jaxdem.rl#
JaxDEM reinforcement learning (RL) module. It provides models, environments, and trainers with RL algorithms such as PPO.
- class jaxdem.rl.ActionSpace#
Bases:
FactoryRegistry/namespace for action-space constraints implemented as
distrax.Bijectorobjects.Wrap these bijectors around a base policy distribution (e.g.,
MultivariateNormalDiag) withdistrax.Transformed. The bijector’s ‘forward_and_log_det’ / ‘inverse_and_log_det’ methods then adjust sampling and log-probabilities correctly. See the Distrax/TFP bijector interface for details on shape semantics and ‘event_ndims_in/out’.Example:#
To define a custom action space, inherit from
distrax.BijectorandActionSpaceand implement its abstract methods:>>> @ActionSpace.register("myCustomActionSpace") >>> class MyCustomActionSpace(distrax.Bijector, ActionSpace): ...
- class jaxdem.rl.Environment(state: State, system: System, env_params: dict[str, Any])#
Bases:
Factory,ABCDefines the interface for reinforcement-learning environments.
Let A be the number of agents (A ≥ 1). Single-agent environments still use A=1.
The environment flattens observations and actions per agent to fixed sizes. Use
action_space_shapeto reshape inside the environment if needed.
Required shapes
Observation:
(A, observation_space_size)Action (input to
step()):(A, action_space_size)Reward:
(A,)Done: scalar boolean for the whole environment
Todo: - Truncated data field: per-agent termination flag - Render method
Example:#
To define a custom environment, inherit from
Environment. Implement the abstract methods:>>> @Environment.register("MyCustomEnv") >>> @jax.tree_util.register_dataclass >>> @dataclass(slots=True) >>> class MyCustomEnv(Environment): ...
- env_params: dict[str, Any]#
Environment-specific parameters.
- classmethod Create(dim: int = 2) Environment[source]#
- static reset(env: Environment, key: Array | ndarray | bool | number | bool | int | float | complex) Environment[source]#
Initialize the environment to a valid start state.
- Parameters:
env ('MyCustomEnv') – The current environment.
key (jax.random.PRNGKey) – JAX random number generator key.
- Returns:
The initialized environment.
- Return type:
- static reset_if_done(env: Environment, done: Array, key: Array | ndarray | bool | number | bool | int | float | complex) Environment[source]#
Reset the environment when
doneis True.When
doneis True, this method calls the environment’sresetmethod. Whendoneis False, it returns the current environment unchanged.- Parameters:
env (Environment) – The current environment.
done (jax.Array) – A boolean flag that is True when the environment has reached a terminal state.
key (jax.random.PRNGKey) – JAX random number generator key used for the reset.
- Returns:
The reset environment when
doneis True, otherwise the unchanged environment.- Return type:
- static step(env: Environment, action: Array) Environment[source]#
Advance the simulation by one step using per-agent actions.
- Parameters:
env (Environment) – The current environment.
action (jax.Array) – The per-agent action vectors.
- Returns:
The updated environment state.
- Return type:
- static observation(env: Environment) Array[source]#
Return the per-agent observation vector.
- Parameters:
env (Environment) – The current environment.
- Returns:
Per-agent observations, shape
(A, observation_space_size).- Return type:
jax.Array
- static reward(env: Environment) Array[source]#
Return the per-agent immediate rewards.
- Parameters:
env (Environment) – The current environment.
- Returns:
Per-agent rewards for the current state, shape
(A,).- Return type:
jax.Array
- static done(env: Environment) Array[source]#
Return whether the episode has ended.
- Parameters:
env (Environment) – The current environment.
- Returns:
A bool that is True when the episode has ended.
- Return type:
jax.Array
- static info(env: Environment) dict[str, Any][source]#
Return auxiliary diagnostic information.
The default is an empty dict. Subclasses can override this method to provide environment-specific information.
- Parameters:
env (Environment) – The current environment.
- Returns:
A dictionary with additional information about the environment.
- Return type:
Dict
- property action_space_size: int[source]#
Flattened action size per agent. Actions passed to
step()have shape(A, action_space_size).
- property action_space_shape: tuple[int][source]#
Original per-agent action shape (useful for reshaping inside the environment).
- property observation_space_size: int[source]#
Flattened observation size per agent.
observation()returns shape(A, observation_space_size).
- class jaxdem.rl.Model(*args: Any, **kwargs: Any)#
Bases:
Factory,Module,ABCBase interface for reinforcement learning models. Acts as a namespace.
Models map observations to an action distribution and a value estimate.
Example:#
To define a custom model, inherit from
Modeland implement its abstract methods:>>> @Model.register("myCustomModel") >>> class MyCustomModel(Model): ...
- discrete: bool#
Whether the model emits a categorical (discrete) action distribution. Subclasses set this in
__init__before calling_init_policy_head().
- class jaxdem.rl.Trainer(env: Environment, graphdef: nnx.GraphDef[Any], graphstate: nnx.GraphState, key: ArrayLike, advantage_gamma: jax.Array, advantage_lambda: jax.Array, advantage_rho_clip: jax.Array, advantage_c_clip: jax.Array)#
Bases:
Factory,ABCBase class for reinforcement learning trainers.
This class holds the environment and model state (Flax NNX GraphDef/GraphState). It provides rollout utilities (
step(),trajectory_rollout()) and a general advantage method (compute_advantages()). Subclasses must implement algorithm-specific training logic inepoch().Example:#
To define a custom trainer, inherit from
Trainerand implement its abstract methods:>>> @Trainer.register("myCustomTrainer") >>> @jax.tree_util.register_dataclass >>> @dataclass(slots=True) >>> class MyCustomTrainer(Trainer): ...
- env: Environment#
Environment object.
- graphdef: nnx.GraphDef[Any]#
Static graph definition of the model/optimizer.
- graphstate: nnx.GraphState#
Mutable state (parameters, optimizer state, RNGs, etc.).
- key: ArrayLike#
PRNGKey used to sample actions and for other stochastic operations.
- advantage_gamma: jax.Array#
Discount factor \(\gamma \in [0, 1]\).
- advantage_lambda: jax.Array#
Generalized Advantage Estimation parameter \(\lambda \in [0, 1]\).
- advantage_rho_clip: jax.Array#
V-trace \(\bar{\rho}\) (importance weight clip for the TD term).
- advantage_c_clip: jax.Array#
V-trace \(\bar{c}\) (importance weight clip for the recursion/trace term).
- static step(env: Environment, graphdef: nnx.GraphDef[Any], graphstate: nnx.GraphState, key: jax.Array, skip_frames: int = 0) tuple[tuple[Environment, nnx.GraphState, jax.Array], TrajectoryData][source]#
Take one environment step and record a single-step trajectory. Repeats the action for
skip_framesextra frames when set.- Parameters:
env (Environment) – The (vectorized) environment to step.
graphdef (nnx.GraphDef) – Python part of the nnx model.
graphstate (nnx.GraphState) – State of the nnx model.
key (jax.Array) – Jax random key.
skip_frames (int) – Number of additional frames to repeat the action.
- Returns:
Updated state and the new single-step trajectory. Trajectory data is shaped (N_envs, N_agents, …).
- Return type:
Tuple[Tuple[Environment, nnx.GraphState, jax.Array], TrajectoryData]
Notes
This method does not reset finished environments. Under
vmap, a per-step conditional reset would run the reset computation for every environment on every step. Instead, the method marks terminal framesdone, which masks them in the advantage calculation, and the trainer resets the environments once per epoch, before the rollout. The method does clear the recurrent carry of done environments each step with a cheap masked zeroing, so episodes do not bleed into each other.
- static trajectory_rollout(env: Environment, graphdef: nnx.GraphDef[Any], graphstate: nnx.GraphState, key: jax.Array, num_steps_epoch: int, unroll: int = 8, skip_frames: int = 0) tuple[Environment, nnx.GraphState, jax.Array, TrajectoryData][source]#
Roll out \(T = \text{num\_steps\_epoch}\) environment steps using
jax.lax.scan().- Parameters:
env (Environment) – The (vectorized) environment to roll out.
graphdef (nnx.GraphDef) – Python part of the nnx model.
graphstate (nnx.GraphState) – State of the nnx model.
key (jax.Array) – Jax random key.
num_steps_epoch (int) – Number of steps to roll out.
unroll (int) – Number of loop iterations to unroll for compilation speed.
skip_frames (int) – Number of frames to skip (repeat action) per observation.
- Returns:
The final environment, graph state, and PRNG key, plus a
TrajectoryDatainstance whose fields are stacked along time (leading dimension \(T = \text{num_steps_epoch}\)).- Return type:
Tuple[Environment, nnx.GraphState, jax.Array, TrajectoryData]
- static compute_advantages(value: Array, reward: Array, ratio: Array, done: Array, advantage_rho_clip: Array, advantage_c_clip: Array, advantage_gamma: Array, advantage_lambda: Array, last_value: Array | None = None, unroll: int = 8) tuple[Array, Array][source]#
Compute V-trace/GAE advantages and return targets.
Given a policy \(\pi\), define per-step importance ratios:
\[\rho_t = \exp\big( \log \pi_\theta(a_t \mid s_t) - \log \pi_{\theta_\text{old}}(a_t \mid s_t) \big)\]and their clipped versions \(\hat{\rho}, \hat{c}\):
\[\hat{\rho}_t = \min(\rho_t, \bar{\rho}), \quad \hat{c}_t = \min(\rho_t, \bar{c}).\]We form a TD-like residual with an off-policy correction:
\[\delta_t = \hat{\rho}_t \big( r_t + \gamma V(s_{t+1})(1 - \text{done}_t) - V(s_t) \big)\]and propagate a GAE-style trace using \(\hat{c}_t\):
\[A_t = \delta_t + \gamma \lambda (1 - \text{done}_t) \hat{c}_t A_{t+1}\]Finally, the return targets are:
\[\text{returns}_t = A_t + V(s_t)\]Notes
When \(\pi_\theta = \pi_{\theta_\text{old}}\) (i.e.
ratio==1) and \(\bar{\rho} = \bar{c} = 1\), this function reduces to standard GAE.- Parameters:
last_value (jax.Array | None) – Bootstrap value \(V(s_T)\) evaluated on the post-rollout observation. If
None, falls back tovalue[-1](i.e. \(V(s_{T-1})\)), which biases the advantage of the last transition; callers should provide it whenever possible.- Returns:
Computed advantage and returns.
- Return type:
Tuple[jax.Array, jax.Array]
References
Schulman et al., High-Dimensional Continuous Control Using Generalized Advantage Estimation, 2015/2016
Espeholt et al., IMPALA: Scalable Distributed Deep-RL with Importance Weighted Actor-Learner Architectures, 2018
- class jaxdem.rl.TrajectoryData(*, obs: Array, action: Array, value: Array, log_prob: Array, ratio: Array, reward: Array, done: Array)#
Bases:
objectContainer for rollout data (single step or stacked across time).
- obs: Array#
Observations.
- action: Array#
Actions sampled from the policy.
- value: Array#
Baseline value estimates \(V(s_t)\).
- log_prob: Array#
Behavior-policy log-probabilities \(\log \pi_b(a_t \mid s_t)\) at collection time.
- ratio: Array#
Probability ratio between the current policy and the old policy: \(\exp\big( \log \pi_\theta(a_t \mid s_t) - \log \pi_{\theta_\text{old}}(a_t \mid s_t) \big)\).
- reward: Array#
Immediate rewards \(r_t\).
- done: Array#
Episode-termination flags (boolean).
- jaxdem.rl.clip_action_env(env: Environment, min_val: float = -1.0, max_val: float = 1.0) Environment[source]#
Wrap an environment so that its step method clips the action to [min_val, max_val] before calling the original step.
- jaxdem.rl.is_wrapped(env: Environment) bool[source]#
Check whether an environment instance is a wrapped environment.
- Parameters:
env (Environment) – The environment instance to check.
- Returns:
True if the environment is wrapped, False otherwise.
- Return type:
bool
- jaxdem.rl.unwrap(env: Environment) Environment[source]#
Unwrap an environment to its original base class. Keeps all current field values.
- Parameters:
env (Environment) – The wrapped environment instance.
- Returns:
A new instance of the original base environment class with the same field values as the wrapped instance.
- Return type:
- jaxdem.rl.vectorise_env(env: Environment, n: int | None = None) Environment[source]#
Promote an environment instance to a parallel version by applying jax.vmap(…) to its static methods.
- Parameters:
env (Environment) – The environment to vectorize. May already carry a leading batch dimension (e.g. produced with
jax.vmap).n (int, optional) – When given, the wrapper first broadcasts the (scalar) environment to a batch of
nidentical copies, so callers do not need to writejax.vmap(lambda _: env)(jnp.arange(n))themselves. Then callenv.reset(env, keys)to randomize each copy.
Example
>>> env = vectorise_env(env, n=32) >>> env = env.reset(env, jax.random.split(key, 32))
Modules
Interface for bijectors that constrain the policy probability distribution. |
|
Wrappers that modify RL environments. |
|
Reinforcement-learning environment interface. |
|
Interface for defining reinforcement learning models. |
|
Interface for defining reinforcement learning model trainers. |