We implement PPO fully in JAX, following the algorithm described in PPO. Our implementation is single-process, but it is efficiently vectorized using JAX's vmap for batching.
In our implementation, we took inspiration from Brax simulator and it's implementation of PPO and environments. In particular, we use:
- Similar datatypes for
State,Transition, etc. See src/training/types.py and src/envs/types.py for more details. - Similar training loop and environment logic. See src/training/ppo.py and src/envs/env.py for more details.
- Same observation normalization code. See src/training/running_statistics.py for more details.
- Same code for parametric distribution. See src/training/distribution.py for more details.
However, there are some differences:
- We reimplement the wrappers for vectorization, and truncation and termination handling. This also affects GAE computations. See src/envs/wrappers.py for more details.
- We use different metrics collection logic.
- We do not use Brax's simulation or MuJoCo directly and implement our own simulation.
Below we describe some of the most important parts of the implementation in more detail. Some parts of this document are based on the The 37 Implementation Details of Proximal Policy Optimization blog post, which we highly recommend to read. For example, vectorized approach to training is described there in "Implementation detail 1", but here we provide a graphical explanation to unrolls, episodes and truncation handling.
We start with this topic because we found it to be the most complex part of the implementation. This is because proper handling of truncation and termination is important in GAE computation, environment resets and metrics computation.
Suppose that we have a simple case of 1 environment, where we perform environment steps until the episode naturally terminates at step
So, we perform a trajectory unroll until termination:
In code, we use Transition type to represent the transition between two states.
new_state = env.step(action)
transition = Transition(
observation=state.obs,
action=action,
reward=new_state.reward,
discount=1.0 - new_state.done,
next_observation=new_state.obs,
)The following logic e.g. for GAE computation uses Transition object, where each element inside stores a batch of values for the whole trajectory. So, it is important that reward in the Transition object corresponds to "reward for the transition to the next state".
Now, let's consider a case of
If we were to implement this like shown above, we would have to do so in a loop, because each trajectory ends at different time steps
The first logical step from here is to somehow ensure that environment will reach natural termination at some time step
We have two problems here:
- Environment may not reach natural termination at
$T_{\text{max}}$ . - This approach cannot be properly vectorized, because some environments may terminate earlier, and here we do not store any information after environment terminates, so we cannot transition data in a single tensor.
We will address the first problem later (this is the truncation case). To address the second problem, we will continue to step terminated environments until
Note that in code, we mask out the transitions (and particular attributes of Transition object), instead of the states.
This concludes the description of the parallel environments case with the assumption that environment will reach natural termination at some time step
Let's make a step further in our discussion. For now, we still assume that environment will reach natural termination at some time step
In RL, environments are different in terms of how long episodes are. For example, we may have a game or simulation with small number of steps (e.g.
- Too short episodes at the start of training.
- Too long episodes at the end of training.
In general, suppose we have 1M steps in a single episode. This limits our capabilities for vectorized training, because we have to:
- Use smaller batch sizes.
- Collect the trajectories in full (until natural termination).
The second point means that:
- We will update the policy less frequently.
- We will need to store more data in memory.
To address this, we can use the concept of episodes and unrolls:
-
Unroll is a trajectory of a single environment of fixed length
$U$ . -
Episode is a trajectory of a single environment from the start to the maximum length
$T_{\text{max}}$ .
So, the unroll is a part of the episode. Consider a case of parallel environments:
Here we have
- Starting state. This state represents that the environment is reset when the unroll starts.
- Ordinary state. This is just an intermediate state of the unroll.
- Terminal state. This state represents that the environment has naturally terminated.
- Masked state. This state represents that the environment has reached terminal state in this unroll, and the information from these states is masked out in subsequent computations.
- Last state. This state represents that the unroll has reached the maximum length
$U$ without termination.
Note
If the environment reached "last" state, this state will be used as a starting state for the next unroll, without resetting the environment. This is useful because it allows the agent to explore the environment longer.
Important
Here we still assume that the episode will not end in "last" state, so there is no truncation.
Wrappers are applied to the environment to add functionality. For the code to work correctly, we require the wrappers below to be used with wrap method.
To work with parallel environments, we use VmapWrapper, which reimplements reset and step methods by applying jax.vmap to corresponding unwrapped environment methods. The only thing that we need to ensure is proper initialization shape of the observation in order to create a policy and value networks. This allows to use object of class Env as a handler for several parallel environments without the need to vmap any other methods.
Note
This wrapper was adapted from Brax implementation.
However, we additionally implement reset_done method, which is used to reset environments that have reached terminal state.
This is done deliberately, because Brax resets environments in the middle of the unroll, which may mix two "true" unrolls together.
To work with episodes, we implement EpisodeWrapper, which maintains episode metrics and sets truncation and termination signals. Episode wrapper has to be used, otherwise the code for advantage estimation and trajectory unrolling will not work correctly. Here we will describe the wrapper in full, with truncation signal, which is described below.
EpisodeWrapper records the following:
- Total number of steps made in the episode.
- Number of steps in the episode until termination or truncation.
- Termination signal.
- Truncation signal.
- Episode done signal.
- Post-done signal.
- Total reward in the episode until post done signal.
- Total episode metrics, until post done signal.
Truncation is the case of unnatural termination, when the episode is terminated by EpisodeWrapper (in our case), but not by the environment. It is useful to truncate the episode for it not to progress for large number of steps.
Below is the illustration of the wrapper in action. Here we only show unrolls for a single environment until truncation.
By definition, we set:
- Termination signal to 1.0 when natural termination occurs.
- Truncation signal to 1.0 when env reaches episode length defined by
EpisodeWrapper. - Post-done signal to 1.0 for all transitions after episode is done (truncated or terminated).
Important
Unroll is a part of the episode, but each unroll is independent, because weight optimization happens after each unroll.
We may reuse the last state of the unroll as a starting state for the next unroll, but if environment is reset, all variables stored by EpisodeWrapper are also reset. This means that in case of parallel environments, one episode may end in environment
Advantage estimation can be done using different methods. Here we use Generalized Advantage Estimation (GAE), described in PPO.
We use the following formula for GAE:
where
Our implementation of GAE computation is based on the code from Brax,
but due to a different EpisodeWrapper implementation, we have a slightly different notation.
Notice that:
- Episode length
$T$ is different in all environments, therefore we need to use a mask to account for that. - In case terminal state is reached at
$s_T$ , we do include the reward of the last transition, but we don't use the value estimate$s_{T+1}$ , because there is no valid next state (there is, but only because we continue to step the environment after termination). - In case episode is truncated at
$s_T$ , we need to bootstrap the value, i.e. use value estimate$V_{\phi}(s_{T+1})$ , because there should be a next state, we just truncated it. - Ordinary states are treated in a same way as truncated states, meaning that we use value estimate
$V_{\phi}(s_{t+1})$ for all$t < T$ .
From these observations, we can derive the method for GAE computation:
- Compute the TD error
$\delta_t$ for all time steps. - In TD computation, use termination mask to exclude value estimate
$V_{\phi}(s_{T+1})$ . - Do not use the truncation mask, since truncated state is treated as an ordinary state.
- Apply a post done mask to
$\delta_t$ to exclude TD error for masked states. - Compute advantage estimates in reverse order without masking, because masked states are already excluded from
$\delta_t$ .
See src/training/loss.py for the implementation.
We use the following hyperparameters:
- Number of environments (batch size):
$B = 2048$ - Number of minibatches:
$N_{\text{minibatches}} = 8$ - Number of unroll steps:
$U = 16$ - Episode length:
$T_{\text{max}} = 200$ - Number of optimizer steps per batch:
$K = 2$ - Policy and value learning rate:
$\eta = 1e-2$ , with one-cycle learning rate schedule. - Discount factor:
$\gamma = 0.99$ - GAE parameter:
$\lambda = 0.95$ - Clip parameter:
$\epsilon = 0.2$ - Value function coefficient:
$\alpha = 0.5$ - Entropy coefficient:
$\beta = 0.01$ - Number of training epochs is calculated from the total number of environment steps we want to train for.
- Policy and value max gradient norm:
$0.5$ . - Number of layers in policy network: 4.
- Policy hidden size: 128.
- Number of layers in value network: 5.
- Value hidden size: 512.
Regarding the architecture, we follow the recommendations from [1] and use orthogonal weight initialization. But, we use different number of layers, hidden sizes, weight initialization sigmas and different activation functions. In particular, we use swish instead of tanh as we found the training to be more stable.
Some observations in our experiments:
- Without entropy loss, the agent does not reach the target.
-
$K = 2$ works better than$K = 1, 4, 8$ . -
$U = 16$ (with usual number of steps to reach the target around 30) works better than$U = 32$ and higher values. - Observation and reward normalization increases the performance of the agent.
Important
We also found that even though the agent is able to reach the target in some episodes, usually it selects a suboptimal strategy of moving for all allowed time steps
This solution may be suboptimal and highly dependent on particular environment and reward scaling, but it serves our use case.



