Provable Representation Learning for Imitation with Contrastive Fourier Features
Ofir Nachum, Mengjiao Yang
Introduction
In the field of sequential decision making one aims to learn a behavior policy to act in an environment to optimize some criteria. The well-known field of reinforcement learning (RL) corresponds to one aspect of sequential decision making, where the aim is to learn how to act in the environment to maximize cumulative returns via trial-and-error experience . In this work, we focus on imitation learning, where the aim is to learn how to act in the environment to match the behavior of some unknown target policy . This focus puts us closer to the supervised learning regime, and, indeed, a common approach to imitation learning – known as behavioral cloning (BC) – is to perform max-likelihood training on a collected a set of target demonstrations composed of state-action pairs sampled from the target policy .
Since the learned behavior policy produces predictions (actions) conditioned on observations (states), the amount of demonstrations needed to accurately match the target policy typically scales with the state dimension, and this can limit the applicability of imitation learning to settings where collecting large amounts of demonstrations is expensive, \eg, in health and robotics applications. The limited availability of target demonstrations stands in contrast to the recent proliferation of large offline datasets for sequential decision making . These datasets may exhibit behavior far from the target policy and so are not directly relevant to imitation learning via max likelihood training. Nevertheless, the offline datasets provide information about the unknown environment, presenting samples of environment reward and transition dynamics. It is therefore natural to wonder, is it possible to use such offline datasets to improve the sample efficiency of imitation learning?
Recent empirical work suggests that this is possible , by using the offline datasets to learn a low-dimensional state representation via unsupervised training objectives. While these empirical successes are clear, the theoretical foundation for these results is less obvious. The main challenge in providing theoretical guarantees for such techniques is that of aliasing. Namely, even if environment rewards or dynamics exhibit a low-dimensional structure, the target policy and its demonstrations may not. If the target policy acts differently in states which the representation learning objective maps to the same low-dimensional representation, the downstream behavioral cloning objective may end up learning a policy which “averages” between these different states in unpredictable ways.
In this work, we aim to bridge the gap between practical objectives and theoretical understanding. We derive an offline objective that learns low-dimensional representations of the environment dynamics and, if available, rewards. We show that minimizing this objective in conjunction with a downstream behavioral cloning objective corresponds to minimizing an upper bound on the performance difference between the learned low-dimensional BC policy and the unknown and possibly high-dimensional target policy. The form of our bound immediately makes clear that, as long as the learned policy is sufficiently expressive on top of the low-dimensional representations, the implicit “averaging” occurring in the BC objective due to any aliasing is irrelevant, and a learned policy can match the target regardless of whether the target policy itself is low-dimensional.
Extending our results to policies with limited expressivity, we consider the commonly used parameterization of setting the learned policy to be log-linear with respect to the representations (\ie, a softmax of a linear transformation). In this setting, we show that it is enough to use the same offline representation learning objective, but with linearly parameterized dynamics and rewards, and this again leads to an upper bound showing that the downstream BC policy can match the target policy regardless of whether the target is low-dimensional or log-linear itself. We compare the form of our representation learning objective to “latent space model” approaches based on bisimulation principles, popular in the RL literature , and show that these objectives are, in contrast, very liable to aliasing issues even in simple scenarios, explaining their poor performance in recent empirical studies .
We continue to the practicality of our own objective, and show that it can be implemented as a contrastive learning objective that implicitly learns an energy based model, which, in many common cases, corresponds to a linear model with respect to representations given by random Fourier features . We evaluate our objective in both tabular synthetic domains and high-dimensional Atari game environments . We find that our representation learning objective effectively leverages offline datasets to dramatically improve performance of behavioral cloning.
Related Work
Representation learning in sequential decision making has traditionally focused on learning representations for improved RL rather than imitation. While some works have proposed learning action representations , our work focuses on state representation learning, which is more common in the literature, and whose aim is generally to distill aspects of the observation relevant to control from those relevant only to measurement ; see for a review. Of these approaches, bisimulation is the most theoretically mature , and several recent works apply bisimulation principles to derive practical representation learning objectives . However, the existing theoretical results for bisimulation fall short of the guarantees we provide. For one, many of the bisimulation results rely on defining a representation error which holds globally on all states and actions . On the other hand, theoretical bisimulation results that define a representation error in terms of an expectation are inapplicable to imitation learning, as they only provide guarantees bounding the performance difference between policies that are “close” (\eg, Lipschitz) in the representation space and say nothing regarding whether an arbitrary target policy in the true MDP can be represented in the latent MDP . In Section 4.4, we will show that these shortcomings of bisimulation fundamentally limit its applicability to an imitation learning setting.
In contrast to RL, there are comparatively fewer theoretical works on representation learning for imitation learning. One previous line of research in this vein is given by , which considers learning a state representation using a dataset of multiple demonstrations from multiple target policies. Accordingly, this approach requires that each target policy admits a low-dimensional representation. In contrast, our own work makes no assumption on the form of the target policy, and, in fact, this is one of the central challenges of representation learning in this setting.
As imitation learning is close to supervised learning, it is an interesting avenue for future work to extend our results to more common supervised learning domains. We emphasize that our own contrastive objectives are distinct from typical approaches in image domains and popular in image-based RL , which use prior knowledge to generate pairs of similar images (\eg, via random cropping). We avoid any such prior knowledge of the task, and our losses are closer to temporal contrastive learning, more common in NLP .
Background
We begin by introducing the notation and concepts we will build upon in the later sections.
The performance associated with a policy is its expected future discounted reward when acting in the manner described above:
The visitation distribution of is the state distribution induced by the sequential process:
Behavioral Cloning (BC)
In imitation learning, one wishes to recover an unknown target policy with access to demonstrations of acting in the environment. More formally, the demonstrations are given by a dataset where . A popular approach to imitation learning is behavioral cloning (BC), which suggests to learn a policy to approximate via max-likelihood optimization. That is, one wishes to use the samples to approximately minimize the objective
In this work, we consider using a state representation function to simplify this objective. Namely, we consider a function . Given this representation, one no longer learns a policy , but rather a policy . The BC loss with representation becomes
Offline Data
Learning Goal
Similar to related work , we will measure the discrepancy between a candidate and the target via the performance difference:
For any , the performance difference may be bounded as
Notice that the guarantee above for vanilla BC includes a quadratic dependence on horizon in the form of , and this quadratic dependence is maintained in all our subsequent bounds. While there exists a number of imitation learning works that aim to reduce this dependence, the specific problem our paper focuses on – aliasing in the context of learning state representations – is an orthogonal problem to quadratic dependence on horizon. Indeed, if some representation maps two very different raw observations to the same latent state, no downstream imitation learning algorithm (regardless of sample complexity) will be able to learn a good policy. Still, extending our representation learning bounds to more sophisticated algorithms with potentially smaller dependence on horizon, like DAgger , is a promising direction for future work.
Representation Learning with Performance Bounds
We now continue to our contributions, beginning by presenting performance difference bounds analogous to Lemma 1 but with respect to a specific representation . The bounds will necessarily depend on quantities which correspond to how “good” the representation is, and these quantities then form the representation learning objective for learning ; ideally these quantities are independent of , which is unknown.
Consider a representation function and models as defined above. Denote the representation error as
Then the performance difference in between and a latent policy may be bounded as,
2 Log-linear Policies
Theorem 2 establishes a connection between the performance difference and behavioral cloning over representations given by . The optimal latent policy for BC is , and this is the same policy which achieves minimal performance difference. Whether we can find depends on how we parameterize our latent policy. If is tabular or if is represented as a sufficiently expressive neural network, then the approximation error is effectively zero. But what about in other cases?
The statement of Theorem 3 makes it clear that realizability of is irrelevant for log-linear policies. It is enough to only have the gradient with respect to learned be close to zero, which is a guarantee of virtually all gradient-based algorithms. Thus, in these settings performing BC on top of learned representations is provably optimal regardless of both the form of and the form of .
It is possible to extend the statement of Theorem 3 to generalized linear dynamics and reward models based on kernels by replacing the gradient in the bound with the functional gradient with respect to the kernel .
3 Sample Efficiency
4 Comparison to Bisimulation
The form of our representation learning objectives – learning to be predictive of rewards and next state dynamics – recalls similar ideas in the bisimulation literature . However, a key difference is that in bisimulation the divergence over next state dynamics is measured in the latent representation space; \ie, a divergence between and for some “latent space model” , whereas our proposed representation error is between and . We find that this difference is crucial, and in fact there exist no theoretical guarantees for bisimulation similar to those in Theorems 2 and 3. Indeed, one can construct a simple example where the use of latent space models leads to a complete failure, see Figure 1.
Learning the Representations in Practice
One may recover a contrastive learning objective by parameterizing as an energy-based model. Namely, consider parameterizing as
Similar contrastive learning objectives have appeared in related works , and so our theoretical bounds can be used to explain these previous empirical successes.
2 Linear Models with Contrastive Fourier Features
While the connection between temporal contrastive learning and approximate dynamics models has appeared in previous works , it is not immediately clear how one should learn the approximate linear dynamics required by Theorem 3. In this section, we show how the same contrastive learning objective can be used to learn approximate linear models, thus illuminating a new connection between contrastive learning and near-optimal sequential decision making; see Appendix A for pseudocode.
Experiments
We now empirically verify the performance benefits of the proposed representation learning objective in both tabular and Atari game environments. See environment details in Appendix D.
We learn contrastive Fourier features as described in Section 5.2 using tabular . We then fix these representations and train a log-linear policy on the target demonstrations. For the baseline, we learn a vanilla BC policy with tabular parametrization directly on target demonstrations.
We also experiment with representations given by singular value decomposition (SVD) of the empirical transition matrix, which is another form of learning factored linear dynamics. Figure 2 shows the performance achieved by the learned policy with and without representation learning. Representation learning consistently yields significant performance gains, especially with few target demonstrations. SVD performs similar to contrastive Fourier features when the offline data is abundant with respect to the state space size, but degrades as the offline data size reduces or the state space grows.
2 Atari 2600 with Deep Neural Networks
We now study the practical benefit of the proposed contrastive learning objective to imitation learning on Atari 2600 games , taking for the offline dataset the DQN Replay Dataset , which for each game provides 50M steps collected during DQN training. For the target demonstrations, we take k single-step transitions from the last 1M steps of each dataset, corresponding to the data collected near the end of the DQN training. For the offline data, we use all M transitions of each dataset.
For our learning agents, we extend the implementations found in Dopamine . We use the standard Atari CNN architecture to embed image-based inputs to vectors of dimension . In the case of vanilla BC, we pass this embedding to a log-linear policy and learn the whole network end-to-end with behavioral cloning. For contrastive Fourier features, we use separate CNNs to parameterize in the objective in Section 5.2, and the representation is given by the random Fourier feature procedure described in Section 5.2. A log-linear policy is then trained on top of this representation, but without passing any BC gradients through . This corresponds to the setting of Theorem 3. We also experiment with the setting of Theorem 2; in this case the setup is same as for contrastive Fourier features, only that we define (\ie, is an energy-based dynamics model) and we parameterize as a more expressive single-hidden-layer softmax policy on top of .
We compare contrastive learning with Fourier features and energy-based models to two latent space models, DeepMDP and Deep Bisimulation for Control (DBC) , in Figure 3. Both linear (Fourier features) and energy-based parametrization of contrastive learning achieve dramatic performance gains () on over half of the games. DeepMDP and DBC, on the other hand, achieve little improvement over vanilla BC when presented as a separate loss from behavioral cloning. Enabling end-to-end learning of the latent space models as an auxiliary loss to behavioral cloning leads to better performance, but DeepMDP and DBC still underperform contrastive learning. See Appendix D for further ablations.
Conclusion
We have derived an offline representation learning objective which, when combined with BC, provably minimizes an upper bound on the performance difference from the target policy. We further showed that the proposed objective can be implemented as contrastive learning with an optional projection to Fourier features. Interesting avenues for future work include (1) extending our theory to multi-step contrastive learning, popular in practice , (2) deriving similar results for policy learning in offline and online RL settings, and (3) reducing the effect of offline distribution shifts. We also note that our use of contrastive Fourier features for learning a linear dynamics model may be of independent interest, especially considering that a number of theoretical RL works rely on such an approximation (\eg, ), while to our knowledge no previous work has demonstrated a practical and scalable learning algorithm for linear dynamics approximation. Determining if the technique of learning contrastive Fourier features works well for these settings offers another interesting direction to explore.
Acknowledgments and Disclosure of Funding
We thank Bo Dai, Rishabh Agarwal, Mohammad Norouzi, Pablo Castro, Marlos Machado, Marc Bellemare, and the rest of the Google Brain team for fruitful discussions and valuable feedback.
References
Appendix A Pseudocode
We present basic pseudocode of feature learning below.
For the Fourier feature representation, we normalize the components of , which, when the inputs are normally distributed with some unknown mean and variance, may be interpreted as focusing the sampling distribution of on the most informative Fourier features; mathematically, one may show this still approximates the kernel as by using importance sampling with importance weights placed on instead of .
Appendix B Proofs
We first present a basic performance difference lemma:
If and are two policies in , then
Note the second quantity above is a scaled TV-divergence between and . When the MDP reward is action-independent, \ie for all , the same bound holds with the reward term removed.
Following similar derivations in , we express the performance difference in linear operator notation:
where are linear operators such that . Notice that may be expressed in this notation as . We split the expression above into two parts:
Using matrix norm inequalities, we bound the above by
and so we immediately achieve the desired bound in (20).
In the case of action-independent rewards, one may follow the same derivation for the first part of (24), starting with
Now we incorporate a representation function , showing how the errors above may be further reduced in the special case of :
The result follows from straightforward algebraic manipulation via the definitions of and triangle inequality. For the reward error, we have,
as desired. The bound for the transition error may be derived analogously. ∎
In the case of linear reward and transition models, we may derive a variant of Lemma 7, which will be useful in the proof of Theorem 3:
where and .
The crux of the proof is noting that the gradient above for a specific column of may be expressed as
as desired. The bound for the transition error may be derived analogously. ∎
Our final lemma will be used to translate on-policy bounds to off-policy.
For two distributions with , we have,
The lemma is a straightforward consequence of Cauchy-Schwartz:
B.2 Proof of Theorem 2
With the lemmas in the above section, we are now prepared to prove Theorem 2. We begin with the following on-policy version of Theorem 2 and then derive the off-policy bound:
Then the performance difference between and a latent policy may be bounded as,
We combine Lemmas 5, 6, and 7 with and to yield the following bound:
To yield the desired bound, we simply apply Pinsker’s inequality and the concavity of the square-root function:
B.3 Proof of Theorem 3
The proof of Theorem 3 is derived analogously to that of Theorem 2 above, except using Lemma 8 in place of Lemma 7.
B.4 Proof of Lemma 1
Lemma 1 may be immediately derived from Theorem 2, using and with and .
Appendix C Sample Efficiency
Let be a distribution with finite support. Let denote the empirical estimate of from i.i.d. samples . Then,
The first inequality is Lemma 8 in while the second inequality is due to the concavity of the square root function. ∎
Let be i.i.d. samples from a factored distribution for . Let be the empirical estimate of in and be the empirical estimate of in . Then,
Let be the empirical estimate of in . We have,
Finally, the bound in the lemma is achieved by application of Lemma 11 to each of the TV divergences. ∎
To prove Theorem 4, we first present the following variant of the bound in Theorem 3, which maintains the BC loss as a TV divergence rather than a KL divergence and whose validity is clear from (63):
The result in Theorem 4 is then derived by setting and using the result of Lemma 12, noting that learning a tabular with BC on corresponds to setting to be the empirical conditional distribution appearing in with respect to .
Appendix D Experiment Details
The three-level binary decision tree used in Section 6 has a chance of landing on the intended child node and a chance of random exploration following the Dirichlet distribution (). Each node in the tree has two rewards associated with taking left or right actions. The mean of the rewards are generated uniformly between and and are powered to the third and normalized to have a maximum of at each step. A unit Gaussian noise is then applied to the rewards. An optimal policy and a random policy achieve near and average per-step reward respectively.
D.2 Experiment Details
For tree experiments in Figure 2, we set , when ablating over , , when ablating over , , , when ablating over , and , , when ablating over . For the representation dimensions, we use SVD features or Fourier features. We use the Adam optimizer with learning rate for both representation learning and behavioral cloning.
Atari
For Atari experiments in Figure 3, we set state embedding size to 256, Fourier features size to and use all M transitions as the offline data by default. When sampling batches, to encourage better negative samples in the contrastive learning, we sample 4 sequences of length 64 from the replay buffer, making a total batch size of 256 single-step transitions. We use the standard Atari CNN architecture with three convolutional layers interleaved with ReLU activation followed by two fully connected layers with units each to output state representations. For the energy-based parametrization in Section D.4, we is a single-hidden-layer NN with units. All networks are trained using the Adam optimizer with learning rate . We train the representations and behavior cloning concurrently with separate losses. Training is conducted on NVIDIA P100 GPUs.
D.3 Improvements on individual Atari games
D.4 Ablation study of contrastive learning
We further investigate various factors that potentially affect the benefit of contrastive representation learning. We consider:
Size of the state representations: Either 256 or 512.
Parameterization of the approximate dynamics model: Either using Fourier features with a log-linear policy (corresponding to Section 5.2 and Theorem 3) or using contrastive learning of an energy-based dynamics model with a more expressive single-hidden-layer neural network for (corresponding to Section 5.1 and Theorem 2).
How far the offline distribution differs from target demonstrations: Using either the last 10M, 25M, or whole 50M transitions in the offline replay dataset.
Figure 5 shows the quantile of the normalized improvements among the Atari games. Overall, the benefit of representation learning is robust to these factors, and is slightly more pronounced with larger state embedding size and smaller difference between the offline distribution and target demonstrations.