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 ¶
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.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.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.ActorCritic ¶
Bases: Module
src.rl.ippo.validate_ippo_config ¶
Merge and validate IPPO configuration, rejecting misspelled options.
src.rl.parameterized_ippo.ParameterizedIPPO ¶
IPPO over factored categorical/raw-Gaussian augmented shield actions.
src.rl.parameterized_ippo.ParameterizedActorCritic ¶
Bases: Module
src.rl.parameterized_ippo.ParameterizedPolicyAdapter
dataclass
¶
src.rl.parameterized_ippo.validate_parameterized_ippo_config ¶
src.rl.ippo_lag.ActorCriticWithCost ¶
Bases: Module
src.rl.ICPO.CPO ¶
src.rl.ICPO.CategoricalActor ¶
CategoricalActor(
obs_dim: int,
action_dim: int,
*,
hidden_sizes: tuple[int, ...],
activation: str,
)
Bases: Module
src.rl.ICPO.ValueNetwork ¶
Bases: Module