Skip to content

RL API

RL APIs cover baseline learners, shared training data structures, policy adapters, and model export helpers.

Shared Training Helpers

src.rl.base.AgentSpec dataclass

AgentSpec(
    agent_id: str,
    obs_dim: int,
    action_dim: int,
    action_kind: str,
    action_nvec: tuple[int, ...],
    action_shape: tuple[int, ...],
    action_start: tuple[int, ...],
)

src.rl.base.TrainResult dataclass

TrainResult(
    history: list[Any],
    metrics: dict[str, list[dict[str, float]]],
    policies: dict[str, "PolicyAdapter"],
    agent_states: dict[str, Any],
)

src.rl.base.EpisodicCostTracker dataclass

EpisodicCostTracker(
    gamma: float = 1.0,
    initial_return: float = 0.0,
    running_return: float = 0.0,
    running_mass: float = 0.0,
    discount: float = 1.0,
    last_return: float | None = None,
    last_scale: float = 1.0,
)

Track returns from true episode starts across rollout boundaries.

src.rl.base.PolicyAdapter dataclass

PolicyAdapter(
    agent_id: str,
    params: Any,
    apply_fn: Callable[..., Any],
    obs_dim: int,
    action_kind: str = "discrete",
    action_nvec: tuple[int, ...] = (),
    action_shape: tuple[int, ...] = (),
    action_start: tuple[int, ...] = (),
)

component_action_probs

component_action_probs(obs: Any) -> tuple[np.ndarray, ...]

Return one categorical distribution per discrete action component.

src.rl.base.merge_configs

merge_configs(
    defaults: Mapping[str, Any],
    overrides: Mapping[str, Any] | None,
) -> dict[str, Any]

src.rl.base.infer_agent_specs

infer_agent_specs(env: ParallelEnv) -> list[AgentSpec]

src.rl.base.compute_gae

compute_gae(
    rewards: ndarray,
    bootstrap_dones: ndarray,
    episode_ends: ndarray,
    values: ndarray,
    next_values: ndarray,
    gamma: float,
    gae_lambda: float,
) -> tuple[jnp.ndarray, jnp.ndarray]

src.rl.base.normalize_advantages

normalize_advantages(advantages: ndarray) -> jnp.ndarray

src.rl.base.num_optimizer_steps_for_timesteps

num_optimizer_steps_for_timesteps(
    timesteps: int,
    rollout_length: int,
    learning_epochs: int,
    requested_mini_batches: int,
) -> int

Return the exact number of optimizer applications for a training run.

Algorithms

src.rl.ippo.IPPO

IPPO(cfg: dict[str, Any] | None = None, *, seed: int = 0)

src.rl.ippo.ActorCritic

Bases: Module

src.rl.ippo.validate_ippo_config

validate_ippo_config(
    overrides: dict[str, Any] | None = None,
) -> dict[str, Any]

Merge and validate IPPO configuration, rejecting misspelled options.

src.rl.parameterized_ippo.ParameterizedIPPO

ParameterizedIPPO(
    cfg: dict[str, Any] | None = None, *, seed: int = 0
)

IPPO over factored categorical/raw-Gaussian augmented shield actions.

src.rl.parameterized_ippo.ParameterizedActorCritic

Bases: Module

src.rl.parameterized_ippo.ParameterizedPolicyAdapter dataclass

ParameterizedPolicyAdapter(
    agent_id: str,
    params: Any,
    apply_fn: Callable[..., Any],
    obs_dim: int,
    action_kind: str = "discrete",
    action_nvec: tuple[int, ...] = (),
    action_shape: tuple[int, ...] = (),
    action_start: tuple[int, ...] = (),
    continuous_dim: int = 0,
)

Bases: PolicyAdapter

Inference adapter that preserves the complete learned augmented action.

src.rl.parameterized_ippo.validate_parameterized_ippo_config

validate_parameterized_ippo_config(
    overrides: dict[str, Any] | None = None,
) -> dict[str, Any]

src.rl.ippo_lag.IPPO_Lagrangian

IPPO_Lagrangian(cfg: dict[str, Any], *, seed: int = 0)

src.rl.ippo_lag.ActorCriticWithCost

Bases: Module

src.rl.ICPO.ICPO

ICPO(cfg: dict[str, Any], *, seed: int = 0)

src.rl.ICPO.CPO

CPO(cfg: dict[str, Any] | None = None, *, seed: int = 0)

Bases: ICPO

Plain CPO preset without ICPO's learned cost shaping extension.

src.rl.ICPO.CategoricalActor

CategoricalActor(
    obs_dim: int,
    action_dim: int,
    *,
    hidden_sizes: tuple[int, ...],
    activation: str,
)

Bases: Module

src.rl.ICPO.ValueNetwork

ValueNetwork(
    obs_dim: int,
    *,
    hidden_sizes: tuple[int, ...],
    activation: str,
)

Bases: Module

Model Export

src.rl.model_io.save_models

save_models(
    env: Any,
    models: Mapping[Any, Mapping[str, Module]],
    folder: str | Path,
) -> None