Augmented World Models Facilitate Zero-Shot Dynamics Generalization From a Single Offline Environment

Philip J. Ball, Cong Lu, Jack Parker-Holder, Stephen Roberts

Introduction

Offline reinforcement learning (RL) describes the problem setting where RL agents learn policies solely from previously collected experience without further interaction with the environment (Fujimoto et al.,, 2019; Levine et al.,, 2020). This could have tremendous implications for real world problems (Dulac-Arnold et al.,, 2019), with the potential to leverage rich datasets of past experience where exploration is either not feasible (e.g. a Mars Rover) or unsafe (e.g. in medical settings). As such, interest in offline RL has surged in recent times.

This work focuses on model-based offline RL, which has achieved state-of-the-art performance through the use of uncertainty penalized updates (Yu et al.,, 2020; Kidambi et al.,, 2020). However, existing work only addresses the issue of transferring from different behavior policies in the same environment, ignoring any possibility of distribution shift. Consider the case where it is expensive to collect data, and we have access to a single dataset from a robot. Using existing methods we would be unable to make any changes that impact the dynamics, such as using a newer model of the robot or deploying it in a different room.

A related setting is the Sim2Real problem which considers transferring an agent from a simulated environment to the real world. A popular recent approach is domain randomization (Tobin et al.,, 2017; James et al.,, 2017), the process of randomizing non-essential regions of the observation space to make agents robust to ‘observational overfitting’ (Song et al.,, 2020). Indeed, methods seeking to generalize to novel dynamics have also shown promise (Peng et al.,, 2018), by randomizing physical properties such as the mass of the agent. A significant limitation of these approaches is the requirement for a simulator, which may not be available.

In this work we take inspiration from Sim2Real to generalize solely from an offline dataset, in a learned simulator or World Model (WM). We therefore describe our problem setting as follows: an agent must learn to generalize to unseen test-time dynamics whilst having access to offline data from only a single environment; we call this “dynamics generalization from a single offline environment”.

In this paper we concentrate on the zero-shot performance of our agents to unseen dynamics, as it may not be practical (nor safe) to perform multiple rollouts at test time. To tackle this problem, we propose a novel form of data augmentation: rather than augment observations, we focus on augmenting the dynamics. We first learn a world model of the environment, and then augment the transition function at policy training time, making the agent train under different imagined dynamics. In addition, our agent is given access to the augmentation itself as part of the observations, allowing it to consider the context of modified dynamics.

At test time we propose a simple, yet surprisingly effective, self-supervised approach to learning an agent’s augmentation context. We learn a linear dynamics model which is then used to approximate the dynamics augmentation induced by the modified environment. This context is then given to the agent, allowing it to adapt on the fly to the new dynamics within a single episode (i.e. zero-shot). We show that our approach is capable of training agents that can vastly outperform existing Offline RL methods on the “dynamics generalization from a single offline environment” problem. We also note that this approach does not require access to environment rewards at test time. This facilitates application to Sim2Real problems whereby test time rewards may not be available.

Our contributions are twofold: 1) As far as we are aware, we are the first to propose dynamics augmentation for model based RL, allowing us to generalize to changing dynamics despite only training on a single setting. We do this without access to any environment parameters or prior knowledge. 2) We propose a simple self-supervised context adaptation reward-free algorithm, which allows our policy to use information from interactions in the environment to vary its behavior in a single episode, increasing zero-shot performance. We believe both of these approaches are not only novel, but offer significant improvement v.s. state-of-the-art methods, improving generalization and providing a promising approach for using offline RL in the real world.

Related Work

In this work we focus on Model Based RL (MBRL). A key challenge in MBRL is that an inaccurate model can be exploited by the policy, leading to behaviors that fail to transfer to the real environment. As such, a swathe of recent works have made use of model ensembles to improve robustness (Kurutach et al.,, 2018; Chua et al.,, 2018; Clavera et al.,, 2018; Janner et al.,, 2019; Ball et al.,, 2020). With increased accuracy, MBRL has recently been shown to be competitive with model free methods in continuous control (Ha and Schmidhuber,, 2018; Chua et al.,, 2018; Janner et al.,, 2019) and games (Schrittwieser et al.,, 2019; Kaiser et al.,, 2020). We make use of an ensemble of probabilistic dynamics models, first introduced in Lakshminarayanan et al., (2017) and subsequently used in Chua et al., (2018).

In this paper we focus on Model-Based offline RL, where MOPO (Yu et al.,, 2020) and MOReL (Kidambi et al.,, 2020) have recently demonstrated the effectiveness of learned dynamics models, using model uncertainty to constrain policy optimization. We build upon this approach for zero-shot dynamics generalization from offline data. There have also been successes in off policy methods for offline RL (Wu et al.,, 2019; Fujimoto et al.,, 2019; Kumar et al.,, 2020; Rudner et al.,, 2021) and context based approaches (Ajay et al.,, 2021), although these works only consider tasks within the support of the offline dataset. Finally, MBOP (Argenson and Dulac-Arnold,, 2021) addresses the problem of goal-conditioned zero-shot transfer from offline datasets. However, their goal-conditioning relies on unchanged dynamics in the test environment.

In online RL, recent work has achieved strong dynamics generalization with a learned model (Seo et al.,, 2020). However, this required training under varied dynamics, assigning different experiences to models. In addition, this work used MPC whereas we train a policy inside the model, which is significantly faster at deployment time. Also related are Clavera et al., (2019); Nagabandi et al., (2019), where the model is trained to quickly to adapt to new dynamics P(s′∣s,a)P(s^{\prime}|s,a), however both these works place more emphasis on model-adaption rather than zero-shot policy performance. Furthermore, access to an underlying task distribution is required, something we do not have in our offline setting. Also similar to our work is the recently proposed Policy Adaptation during Deployment (PAD, (Hansen et al.,, 2021)) approach. Our approach differs in that we learn a context, whereas PAD uses a auxiliary objective to adapt its features. In addition, PAD considers the online model free setting, while our method is offline and model based.

Sim2Real is the setting where an agent trained in a simulator must transfer to the real world. A common approach to solve this problem is through domain randomization (Tobin et al.,, 2017; James et al.,, 2017), whereby parameters in the simulator are varied during training. This has shown to be effective for dynamics generalization (Andrychowicz et al.,, 2020; Antonova et al.,, 2017; Peng et al.,, 2018; Yu et al.,, 2017; Zhou et al.,, 2019; OpenAI et al.,, 2019), but requires access to a simulator which we do not have. Another form of domain randomization, data augmentation, has proved to be effective for training RL policies (Laskin et al., 2020a, ; Laskin et al., 2020b, ; Kostrikov et al.,, 2021; Raileanu et al.,, 2020), resulting in improved efficiency and generalization. So far, these works have focused on online model free methods, and used data augmentation on the state space, reducing observational overfitting (Song et al.,, 2020). In contrast, we focus on offline MBRL and instead augment the dynamics.

We also note clear links to contextual MDPs (Hallak et al.,, 2015; Modi et al.,, 2018) and hidden parameter MDPs (HiP-MDP) (Doshi-Velez and Konidaris,, 2016; Killian et al.,, 2017; Zhang et al.,, 2021) settings, whereby our self-supervised dynamics embedding can be considered as a context/hidden parameter. However in these settings the embedding is chosen at the beginning of each episode and is fixed throughout, whereas our embedding varies per timestep.

We are not the first to propose data augmentation in the MBRL setting, Pitis et al., (2020) proposed Counterfactual data augmentation for improving performance in the context of locally factored tasks. Approaches to ensuring adversarial robustness can include data augmentations that assist with out-of-domain generalization, as opposed to observational overfitting. In Volpi et al., (2018) this is done without a simulator and from a single source of data, however they only work on supervised learning problems and require an adversary to be learned, adding computational complexity. Finally, Wellmer and Kwok, (2021) concurrently explore the idea of augmenting world model dynamics for improved test-time transferability, however they focus on in-domain generalization, and do not infer context at test time.

Preliminaries

When training a model, we follow MBPO (Janner et al.,, 2019) and MOPO (Yu et al.,, 2020) and train an ensemble of NN probabilistic dynamics models (Nix and Weigend,, 1994). Each model learns to predict both next state s′s^{\prime} and reward rr from a state-action pair, using Denv\mathcal{D}_{env} in a supervised fashion. Furthermore, each model outputs a Gaussian P^i(st+1,rt∣st,at)=N(μ(st,at),Σ(st,at))\widehat{P}_{i}(s_{t+1},r_{t}|s_{t},a_{t})=\mathcal{N}(\mu(s_{t},a_{t}),\Sigma(s_{t},a_{t})). The resulting model P^\widehat{P} defines a model MDP M^=(S,A,P^,R^,ρ0,γ)\widehat{M}=(\mathcal{S},\mathcal{A},\widehat{P},\widehat{R},\rho_{0},\gamma), where R^\widehat{R} refers to the learned reward model.

While MOPO and MOReL have addressed the issue of training a policy in Denv\mathcal{D}_{env}, and transferring to the true environment MM, they only consider where the data in Denv\mathcal{D}_{env} is actually drawn from PP. However, sometimes this may not be sufficient for deployment. For example, a robot could fail to walk when learning from data that was collected by a different version of the robot (with different mass), or if the same robot collected data but in a different room to deployment (with varied friction). It is this setting, where dynamics may vary at test time, that is the focus of our work. To learn successfully we propose a novel approach to training robust context-dependent policies.

Augmented World Models with Self Supervised Policy Adaptation

In this section we introduce our algorithm: Augmented World Models (AugWM). We first discuss our training procedure (Fig. 2) before moving onto our self-supervised approach to online context learning (Fig. 3).

for all s,as,a, some small ϵ>0\epsilon>0, and suitable distance/divergence metric DD. We consider several augmentations, beginning with existing works before moving to new approaches which specifically target the problem of dynamics generalization. We begin with Random Amplitude Scaling as in Laskin et al., 2020a , which we refer to as RAD. RAD scales both sts_{t} and st+1s_{t+1} as follows:

One crucial addition to our method is the use of context. Concretely, when we are optimizing the policy using a batch of data, we concatenate the next state with the augmentation vector zz. This allows our policy to be informed of the specific augmentation that was applied to the environment and thus behave accordingly. However, at test time we do not know zz, so what can we use? Next we propose a solution to this problem, learning the context on the fly.

2 Self-Supervised Policy Selection

In the meta-learning literature there have been many recent successes making use of a learned context to adapt a policy at test time to a new environment (Rakelly et al.,, 2019; Zintgraf et al.,, 2020, 2021), typically using a blackbox model with a latent state. Crucially, these approaches require several episodes to adapt at test time, making them unfeasible in our zero-shot setting. What makes our setting unique is we explicitly know what the context represents: a linear transformation of sts_{t}, or δt\delta_{t}. Using this insight, we are able to learn an effective context on the fly at test time. Concretely, we observe that from a state sts_{t} drawn from M⋆M^{\star}, we can sample an action at∼πa_{t}\sim\pi and then compute an approximate s^t+1\widehat{s}_{t+1} using our model P^\widehat{P}. With s^t+1\widehat{s}_{t+1}, we have a sample estimate of the state change under P^\widehat{P}, i.e. δ^t=s^t+1−st\widehat{\delta}_{t}=\widehat{s}_{t+1}-s_{t}. We can make this approximation of the next state without interacting with the environment, but once we do take the action ata_{t} in the environment, we then receive the true next state st+1s_{t+1} and can store the true difference δt=st+1−st\delta_{t}=s_{t+1}-s_{t}. Using the DAS augmentation, we can approximate zz as \nicefracδtδ^t\nicefrac{{\delta_{t}}}{{\hat{\delta}_{t}}}.

This however is retrospective; we can only approximate zz having already seen the next state, by which time our agent has already acted. Furthermore, we believe under changed dynamics the true zz likely depends on ss, thus we cannot use a previous zz for future timesteps. Therefore, we learn a forward dynamics model using the data collected during the test rollout. After hh timesteps in the environment, we have the following dataset: D={(s1,s2),…,(sh−1,sh)}\mathcal{D}=\{(s_{1},s_{2}),\dots,(s_{h-1},s_{h})\}. This allows us to learn a simple dynamics model fψ:(st)↦δt=st+1−stf_{\psi}:(s_{t})\mapsto\delta_{t}=s_{t+1}-s_{t}, by minimizing the mean squared error LMSE(ψ,D)\mathcal{L}_{\text{MSE}}(\psi,\mathcal{D}). Notably, since we never actually plan with this model, it does not need to be as accurate as a typical dynamics model in MBRL. Instead, it is crucial that the model learns quickly enough such that we can use it in a zero-shot evaluation. Thus, we choose to use a simple linear model for ff. To show the effectiveness of this, in Fig. 4 we show the mean R-squared of linear models learned on the fly during evaluation rollouts.

We observe that in less than 100100 timesteps the linear model achieves high accuracy on the test data. Subsequently, equipped with fψf_{\psi}, we can approximate δt\delta_{t}, and predict the augmentation as z^t=\nicefracδtδ^t\widehat{z}_{t}=\nicefrac{{\delta_{t}}}{{{\hat{\delta}_{t}}}}. We then provide the agent with z^t\widehat{z}_{t} to compute action ata_{t}. The full procedure is shown in Algorithm 2.

Experiments

In our experiments we aim to investigate the effectiveness of our approach for zero-shot dynamics generalization from a single offline dataset. To assess this, we will answer a series of questions, beginning with a question on the necessity of our method:

Do we really need to develop methods specifically for dynamics generalization?

To answer this, we train MOPO using offline data from a single environment, and test it under changed dynamics. We consider the HalfCheetah environment from the OpenAI Gym (Brockman et al.,, 2016), using offline data from D4RL (Fu et al.,, 2021). We train a MOPO agent using the mixed dataset, using our own implementation of the algorithm (but using the same hyperparameters as the original authors). To test the trained policy, we vary both the mass of the agent and damping coefficient by a multiplicative factor. The standard environment (both in Gym and D4RL) corresponds to both these values being set to 1.01.0. In this work we consider a grid of the following values for HalfCheetah: {0.25,0.5,0.75,1.0,1.25,1.5,1.75}\{0.25,0.5,0.75,1.0,1.25,1.5,1.75\} and {0.5,0.75,1.0,1.25,1.5}\{0.5,0.75,1.0,1.25,1.5\} for Walker2d, representing a significant out-of-distribution shift.

The results (Fig. 5), show that MOPO performance is clearly impacted by changing dynamics. We see in the central cell, that performance for our version of MOPO matches the author results (Yu et al.,, 2020), and in some cases we even see small gains (e.g. mass =0.75=0.75, damping =1.0=1.0). However, on the top left we see dramatically weaker performance, often below 11k, indicating the robot is failing to achieve locomotion. Before evaluating AugWM, we first test whether training with the “correct” augmentation improves generalization performance. In short, we ask:

To answer this, we train SAC for 1×1051\times 10^{5} steps and save the states visited in the ‘true’ environment. We then use these starting states to train a policy using an offline MBPOSince we have access to the true environment, there is no need for the MOPO penalty. approach with AugWM. However, instead of sampling zt∼Zz_{t}\sim\mathcal{Z}, we provide the actual z=\nicefracδ⋆δ^z=\nicefrac{{\delta^{\star}}}{{\hat{\delta}}} as we have access to the ‘true’ and ‘modified’ environments; we refer to this as an oracle version of our method, and is designed to assess the viability of our approach. Note that we do not augment the ‘true’ environment rewards. We consider two baselines: a) offline MBPO in the ‘true’ environment; b) online SAC in the ‘modified’ environment. We train MBPO until convergence, and SAC for 1×1051\times 10^{5} steps.

As shown Fig. 6, when provided with the true zz, AugWM outperforms both baselines. The SAC result is surprising: the baseline agent was trained directly on the ‘modified’ environment for the same number of steps as the policy that generated our oracle starting states. One explanation is the greater exploration induced by the ‘easier’ dynamics of the ‘true’ environment. This validates our approach; if we augment the dynamics P^\widehat{P} from a model correctly, we can generalize to unseen dynamics. In other words, neither the starting states nor rewards need to be from the ‘modified’ environment. With this in mind, our next question is a simple one:

Which augmentation strategy is most effective?

To test this, we train as in Algorithm 1, without context, to isolate the effectiveness of the training process. We use the HalfCheetah Mixed dataset and train a MOPO agent, augmenting either both ss and s′s^{\prime} (RAD), just s′s^{\prime} (RANS) or just δ\delta (DAS).

The results are shown in Fig. 7. As we see, the RAD augmentation fails to improve dynamics generalization, actually leading to worse performance overall. RANS does improve performance on unseen dynamics, as we are influencing the dynamics, not just the observation. However, DAS clearly provides the strongest performance. As a result, we use DAS for AugWM. Our final algorithm design question is as follows:

Does training with context improve performance?

To answer this question, we return to the HalfCheetah Mixed setting from Fig. 7, taking the policy trained with DAS. We now train two additional agents: 1) Default Context: at train time the agent is provided with the DAS augmentation zz as context, at test time it is provided with a vector of ones, 1∣S∣\mathbf{1}^{|\mathcal{S}|}; 2) Learned Context: trained as in 1), but context is learned online using Algorithm 2.

The results are shown in in Fig. 8. We observe that training with context (orange) improves performance on average, while adapting the context on the fly (green) leads to further gains (+80 on average). These methods combine to produce our AugWM algorithm. We are now ready for the final question:

Can Augmented World Models improve zero-shot generalization?

To answer this question we perform a rigorous analysis, using multiple benchmarks from the D4RL dataset (Fu et al.,, 2021). Namely, we consider the Random, Medium, Mixed and Medium-Expert datasets for both Walker2d and HalfCheetah. In each setting, we compare AugWM against base MOPO on zero-shot performance, training entirely on the data provided, but not seeing the test environment until evaluation. The results are shown as a change v.s. MOPO, averaged over one dimension in Fig. 11, and as a total return number averaged over both dimensions in Table 1. For additional implementation details (e.g., hyperparameters) see Appendix B. AugWM provides statistically significant improvements in zero-shot performance v.s. MOPO in many cases, achieving successful policies where MOPO fails.

By now we have provided significant evidence that AugWM can significantly improve performance for HalfCheetah and Walker2d with varied mass and damping. However, this is only a small subset of possible dynamics changes. We next consider several significantly harder settings. We test increased dimensionality, using the Ant environment from MOPO (Yu et al.,, 2020), and consider additional types of dynamics changes (e.g., Ant with crippled legs, HalfCheetah with changed limb sizes from Henderson et al., (2017)). We illustrate the impact of the crippled leg Ant environment on baseline agent performance, and the improvement provided by AugWM, in Fig. 12. We show the mean results over each of these factors of variation in Table 2, where again AugWM provides a non-trivial improvement over a strong baseline. For more details see Appendix B.

Finally, we note that dynamics may change during an episode; consider a robot that suffers a motor fault, reducing the power delivered to its joints. Evidently the underlying dynamics have been altered, and being robust to such changes when only training from a single dataset of offline experience is challenging. To illustrate this, we perform a 15001500 step rollout in the HalfCheetah environment, starting with offline dynamics (mass/damping = 11), before changing to mass = 0.750.75, damping = 0.50.5 after 500500 steps; performance is shown in Fig. 9. Observe that after 500500 steps, MOPO performance is dramatically reduced. This is because the agent continues to apply the same force and thus falls forward with lighter mass. For our AugWM agent, performance initially drops, then when the new context is learned we achieve higher performance than before, making use of the lighter torso.

Discussion We believe that our experiments provide significant support to the claim that training with AugWM improves zero-shot dynamics generalization. In a broad set of commonly used datasets, and with a wide range of out-of-distribution dynamics, our algorithm learns good policies where a state-of-the-art baseline fails.For videos see: https://sites.google.com/view/augmentedworldmodels/ This is due to a number of novel contributions: 1) using dynamics augmentation rather than observation augmentation; 2) training and testing with a context-based policy. Regarding limitations, we note that training inside the WM with context generally takes longer to converge (Appendix B). Furthermore, in more nonlinear settings such as the HalfCheetah modified body part setting we saw a reduced performance for the learned context. This could be because the dynamics changes are out of the distribution of DAS augmentations (violating Eqn. 1), or due to the difficulty of modeling the task with a linear model. We note that linear models have achieved success in planning (Gu et al.,, 2016) and meta learning (Peng et al.,, 2021), and are effective in our case due to their data efficiency, but can be replaced by more flexible models to deal with different augmentations. Indeed, given our work is the first of its kind, we believe significant improvements are possible, such as using more complex and problem-specific augmentations.

Conclusion and Future Work

In this paper we propose Augmented World Models (AugWM), which we show sufficiently simulates changes in dynamics such that agents can generalize in a zero-shot manner. We believe that we are the first to propose this problem setting, and our results show a significant improvement over existing state-of-the-art methods which ignore this problem.

A promising line of future work would be to meta-train a policy over AugWM such that it can quickly adapt to new dynamics in the few-shot setting. There is evidence that data augmentation can improve robustness in meta-learning (Rajendran et al.,, 2020), and could extend to strong performance in out-of-distribution tasks. We also wish to consider varying goals at test time, and other potential sources of non-stationarity which may impact policy performance (Igl et al.,, 2021). It may also be possible to extend AugWM to pixel-based tasks, which have received a great deal of recent attention (Hafner et al.,, 2019, 2020). We believe that our transition based augmentations will be applicable to a latent representation, as commonly used in state-of-the-art vision MBRL approaches. Thus we think that extending our work to this setting, while a considerable feat of engineering, should not require significant methodological changes.

Acknowledgments

Philip Ball is funded through the Willowgrove Studentship. Cong Lu is funded by the Engineering and Physical Sciences Research Council (EPSRC). We are grateful to Taylor Killian for useful discussions on contextual/HiP-MDPs, and to Vitaly Kurin for his feedback on an earlier version of this paper via his ‘Paper Notes’. The authors would also like to thank the anonymous ICLR SSL-RL Workshop + ICML reviewers and the area chair for constructive feedback which helped us in improving the paper.

Changes From ICML 2021 Proceedings

We added additional related work that we were not originally aware of, updated the Acknowledgments section, and generally tidied up the formatting.

References

Appendix

Appendix A Additional Experiments

In this section we show the performance for Augmented World Models with different training ranges for the DAS augmentation (zz train in Table 4). We train with adaptive context on the HalfCheetah mixed dataset, and present the results in Fig. 13. As we see, [0.75,1.25][0.75,1.25] and [0.5,1.5][0.5,1.5] perform the best. Based on this, we use [0.5,1.5][0.5,1.5] for our experiments as we believe this helps us sample a wider set of dynamics, helping us generalize better across all environments and data sets.

Appendix B Implementation Details

Our algorithm is based on MOPO (Yu et al.,, 2020) with values for the rollout length hh and penalty coefficient λ\lambda shown in Table 3.

AugWM specific hyperparameters are listed in Table 4. For each evaluation rollout, we clear the buffer of stored true modified environment transitions to measure zero-shot performance. We adapt using the context after a set number of steps, kk, in the environment to train the linear model. The two ranges used for the context zz during training and test time are different. At test time, the estimated context is clipped to remain within the given bounds.

B.2 D4RL dataset

We evaluate our method on D4RL (Fu et al.,, 2021) datasets based on the MuJoCo continuous control tasks (halfcheetah and walker2d). The four dataset types we evaluate on are:

random: roll out a randomly initialized policy for 1M steps.

medium: partially train a policy using SAC, then roll it out for 1M steps.

mixed: train a policy using SAC until a certain (environment-specific) performance threshold is reached, and take the replay buffer as the batch.

medium-expert: combine 1M samples of rollouts from a fully-trained policy with another 1M samples of rollouts from a partially trained policy or a random policy.

B.3 Ant Environment

For the Ant experiments, we follow the Ant Changed Direction approach in MOPO (Yu et al.,, 2020). Since this offline dataset is not provided in the authors’ code, nor is it in the standard D4RL library (Fu et al.,, 2021), we were required to generate our own offline Ant dataset. Since the authors’ did not outline certain details in their experiment, we found the following was required to match their performance with our codebase: 1) Training our SAC policy for 1×1061\times 10^{6} timesteps in the Ant environment provided by the authors’ code in (Yu et al.,, 2020); 2) relabelling each reward in the buffer using the new direction, without the living reward; 3) training a world model over this offline dataset; 4) training a policy in the world model, adding in living reward post-hoc; 5) evaluating the policy with the living reward.

B.4 HalfCheetah Modified Agent

We use the modified HalfCheetah environments from (Henderson et al.,, 2017). In each setting one body part of the agent is changed, from following set: {Foot, Leg, Thigh, Torso, Head}. The body part can either be “Big” or “Small”, where Big bodyparts involve scaling the mass and width of the limb by 1.25 and Small bodyparts are scaled by 0.75. In Table 2 we show the mean over each of these five body parts, for agents trained on each of the D4RL datasets, repeated for five seeds.