Image Augmentation Is All You Need: Regularizing Deep Reinforcement Learning from Pixels
Ilya Kostrikov, Denis Yarats, Rob Fergus
Introduction
Sample-efficient deep reinforcement learning (RL) algorithms capable of directly training from image pixels would open up many real-world applications in control and robotics. However, simultaneously training a convolutional encoder alongside a policy network is challenging when given limited environment interaction, strong correlation between samples and a typically sparse reward signal. Naive attempts to use a large capacity encoder result in severe over-fitting (see Figure 1a) and smaller encoders produce impoverished representations that limit task performance.
Limited supervision is a common problem across AI and a number of approaches are adopted: (i) pre-training with self-supervised learning (SSL), followed by standard supervised learning; (ii) supervised learning with an additional auxiliary loss and (iii) supervised learning with data augmentation. SSL approaches are highly effective in the large data regime, e.g. in domains such as vision and NLP where large (unlabeled) datasets are readily available. However, in sample-efficient RL, training data is more limited due to restricted interaction between the agent and the environment, resulting in only – transitions from a few hundred trajectories. While there are concurrent efforts exploring SSL in the RL context , in this paper we take a different approach, focusing on data augmentation.
A wide range of auxiliary loss functions have been proposed to augment supervised objectives, e.g. weight regularization, noise injection , or various forms of auto-encoder . In RL, reconstruction objectives or alternate tasks are often used . However, these objectives are unrelated to the task at hand, thus have no guarantee of inducing an appropriate representation for the policy network.
Data augmentation methods have proven highly effective in vision and speech domains, where output-invariant perturbations can easily be applied to the labeled input examples. Surprisingly, data augmentation has received relatively little attention in the RL community, and this is the focus of this paper. The key idea is to use standard image transformations to peturb input observations, as well as regularizing the -function learned by the critic so that different transformations of the same input image have similar -function values. No further modifications to standard actor-critic algorithms are required, obviating the need for additional losses, e.g. based on auto-encoders , dynamics models , or contrastive loss terms .
The paper makes the following contributions: (i) we demonstrate how straightforward image augmentation, applied to pixel observations, greatly reduces over-fitting in sample-efficient RL settings, without requiring any change to the underlying RL algorithm. (ii) exploiting MDP structure, we introduce two simple mechanisms for regularizing the value function which are generally applicable in the context of model-free off-policy RL. (iii) Combined with vanilla SAC and using hyper-parameters fixed across all tasks, the overall approach obtains state-of-the-art performance on the DeepMind control suite . (iv) Combined with a DQN-like agent, the approach also obtains state-of-the-art performance on the Atari 100k benchmark. (v) It is thus the first effective approach able to train directly from pixels without the need for unsupervised auxiliary losses or a world model. (vi) We also provide a PyTorch implementation of the approach combined with SAC and DQN.
Background
Soft Actor-Critic The Soft Actor-Critic (SAC) learns a state-action value function , a stochastic policy and a temperature to find an optimal policy for an MDP by optimizing a -discounted maximum-entropy objective . is used generically to denote the parameters updated through training in each part of the model.
Deep Q-learning DQN also learns a convolutional neural net to approximate Q-function over states and actions. The main difference is that DQN operates on discrete actions spaces, thus the policy can be directly inferred from Q-values. In practice, the standard version of DQN is frequently combined with a set of refinements that improve performance and training stability, commonly known as Rainbow . For simplicity, the rest of the paper describes a generic actor-critic algorithm rather than DQN or SAC in particular. Further background on DQN and SAC can be found in Appendix A.
Sample Efficient Reinforcement Learning from Pixels
This work focuses on the data-efficient regime, seeking to optimize performance given limited environment interaction. In Figure 1a we show a motivating experiment that demonstrates over-fitting to be a significant issue in this scenario. Using three tasks from the DeepMind control suite , SAC is trained with the same policy network architecture but using different image encoder architectures, taken from the following RL approaches: NatureDQN , Dreamer , Impala , SAC-AE (also used in CURL ), and D4PG . The encoders vary significantly in their capacity, with parameter counts ranging from 220k to 2.4M. The curves show that performance decreases as parameter count increases, a clear indication of over-fitting.
A range of successful image augmentation techniques to counter over-fitting have been developed in computer vision . These apply transformations to the input image for which the task labels are invariant, e.g. for object recognition tasks, image flips and rotations do not alter the semantic label. However, tasks in RL differ significantly from those in vision and in many cases the reward would not be preserved by these transformations. We examine several common image transformations from in Appendix E and conclude that random shifts strike a good balance between simplicity and performance, we therefore limit our choice of augmentation to this transformation.
Figure 1b shows the results of this augmentation applied during SAC training. We apply data augmentation only to the images sampled from the replay buffer and not for samples collection procedure. The images from the DeepMind control suite are . We pad each side by pixels (by repeating boundary pixels) and then select a random crop, yielding the original image shifted by pixels. This procedure is repeated every time an image is sampled from the replay buffer. The plots show overfitting is greatly reduced, closing the performance gap between the encoder architectures. These random shifts alone enable SAC to achieve competitive absolute performance, without the need for auxiliary losses.
2 Optimality Invariant Image Transformations
While the image augmentation described above is effective, it does not fully exploit the MDP structure inherent in RL tasks. We now introduce a general framework for regularizing the value function through transformations of the input state. For a given task, we define an optimality invariant state transformation as a mapping that preserves the -values
where are the parameters of , drawn from the set of all possible parameters . One example of such transformations are the random image translations successfully applied in the previous section.
For every state, the transformations allow the generation of several surrogate states with the same -values, thus providing a mechanism to reduce the variance of -function estimation. In particular, for an arbitrary distribution of states and policy , instead of using a single sample , estimation of the following expectation
we can instead generate samples via random transformations and obtain an estimate with lower variance
This suggests two distinct ways to regularize -function. First, we use the data augmentation to compute the target values for every transition tuple as
where corresponds to a transformation parameter of . Then the Q-function is updated using these targets through an SGD update using learning rate
In tandem, we note that the same target from Equation 1 can be used for different augmentations of , resulting in the second regularization approach
When both regularization methods are used, and are drawn independently.
3 Our approach: Data-regularized Q (DrQ)
Our approach, DrQ, is the union of the three separate regularization mechanisms introduced above:
transformations of the input image (Section 3.1).
averaging the target over K image transformations (Equation 1).
averaging the function itself over M image transformations (Equation 3).
Algorithm 1 details how they are incorporated into a generic pixel-based off-policy actor-critic algorithm. If [K=1,M=1] then DrQ reverts to image transformations alone, this makes applying DrQ to any model-free RL algorithm straightforward as it does not require any modifications to the algorithm itself. Note that DrQ [K=1,M=1] also exactly recovers the concurrent work of RAD , up to a particular choice of hyper-parameters and data augmentation type.
For the experiments in this paper, we pair DrQ with SAC and DQN , popular model-free algorithms for control in continuous and discrete action spaces respectively. We select image shifts as the class of image transformations , with , as explained in Section 3.1. For target Q and Q augmentation we use [K=2,M=2] respectively. Figure 2 shows DrQ and ablated versions, demonstrating clear gains over unaugmented SAC. A more extensive ablation can be found in Appendix F.
Experiments
In this section we evaluate our algorithm (DrQ) on the two commonly used benchmarks based on the DeepMind control suite , namely the PlaNet and Dreamer setups. Throughout these experiments all hyper-parameters of the algorithm are kept fixed: the actor and critic neural networks are trained using the Adam optimizer with default parameters and a mini-batch size of . For SAC, the soft target update rate is , initial temperature is , and target network and the actor updates are made every critic updates (as in ). We use the image encoder architecture from SAC-AE and follow their training procedure. The full set of parameters is in Appendix B.
Following , the models are trained using different seeds; for every seed the mean episode returns are computed every environment steps, averaging over episodes. All figures plot the mean performance over the seeds, together with 1 standard deviation shading. We compare our DrQ approach to leading model-free and model-based approaches: PlaNet , SAC-AE , SLAC , CURL and Dreamer . The comparisons use the results provided by the authors of the corresponding papers.
PlaNet Benchmark consists of six challenging control tasks from with different traits. The benchmark specifies a different action-repeat hyper-parameter for each of the six tasksThis means the number of training observations is a fraction of the environment steps (e.g. an episode of steps with action-repeat results in training observations).. Following common practice , we report the performance using true environment steps, thus are invariant to the action-repeat hyper-parameter. Aside from action-repeat, all other hyper-parameters of our algorithm are fixed across the six tasks, using the values previously detailed.
Figure 3 compares DrQ [K=2,M=2] to PlaNet , SAC-AE , CURL , SLAC , and an upper bound performance provided by SAC that directly learns from internal states. We use the version of SLAC that performs one gradient update per an environment step to ensure a fair comparison to other approaches. DrQ achieves state-of-the-art performance on this benchmark on all the tasks, despite being much simpler than other methods. Furthermore, since DrQ does not learn a model or any auxiliary tasks , the wall clock time also compares favorably to the other methods. In Table 1 we also compare performance given at a fixed number of environment interactions (e.g. k and k). Furthermore, in Appendix G we demonstrate that DrQ is robust to significant changes in hyper-parameter settings.
Dreamer Benchmark is a more extensive testbed that was introduced in Dreamer , featuring a diverse set of tasks from the DeepMind control suite. Tasks involving sparse reward were excluded (e.g. Acrobot and Quadruped) since they require modification of SAC to incorporate multi-step returns , which is beyond the scope of this work. We evaluate on the remaining tasks, fixing the action-repeat hyper-parameter to , as in Dreamer .
We compare DrQ [K=2,M=2] to Dreamer and the upper-bound performance of SAC from statesNo other publicly reported results are available for the other methods due to the recency of the Dreamer benchmark.. Again, we keep all the hyper-parameters of our algorithm fixed across all the tasks. In Figure 4, DrQ demonstrates the state-of-the-art results by collectively outperforming Dreamer , although Dreamer is superior on of the tasks (Walker Run, Cartpole Swingup Sparse and Pendulum Swingup). On many tasks DrQ approaches the upper-bound performance of SAC trained directly on states.
2 Atari 100k Experiments
We evaluate DrQ [K=1,M=1] on the recently introduced Atari 100k benchmark – a sample-constrained evaluation for discrete control algorithms. The underlying RL approach to which DrQ is applied is a DQN, combined with double Q-learning , n-step returns , and dueling critic architecture . As per common practice , we evaluate our agent for 125k environment steps at the end of training and average its performance over random seeds. Figure 5 shows the median human-normalized episode returns performance (as in ) of the underlying model, which we refer to as Efficient DQN, in pink. When DrQ is added there is a significant increase in performance (cyan), surpassing OTRainbow and Data Efficient Rainbow . DrQ is also superior to CURL that uses an auxiliary loss built on top of a hybrid between OTRainbow and Efficient rainbow. DrQ combined with Efficient DQN thus achieves state-of-the-art performance, despite being significantly simpler than the other approaches. The experimental setup is detailed in Appendix C and full results can be found in Appendix D.
Related Work
Computer Vision Data augmentation via image transformations has been used to improve generalization since the inception of convolutional networks . Following AlexNet , they have become a standard part of training pipelines. For object classification tasks, the transformations are selected to avoid changing the semantic category, i.e. translations, scales, color shifts, etc. Perturbed versions of input examples are used to expand the training set and no adjustment to the training algorithm is needed. While a similar set of transformations are potentially applicable to control tasks, the RL context does require modifications to be made to the underlying algorithm.
Data augmentation methods have also been used in the context of self-supervised learning. use per-exemplar perturbations in a unsupervised classification framework. More recently, a several approaches have used invariance to imposed image transformations in contrastive learning schemes, producing state-of-the-art results on downstream recognition tasks. By contrast, our scheme addresses control tasks, utilizing different types of invariance.
Generalization between Tasks and Domains A range of datasets have been introduced with the explicit aim of improving generalization in RL through deliberate variation of the scene colors/textures/backgrounds/viewpoints. These include Robot Learning in Homes , Meta-World , the ProcGen benchmark . There are also domain randomization techniques which synthetically apply similar variations, but assume control of the data generation procedure, in contrast to our method. Furthermore, these works address generalization between domains (e.g. synthetic-to-real or different game levels), whereas our work focuses on a single domain and task. In concurrent work, RAD also demonstrates that image augmentation can improve sample efficiency and generalization of RL algorithms. However, RAD represents a specific instantiation of our algorithm when [K=1,M=1] and different image augmentations are used.
Continuous Control from Pixels There are a variety of methods addressing the sample-efficiency of RL algorithms that directly learn from pixels. The most prominent approaches for this can be classified into two groups, model-based and model-free methods. The model-based methods attempt to learn the system dynamics in order to acquire a compact latent representation of high-dimensional observations to later perform policy search . In contrast, the model-free methods either learn the latent representation indirectly by optimizing the RL objective or by employing auxiliary losses that provide additional supervision . Our approach is complementary to these methods and can be combined with them to improve performance.
Conclusion
We have introduced a simple regularization technique that significantly improves the performance of SAC trained directly from image pixels on standard continuous control tasks. Our method is easy to implement and adds a negligible computational burden. We compared our method to state-of-the-art approaches on both DeepMind control suite, where we demonstrated that it outperforms them on the majority of tasks, and Atari 100k benchmarks, where it outperforms other methods in the median metric. Furthermore, we demonstrate the method to be robust to the choice of hyper-parameters.
Acknowledgements
We would like to thank Danijar Hafner, Alex Lee, and Michael Laskin for sharing performance data for the Dreamer and PlaNet , SLAC , and CURL baselines respectively. Furthermore, we would like to thank Roberta Raileanu for helping with the architecture experiments. Finally, we would like to thank Ankesh Anand for helping us finding an error in our evaluation script for the Atari 100k benchmark experiments.
References
Appendix
Appendix A Extended Background
Soft Actor-Critic
The policy evaluation step learns the critic network by optimizing a single-step of the soft Bellman residual
where is a replay buffer of transitions, is an exponential moving average of the weights as done in . SAC uses clipped double-Q learning , which we omit from our notation for simplicity but employ in practice.
The policy improvement step then fits the actor policy network by optimizing the objective
Finally, the temperature is learned with the loss
Deep Q-learning
DQN also learns a convolutional neural net to approximate Q-function over states and actions. The main difference is that DQN operates on discrete actions spaces, thus the policy can be directly inferred from Q-values. The parameters of DQN are updated by optimizing the squared residual error
In practice, the standard version of DQN is frequently combined with a set of tricks that improve performance and training stability, wildly known as Rainbow .
Appendix B The DeepMind Control Suite Experiments Setup
Our PyTorch SAC implementation is based off of .
We employ clipped double Q-learning for the critic, where each -function is parametrized as a 3-layer MLP with ReLU activations after each layer except of the last. The actor is also a 3-layer MLP with ReLUs that outputs mean and covariance for the diagonal Gaussian that represents the policy. The hidden dimension is set to for both the critic and actor.
B.2 Encoder Network
We employ an encoder architecture from . This encoder consists of four convolutional layers with kernels and channels. The ReLU activation is applied after each conv layer. We use stride to everywhere, except of the first conv layer, which has stride . The output of the convnet is feed into a single fully-connected layer normalized by LayerNorm . Finally, we apply tanh nonlinearity to the dimensional output of the fully-connected layer. We initialize the weight matrix of fully-connected and convolutional layers with the orthogonal initialization and set the bias to be zero.
The actor and critic networks both have separate encoders, although we share the weights of the conv layers between them. Furthermore, only the critic optimizer is allowed to update these weights (e.g. we stop the gradients from the actor before they propagate to the shared conv layers).
B.3 Training and Evaluation Setup
Our agent first collects seed observations using a random policy. The further training observations are collected by sampling actions from the current policy. We perform one training update every time we receive a new observation. In cases where we use action repeat, the number of training observations is only a fraction of the environment steps (e.g. a steps episode at action repeat will only results into training observations). We evaluate our agent every true environment steps by computing the average episode return over evaluation episodes. During evaluation we take the mean policy action instead of sampling.
B.4 PlaNet and Dreamer Benchmarks
We consider two evaluation setups that were introduced in PlaNet and Dreamer , both using tasks from the DeepMind control suite . The PlaNet benchmark consists of six tasks of various traits. Importantly, the benchmark proposed to use a different action repeat hyper-parameter for each task, which we summarize in Table 2.
The Dreamer benchmark considers an extended set of tasks, which makes it more difficult that the PlaNet setup. Additionally, this benchmark requires to use the same set hyper-parameters for each task, including action repeat (set to ), which further increases the difficulty.
B.5 Pixels Preprocessing
We construct an observational input as an -stack of consecutive frames , where each frame is a RGB rendering of size from the th camera. We then divide each pixel by to scale it down to $$ range.
B.6 Other Hyper Parameters
Due to computational constraints for all the continuous control ablation experiments in the main paper and appendix we use a minibatch size of , while for the main results we use minibatch of size . In Table 3 we provide a comprehensive overview of all the other hyper-parameters.
Appendix C The Atari 100k Experiments Setup
For ease of reproducibility in Table 4 we report the hyper-parameter settings used in the Atari 100k experiments. We largely reuse the hyper-parameters from OTRainbow , but adapt them for DQN . Per common practise, we average performance of our agent over random seeds. The evaluation is done for 125k environment steps at the end of training for 100k environment steps.
Appendix D Full Atari 100k Results
Besides reporting in Figure 5 median human-normalized episode returns over the Atari games used in , we also provide the mean episode return for each individual game in Table 5. Human/Random scores are taken from to be consistent with the established setup.
Appendix E Image Augmentations Ablation
Following , we evaluate popular image augmentation techniques, namely random shifts, cutouts, vertical and horizontal flips, random rotations and imagewise intensity jittering. Below, we provide a comprehensive overview of each augmentation. Furthermore, we examine effectiveness of these techniques in Figure 6.
We bring our attention to random shifts that are commonly used to regularize neural networks trained on small images . In our implementation of this method images of size are padded each side by pixels (by repeating boundary pixels) and then randomly cropped back to the original size.
Cutout
Cutouts introduced in represent a generalization of Dropout . Instead of masking individual pixels cutouts mask square regions. Since image pixels can be highly correlated, this technique is proven to improve training of neural networks.
Horizontal/Vertical Flip
This technique simply flips an image either horizontally or vertically with probability .
Rotate
Here, an image is rotated by degrees, where is uniformly sampled from $$.
Intensity
Implementation
Finally, we provide Python-like implementation for the aforementioned augmentations powered by Kornia .
Appendix F K and M Hyper-parameters Ablation
We further ablate the K,M hyper-parameters from Algorithm 1 to understand their effect on performance. In Figure 7 we observe that increase values of K,M improves the agent’s performance. We choose to use the [K=2,M=2] parametrization as it strikes a good balance between performance and computational demands.
Appendix G Robustness Investigation
To demonstrate the robustness of our approach , we perform a comprehensive study on the effect different hyper-parameter choices have on performance. A review of prior work shows consistent values for discount and target update rate parameters, but variability on network architectures, mini-batch sizes, learning rates. Since our method is based on SAC , we also check whether the initial value of the temperature is important, as it plays a crucial role in the initial phase of exploration. We omit search over network architectures since Figure 1b shows our method to be robust to the exact choice. We thus focus on three hyper-parameters: mini-batch size, learning rate, and initial temperature.
Due to computational demands, experiments are restricted to a subset of tasks from : Walker Walk, Cartpole Swingup, and Finger Spin. These were selected to be diverse, requiring different behaviors including locomotion and goal reaching. A grid search is performed over mini-batch sizes , learning rates , and initial temperatures . We follow the experimental setup from Appendix B, except that only seeds are used due to the computation limitations, but since variance is low the results are representative.
Figure 8 shows performance curves for each configuration as well as a heat map over the mean performance of the final evaluation episodes, similar to . Our method demonstrates good stability and is largely invariant to the studied hyper-parameters. We emphasize that for simplicity the experiments in Section 4 use the default learning rate of Adam (0.001), even though it is not always optimal.
Appendix H Improved Data-Efficient Reinforcement Learning from Pixels
Our method allows to generate many various transformations from a training observation due to the data augmentation strategy. Thus, we further investigate whether performing more training updates per an environment step can lead to even better sample-efficiency. Following we compare a single update with a mini-batch of transitions with updates with different mini-batches of size samples each. Performing more updates per an environment step leads to even worse over-fitting on some tasks without data augmentation (see Figure 9a), while our method DrQ, that takes advantage of data augmentation, demonstrates improved sample-efficiency (see Figure 9b).