Learning Invariant Representations for Reinforcement Learning without Reconstruction

Amy Zhang, Rowan McAllister, Roberto Calandra, Yarin Gal, Sergey Levine

Introduction

Learning control from images is important for many real world applications. While deep reinforcement learning (RL) has enjoyed many successes in simulated tasks, learning control from real vision is more complex, especially outdoors, where images reveal detailed scenes of a complex and unstructured world. Furthermore, while many RL algorithms can eventually learn control from real images given unlimited data, data-efficiency is often a necessity in real trials which are expensive and constrained to real-time. Prior methods for data-efficient learning of simulated visual tasks typically use representation learning. Representation learning summarizes images by encoding them into smaller vectored representations better suited for RL. For example, sequential autoencoders aim to learn lossless representations of streaming observations—sufficient to reconstruct current observations and predict future observations—from which various RL algorithms can be trained (Hafner et al., 2019; Lee et al., 2020; Yarats et al., 2021). However, such methods are task-agnostic: the models represent all dynamic elements they observe in the world, whether they are relevant to the task or not. We argue such representations can easily “distract” RL algorithms with irrelevant information in the case of real images. The issues of distraction is less evident in popular simulation MuJoCo and Atari tasks, since any change in observation space is likely task-relevant, and thus, worth representing. By contrast, visual images that autonomous cars observe contain predominately task-irrelevant information, like cloud shapes and architectural details, illustrated in Figure 1.

Rather than learning control-agnostic representations that focus on accurate reconstruction of clouds and buildings, we would rather achieve a more compressed representation from a lossy encoder, which only retains state information relevant to our task. If we would like to learn representations that capture only task-relevant elements of the state and are invariant to task-irrelevant information, intuitively we can utilize the reward signal to help determine task-relevance, as shown by Jonschkowski & Brock (2015). As cumulative rewards are our objective, state elements are relevant not only if they influence the current reward, but also if they influence state elements in the future that in turn influence future rewards. This recursive relationship can be distilled into a recursive task-aware notion of state abstraction: an ideal representation is one that is predictive of reward, and also predictive of itself in the future.

We propose learning such an invariant representation using the bisimulation metric, where the distance between two observation encodings correspond to how “behaviourally different” (Ferns & Precup, 2014) both observations are. Our main contribution is a practical representation learning method based on the bisimulation metric suitable for downstream control, which we call deep bisimulation for control (DBC). We additionally provide theoretical analysis that proves value bounds between the optimal value function of the true MDP and the optimal value function of the MDP constructed by the learned representation. Empirical evaluations demonstrate our non-reconstructive approach using bisimulation is substantially more robust to task-irrelevant distractors when compared to prior approaches that use reconstruction losses or contrastive losses. Our initial experiments insert natural videos into the background of MoJoCo control task as complex distraction. Our second setup is a high-fidelity highway driving task using CARLA (Dosovitskiy et al., 2017), showing that our representations can be trained effectively even on highly realistic images with many distractions, such as trees, clouds, buildings, and shadows. For example videos see https://sites.google.com/view/deepbisim4control. Code is available at https://github.com/facebookresearch/deep_bisim4control.

Related Work

Our work builds on the extensive prior research on bisimulation in MDP state aggregation.

Contrastive-based Representations. Contrastive losses are a self-supervised approach to learn useful representations by enforcing similarity constraints between data (van den Oord et al., 2018; Chen et al., 2020). Similarity functions can be provided as domain knowledge in the form of heuristic data augmentation, where we maximize similarity between augmentations of the same data point (Laskin et al., 2020) or nearby image patches (Hénaff et al., 2020), and minimize similarity between different data points. In the absence of this domain knowledge, contrastive representations can be trained by predicting the future (van den Oord et al., 2018). We compare to such an approach in our experiments, and show that DBC is substantially more robust. While contrastive losses do not require reconstruction, they do not inherently have a mechanism to determine downstream task relevance without manual engineering, and when trained only for prediction, they aim to capture all predictable features in the observation, which performs poorly on real images for the same reasons world models do. A better method would be to incorporate knowledge of the downstream task into the similarity function in a data-driven way, so that images that are very different pixel-wise (e.g. lighting or texture changes), can also be grouped as similar w.r.t. downstream objectives.

Bisimulation. Various forms of state abstractions have been defined in Markov decision processes (MDPs) to group states into clusters whilst preserving some property (e.g. the optimal value, or all values, or all action values from each state) (Li et al., 2006). The strictest form, which generally preserves the most properties, is bisimulation (Larsen & Skou, 1989). Bisimulation only groups states that are indistinguishable w.r.t. reward sequences output given any action sequence tested. A related concept is bisimulation metrics (Ferns & Precup, 2014), which measure how “behaviorally similar” states are. Ferns et al. (2011) defines the bisimulation metric with respect to continuous MDPs, and propose a Monte Carlo algorithm for learning it using an exact computation of the Wasserstein distance between empirically measured transition distributions. However, this method does not scale well to large state spaces. Taylor et al. (2009) relate MDP homomorphisms to lax probabilistic bisimulation, and define a lax bisimulation metric. They then compute a value bound based on this metric for MDP homomorphisms, where approximately equivalent state-action pairs are aggregated. Most recently, Castro (2020) propose an algorithm for computing on-policy bisimulation metrics, but does so directly, without learning a representation. They focus on deterministic settings and the policy evaluation problem. We believe our work is the first to propose a gradient-based method for directly learning a representation space with the properties of bisimulation metrics and show that it works in the policy optimization setting.

Preliminaries

We start by introducing notation and outlining realistic assumptions about underlying structure in the environment. Then, we review state abstractions and metrics for state similarity.

Bisimulation is a form of state abstraction that groups states si\mathbf{s}_{i} and sj\mathbf{s}_{j} that are “behaviorally equivalent” (Li et al., 2006). For any action sequence a0:∞\mathbf{a}_{0:\infty}, the probabilistic sequence of rewards from si\mathbf{s}_{i} and sj\mathbf{s}_{j} are identical. A more compact definition has a recursive form: two states are bisimilar if they share both the same immediate reward and equivalent distributions over the next bisimilar states (Larsen & Skou, 1989; Givan et al., 2003).

Given an MDP M\mathcal{M}, an equivalence relation BB between states is a bisimulation relation if, for all states si,sj∈S\mathbf{s}_{i},\mathbf{s}_{j}\in\mathcal{S} that are equivalent under BB (denoted si≡Bsj\mathbf{s}_{i}\equiv_{B}\mathbf{s}_{j}) the following conditions hold:

where SB\mathcal{S}_{B} is the partition of S\mathcal{S} under the relation BB (the set of all groups GG of equivalent states), and P(G∣s,a)=∑s′∈GP(s′∣s,a).\mathcal{P}(G|\mathbf{s},\mathbf{a})=\sum_{\mathbf{s}^{\prime}\in G}\mathcal{P}(\mathbf{s}^{\prime}|\mathbf{s},\mathbf{a}).

Defining a distance dd between states requires defining both a distance between rewards (to soften Equation 1), and distance between state distributions (to soften Equation 2). Prior works use the Wasserstein metric for the latter, originally used in the context of bisimulation metrics by van Breugel & Worrell (2001). The pthp^{\text{th}} Wasserstein metric is defined between two probability distributions Pi\mathcal{P}_{i} and Pj\mathcal{P}_{j} as Wp(Pi,Pj;d)=(inf⁡γ′∈Γ(Pi,Pj)∫S×Sd(si,sj)p dγ′(si,sj))1/pW_{p}(\mathcal{P}_{i},\mathcal{P}_{j};d)=(\inf_{\gamma^{\prime}\in\Gamma(\mathcal{P}_{i},\mathcal{P}_{j})}\int_{\mathcal{S}\times\mathcal{S}}d(\mathbf{s}_{i},\mathbf{s}_{j})^{p}\,\textnormal{d}\gamma^{\prime}(\mathbf{s}_{i},\mathbf{s}_{j}))^{1/p}, where Γ(Pi,Pj)\Gamma(\mathcal{P}_{i},\mathcal{P}_{j}) is the set of all couplings of Pi\mathcal{P}_{i} and Pj\mathcal{P}_{j}. This is known as the “earth mover” distance, denoting the cost of transporting mass from one distribution to another (Villani, 2003). Finally, the bisimulation metric is the reward difference added to the Wasserstein distance between transition distributions:

From Theorem 2.6 in Ferns et al. (2011) with c∈[0,1)c\in[0,1):

Learning Representations for Control with Bisimulation Metrics

Incorporating control. We combine our representation learning approach (Algorithm 1) with the soft actor-critic (SAC) algorithm (Haarnoja et al., 2018) to devise a practical reinforcement learning method. We modified SAC slightly in Algorithm 2 to allow the value function to backprop to our encoder, which can improve performance further (Yarats et al., 2021; Rakelly et al., 2019). Although, in principle, our method could be combined with any RL algorithm, including the model-free DQN (Mnih et al., 2015), or model-based PETS (Chua et al., 2018). Implementation details and hyperparameter values of DBC are summarized in the appendix, Table 2. We train DBC by iteratively updating three components in turn: a policy π\pi (in this case SAC), an encoder ϕ\phi, and a dynamics model P^\hat{\mathcal{P}} (lines 7–9, Algorithm 1). We found a single loss function was less stable to train. The inputs of each loss function J(⋅)J(\cdot) in Algorithm 1 represents which components are updated. After each training step, the policy π\pi is used to step in the environment, the data is collected in a replay buffer D\mathcal{D}, and a batch is randomly selected to repeat training.

Generalization Bounds and Links to Causal Inference

While DBC enables representation learning without pixel reconstruction, it leaves open the question of how good the resulting representations really are. In this section, we present theoretical analysis that bounds the suboptimality of a value function trained on the representation learned via DBC.

First, we show that our π∗\pi^{*}-bisimulation metric converges to a fixed point, starting from the initialized policy π0\pi_{0} and converging to an optimal policy π∗\pi^{*}.

Let met\mathfrak{met} be the space of bounded pseudometrics on S\mathcal{S} and π\pi a policy that is continuously improving. Define F:met↦met\mathcal{F}:\mathfrak{met}\mapsto\mathfrak{met} by

Given an MDP Mˉ\bar{\mathcal{M}} constructed by aggregating states in an ϵ\epsilon-neighborhood, and an encoder ϕ\phi that maps from states in the original MDP M\mathcal{M} to these clusters, the optimal value functions for the two MDPs are bounded as

MDP dynamics have a strong connection to causal inference and causal graphs, which are directed acyclic graphs (Jonsson & Barto, 2006; Schölkopf, 2019; Zhang et al., 2020). Specifically, the state and action at time tt causally affect the next state at time t+1t+1. In this work, we care about the components of the state space that causally affect current and future reward. Deep bisimulation for control representations connect to causal feature sets, or the minimal feature set needed to predict a target variable (Zhang et al., 2020).

If we partition observations using the bisimulation metric, those clusters (a bisimulation partition) correspond to the causal feature set of the observation space with respect to current and future reward.

This connection tells us that these features are the minimal sufficient statistic of the current and future reward, and therefore consist of (and only consist of) the causal ancestors of the reward variable rr.

In a causal graph where nodes correspond to variables and directed edges between a parent node PP and child node CC are causal relationships, the causal ancestors AN(C)AN(C) of a node are all nodes in the path from CC to a root node.

If there are interventions on distractor variables, or variables that control the rendering function qq and therefore the rendered observation but do not affect the reward, the causal feature set will be robust to these interventions, and correctly predict current and future reward in the linear function approximation setting (Zhang et al., 2020). As an example, in autonomous driving, an intervention can be a change from day to night which affects the observation space but not the dynamics or reward. Finally, we show that a representation based on the bisimulation metric generalizes to other reward functions with the same causal ancestors.

Proof in appendix. This result shows that the learned representation will generalize to unseen reward functions, as long as the new reward function has a subset of the same causal ancestors. As an example, a representation learned for a robot to walk will likely generalize to learning to run, because the reward function depends on forward velocity and all the factors that contribute to forward velocity. However, that representation will not generalize to picking up objects, as those objects will be ignored by the learned representation, since they are not likely to be causal ancestors of a reward function designed for walking. Theorem 4 shows that the learned representation will be robust to spurious correlations, or changes in factors that are not in AN(R)AN(R). This complements Theorem 5, that the representation is a minimal sufficient statistic of the optimal value function, improving generalization over non-minimal representations.

See Theorem 5.1 in Ferns et al. (2004) for proof. We show empirical validation of these findings in Section 6.2.

Experiments

Our central hypothesis is that our non-reconstructive bisimulation based representation learning approach should be substantially more robust to task-irrelevant distractors. To that end, we evaluate our method in a clean setting without distractors, as well as a much more difficult setting with distractors. We compare against several baselines. The first is Stochastic Latent Actor-Critic (SLAC, Lee et al. (2020)), a state-of-the-art method for pixel observations on DeepMind Control that learns a dynamics model with a reconstruction loss. The second is DeepMDP (Gelada et al., 2019), a recent method that also learns a latent representation space using a latent dynamics model, reward model, and distributional Q learning, but for which they needed a reconstruction loss to scale up to Atari. Finally, we compare against two methods using the same architecture as ours but exchange our bisimulation loss with (1) a reconstruction loss (“Reconstruction”) and (2) contrastive predictive coding (Oord et al., 2018) (“Contrastive”) to ground the dynamics model and learn a latent representation.

In this section, we benchmark DBC and the previously described baselines on the DeepMind Control (DMC) suite (Tassa et al., 2018) in two settings and nine environments (Figure 3), finger_spin, cheetah_run, and walker_walk and additional environments in the appendix.

Default Setting. Here, the pixel observations have simple backgrounds as shown in Figure 3 (top row) with training curves for our DBC and baselines. We see SLAC, a recent state-of-the-art model-based representation learning method that uses reconstruction, generally performs best.

Simple Distractors Setting. Next, we include simple background distractors, shown in Figure 3 (middle row), with easy-to-predict motions. We use a fixed number of colored circles that obey the dynamics of an ideal gas (no attraction or repulsion between objects) with no collisions. Note the performance of DBC remains consistent, as other methods start decreasing.

Natural Video Setting. Then, we incorporate natural video from the Kinetics dataset (Kay et al., 2017) as background (Zhang et al., 2018), shown in Figure 3 (bottom row). The results confirm our hypothesis: although a number of prior methods can learn effectively in the absence of distractors, when complex distractions are introduced, our non-reconstructive bisimulation based method attains substantially better results.

To visualize the representation learned with our bisimulation metric loss function in Equation 4, we use a t-SNE plot (Figure 4). We see that even when the background looks drastically different, our encoder learns to ignore irrelevant information and maps observations with similar robot configurations near each other. See Appendix D for another visualization.

2 Generalization Experiments

We test generalization of our learned representation in two ways. First, we show that the learned representation space can generalize to different types of distractors, by training with simple distractors and testing on the natural video setting. Second, we show that our learned representation can be useful reward functions other than those it was trained for.

Generalizing over backgrounds. We first train on the simple distractors setting and evaluate on natural video. Figure 5 shows an example of the simple distractors setting and performance during training time of two experiments, blue being the zero-shot transfer to the natural video setting, and orange the baseline which trains on natural video. This result empirically validates that the representations learned by DBC are able to effectively learn to ignore the background, regardless of what the background contains or how dynamic it is.

Generalizing over reward functions. We evaluate (Figure 5) the generalization capabilities of the learned representation by training SAC with new reward functions walker_stand and walker_run using the fixed representation learned from walker_walk. This is empirical evidence that confirms Theorem 4: if the new reward functions are causally dependent on a subset of the same factors that determine the original reward function, then our representation is sufficient.

3 Comparison with other Bisimulation Encoders

Even though the purpose of bisimulation metrics by Castro (2020) is learning distances dd, not representation spaces Z\mathcal{Z}, it nevertheless implements dd with function approximation: d(\mathbf{s}_{i},\mathbf{s}_{j})=\psi\big{(}\phi(\mathbf{s}_{i}),\phi(\mathbf{s}_{j})\big{)} by encoding observations with ϕ\phi before computing distances with ψ\psi, trained as:

4 Autonomous Driving with Visual Redundancy

Results in Figure 9 compare the same baselines as before, except for SLAC which is easily distracted (Figure 3). Instead we used SAC, which does not explicitly learn a representation, but performs surprisingly well from raw images. DeepMDP performs well too, perhaps given its similarly to bisimulation. But, Reconstruction and Contrastive methods again perform poorly with complex images. More intuitive metrics are in Table 1 and Figure 8 depicts the representation space as a t-SNE with corresponding observations. Each run took 12 hours on a GTX 1080 GPU.

Discussion

This paper presents Deep Bisimulation for Control: a new representation learning method that considers downstream control. Observations are encoded into representations that are invariant to different task-irrelevant details in the observation. We show this is important when learning control from outdoor images, or otherwise images with background “distractions”. In contrast to other bisimulation methods, we show performance gains when distances in representation space match the bisimulation distance between observations.

Future work: Several options exist for future work. First, our latent dynamics model P^\hat{\mathcal{P}} was only used for training our encoder in Equation 4, but could also be used for multi-step planning in latent space. Second, estimating uncertainty could also be important to produce agents that can work in the real world, perhaps via an ensemble of models {P^k}k=1K\{\hat{\mathcal{P}}_{k}\}_{k=1}^{K}, to detect—and adapt to—distributional shifts between training and test observations. Third, an undressed issue is that of partially observed settings (that assumed approximately full observability by using stacked images), possibly using explicit memory or implicit memory such as an LSTM. Finally, investigating which metrics (L1 or L2) and dynamics distributions (Gaussians or not) would be beneficial.

References

Appendix A Additional Theorems and Proofs

Let met\mathfrak{met} be the space of bounded pseudometrics on SS and π∈Π\pi\in\Pi a policy that is continuously improving in the space of policies Π\Pi. Define F:met×Π↦met\mathcal{F}:\mathfrak{met}\times\Pi\mapsto\mathfrak{met} by

Ideally, to prove this theorem we show that F\mathcal{F} is monotonically increasing and continuous, and apply Fixed Point Theorem to show the existence of a fixed point that F\mathcal{F} converges to. Unfortunately, we can show that F\mathcal{F} under π\pi as π\pi monotonically converges to π∗\pi^{*} is not also monotonic, unlike the original bisimulation metric setting (Ferns et al., 2004) and the policy evaluation setting (Castro, 2020). We start the iterates Fn\mathcal{F}^{n} from bottom ⊥\perp, denoted as Fn(⊥)\mathcal{F}^{n}(\perp). In Ferns et al. (2004) the max⁡a∈A\max_{\mathbf{a}\in\mathcal{A}} can be thought of as learning a policy between every two pairs of states to maximize their distance, and therefore this distance can only stay the same or grow over iterations of F\mathcal{F}. In Castro (2020), π\pi is fixed, and under a deterministic MDP it can also be shown that distance between states dn(si,sj)d_{n}(\mathbf{s}_{i},\mathbf{s}_{j}) will only expand, not contract as nn increases. In the policy iteration setting, however, with π\pi starting from initialization π0\pi_{0} and getting updated:

Instead, we show that using the policy improvement theorem which gives us

π\pi will converge to a fixed point using the Fixed Point Theorem, and taking the result by Castro (2020) that Fπ\mathcal{F}^{\pi} has a fixed point for every π∈Π\pi\in\Pi, we can show that a fixed point bisimulation metric will be found with policy iteration. ∎

Given a new aggregated MDP Mˉ\bar{\mathcal{M}} constructed by aggregating states in an ϵ\epsilon-neighborhood, and an encoder ϕ\phi that maps from states in the original MDP M\mathcal{M} to these clusters, the optimal value functions for the two MDPs are bounded as

From Theorem 5.1 in Ferns et al. (2004) we have:

We assume a MDP with a state space S:={S1,...,SK}\mathcal{S}:=\{\mathcal{S}^{1},...,\mathcal{S}^{K}\} that can be factorized into KK variables with 1-step causal transition dynamics described by a causal graph G\mathcal{G} (example in Figure 10). We break the proof up into two parts: 1) show that if a factor Si∉AN(R)\mathcal{S}^{i}\notin AN(R) changes, the bisimulation distance between the original state s\mathbf{s} and the new state s′\mathbf{s}^{\prime} is 0. and 2) show that if a factor Sj∈AN(R)\mathcal{S}^{j}\in AN(R) changes, the bisimulation distance can be >0>0.

1) If Si∉AN(R)\mathcal{S}^{i}\notin AN(R), an intervention on that factor does not affect current or future reward.

If Si\mathcal{S}^{i} does not affect future reward, then states si\mathbf{s}_{i} and sj\mathbf{s}_{j} will have the same future reward conditioned on all future actions. This gives us

Appendix B Definition of State

Since we are concerned primarily with learning from image observations, we could explicitly distinguish the image observation space O\mathcal{O} from an unknown state space S\mathcal{S}. However, since we are not tackling the general POMDP problem, we consider the Block MDP (Du et al., 2019), which assumes the state space is latent, and that we are instead given access to an observation space O\mathcal{O} and rendering function q:S↦Oq:\mathcal{S}\mapsto\mathcal{O}. The crucial assumption that distinguishes the Block MDP from partially observable MDPs is the following:

Each observation o\mathbf{o} uniquely determines its generating state s\mathbf{s}. That is, the observation space O\mathcal{O} can be partitioned into disjoint blocks Os\mathcal{O}_{s}, each containing the support of the conditional distribution q(o∣s)q(\mathbf{o}|\mathbf{s}).

This assumption gives us the Markov property in the observation space o∈O\mathbf{o}\in\mathcal{O}. As an example, one can think of the proprioceptive state consisting of positions and velocities of actuators as the underlying state, and stacked pixel observations from a specific camera angle as a particular rendering function and corresponding observation space.

Appendix C Additional DMC Results

In Figure 11 we show performance on the default setting on 9 different environments from DMC. Figures 12 and 13 give performance on the simple distractors and natural video settings for all 9 environments.

Appendix D Additional Visualizations

In addition to Figure 4, we also took 10 nearby points in the t-SNE plot and average the observations, shown on the far left of Figure 14. Note the robot agent is quite crisp, which means neighboring points encode the agent in similar positions, but the backgrounds are very different, and so are blurry when averaged.

Appendix E Implementation Details

We use the same encoder architecture as in Yarats et al. (2021), which is an almost identical encoder architecture as in Tassa et al. (2018), with two more convolutional layers to the convnet trunk. The encoder has kernels of size 3×33\times 3 with 3232 channels for all the convolutional layers and set stride to 11 everywhere, except of the first convolutional layer, which has stride 22, and interpolate with ReLU activations. Finally, we add tanh nonlinearity to the 5050 dimensional output of the fully-connected layer.

For the reconstruction method, the decoder consists of a fully-connected layer followed by four deconvolutional layers. We use ReLU activations after each layer, except the final deconvolutional layer that produces pixels representation. Each deconvolutional layer has kernels of size 3×33\times 3 with 3232 channels and stride 11, except of the last layer, where stride is 22.

The dynamics and reward models are both MLPs with two hidden layers with 200 neurons each and ReLU activations.

Soft Actor Critic (SAC) (Haarnoja et al., 2018) is an off-policy actor-critic method that uses the maximum entropy framework for soft policy iteration. At each iteration, SAC performs soft policy evaluation and improvement steps. The policy evaluation step fits a parametric soft Q-function Q(st,at)Q(\mathbf{s}_{t},\mathbf{a}_{t}) using transitions sampled from the replay buffer D\mathcal{D} by minimizing the soft Bellman residual,

The target value function Vˉ\bar{V} is approximated via a Monte-Carlo estimate of the following expectation,

where Qˉ\bar{Q} is the target soft Q-function parameterized by a weight vector obtained from an exponentially moving average of the Q-function weights to stabilize training. The policy improvement step then attempts to project a parametric policy π(at∣st)\pi(\mathbf{a}_{t}|\mathbf{s}_{t}) by minimizing KL divergence between the policy and a Boltzmann distribution induced by the Q-function, producing the following objective,

We modify the Soft Actor-Critic PyTorch implementation by Yarats & Kostrikov (2020) and augment with a shared encoder between the actor and critic, the general model fsf_{s} and task-specific models fηef_{\eta}^{e}. The forward models are multi-layer perceptions with ReLU non-linearities and two hidden layers of 200 neurons each. The encoder is a linear layer that maps to a 50-dim hidden representation. The hyperparameters used for the RL experiments are in Table 2.