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.