Skip to content

Latest commit

 

History

History
229 lines (153 loc) · 13.9 KB

File metadata and controls

229 lines (153 loc) · 13.9 KB

Implementation details

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:

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.

Handling truncation and termination

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.

Single environment

Suppose that we have a simple case of 1 environment, where we perform environment steps until the episode naturally terminates at step $T$.

Simple transition

So, we perform a trajectory unroll until termination:

$$s_0, a_0, r_1, s_1, a_1, r_2, \ldots, s_{T-1}, a_{T-1}, r_T, s_T.$$

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".

Parallel environments

Now, let's consider a case of $B$ parallel environments.

Parallel transitions

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 $T_i, ~i = 1, \ldots, B$, which is not efficient.

The first logical step from here is to somehow ensure that environment will reach natural termination at some time step $T_{\text{max}}$ (which may be done by environment itself).

Parallel transitions with max length

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 $T_{\text{max}}$ and store the data in a single tensor. We will assume that environment always reaches natural termination at step $T \leq T_{\text{max}}$. Later, to correctly account for the termination at different time steps, we will use a mask.

Parallel transitions

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 $T_{\text{max}}$.

Unrolls and episodes

Let's make a step further in our discussion. For now, we still assume that environment will reach natural termination at some time step $T \leq T_{\text{max}}$.

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. $T_{\text{max}} = 100$), or a long-running simulation (e.g. 1000000 steps). We won't have any problems with the first case using approach to parallel environments described above. However, in the second case, we may have:

  • 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:

  1. Use smaller batch sizes.
  2. Collect the trajectories in full (until natural termination).

The second point means that:

  1. We will update the policy less frequently.
  2. 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:

Parallel unrolls

Here we have $B$ parallel environments, and each env performs $N$ unrolls of length $U$. We define several types of states:

  • 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

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.

Episode wrapper

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 $n$, and another may start in environment $k$.

Advantage estimation

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:

$$\hat{A}_t = \delta_t + (\gamma \lambda) \delta_{t+1} + \cdots + (\gamma \lambda)^{T-t} \delta_{T-1},$$

where $\delta_t = r_t + \gamma V_{\phi}(s_{t+1}) - V_{\phi}(s_t)$ is the TD error. $T$ is the length of the episode.

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:

  1. Episode length $T$ is different in all environments, therefore we need to use a mask to account for that.
  2. 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).
  3. 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.
  4. 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:

  1. Compute the TD error $\delta_t$ for all time steps.
  2. In TD computation, use termination mask to exclude value estimate $V_{\phi}(s_{T+1})$.
  3. Do not use the truncation mask, since truncated state is treated as an ordinary state.
  4. Apply a post done mask to $\delta_t$ to exclude TD error for masked states.
  5. 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.

Training

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 $T_{\text{max}}$, without going in penalizing terminal states (collision and out of bounds). This is not what we want, because our goal is to reach a target rather than train the agent to fly a drone. Therefore, we propose a solution to apply a penalty of 50.0 to the reward at truncated state, so the agent is discouraged from moving for all allowed time steps $T_{\text{max}}$.

This solution may be suboptimal and highly dependent on particular environment and reward scaling, but it serves our use case.

References

  1. The 37 Implementation Details of Proximal Policy Optimization