WorldSense 技术笔记

DreamerV3 训练工程实践:从 GPU 配置到超参调优

2026年8月28日 · 阅读约34分钟 · DreamerV3, 训练技巧, GPU, 超参数, 工程实践, Dreamer系列
目录

Dreamer 系列 · 第 3 篇

系列目录(当前在第 3 篇):

  1. (一)读懂 Dreamer:世界模型是怎么学会’想象’的?
  2. (二)Dreamer 的 Actor-Critic:想象空间里的策略优化
  3. (三)DreamerV3 训练工程实践:从 GPU 配置到超参调优

前两篇文章从理论层面讲清楚了 Dreamer 的架构设计和 Actor-Critic 工作原理。但要把 DreamerV3 跑起来、训得好,还需要解决一堆工程问题:GPU 显存不够怎么办?超参数怎么调?训练不稳定怎么排查?

这篇文章从实战角度出发,总结 DreamerV3 训练中的工程经验和常见坑。内容基于 danijar/dreamerv3@e3f02248 JAX reference implementation。

一、环境配置:JAX + GPU

JAX 的 GPU 支持

DreamerV3 使用 JAX 实现。JAX 的 GPU 支持依赖 CUDA 和 cuDNN,安装时需要注意版本匹配:

# 推荐安装方式(以 CUDA 12.x 为例)
pip install --upgrade "jax[cuda12]" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html

验证 GPU 是否被正确识别:

import jax
print(jax.devices())  # 应该看到 GPU 设备

如果输出只有 [CpuDevice],说明 JAX 没有正确检测到 GPU。常见原因包括 CUDA 版本不匹配、cuDNN 未安装、或者环境变量配置问题。

需要注意的是:nvidia-smi 能看到 GPU 不等于 JAX 一定可用——CUDA driver 版本必须 ≥ runtime 版本,而且多版本 CUDA 环境容易互相污染。建议用更详细的诊断命令:

import jax
jax.print_environment_info()  # 打印 JAX、CUDA、cuDNN 版本信息

显存管理

JAX 默认会预分配几乎所有 GPU 显存。在多任务训练或调试时,这可能造成问题。可以通过环境变量控制:

变量作用
XLA_PYTHON_CLIENT_PREALLOCATE=false关闭启动时大块预分配
XLA_PYTHON_CLIENT_MEM_FRACTION=0.9限制预分配比例为 90%
# 关闭预分配(按需分配)
export XLA_PYTHON_CLIENT_PREALLOCATE=false

# 或者限制预分配比例(仍预分配,但只占 90%)
export XLA_PYTHON_CLIENT_MEM_FRACTION=0.9

注意:MEM_FRACTION=0.9 不是"按需分配",而是"预分配最多 90% 显存"。真正关闭预分配需要用 PREALLOCATE=false

对于 24GB 显存的 GPU(如 RTX 3090/4090),DreamerV3 的默认配置通常可以跑通。如果显存紧张,需要调整 batch size 或 imagination length。

二、关键超参数

DreamerV3 的超参数在 configs.yaml 中定义。理解每个参数的作用,才能在具体任务上调得好。

世界模型相关

# RSSM 结构
rssm:
  deter: 4096      # deterministic state 维度
  stoch: 32        # categorical variables 数量
  classes: 32      # 每个 categorical variable 的类别数
  hidden: 4096     # RSSM 内部 MLP 隐藏层维度

# imagination
imag_length: 15    # 想象轨迹长度

注意:stoch=32classes=32 表示 stochastic state 由 32 个 categorical variables 组成,每个变量有 32 个类别。实际 latent 是 32 × 32 的 categorical distribution,而不是 32 维的连续向量。

KL divergence 在这个 categorical distribution 上以两种形式出现:

两者的主要区别不仅在 KL 方向,还在于梯度更新对象不同:dynamics loss 主要更新 prior 网络,representation loss 主要更新 posterior encoder。这种分离的梯度流设计是 DreamerV3 防止 latent collapse 和保持预测能力的重要机制。

imag_length 是最重要的超参数之一。15 步是 DreamerV3 的默认值,在大多数任务上表现良好。

但要注意一个常见误解:长 imagination 不等于更好的 long horizon credit assignment。世界模型的预测误差 $p(s_{t+k})$ 随步数 $k$ 增大而快速累积,所以:

如果任务需要更长的 horizon 才能看到奖励信号,可以适当增加到 20-25,但前提是世界模型本身足够准确。否则,长 imagination 反而会把策略带偏。

Actor-Critic 相关

# Actor
policy:
  layers: 3
  units: 1024
  act: silu
  norm: rms
  minstd: 0.1      # 连续动作最小标准差
  maxstd: 1.0      # 连续动作最大标准差
  outscale: 0.01   # 输出层权重缩放
  unimix: 0.01     # 离散动作均匀混合系数

# Critic
value:
  layers: 3
  units: 1024
  output: symexp_twohot
  bins: 255        # two-hot distribution 的 bin 数量

# 训练
imag_loss:
  lam: 0.95        # lambda-return 的 lambda
  actent: 3e-4     # 熵正则系数
  slowtar: False   # 是否使用 slow value 作为 target
  slowreg: 1.0     # slow value network 正则权重

slowvalue:
  rate: 0.02       # EMA 更新速率
  every: 1         # 每步更新

actent 控制探索程度。默认 3e-4 在大多数任务上表现良好。但要注意,DreamerV3 中探索不是简单靠 actent 一个参数控制的——Actor loss 是 imagination return + entropy,而 policy distribution 本身受 minstd/maxstd(连续动作)和 unimix(离散动作)控制。

如果需要调整探索行为,经验优先级(非官方推荐顺序,不同任务可能不同):

  1. minstd / maxstd(连续动作影响最大):直接改变策略分布的标准差范围
  2. unimix(离散动作):控制均匀混合系数,影响离散动作的随机性
  3. actent(entropy coefficient):作为最后微调手段

很多情况下,默认 actent 已经足够,问题往往出在 minstd/maxstd 的设置上。

lam 控制 TD 和 Monte Carlo 之间的权衡。0.95 偏向 Monte Carlo(多步回报),适合 imagination 比较准确的场景。如果世界模型不够准确,可以减小到 0.9 或 0.85,增加 TD bootstrap 的比重来降低方差。

优化器相关

lr: 4e-5           # 学习率
opt: {eps: 1e-20, clip: 1000.0}  # Adam 优化器参数

DreamerV3 使用较小的学习率 4e-5,这是训练稳定的关键之一。clip: 1000.0 是梯度裁剪阈值,防止梯度爆炸。

三、显存优化实战

问题:OOM(Out of Memory)

当 batch size 较大或 imagination length 较长时,容易遇到显存不足。以下是几种解决方案:

方案 1:减小 batch size

batch_size: 8      # 从 16 减到 8
batch_length: 512  # 序列长度也可以适当减小

方案 2:减小 imagination length

imag_length: 10    # 从 15 减到 10

这会减少 imagination 阶段的计算量,但可能影响长期信用的分配。

方案 3:使用梯度检查点

JAX 支持梯度检查点(gradient checkpointing),用计算换显存:

# 在关键函数上添加 gradient checkpointing
import jax

@jax.remat
def expensive_function(x):
    # ...

但这会增加约 20-30% 的训练时间。

方案 4:混合精度训练(谨慎使用)

DreamerV3 默认使用 float32。如果需要混合精度,强烈推荐 bfloat16 而非 float16,因为 bfloat16 的指数范围更大,数值更稳定。

但即使使用 bfloat16,也要谨慎:DreamerV3 对数值稳定高度敏感,涉及 KL balancing、categorical logits、two-hot value distribution、symlog/symexp 等组件。

推荐策略:

# JAX 中设置 matmul 精度(谨慎使用)
from jax import config
config.update("jax_default_matmul_precision", "bfloat16")

注意:jax_default_matmul_precision 控制的是矩阵乘法内部精度选择,不是模型 dtype。它不会自动完成 dtype casting,不等同于 PyTorch 的 AMP(Automatic Mixed Precision)。JAX 的精度控制(X64/X32/bf16)是全局或逐操作的显式设置,而 PyTorch AMP 通过 autocast 自动为不同操作选择合适精度。开了 jax_default_matmul_precision 不等于"开启了 bf16 training",实际效果仅限于 matmul 运算的内部精度。

注意:float16 在 DreamerV3 中很容易导致 KL collapse、value explosion 或 NaN,不建议使用。

显存监控

训练时监控显存使用情况:

# 实时监控
watch -n 1 nvidia-smi

# 或者使用更详细的工具
nvtop

如果发现显存使用量在训练过程中逐渐增加,可能存在显存泄漏。JAX 的显存管理通常比较稳定,但某些自定义操作可能导致问题。

四、训练稳定性排查

问题 1:Reward 不增长

如果训练过程中,真实环境的累积奖励不增长,可能原因:

世界模型没有学好

检查世界模型的 loss 曲线:

如果世界模型 loss 不下降,可能是学习率太大或太小,或者 encoder/decoder 结构不适合当前任务。

Actor-Critic 没有从想象中提取到有效信号

检查 Actor-Critic 的 loss:

如果 advantage 方差很大,可能是 value network 没有学好,可以检查 slow value regularization(slowreg)、lam、reward scale 等因素。slowreg 不是第一调节项——先排查 reward scale 和 lambda 设置是否合理。

问题 2:训练突然崩溃

训练前期正常,但某个 step 突然 loss 变成 NaN 或 Inf:

梯度爆炸

检查梯度范数。DreamerV3 有梯度裁剪,但如果 clip 值设置太大,可能无法有效防止爆炸。可以尝试:

opt: {clip: 100.0}  # 从 1000 减到 100

数值溢出

two-hot distribution 的 bins 覆盖的是 symlog 空间,映射回原始尺度后覆盖极大范围(约 ±4.8×10⁸),但训练主要发生在压缩后的 symlog 空间。DreamerV3 通过 symlog 压缩动态范围,使得 value prediction 不需要直接预测巨大的原始 return。

如果 value prediction 出现 NaN,检查:

imagination 起点质量差

如果 replay buffer 中某些序列质量很差(比如 episode 很短,或者 reward 异常大),可能导致 imagination 产生异常值。可以检查 replay buffer 中序列的分布。

问题 3:不同 seed 结果差异大

如果换随机 seed 后训练结果差异很大,说明训练不稳定:

增大 replay buffer

replay_size: 5e6   # 从 1e6 增大到 5e6

更大的 replay buffer 可以平滑采样分布,减少方差。

减小学习率

lr: 2e-5           # 从 4e-5 减到 2e-5

增大 batch size

batch_size: 32     # 从 16 增大到 32

更大的 batch size 可以提供更稳定的梯度估计。

多 seed 评估

RL 实验的 seed 敏感性是已知问题。如果你在做严肃的实验对比,建议:

问题 4:Replay buffer 采样与训练节奏

Dreamer 对 replay buffer 的采样策略非常敏感,但这一点经常被忽略。只关注 replay_size 是不够的,还需要注意:

Replay ratio(训练/采集比)

Dreamer 不是"先训世界模型再训 Actor",而是联合更新

采集数据 → 联合更新:世界模型 + Actor-Critic → 采集更多数据 → ...

DreamerV3 通常采用较高 replay ratio,但具体取决于环境速度和配置——对于仿真速度快的环境,ratio 可能远大于 1;对于真实机器人等慢速环境,ratio 可能接近 1。这个比例直接影响训练效率:

DreamerV3 默认的 ratio 在大多数任务上已经比较合理,但如果训练效率不理想,可以调整。

Warmup 阶段

训练初期,replay buffer 中的数据量很少,世界模型还没有学好。这个阶段需要注意:

Train/eval ratio

训练和评估的比例也很重要。评估太频繁会浪费训练时间,评估太少则无法及时发现问题。建议每 1000-5000 步评估一次。

问题 5:世界模型坍缩诊断

世界模型坍缩(world model collapse)是 Dreamer 训练中最隐蔽的失败模式。与 reward 不增长不同,坍缩时 loss 曲线可能看起来完全正常。

什么是坍缩?

RSSM 的 latent space 逐渐失去信息表征能力。具体表现为:

诊断方法:

  1. 监控 posterior entropyprior entropy 的变化趋势
  2. 如果两者都持续下降且差距缩小,可能正在坍缩
  3. Latent embedding visualization(如 PCA/t-SNE)仅作为辅助参考——RSSM 的 categorical latent 和 temporal structure 使得 t-SNE 不一定可靠。更推荐的诊断手段是 latent probing、linear probe、或 rollout reconstruction

修复策略:

五、任务特定的调优建议

Atari 游戏

Atari 游戏的 observation 是像素图像,reward 通常是稀疏的(只有得分变化时才有奖励)。

推荐配置调整:

# 可以尝试增大 imagination length(需要确认 world model prediction quality)
imag_length: 20    # 先检查 imagined rollout 质量再决定是否增大

# 使用更激进的探索
policy:
  unimix: 0.05     # 增大均匀混合,增加随机性

# 调整 lambda,因为世界模型在像素空间可能不够准确
imag_loss:
  lam: 0.9         # 减小 lambda,增加 TD 比重

注意事项:

MuJoCo 控制

MuJoCo 是连续控制任务,observation 是低维状态向量,action 也是连续的。

推荐配置调整:

# MuJoCo 通常不需要太长的 imagination
imag_length: 15    # 默认值通常够用

# 连续动作的标准差范围
policy:
  minstd: 0.1
  maxstd: 1.0

# 默认 entropy 通常足够,只有在探索不足时才增大
imag_loss:
  actent: 3e-4     # 默认值,如果探索不足再考虑增大到 1e-3

注意事项:

机器人操控

机器人操控任务通常有较高的状态维度和较复杂的动力学。

推荐配置调整:

# 机器人任务可能需要更长的 imagination 来理解因果
imag_length: 20

# 增大 RSSM 容量
rssm:
  deter: 4096
  hidden: 4096

# 如果 episode 经常提前终止,continuation model 很重要
# 确保 continuation prediction loss 正常下降

注意事项:

机器人任务特有的工程问题:

六、训练监控与日志

关键指标

训练时应该监控以下指标:

世界模型指标:

latent space 健康度(关键!):

很多 Dreamer 训练失败不是 reward loss 不下降,而是 latent collapse——latent space 的信息量逐渐坍缩到接近零,posterior entropy 持续下降。这种情况下,世界模型 loss 可能看起来正常,但策略已经无法从 latent state 中提取有用信息。

如果发现 latent collapse 迹象:

Actor-Critic 指标:

系统指标:

进阶诊断指标:

使用 TensorBoard

DreamerV3 支持 TensorBoard 日志:

# 启动训练时指定日志目录
python dreamerv3/train.py --logdir ./logs/my_experiment

# 启动 TensorBoard
tensorboard --logdir ./logs

在 TensorBoard 中可以直观地看到各项指标的变化趋势,方便调试。

定期保存 checkpoint

训练过程中定期保存模型,防止意外中断:

# configs.yaml
save_every: 10000  # 每 10000 步保存一次

保存的 checkpoint 可以用于:

七、常见坑与解决方案

坑 1:JAX 编译很慢

JAX 使用 XLA 编译器,第一次运行时会编译计算图,可能需要几分钟。这是正常的,后续运行会快很多。

解决方案:

坑 2:多 GPU 训练

DreamerV3 的官方 JAX 参考实现默认配置主要针对单个 accelerator,但这并不意味着"不支持"多 GPU。JAX 本身提供了 jax.pmapjax.shard_map、PJRT 等并行机制,理论上可以进行多 GPU 扩展。

实际情况:

坑 3:Replay buffer 占用大量内存

DreamerV3 的 replay buffer 保存的是原始 trajectory 数据(observation、action、reward、continuation、episode 信息),而不是训练后的 latent sequence。训练时数据经过 encoder 才得到 latent。但如果存储原始 observation(尤其是像素图像),内存占用仍然可能很大。

解决方案:

replay_size: 1e6   # 控制 replay buffer 大小

坑 4:训练速度慢

DreamerV3 的训练速度受多个因素影响:

可能原因:

解决方案:

坑 5:symlog/symexp 数值问题

symlog 和 symexp 是 DreamerV3 处理尺度的关键,但如果输入值极端,可能导致数值问题。

解决方案:

八、实战经验总结

调参优先级

当需要调优时,建议按以下优先级尝试:

  1. 环境接口正确性:reward scale、action normalization、observation preprocessing、continuation 设置——这些错误会直接导致训练失败
  2. Replay / warmup 配置:确保 replay buffer 正常填充,warmup 阶段合理
  3. Learning rate:影响训练稳定性,DreamerV3 默认 4e-5 通常不需要大改
  4. Batch size:影响梯度估计的稳定性
  5. RSSM 容量deterhidden 维度,影响世界模型表征能力
  6. Imagination length:影响长期信用的分配,但很多失败并不是 imagination length 的问题
  7. Entropy / 探索:通过 minstd/maxstd → unimix → actent 的顺序调整

很多训练失败的根因是 reward scale 错误、action normalization 错误、continuation 配置错误等基础问题,而不是 imagination length 不够长。先排查基础配置,再调超参数。

训练 checklist

开始训练前,检查以下事项:

何时停止训练

训练不是越久越好。以下情况可以考虑停止:

九、实验复现指南

Dreamer 类算法非常依赖实验条件,复现性是一个常见问题。做严肃实验时,建议建立以下习惯:

Seed 管理

Config 保存

Checkpoint 命名

Git commit 记录

Hardware 记录

十、把之前的文章串起来

世界模型入门 → RSSM 深度解析 → RSSM 代码系列(6篇)
                              Dreamer 系列 #1:整体架构
                              Dreamer 系列 #2:Actor-Critic
                              Dreamer 系列 #3:训练技巧(本篇)
                              GPU 选型指南

如果你还没读过前两篇,建议先看 Dreamer 整体架构Actor-Critic 详解,再来读这篇训练技巧,会更有收获。

十一、总结

DreamerV3 的训练工程实践可以概括为:

DreamerV3 的设计已经非常鲁棒,默认配置在大多数任务上都能工作。但如果想要达到最佳性能,或者遇到训练问题,就需要深入理解每个组件的作用,才能有针对性地调优。

希望这篇实战指南能帮助你更顺利地训练 DreamerV3。如果遇到问题,欢迎在评论区讨论。


← Dreamer 的 Actor-Critic:想象空间里的策略优化是怎么工作的? DreamerV3 GPU 选型指南:从显存需求到性价比分析 →

评论

W
侯晓琴

西北工业大学硕士,十余年自动化与 AI 工程经验。著有《Visual C++入门很容易》《C++程序设计经典300例》。目前聚焦世界模型与具身智能方向,记录从传统自动化到机器人 AI 的转型之路。