GPI-Linear Support (Jax)¶
- class morl_baselines.multi_policy.gpi_ls_jax.gpi_ls_jax.GPILS(env, learning_rate: float = 0.0003, initial_epsilon: float = 0.01, final_epsilon: float = 0.01, epsilon_decay_steps: int = None, target_net_update_freq: int = 1000, buffer_size: int = 1000000, net_arch: List = [256, 256, 256, 256], num_nets: int = 2, batch_size: int = 128, learning_starts: int = 100, gradient_updates: int = 20, gamma: float = 0.99, use_gpi: bool = True, gpi_type: str = 'gpi', pessimism: float = 0.99, per: bool = False, alpha_per: float = 0.6, min_priority: float = 0.01, drop_rate: float = 0.01, layer_norm: bool = True, project_name: str = 'MORL-Baselines', experiment_name: str = 'GPI-LS - Jax', wandb_entity: str | None = None, log: bool = True, seed: int | None = None)¶
GPI-LS Algorithm in Jax.
Alegre, L.N., Bazzan, A.L.C., Roijers, D.M. et al. Generalized policy improvement for efficient and robust multi-objective reinforcement learning. Autonomous Agents and Multi-Agent Systems 40, 12 (2026). https://doi.org/10.1007/s10458-026-09736-w
Initialize the GPI-LS algorithm.
- Parameters:
env – The environment to learn from.
learning_rate – The learning rate.
initial_epsilon – The initial epsilon value.
final_epsilon – The final epsilon value.
epsilon_decay_steps – The number of steps to decay epsilon.
target_net_update_freq – The target network update frequency.
buffer_size – The size of the replay buffer.
net_arch – The network architecture.
num_nets – The number of networks.
batch_size – The batch size.
learning_starts – The number of steps before learning starts.
gradient_updates – The number of gradient updates per step.
gamma – The discount factor.
use_gpi – Whether to use GPI.
gpi_type – “gpi” or “ugpi” for uncertainty-aware GPI.
pessimism – Pessimism level when using ugpi.
per – Whether to use PER.
alpha_per – The alpha parameter for PER.
min_priority – The minimum priority for PER.
drop_rate – The dropout rate.
layer_norm – Whether to use layer normalization.
project_name – The name of the project.
experiment_name – The name of the experiment.
wandb_entity – The name of the wandb entity.
log – Whether to log.
seed – The seed for random number generators.
- eval(obs: ndarray, w: ndarray) int¶
Evaluate the policy.
- get_config()¶
Return the configuration of the agent.
- static gpi_action(q_net: VectorQNetwork, q_state: TrainState, obs: Array, w: Array, M: Array, key: PRNGKey) Array¶
Generalized Policy Improvement (GPI).
- load(path: str, step: int | None = None)¶
Load the model parameters.
- static max_action(q_net: VectorQNetwork, q_state: TrainState, obs: Array, w: Array, key: PRNGKey) int¶
Select the action with the maximum Q-value.
- static one_update(q_net: VectorQNetwork, q_state: TrainState, weight: Array, weight_support: Array, data: ReplayBufferSamplesNp, gamma: float, min_priority: float, gradient_updates: int, key: PRNGKey) Tuple[TrainState, Array, Array, Array, PRNGKey]¶
Perform a single update step.
- save(save_dir='weights/', filename=None)¶
Save the model parameters.
- set_weight_support(M: List[ndarray])¶
Set the weight support set.
- train(total_timesteps: int, eval_env, ref_point: ndarray, known_pareto_front: List[ndarray] | None = None, num_eval_weights_for_front: int = 100, num_eval_episodes_for_front: int = 5, num_eval_weights_for_eval: int = 50, timesteps_per_iter: int = 10000, weight_selection_algo: str = 'gpi-ls', eval_freq: int = 1000, eval_mo_freq: int = 10000, checkpoints: bool = True)¶
Train agent.
- Parameters:
total_timesteps (int) – Number of timesteps to train for.
eval_env (gym.Env) – Environment to evaluate on.
ref_point (np.ndarray) – Reference point for hypervolume calculation.
known_pareto_front (Optional[List[np.ndarray]]) – Optimal Pareto front if known.
num_eval_weights_for_front – Number of weights to evaluate for the Pareto front.
num_eval_episodes_for_front – number of episodes to run when evaluating the policy.
num_eval_weights_for_eval (int) – Number of weights use when evaluating the Pareto front, e.g., for computing expected utility.
timesteps_per_iter (int) – Number of timesteps to train for per iteration.
weight_selection_algo (str) – Weight selection algorithm to use.
eval_freq (int) – Number of timesteps between evaluations.
eval_mo_freq (int) – Number of timesteps between multi-objective evaluations.
checkpoints (bool) – Whether to save checkpoints.
- train_iteration(total_timesteps: int, weight: ndarray, weight_support: List[ndarray], change_w_every_episode: bool = True, reset_num_timesteps: bool = True, eval_env: Env | None = None, eval_freq: int = 1000, reset_learning_starts: bool = False)¶
Train the agent for one iteration.
- Parameters:
total_timesteps (int) – Number of timesteps to train for
weight (np.ndarray) – Weight vector
weight_support (List[np.ndarray]) – Weight support set
change_w_every_episode (bool) – Whether to change the weight vector at the end of each episode
reset_num_timesteps (bool) – Whether to reset the number of timesteps
eval_env (Optional[gym.Env]) – Environment to evaluate on
eval_freq (int) – Number of timesteps between evaluations
reset_learning_starts (bool) – Whether to reset the learning starts
- static ugpi_action(q_net: VectorQNetwork, q_state: TrainState, obs: Array, w: Array, M: Array, pessimism: float, key: PRNGKey) Array¶
Uncertainty-Aware GPI (uGPI).
- update(weight: Array)¶
Update the parameters of the networks.