Mastering Atari Games with Limited Data

Weirui Ye, Shaohuai Liu, Thanard Kurutach, Pieter Abbeel, Yang Gao

Introduction

Reinforcement learning has achieved great success on many challenging problems. Notable work includes DQN , AlphaGo and OpenAI Five . However, most of these works come at the cost of a large number of environmental interactions. For example, AlphaZero needs to play 21 million games at training time. On the contrary, a professional human player can only play around 5 games per day, meaning it would take a human player 11,500 years to achieve the same amount of experience. The sample complexity might be less of an issue when applying RL algorithms in simulation and games. However, when it comes to real-world problems, such as robotic manipulation, healthcare, and advertisement recommendation systems, achieving high performance while maintaining low sample complexity is the key to viability.

People have made a lot of progress in sample efficient RL in the past years . Among them, model-based methods have attracted a lot of attention, since both the data from real environments and the “imagined data” from the model can be used to train the policy, making these methods particularly sample-efficient . However, most of the successes are in state-based environments. In image-based environments, some model-based methods such as MuZero and Dreamer V2 achieve super-human performance, but they are not sample efficient; other methods such as SimPLe is quite efficient but achieve inferior performance (0.144 human normalized median scores). Recently, data-augmented and self-supervised methods applied to model-free methods have achieved more success in the data-efficient regime . However, they still fail to achieve the levels which can be expected of a human.

Therefore, for improving the sample efficiency as well as keeping superior performance, we find the following three components are essential to the model-based visual RL agent: a self-supervised environment model, a mechanism to alleviate the model compounding error, and a method to correct the off-policy issue. In this work, we propose EfficientZero, a model-based RL algorithm that achieves high performance with limited data. Our proposed method is built on MuZero. We make three critical changes: (1) use self-supervised learning to learn a temporally consistent environment model, (2) learn the value prefix in an end-to-end manner, thus helping to alleviate the compounding error in the model, (3) use the learned model to correct off-policy value targets.

As illustrated as Figure 1, our model achieves state-of-the-art performance on the widely used Atari 100k benchmark and it achieves super-human performance with only 2 hours of real-time gameplay. More specifically, our model achieves 194.3% mean human normalized performance and 109.0% median human normalized performance. As a reference, DQN achieves 220% mean human normalized performance, and 96% median human normalized performance, at the cost of 500 times more data (200 million frames). To further verify the effectiveness of EfficientZero, we conduct experiments on some simulated robotics environments of the DeepMind Control (DMControl) suite. It achieves state-of-the-art performance and outperforms the state SAC which directly learns from the ground truth states. Our sample efficient and high-performance algorithm opens the possibility of having more impact on many real-world problems.

Related Work

Sample efficiency has attracted significant work in the past. In RL with image inputs, model-based approaches which model the world with both a stochastic and a deterministic component, have achieved promising results for simulated robotic control. Kaiser et al. propose to use an action-conditioned video prediction model, along with a policy learning algorithm. It achieves the first strong performance on Atari games with as little as 400k frames. However, Kielak and van Hasselt et al. argue that this is not necessary to achieve strong results with model-based methods, and they show that when tuned appropriately, Rainbow can achieve comparable results.

Recent advances in self-supervised learning, such as SimCLR , MoCo , SimSiam and BYOL have inspired representation learning in image-based RL. Srinivas et al. propose to use contrastive learning in RL algorithms and their work achieves strong performance on image-based continuous and discrete control tasks. Later, Laskin et al. and Kostrikov et al. find that contrastive learning is not necessary, but with data augmentations alone, they can achieve better performance. Some researchers propose to use contrastive learning to enforce action equivariance on the learned representations . Schwarzer et al. propose a temporal consistency loss, which is combined with data augmentations and achieves state-of-the-art performance. Notably, our self-supervised consistency loss is quite similar to SPR, except we use SimSiam while they use BYOL as the base self-supervised learning framework. However, they only apply the learned representations in a model-free manner, while we combine the learned model with model-based exploration and policy improvement, thus leading to more efficient use of the environment model.

Despite the recent progress in the sample-efficient RL, today’s RL algorithms are still well behind human performance when the amount of data is limited. Although traditional model-based RL is considered more sample efficient than model-free ones, current model-free methods dominate in terms of performance for image-input settings. In this paper, we propose a model-based RL algorithm that for the first time, achieves super-human performance on Atari games with limited data.

2 Reinforcement Learning with MCTS

Temporal difference learning and policy gradient based methods are two types of popular reinforcement learning algorithms. Recently, Silver et al. propose to use MCTS as a policy improvement operator and has achieved great success in many board games, such as Go, Chess, and Shogi . Later, the algorithm is adapted to learn the world model at the same time . It has also been extended to deal with continuous action spaces and offline data . These MCTS RL algorithms are a hybrid of model-based learning and model-free learning.

However, most of them are trained with a lot of environmental samples. Our method is built on top of MuZero , and we demonstrate that our method can achieve higher sample efficiency while still achieving competitive performance on the Atari 100k benchmark. de Vries et al. have studied the potential of using auxiliary loss similar to our self-supervised consistency loss. However, they only test on two low dimensional state-based environments and find the auxiliary loss has mixed effects on the performance. On the contrary, we find that the consistency loss is critical in most environments with high dimensional observations and limited data.

3 Multi-Step Value Estimation

In Q-learning , the target Q value is computed by one step backup. In practice, people find that incorporating multiple steps of rewards at once, i.e. zt=∑i=0k−1γiut+i+γkvt+kz_{t}=\sum_{i=0}^{k-1}\gamma^{i}u_{t+i}+\gamma^{k}v_{t+k}, where ut+iu_{t+i} is the reward from the replay buffer, vt+kv_{t+k} is the value estimation from the target network, to compute the value target ztz_{t} leads to faster convergence . However, the use of multi-step value has off-policy issues, since ut+iu_{t+i} are not generated by the current policy. In practice, this issue is usually ignored when there is a large amount of data since the data can be thought as approximately on-policy. TD(λ\lambda) and GAE improve the value estimation by better trading off the bias and the variance, but they do not deal with the off-policy issue. Recently, image input model-based algorithms such as Kaiser et al. and Hafner et al. use model imaginary rollouts to avoid the off-policy issue. However, this approach has the risk of model exploitation. Asadi et al. proposed a multi-step model to combat the compounding error. Our proposed model-based off-policy correction method starts from the rewards in the real-world experience and uses model-based value estimate to bootstrap. Our approach balances between the off-policy issue and model exploitation.

Background

Our method is built on top of the MuZero Reanalyze algorithm. For brevity, we refer to it as MuZero throughout the paper. MuZero is a policy learning method based on the Monte-Carlo Tree Search (MCTS) algorithm. The MCTS algorithm operates with an environment model, a prior policy function, and a value function. The environment model is represented as the reward function R\mathcal{R} and the dynamic function G\mathcal{G}: rt=R(st,at)r_{t}=\mathcal{R}(s_{t},a_{t}), s^t+1=G(st,at)\hat{s}_{t+1}=\mathcal{G}(s_{t},a_{t}), which are needed when MCTS expands a new node. In MuZero, the environment model is learned. Thus the reward and the next state are approximated. Besides, the predicted policy pt=p_{t}= acts as a search prior over actions of a node. It helps the MCTS focus on more promising actions when expanding the node. MCTS also needs a value function V(st)\mathcal{V}(s_{t}) that measures the expected return of the node sts_{t}, which provides a long-term evaluation of the tree’s leaf node without further search. MCTS will output an action visit distribution πt\pi_{t} over the root node, which is potentially a better policy, compared to the current neural network. Thus, the MCTS algorithm can be thought of as a policy improvement operator.

In practice, the environment model, policy function, and value function operate on a hidden abstract state sts_{t}, both for computational efficiency and ease of environment modeling. The abstract state is extracted by a representation function H\mathcal{H} on observations oto_{t}: st=H(ot)s_{t}=\mathcal{H}(o_{t}). All of the mentioned models above are usually represented as neural networks. During training, the algorithm collects roll-out data in the environment using MCTS, resulting in potentially higher quality data than the current neural network policy. The data is stored in a replay buffer. The optimizer minimizes the following loss on the data sampled from the replay buffer:

Here, utu_{t} is the reward from the environment, rt=R(st,at)r_{t}=\mathcal{R}(s_{t},a_{t}) is the predicted reward, πt\pi_{t} is the output visit count distribution of the MCTS, pt=P(st)p_{t}=\mathcal{P}(s_{t}) is the predicted policy, zt=∑i=0k−1γiut+i+γkvt+kz_{t}=\sum_{i=0}^{k-1}\gamma^{i}u_{t+i}+\gamma^{k}v_{t+k} is the bootstrapped value target and vt=V(st)v_{t}=\mathcal{V}(s_{t}) is the predicted value. Specifically, the reward function R\mathcal{R}, policy function P\mathcal{P}, value function V\mathcal{V}, the representation function H\mathcal{H} and the dynamics function G\mathcal{G} are trainable neural networks. It is worth noting that MuZero does not explicitly learn the environment model. Instead, it solely relies on the reward, value, and policy prediction to learn the model.

2 Monte-Carlo Tree Search

Monte-Carlo Tree Search , or MCTS, is a heuristic search algorithm. In our setup, MCTS is used to find an action policy that is better than the current neural network policy.

More specifically, MCTS needs an environment model, including the reward function and the dynamics function. It also needs a value function and a policy function, which act as heuristics during search. MCTS operates by expanding a search tree from the current node. It saves computation by selectively expanding nodes. To find a high-quality decision, the expansion process has to balance between exploration versus exploitation, i.e. balance between expanding a node that is promising with more visits versus a node with lower performance but fewer visits. MCTS employs the UCT rule, i.e. UCB on trees. At every node expansion step, UCT will select a node as follows :

where, Q(s,a)Q(s,a) is the current estimate of the Q-value, P(s,a)P(s,a) is the current neural network policy for selecting this action, helping the MCTS prioritize exploring promising part of the tree. During training time, P(s,a)P(s,a) is usually perturbed by noises to allow explorations. N(s,a)N(s,a) denotes how many times this state-action pair is visited in the tree search, and N(s,b)N(s,b) denote that of aa’s siblings. Thus this term will encourage the search to visit the nodes whose siblings are visited often, but itself less visited. Finally, the last term gives a weights to the previous terms.

After expanding the nodes for a pre-defined number of times, the MCTS will return how many times each action under the root node is visited, as the improved policy to the root node. Thus, MCTS can be considered as a policy improvement operator in the RL setting.

EfficientZero

Model-based algorithms have achieved great success in sample-efficient learning from low-dimensional states. However, current visual model-based algorithms either require large amounts of training data or exhibit inferior performance to model-free algorithms in data-limited settings . Many previous works even suspect whether model-based algorithms can really offer data efficiency when using image observations . We provide a positive answer here. We propose the EfficientZero, a model-based algorithm built on the MCTS, that achieves super-human performance on the 100k Atari benchmark, outperforming the previous SoTA to a large degree.

When directly running MCTS-based RL algorithms such as MuZero, we find that they do not perform well on the limited-data benchmark. Through our ablations, we confirm the following three issues which pose challenges to algorithms like MuZero in data-limited settings.

Lack of supervision on environment model. First, the learned model in the environment dynamics is only trained through the reward, value and policy functions. However, the reward is only a scalar signal and in many scenarios, the reward will be sparse. Value functions are trained with bootstrapping, and thus are noisy. Policy functions are trained with the search process. None of the reward, value and policy losses can provide enough training signals to learn the environment model.

Hardness to deal with aleatoric uncertainty. Second, we find that even with enough data, the predicted rewards still have large prediction errors. This is caused by the aleatoric uncertainty of the underlying environment. For example, when the environment is hard to model, the reward prediction errors will accumulate when expanding the MCTS tree to a large depth.

Off-policy issues of multi-step value. Lastly, as for value targets, MuZero uses the multi-step reward observed in the environment. Although it allows rewards to be propagated to the value function faster, it suffers from severe off-policy issues and hinders convergence given limited data.

To address the above issues, we propose the following three critical modifications, which can greatly improve performance when samples are limited.

In previous MCTS RL algorithms, the environment model is either given or only trained with rewards, values, and policies, which cannot provide sufficient training signals due to their scalar nature. The problem is more severe when the reward is sparse or the bootstrapped value is not accurate. The MCTS policy improvement operator heavily relies on the environment model. Thus, it is vital to have an accurate one. We notice that the output s^t+1\hat{s}_{t+1} from the dynamic function G\mathcal{G} should be the same as st+1s_{t+1}, i.e. the output of the representation function H\mathcal{H} with input of the next observation ot+1o_{t+1} (Fig. 2). This can help to supervise the predicted next state s^t+1\hat{s}_{t+1} using the actual st+1s_{t+1}, which is a tensor with at least a few hundred dimensions. This provides s^t+1\hat{s}_{t+1} with much more training signals than the default scalar reward and value.

Notably, in MCTS RL algorithms, the consistency between the hidden states and the predicted states can be shaped through the dynamics function directly without extra models. More specifically, we adopt the recently proposed SimSiam self-supervised framework. SimSiam is a self-supervised method that takes two augmentation views of the same image and pulls the output of the second branch close to that of the first branch, where the first branch is an encoder network without gradient, and the second is the same encoder network with the gradient and a predictor head. The head can simply be an MLP.

Note that SimSiam only learns the representation of individual images, and is not aware of how different images are connected. The learned image representations of SimSiam might not be a good candidate for learning the transition function, since adjacent observations might be encoded to very different representation encodings. We propose a self-supervised method that learns the transition function, along with the image representation function in an end-to-end manner, as Figure 2 shows. Since we aim to learn the transition between adjacent observations, we pull oto_{t} and ot+1o_{t+1} close to each other. The transition function is applied after the representation of oto_{t}, such that sts_{t} is transformed to s^t+1\hat{s}_{t+1}, which now represents the same entity as the other branch. Then both of st+1s_{t+1} and s^t+1\hat{s}_{t+1} go through a common projector network. Since st+1s_{t+1} is potentially a more accurate description of ot+1o_{t+1} compared to s^t+1\hat{s}_{t+1}, we make the ot+1o_{t+1} branch as the target branch. It is common in self-supervised learning that the second or the third layer from the last is chosen as the features for some reason. Here, we choose the outputs from the representation network or the dynamics network as the hidden states rather than those from the projector or the predictor. The two adjacent observations provide two views of the same entity. In practice, we find that applying augmentations to observations on the image helps to further improve the learned representation quality . We also unroll the dynamic function recurrently for 5 further steps and also pull s^t+k\hat{s}_{t+k} close to st+ks_{t+k} (k=1,...,5k=1,...,5). Please see App.A.1 for more details. We note that our temporal consistency loss is similar to SPR , an unsupervised representation method applied on rainbow. However, the consistency loss in our case is applied in a model-based manner, and we use the SimSiam loss function.

2 End-To-End Prediction of the Value Prefix

In model-based learning, the agent needs to predict the future states conditioned on the current state and a series of hypothetical actions. The longer the prediction, the harder to predict it accurately, due to the compounding error in the recurrent rollouts. This is called the state aliasing problem. The environment model plays an important role in MCTS. The state aliasing problem harms the MCTS expansion, which will result in sub-optimal exploration as well as sub-optimal action search.

Predicting the reward from an aliased state is a hard problem. For example, as shown in Figure 3, the right agent loses the ball. If we only see the first observation, along with future actions, it is very hard both for an agent and a human to predict at which exact future timestep the player would lose a point. However, it is easy to predict the agent will miss the ball after a sufficient number of timesteps if he does not move. In practice, a human will never try to predict the exact step that he loses the point but will imagine over a longer horizon and thus get a more confident prediction.

Inspired by this intuition, we propose an end-to-end method to predict the value prefix. We notice that the predicted reward is always used in the estimation of the Q-value Q(s,a)Q(s,a) in UCT of Equation 2

, where rt+ir_{t+i} is the reward predicted from unrolled state s^t+i\hat{s}_{t+i}. We name the sum of rewards ∑i=0k−1γirt+i\sum_{i=0}^{k-1}\gamma^{i}r_{t+i} as the value prefix, since it is used as a prefix in the later Q-value computation.

We propose to predict value prefix from the unrolled states (st,s^t+1,⋯ ,s^t+k−1s_{t},\hat{s}_{t+1},\cdots,\hat{s}_{t+k-1}) in an end-to-end manner, i.e. value-prefix=f(st,s^t+1,⋯ ,s^t+k−1)\text{value-prefix}=f(s_{t},\hat{s}_{t+1},\cdots,\hat{s}_{t+k-1}). Here ff is some neural network architecture that takes in a variable number of inputs and outputs a scalar. We choose the LSTM in our experiment. During the training time, the LSTM is supervised at every time step, since the value prefix can be computed whenever a new state comes in. This per-step rich supervision allows the LSTM can be trained well even with limited data. Compared with the naive per step reward prediction and summation approach, the end-to-end value prefix prediction is more accurate, because it can automatically handle the intermediate state aliasing problem. See Experiment Section 5 for empirical evaluations. As a result, it helps the MCTS to explore better, and thus increases the performance. See App.A.1 for architectural details.

3 Model-Based Off-Policy Correction

In MCTS RL algorithms, the value function fits the value of the current neural network policy. However, in practice as MuZero Reanalyze does, the value target is computed by sampling a trajectory from the replay buffer and computing: zt=∑i=0k−1γiut+i+γkvt+kz_{t}=\sum_{i=0}^{k-1}\gamma^{i}u_{t+i}+\gamma^{k}v_{t+k}. This value target suffers from off-policy issues, since the trajectory is rolled out using an older policy, and thus the value target is no longer accurate. When data is limited, we have to reuse the data sampled from a much older policy, thus exaggerating the inaccurate value target issue.

In previous model-free settings, there is no straightforward approach to fix this issue. On the contrary, since we have a model of the environment, we can use the model to imagine an "online experience". More specifically, we propose to use rewards of a dynamic horizon ll from the old trajectory, where l<kl<k and ll should be smaller if the trajectory is older. This reduces the policy divergence by fewer rollout steps. Further, we redo an MCTS search with the current policy on the last state st+ls_{t+l} and compute the empirical mean value at the root node. This effectively corrects the off policy issue using imagined rollouts with current policy and reduces the increased bias caused by setting ll less than kk. Formally, we propose to use the following value target:

where l<=kl<=k and the older the sampled trajectory, the smaller the ll. νMCTS(st+l)\nu^{\text{MCTS}}(\text{s}_{t+l}) is the root value of the MCTS tree expanded from st+ls_{t+l} with the current policy, as MuZero non-Reanalyze does. See App.A.4 for how to choose ll. In practice, the computation cost of the correction is two times on the reanalyzed side. However, the training will not be affected due to the parallel implementation.

Experiments

In this section, we aim to evaluate the sample efficiency of the proposed algorithm. Here, the sample efficiency is measured by the performance of each algorithm at a common, small amount of environment transitions, i.e. the better the performance, the higher the sample efficiency. More specifically, we use the Atari 100k benchmark. Intuitively, this benchmark asks the agent to learn to play Atari games within two hours of real-world game time. Additionally, we conduct some ablation studies to investigate and analyze each component on Atari 100k. To further show the sample efficiency, we apply EfficientZero to some simulated robotics environments on the DMControl 100k benchmark, which contains the same 100k environment steps.

Atari 100k Atari 100k was first proposed by the SimPLe method, and is now used by many sample-efficient RL works, such as Srinivas et al. , Laskin et al. , Kostrikov et al. , Schwarzer et al. . The benchmark contains 26 Atari games, and the diverse set of games can effectively measure the performance of different algorithms. The benchmark allows the agent to interact with 100 thousand environment steps, i.e. 400 thousand frames due to a frameskip of 4, with each environment. 100k steps roughly correspond to 2 hours of real-time gameplay, which is far less than the usual RL settings. For example, DQN uses 200 million frames, which is around 925 hours of real-time gameplay. Note that the human player’s performance is tested after allowing the human to get familiar with the game after 2 hours as well. We report the raw performance on each game, as well as the mean and median of the human normalized score. The human normalized score is defined as: (scoreagent−scorerandom)/(scorehuman−scorerandom)(\text{score}_{\text{agent}}-\text{score}_{\text{random}})/(\text{score}_{\text{human}}-\text{score}_{\text{random}}).

We compare our method to the following baselines. (1) SimPLe , a model-based RL algorithm that learns an action conditional video prediction model and trains PPO within the learned environment. (2) OTRainbow , which tunes the hyper-parameters of the Rainbow method to achieve higher sample efficiency. (3) CURL , which uses contrastive learning as a side task to improve the image representation quality. (4) DrQ , which adds data augmentations to the input images while learning the original RL objective. (5) SPR , the previous SoTA in Atari 100k which proposes to augment the Rainbow agent with data augmentations as well as a multi-step consistency loss using BYOL-style self-supervision. (6) MuZero with our implementations and the same hyper-parameters as EfficientZero. (7) Random Agent (8) Human performance.

DeepMind Control 100k Tassa et al. propose the DMControl suite, which includes some challenging visual robotics tasks with continuous action space. And some works have benchmarked for the sample efficiency on the DMControl 100k which contains 100k environment steps data. Since the MCTS-based methods cannot deal with tasks with continuous action space, we discretize each dimension into 5 discrete slots in MuZero and EfficientZero. To avoid the dimension explosion, we evaluate EfficientZero in three low-dimensional tasks.

We compare our method to the following baselines. (1) Pixel SAC, which applies SAC directly to pixels. (2) SAC-AE , which combines the SAC and an auto-encoder to handle image-based inputs. (3) State SAC, which applies SAC directly to ground truth low dimensional states rather than the pixels. (4) Dreamer , which learns a world model and is trained in dreamed scenarios. (5) CURL , the previous SoTA in DMControl 100k. (6) MuZero with action discretizations.

2 Results

Table 1 shows the results of EfficientZero on the Atari 100k benchmark. Normalizing our score with the score of human players, EfficientZero achieves a mean score of 1.904 and a median score of 1.160. As a reference, DQN achieves a mean and median performance of 2.20 and 0.959 on these 26 games. However, it is trained with 500 times more data (200 million frames). For the first time, an agent trained with only 2 hours of game data can outperform the human player in terms of the mean and median performance. Among all games, our method outperforms the human in 14 out of 26 games. Compared with the previous state-of-the-art method (SPR ), we are 170% and 180% better in terms of mean and median score respectively. As for more robust results, we record the aggregate metrics in App.A.5 with statistical tools proposed by Agarwal et al. .

Apart from the Atari games, EffcientZero achieves remarkable results in the simulated tasks. As shown in Table 2, EffcientZero outperforms CURL, the previous SoTA, to a considerable degree and keeps a smaller variance but MuZero cannot work well. Notably, EfficientZero achieves comparable results to the state SAC, which consumes the ground truth states and is considered as the oracles.

3 Ablations

In Section 4, we discuss three issues that prevent MuZero from achieving high performance when data is limited: (1) the lack of environment model supervision, (2) the state aliasing issue, and (3) the off-policy target value issue. We propose three corresponding approaches to fix those issues and demonstrate the usefulness of the combination of those approaches on a wide range of 26 Atari games. In this section, we will analyze each component individually.

Each Component Firstly, we do an ablation study by removing the three components from our full model one at a time. As shown in Table 3, we find that removing any one of the three components will lead to a performance drop compared to our full model. Furthermore, the richer learning signals are the aspect Muzero lacks most in the low-data regime as the largest performance drop is from the version without consistency supervision. As for the performance in the high-data regime, We find that the temporal consistency can significantly accelerate the training. The value prefix seems to be helpful during the early learning process, but not as much in the later stage. The off-policy correction is not necessary as it is specifically designed under limited data.

Temporal Consistency As the version without self-supervised consistency cannot work well in most of the games, we attempt to dig into the reason for such phenomenon. We design a decoder D\mathcal{D} to reconstruct the original observations, taking the latent states as inputs. Specifically, the architecture of D\mathcal{D} and the H\mathcal{H} are symmetrical, which means that all the convolutional layers are replaced by deconvolutional layers in D\mathcal{D} and the order of the layers are reversed in D\mathcal{D}. Therefore, H\mathcal{H} is an encoder to obtain state sts_{t} from observation oto_{t} and D\mathcal{D} tries to decode the oto_{t} from sts_{t}. In this ablation, we freeze all parameters of the trained EfficientZero network with or without consistency respectively and the reconstructed results are shown in different columns of Figure 4. We regard the decoder as a tool to visualize the current states and unrolled states, shown in different rows of Figure 4. Here we note that Mcon\mathcal{M}_{\text{con}} is the trained EfficientZero model with consistency and Mnon\mathcal{M}_{\text{non}} is the one without consistency. As shown in Figure 4, as for the current state sts_{t}, the observation is reconstructed well enough in the two versions. However, it is remarkable that the the decoder given Mnon\mathcal{M}_{\text{non}} can not reconstruct images from the unrolled predicted states s^t+k\hat{s}_{t+k} while the one given Mcon\mathcal{M}_{\text{con}} can reconstruct basic observations.

To sum up, there are some distributional shifts between the latent states from the representation network and the states from the dynamics function without consistency. The consistency component can reduce the shift and provide more supervision for training the dynamics network.

Value Prefix We further validate our assumptions in the end-to-end learning of value prefix, i.e. the state aliasing problem will cause difficulty in predicting the reward, and end-to-end learning of value prefix can alleviate this phenomenon.

To fairly compare directly predicting the reward versus end-to-end learning of the value prefix, we need to control for the dataset that both methods are trained on. Since during the RL training, the dataset distribution is determined by the method, we opt to load a half-trained Pong model and rollout total 100k steps as the common static dataset. We split this dataset into a training set and a validation set. Then we run both the direct reward prediction and the value prefix method on the training split.

As shown in Figure 5, we find that the direct reward prediction method has lower losses on the training set. However, the value prefix’s validation error is much smaller when unrolled for 5 steps. This shows that the value prefix method avoids overfitting the hard reward prediction problem, and thus it can reduce the state aliasing problem, reaching a better generalization performance.

Off-Policy Correction To prove the effectiveness of the off-policy correction component, we compare the error between the target values and the ground truth values with or without off-policy correction. Specifically, the ground truth values are estimated by Monte Carlo sampling.

We train a model for the game UpNDown with total 100k training steps, and collect the trajectories at different training stages respectively (20k, 40k, …, 100k steps). Then we calculate the ground truth values with the final model. We choose the trajectories at the same stage (20k) and use the final model to evaluate the target values with or without off-policy correction, following the Equation 4. We evaluate the L1 error of the target values and the ground truth, as shown in Table 4. The error of unrolled next 5 states means the average error of the unrolled 1-5 states with dynamics network from current states. The error is smaller in both current states and the unrolled states with off-policy correction. Thus, the correction component does reduce the bias caused by the off-policy issue.

Furthermore, we also ablate the value error of the trajectories at distinct stages in Table 5. We can find that the value error becomes smaller as the trajectories are fresher. This indicates that the off-policy issue is severe due to the staleness of the data. More significantly, the off-policy correction can provide more accurate target value estimation for the trajectories at distinct time-steps as all the errors with correction shown in the table are smaller than those without correction at the same stage.

Discussion

In this paper, we propose a sample-efficient model-based method EfficientZero. It achieves super-human performance on the Atari games with as little as 2 hours of the gameplay experience and state-of-the-art performance on some DMControl tasks. Apart from the full results, we do detailed ablation studies to examine the effectiveness of the proposed components. This work is one step towards running RL in the physical world with complex sensory inputs. In the future, we plan to extend it to more directions, such as a better design for the continuous action space. And we also plan to study the acceleration of MCTS and how to combine this framework with life-long learning.

Acknowledgments and Disclosure of Funding

This work is supported by the Ministry of Science and Technology of the People’s Republic of China, the 2030 Innovation Megaprojects “Program on New Generation Artificial Intelligence” (Grant No. 2021AAA0150000).

References

Appendix A Appendix

As for the architecture of the networks, there are three parts in our model pipeline: the representation part, the dynamics part, and the prediction part. The architecture of the representation part is as follows:

1 convolution with stride 2 and 32 output planes, output resolution 48x48. (BN + ReLU)

1 residual downsample block with stride 2 and 64 output planes, output resolution 24x24.

Average pooling with stride 2, output resolution 12x12. (BN + ReLU)

Average pooling with stride 2, output resolution 6x6. (BN + ReLU)

, where the kernel size is 3×33\times 3 for all operations.

As for the dynamics network, we follow the architecture of MuZero but reduce the residual blocks from 16 to 1. Furthermore, we add an extra residual link in the dynamics part to keep the information of historical hidden states during recurrent inference. The design of the dynamics network is listed here:

Concatenate the input states and input actions into 65 planes.

1 convolution with stride 2 and 64 output planes. (BN)

A residual link: add up the output and the input states. (ReLU)

In the prediction part, we use two-layer MLPs with batch normalization to predict the reward, value, or policy. Considering the stability of the prediction part, we set the weights and bias of the last layer to zero in prediction networks. As for the reward prediction network, it predicts the sum of the rewards, namely value prefix: rt,ht+1=R(s^t+1,ht)r_{t},h_{t+1}=\mathcal{R}(\hat{s}_{t+1},h_{t}), where rtr_{t} is the predicted sum of rewards, h0h_{0} is zero-initialized and hidden size of LSTM is 512. The architecture of the value prediction network is as follows:

1 1x1convolution and 16 output planes. (BN + ReLU)

1 fully connected layers and 32 output dimensions. (BN + ReLU)

1 fully connected layers and 601 output dimensions.

The horizontal length of the LSTM during training is limited to the unrolled steps lunroll=5l_{\text{unroll}}=5, but it will be larger in MCTS as the dynamics process can go deeper. Therefore, we reset the hidden state of LSTM after ζ=5\zeta=5 steps of recurrent inference, where ζ\zeta is the valid horizontal length.

The design of the reward and policy prediction networks are the same except for the dimension of the outputs:

1 1x1convolution and 16 output planes. (BN + ReLU)

1 fully connected layers and 32 output dimensions. (BN + ReLU)

1 fully connected layers and DD output dimensions.

, where D=601D=601 in the reward prediction network and DD is equal to the action space in the policy prediction network.

Here is the brief introduction of the training pipeline, taking one-step rollout as an example.

, where H\mathcal{H} is the representation network, G\mathcal{G} is the dynamics network, V\mathcal{V} is the value prediction network, P\mathcal{P} is the policy prediction network, R\mathcal{R} is the reward (value prefix) prediction network. ot,st,ato_{t},s_{t},a_{t} are observations, states and actions. hth_{t} is the hidden states in recurrent neural networks.

Here is the training loss, taking one-step rollout as an example:

, where L\mathcal{L} is the total loss of the unrolled lunrolll_{\text{unroll}} steps, L1\mathcal{L}_{1} is the Cross-Entropy loss, and L2\mathcal{L}_{2} is the negtive cosine similarity loss. Besides, P1P_{1} is a 3-layer MLP while P2P_{2} is a 2-layer MLP. The dimension of the hidden layers is 512 and the dimension of the output layers is 1024. We add batch normalization between every two layers in those MLP except the final layer. sg(P1)sg({P}_{1}) means stopping gradients.

We stack 4 historical frames, with an interval of 4 frames-skip. Thus the input effectively covers 16 frames of the game history. We stack the input images on the channel dimension, resulting in a 96×96×1296\times 96\times 12 tensor. We do not use any extra state normalization besides the batch norm and we choose reward clipping to keep better scales in the searching process.

Generally, compared with MuZero , we reduce the number of residual blocks and the number of planes as we find that there is no capability issue caused by much smaller networks in our EfficientZero with limited data. In another word, such a tiny network can acquire good performance in the limited setting.

For other details, we provide hyper-parameters in Table 6. It is notable that we train the model for 120k steps where we only collect data during the first 100k steps. In this way, the latter trajectories can be fully used in training. Besides, the learning rate will drop after every 100k training steps (from 0.2 to 0.02 at 100k).

A.2 More Ablations

In the experiment section , we list some ablation studies to prove the effectiveness of each component. In this section, we will display more results for the ablation study.

Firstly, the detailed results of the ablation study of each component are listed in Table 7. In this table, We find that the full version of EfficientZero outperforms the others without any one of the components. Furthermore, for those environments EfficientZero can already solve, the performance is similar between the full version and the version without off-policy correction, such as Breakout, Pong, etc. In such a case, the off-policy issue is not severe, which is the reason for this phenomenon. Besides, for some environments with sparse rewards, the value prefix component matters, such as Pong; and for those with dense rewards, the state aliasing problem has less negative effects for the reward signals are sufficient, such as Qbert. As for the version without self-supervised consistency, the results of all the environments are much poorer.

In addition, we do the ablation study for the data augmentation technique in the consistency component to examine the effect of data augmentations. We apply a random small shift of 0-4 pixels as well as the change of the intensity as the augmentation techniques. Here we choose several Atari games and train the model for 100k steps. The results are shown in Table 8. We can find that the version without data augmentation has similar performances while the version without consistency component is worse. This indicates that the improvement of the consistency component is basically from the self-supervised learning loss rather than the data augmentation.

Finally, we also do the ablation study for the MCTS root value and the dynamic horizon in the off-policy correction component. Here we choose several Atari games and train the model for 100k steps. As shown in Table 9, the version without dynamic horizon has poorer results than that without the MCTS root value. In the off-policy correction component, the dynamic horizon seems more important.

A.3 MCTS Details

Our policy searching approach is based on Monte-Carlo tree search (MCTS). We follow the procedure in MuZero , which includes three stages and repeats the searching process for Nsim=50N_{\text{sim}}=50 simulations. Here are some brief introductions for each stage.

Selection In the selection part, it targets at choosing an appropriate unvisited node while balancing exploration and exploitation with UCT:

, where Q(s,a)Q(s,a) is the average Q values after simulations, N(s,a)N(s,a) is the total visit counts at state ss by selecting action aa, and P(s,a)P(s,a) is the policy prior set in the expansion process. In each simulation, the MCTS starts from the root node s0s^{0}. And for each time-step k=1...lk=1...l of the simulation, the algorithm will select the action aka^{k} according to the UCT. Usually, c1=1.25c_{1}=1.25 and c2=19652c_{2}=19652 according to the literature .

However, the default Q value of the unvisted node is set to 0, which indicates the worst state. To give a better Q-value estimation of the unvisited nodes, we evaluate a mean Q value mechanism in each simulation for tree nodes, similar to the implementation of Elf OpenGo .

, where Q^(s)\hat{Q}(s) is the estimated Q value for unvisited nodes to make better selections considering exploration and exploitation. sroots^{\text{root}} is the state of the root node and sparents^{\text{parent}} is the state of the parent node of ss. In experiments, we find that the mean Q value mechanism gives a better exploration than the default one.

Expansion Then the newly selected node will be expanded with the predicted reward and policy as its prior. Furthermore, when the root node is to expand, we apply the Dirichlet noise to the policy prior during the self-play stage and the reanalyzing stage to give more explorations.

, where ND(ξ)\mathcal{N}_{\mathcal{D}}(\xi) is the Dirichlet noise distribution, ρ,ξ\rho,\xi is set to 0.25 and 0.3 respectively. However, we do not use any noise and set ρ\rho to 0 instead for those non-root node or during evaluations.

Backup After selecting and expanding a new node, we need to backup along the current searching trajectory to update the Q(s,a)Q(s,a). Considering the scales of values in distinct environments, we compute a normalized Q-value by using the minimum-maximum values calculated along with all visited tree nodes, which is applied in MuZero. However, when the data is limited, the small difference between the minimum and maximum values will result in overconfidence in UCT calculation. For example, when all the Q-values in those visited tree nodes are in a range of 0 to 10−410^{-4}, the normalized Q-value of 10−510^{-5} and 5×10−55\times 10^{-5} will make a huge difference as one is normalized to 0.10.1 and another is 0.50.5. Therefore, we set a threshold here to reduce overconfidence in such occasions, which is called the soft minimum-maximum updates:

, where ϵ\epsilon, the threshold to give a smooth range of the min-max bound, is set to 0.010.01.

After all the expansions in the MCTS, we will obtain average value and visit count distributions of the root node. Here, the root value can be applied in off-policy correction and the visit count distribution is the target policy distribution:

We decay the temperature of the MCTS output policy distribution here twice during training, at 50% and 75% of the training progress to 0.5 and 0.25 respectively.

A.4 Training Details

In this subsection, we will introduce more training details.

Pipeline As for the code implementation of EfficientZero, we design a paralleled architecture with a double buffering mechanism in Pytorch and Ray, as shown in Figure 6.

Intuitively, we will describe the training process in a synchronized way. Firstly, the data workers called self-play actors are aimed at doing self-play with the given model updated within 600 training steps and then they will send the rolled-out trajectories into the replay buffer. Then the CPU rollout workers attempt to prepare the contexts of those batch transitions sampled from the replay buffer, in which way only CPU resources are required. Afterward, the GPU batch workers reanalyze those past data with the given contexts by the given target model, and most of the time-consuming parts in this procedure are in GPUs. Considering the frequent utilization of CPUs and GPUs in MCTS, the searching process is assigned for those GPU workers. Finally, the learner will obtain the reanalyzed batch and begin to train the agent.

The learner, all the data workers, CPU workers, and GPU workers start in parallel. The data workers and CPU workers share the replay buffer to sample data while the CPU and GPU workers share a context queue for reanalyzing data. Besides, the learner and the GPU workers use a batch queue to communicate. In such a design, we can utilize the CPU and GPU as much as possible.

Self-play During self-play, the priorities of the transition to collect are set to the max of the whole priorities in replay buffer. We also update the priority in EfficientZero according to MuZero : P(i)=piα∑kpkαP(i)=\frac{p_{i}^{\alpha}}{\sum_{k}p_{k}^{\alpha}}, where pip_{i} is the L1 error of the value during training. And the we scale with important sampling ratio wi=(1N×P(i))βw_{i}=(\frac{1}{N\times P(i)})^{\beta}. We set α\alpha to 0.6 and anneal β\beta from 0.40.4 to 1.01.0, following prioritized replay . However, we find the priority mechanism only improves a little with limited data. Considering the long horizons in atari games, we collect the intermediate sequences of 400 moves.

Reanalyze The reanalyzed part is introduced in MuZero , which revisits the past trajectories and re-executes the data with lasted target model to obtain a fresher value and policy with model inference as well as MCTS.

For the off-policy correction, the target values are reanalyzed as follows:

, where kk is the TD steps here, and is set to 5; TcurrentT_{\text{current}} is the current training steps, TstT_{s_{t}} is the training steps of collecting the data sts_{t}, TtotalT_{\text{total}} is the total training steps (100k), and τ\tau is a coefficient which is set to 0.3. Intuitively, ll is to define how fresh the collected data sts_{t} is. When the trajectory is stale, we need to unroll less to estimate the target values for the sake of the gaps between current model predictions and the stale trajectory rollouts. Besides, we replace the predicted value vt+kv_{t+k} with the averaged root value from MCTS νt+lMCTS\nu^{\text{MCTS}}_{t+l} to alleviate the off-policy bias.

Notably, we re-sample Dirichlet noise into the MCTS procedure in reanalyzed part to improve the sample efficiency with a more diverse searching process. Besides, we reanalyze the policy among 99% of the data and reanalyze the value among 100% data.

A.5 Evaluation

We evaluate the EfficientZero on Atari 100k benchmark with a total of 26 games. Here are the evaluation curves during training, as shown in Figure 7.

Besides, we also report the scores for 3 runs (different seeds) with 32 evaluation seeds across the 26 Atari games, which is shown in Table 10.

Recently, Agarwal et al. propose to use statistical tools to present more robust and efficient aggregate metrics. Here we display the corresponding results based on its open-sourced codebase. Figure 8 illustrates that EfficientZero significantly outperforms the other methods on Atari 100k benchmark concerning all the metrics.

A.6 Open Source EfficientZero Implementation

MCTS-based RL algorithms present a promising future research direction: to achieve strong performance with model-based methods. However, two major practical obstacles prevent them from being widely used currently. First, there are no high-quality open-source implementations of these algorithms. Existing implementations can only deal with simple state-based environments, such as CartPole . Accurately scaling to complex image input environments requires non-trivial engineering efforts. Second, MCTS RL algorithms such as MuZero require a large number of computations. For example, MuZero needs 64 TPUs to train 12 hours for one agent on Atari games. The high computational costs pose problems both for the future development of such methods as well as practical applications.

We think our open-source implementation of EfficientZero can drastically accelerate the research in MCTS RL algorithms. Our implementation is computationally friendly. To train an Atari agent for 100k steps, it only needs 4 GPUs to train 7 hours. Our framework could potentially have a large impact on many real-world applications, such as robotics since it requires significantly fewer samples.

Our open-source framework aims to provide an easy way to understand the implementation while keeping relatively high compute efficiency. As shown in Fig. 9, the system is composed of four components: the replay buffer, the experience sampling actor, the reanalyze training target preparation module, and the training component.

To make sure the framework is easy to use, we implement them based on Ray , and the four components are implemented as ray actors which run in parallel. The main computation bottleneck is in the reanalyze module, which samples from the replay, and runs an MCTS search on each observation. To accelerate the reanalyze module, we split the reanalyze computation into the CPU part and the GPU part, such that computation on CPU and GPU are run in parallel. We use a different number of actors between CPU and GPU to match their total throughput. To increase the throughput on GPU, we also collocate multiple batch computation threads on one GPU, as in Tian et al. . We also implement the MCTS in C++ to avoid performance issues with Python on large amounts of atomic computations.

We implement the MCTS by a couple of important techniques, which are quite crucial to improve the efficiency of the MCTS process. On the one hand, we implement batch MCTS to allow the agent to search a batch of trees in parallel, to enlarge the throughput of MCTS during self-play and reanalyzing targets. On the other, we choose C++ in the MCTS process. However, the process of MCTS needs to do searching as well as model inference, which needs to communicate with Pytorch. Therefore, we use Python to do model inference, C++ to do other atomic computations, and Cython to communicate between Python contexts and C++ contexts. In another word, we use pure C++ to do selection, expansion, and backup while using neural networks in Python. Meanwhile, we build a database to store the hidden states in Python while storing the corresponding data index during the searching process in C++. For more details of the implementation, please refer to https://github.com/YeWR/EfficientZero.