DreamerV3 in a nutshell#
DreamerV3 is a model-based reinforcement learning algorithm. It learns a compact model of the environment from replayed experience, then trains an actor and a critic on trajectories generated inside that model. The real environment supplies data for the world model; most policy improvement happens in latent-space imagination.
Paper and maintained implementation#
This page treats the DreamerV3 paper as the source of truth for the algorithm. The author-maintained JAX implementation continues to evolve and its named experiment presets can differ from the protocol reported in the paper. TorchRL documents those presets as separate reproduction targets rather than redefining the paper algorithm around the latest JAX configuration.
Some constructor defaults predate full paper parity and remain for backward compatibility. The runnable DreamerV3 recipes pass the paper-compatible loss settings explicitly; changes to public defaults require the normal deprecation cycle.
The high-level data flow is:
real transition sequences
|
v
encoder -> posterior RSSM state -> reconstruction, reward, continuation
^ |
| v
prior dynamics + action
|
v
imagined trajectories
|
v
actor + online critic
|
v
slow (target) critic
Nomenclature#
Dreamer papers and implementations use several names for closely related objects. In the TorchRL API:
Term |
Meaning |
|---|---|
World model |
The observation encoder, recurrent state-space model (RSSM), observation decoder, reward predictor, and optional continuation predictor. |
Belief or deterministic state ( |
The recurrent hidden state that summarizes history. TorchRL stores it
under the |
State or stochastic state ( |
A sample from the RSSM’s categorical latent variables. TorchRL stores
the flattened straight-through one-hot sample under |
Prior, dynamics, or transition model |
Predicts the next categorical state from the previous state, belief, and
action, without seeing the next observation. See
|
Posterior or representation model |
Corrects the prior using the encoded next observation. It is used while
learning from real sequences. See
|
Imagination |
A latent rollout that uses the prior, reward model, continuation model, and actor, but no real observations. |
Critic or value model |
The online network that predicts lambda returns from the RSSM state and belief. |
Slow critic, target critic, or EMA critic |
A lagged copy of the online critic. It is updated by Polyak averaging and provides a stable auxiliary target for critic regularization. “Slow” refers to its parameter updates, not its optimizer or runtime. |
Continuation |
The learned probability that an imagined trajectory continues. It replaces a fixed survival assumption when weighting returns and losses. |
For acting in a real environment, compose the encoder,
RSSMStateEstimatorV3 and actor with
TensorDictSequential. The estimator resets recurrent
context per stream and samples only the observation-conditioned posterior.
The collector carries its state, belief and action into the next step.
For discrete actions, DreamerV3DiscreteActor provides
an importable one-hot policy with DreamerV3 initialization, uniform probability
mixing and float32 logits under autocast. Its get_dist() method supports
straight-through sampling for imagination, and its input and output keys can
be nested. Network construction and sampling require no recipe imports.
How the RSSM works#
The RSSM splits its latent representation into a deterministic recurrent state and a stochastic categorical state. At each real-data time step:
RSSMPriorV3updates the belief from the previous stochastic state and action, then predicts prior categorical logits.RSSMPosteriorV3combines that belief with the encoded observation and predicts posterior categorical logits.Both networks sample hard categorical states with a straight-through gradient estimator.
unimixcan mix a small uniform component into the probabilities to prevent overconfident categories.RSSMRolloutV3carries the posterior state and belief through a sequence and resets them at episode boundaries.
During imagination there is no observation, so only the prior advances the
latent state. The actor and prediction heads consume both state and
belief. RSSMPriorV3 supports a conventional GRU and the grouped
"block_gru" core used by the full DreamerV3 example. The accompanying
DreamerV3MLP provides the RMS-normalized SiLU MLP
blocks used by the example’s encoder, decoder, actor, critic, and prediction
heads.
For recurrent features outside an RSSM, DreamerV3BlockGRUCell
exposes the same block-diagonal update as a single-step module, while
DreamerV3BlockGRU executes batch-major sequences with
mixed episode resets.
Selecting the sequence backend#
The sequence backend is selected directly on the high-level module:
from torchrl.modules import DreamerV3BlockGRU
gru = DreamerV3BlockGRU(
input_size=512,
hidden_size=512,
recurrent_backend="triton",
).cuda()
The three backends trade portability for speed:
"reference"(default) runs the time loop with ordinary autograd. It works on every supported device, floating dtype, and elementwise activation, and it is the only backend that supports double backward (create_graph=True). It is the slowest option on long sequences."scan"fuses the time loop throughtorch._higher_order_ops.scanand carries only the hidden cotangent in a specialized reverse scan. It runs on CPU and CUDA, requires a recent PyTorch with thehoptorchpackage, and supports the same activations as the reference backend. Mixed input/hidden dtypes are promoted like the reference backend. Its backward consumes saved gate states, so double backward raises instead of silently returning wrong second-order gradients."triton"fuses the complete forward and reverse-time recurrences into one CUDA kernel each, keeping the carry on-chip across the whole horizon. It requires an NVIDIA GPU and Triton 3.3 or newer, supportsnn.SiLU,nn.Tanh, andnn.ReLUdynamics, and runs infloat32orbfloat16(mixed input and hidden dtypes are promoted like the reference backend; other dtypes raise an error). Parameters stay infloat32and accumulation is performed infloat32in both directions. Like the scan backend, double backward raises. Kernels are autotuned, so the first calls for a new sequence-length/width configuration pay a tuning warmup. On DreamerV3-sized workloads it is roughly an order of magnitude faster than the scan backend in both directions.
Select "scan" or "triton" explicitly so missing dependencies or
unsupported devices are reported instead of silently changing execution; the
optimized backends never fall back to another implementation.
To compare the backends on your own shapes and hardware (synchronized forward and backward timings, peak memory, and 95% confidence intervals), run the developer benchmark from a source checkout:
python benchmarks/bench_rnn_backward.py --rnn block_gru \
--backends reference,scan,triton --batches 16 --seq-lens 64,512 \
--hiddens 512 --input-size 512 --projection-size 512 --blocks 8 \
--dtype bfloat16 --warmup 10 --iters 30
Use the batch size, sequence length, widths, block count, dtype, and compile modes from the intended workload: backend performance is hardware- and shape-dependent.
The three objectives#
World model#
DreamerV3ModelLoss trains the model on real
transition sequences. Its components are:
a dynamics KL that trains the prior toward a stopped-gradient posterior;
a representation KL that trains the posterior toward a stopped-gradient prior;
free nats and optional uniform mixing for the categorical distributions;
an L1 or L2 reconstruction loss in symlog space;
a reward loss using symlog-spaced two-hot bins, or symlog MSE; and
an optional binary continuation loss.
kl_mode="separate" exposes the dynamics and representation KL terms
separately, as used by the full example. kl_mode="balanced" provides the
combined balanced-KL form.
Actor#
DreamerV3ActorLoss starts from posterior states
produced by the world model and rolls the actor through a
DreamerEnv. It computes lambda
returns from predicted rewards, values, and optional continuation
probabilities. It supports:
TD(0), TD(1), and TD(lambda) return estimators;
REINFORCE with a stopped-gradient advantage, or reparameterization gradients for suitable continuous policies;
an entropy bonus;
cumulative discount/continuation weighting; and
EMA percentile-range normalization of REINFORCE returns.
Critic and slow critic#
DreamerV3ValueLoss fits the online critic to the
lambda returns produced by the actor loss. The critic can use symlog MSE or a
distributional two-hot cross-entropy loss.
Setting slow_critic_regularization to a positive value creates target
critic parameters inside the value loss. The slow critic is a soft-updated
copy of the online critic:
The slow critic’s stopped-gradient prediction is an additional target for the online critic. In the current TorchRL objective, the online critic still provides the bootstrap values used to form imagined lambda returns; the slow critic regularizes critic learning rather than replacing that bootstrap.
Target updates are deliberately external to the loss. Associate a
SoftUpdate with the value loss and call it after
each critic optimizer step:
from torchrl.objectives import DreamerV3ValueLoss
from torchrl.objectives.utils import SoftUpdate
value_loss = DreamerV3ValueLoss(
value_model,
value_loss="two_hot",
actor_loss=actor_loss,
slow_critic_regularization=1.0,
)
slow_critic_updater = SoftUpdate(value_loss, tau=0.02)
# After loss.backward() and optimizer.step():
slow_critic_updater.step()
Replay critic loss#
The author-maintained JAX implementation also fits the critic on the real
replay sequences, not only on imagined trajectories.
replay_value_loss() computes that
term. Its return at each replay state uses the following replay reward and
bootstraps from the first imagined lambda return of the next state, so the
critic is fitted on real replay states as well as imagined states. The method
reads its
reward, done, terminated and bootstrap entries through
tensor_keys, so
set_keys() can redirect them:
value_loss.set_keys(bootstrap="first_imagined_return")
replay_td = value_loss.replay_value_loss(replay_features)
loss = replay_td["loss_replay_value"]
Because the input features stay attached, this term also trains the RSSM representation when the world-model loss returns live features.
Reconstruction heads#
Vector and image reconstruction can be composed in the public model loss.
Use symlog distance for vector observations and raw distance for normalized
images. Integer images are scaled by 255 when symlog is disabled. Each head
sums its event dimensions before the batch/time average, unless
global_average=True; the resulting head losses are added together.
model_loss = DreamerV3ModelLoss(world_model, reco_symlog=[True, False])
model_loss.set_keys(
pixels=[("sensors", "vector"), ("sensors", "image")],
reco_pixels=["reco_vector", "reco_pixels"],
)
losses, posterior = model_loss(replay_sample)
losses["loss_model_reco"].backward()
Optimization and training loop#
The public DreamerV3Loss composes the model, actor,
critic and replay-value objectives. Its detached replay_context output can
be written back through native replay’s generation-checked update operation.
DreamerV3OptimizationStepper owns the
forward/backward, optimizer and target-update sequence. It can run inside a
Trainer or a custom loop using step(None, sample). It returns scalar
metrics and writes detached posterior features to sample["replay_context"]
for replay updates. Its warmup(sample)
method prepares compilation and capture before collection starts, preserving
normalization buffers, RNG state and captured gradient storage. Keep shared
modules in one compile scope; do not compile the RSSM separately when using
whole-step compilation.
|
Compose DreamerV3 world-model, imagination and replay-value objectives. |
|
Hydra configuration for |
The loss modules do not create optimizers. This keeps optimizer ownership and the update schedule explicit. A typical update cycle is:
Sample contiguous real transition sequences from replay.
Update the world model on KL, reconstruction, reward, and continuation losses.
Detach posterior states from the real sequence and use them as imagination starting points.
Update the actor on imagined lambda returns.
Update the online critic on those same detached returns.
Soft-update the slow critic.
The public DreamerV3Optimizer jointly optimizes
the world model, actor and critic parameters, reproducing the current JAX
implementation’s adaptive gradient clipping, LaProp-style RMS scaling followed
by momentum, and warmup chain.
Those choices belong to the training recipe rather than the loss API, so users
can substitute another optimizer or schedule without changing the objectives.
API map#
Component |
Purpose |
|---|---|
|
Categorical latent dynamics and deterministic recurrent update. |
|
Observation-conditioned categorical representation model. |
|
Sequential prior/posterior filtering over replayed trajectories. |
RMS-normalized MLP building block. |
|
Symlog-spaced categorical scalar encoder, decoder, and loss helper. |
|
World-model objective. |
|
Latent-imagination actor objective and lambda-return construction. |
|
Online and slow-critic objective. |
|
|
External Polyak update for the slow critic. |
Scale-robust scalar transformations for custom heads and losses. |
For a complete training setup, see the DreamerV3 example.
Checkpoint ownership#
Register the composed DreamerV3Loss and
DreamerV3OptimizationStepper separately with
Checkpoint. The loss owns network and normalization
state; the stepper owns optimizer and target-update progress. Restore both after
compile/capture warm-up, which may modify parameters or running statistics.
DreamerV3SeededPolicy includes its seed and next-draw
counter in its module state. Register the policy and
DreamerV3UpdateRatio directly to preserve
policy RNG progress and fractional update scheduling. When replay is omitted,
call the ratio’s reset(record_count) while refilling it to discard owed updates.
Restore global RNG state last, before collection begins.
After loading native replay, call
end_streams() before new environments
append transitions. This closes unfinished tails and resets incomplete streaming
windows without inventing terminal transitions. Environments restart; checkpoint
resume does not guarantee identical trajectories.
Pass saved logger state to get_logger() through
state_dict=.... The logger layer reopens saved local logs or strictly resumes
the saved W&B run ID, then restores its counters. The recipe chooses checkpoint
paths, cadence, optional replay and when to quiesce collection; it does not inspect
component internals to reconstruct their state.