DreamZero (WAM) 模型架构
1. 整体架构概览
DreamZero 是一个世界动作模型 (World Action Model, WAM),通过联合预测未来视频帧与机器人动作序列,实现对未见任务的零样本泛化。其核心创新在于:基于 Wan2.1 视频扩散模型构建 Causal WAN DiT,以 Flow Matching 框架在视频 latent 空间和动作空间上同步去噪,配合 Category-specific MLP 构成的本体接口(state 编码 / action 编码 / action 解码三处)适配不同机器人平台,在 MolmoSpaces 和 RoboArena 双榜均位列第一(截至 2026 年 2 月)。
graph TB
subgraph Input["输入"]
IMG["历史/当前视频帧<br/>[B, 1, C, H, W]<br/>320×176"]
TXT["语言指令<br/>(text prompt)"]
STATE["机器人本体状态 q_t<br/>[B, ·, max_state_dim]"]
EID["机器人本体 ID<br/>(embodiment_id, 一个整数)"]
end
subgraph Encoders["感知编码器"]
CLIP["CLIP ViT-H/14 图像编码器<br/>open-clip-xlm-roberta-large-vit-huge-14<br/>1280 → 5120"]
T5["umt5-xxl 文本编码器<br/>4096 → 5120"]
end
subgraph VAE["视频 VAE (Wan2.1)"]
ENC["VAE 编码器<br/>观测帧像素 → 观测 latent<br/>(仅压缩,无生成能力)"]
DEC["VAE 解码器<br/>预测 latent → 像素帧<br/>(仅解压,可视化用)"]
end
subgraph NoiseIn["双路噪声起点"]
NOISE_V["视频高斯噪声 ε_v<br/>→ 带噪 latent z_τ"]
NOISE_A["动作高斯噪声 ε_a<br/>→ 带噪动作 a_τ"]
end
subgraph Iface["本体专属接口 (CategorySpecific)"]
S_ENC["state_encoder<br/>状态编码器<br/>CategorySpecificMLP"]
A_ENC["action_encoder<br/>动作编码器<br/>MultiEmbodimentActionEncoder"]
A_DEC["action_decoder<br/>动作解码器<br/>CategorySpecificMLP"]
end
subgraph Core["Causal WAN DiT (核心,真正的生成模型)"]
DIT["40层因果扩散 Transformer<br/>dim=5120, 40 heads<br/>视频/动作/状态 token 同序列<br/>Flow Matching + RoPE + Flash Attn 3"]
end
subgraph Sched["Flow Matching 迭代去噪"]
SCH_V["视频 scheduler<br/>FlowUniPCMultistep"]
SCH_A["动作 scheduler<br/>FlowUniPCMultistep"]
end
subgraph Output["输出"]
ACT_OUT["动作序列<br/>[B, 24, action_dim]"]
VID_OUT["未来视频帧<br/>[B, 33, H, W, 3]<br/>(RGB 解码可选)"]
end
IMG --> CLIP
TXT --> T5
IMG --> ENC
STATE --> S_ENC
NOISE_A --> A_ENC
EID -.->|"按 ID 索引出该本体专属的<br/>线性层权重 W_e, b_e (详见 2.4)"| Iface
CLIP -->|"视觉条件(cross-attn)"| DIT
T5 -->|"语言条件(cross-attn)"| DIT
ENC -->|"观测 latent(条件,填KV Cache)"| DIT
NOISE_V -->|"视频 token"| DIT
A_ENC -->|"动作 token"| DIT
S_ENC -->|"状态 token"| DIT
DIT -->|"动作 flow v_a"| A_DEC --> SCH_A
DIT -->|"视频 flow v_z"| SCH_V
SCH_A -->|"循环去噪 N 步"| ACT_OUT
SCH_V -->|"循环去噪 N 步"| DEC --> VID_OUT
style Input fill:#e8f4fd,stroke:#2196F3
style Encoders fill:#fff3e0,stroke:#FF9800
style VAE fill:#f3e5f5,stroke:#9C27B0
style NoiseIn fill:#fff8e1,stroke:#FBC02D
style Iface fill:#ede7f6,stroke:#673AB7
style Core fill:#e8f5e9,stroke:#4CAF50
style Sched fill:#e0f7fa,stroke:#00BCD4
style Output fill:#fce4ec,stroke:#E91E63
关于视频分支与控制的关系:只有「把最终 video latent 用 VAE 解码成 RGB 像素」这一步在部署时可以省略。视频分支本身是动作预测的核心组成部分——带噪视频 latent 与动作 token 处于同一 Transformer 序列、共享 self-attention,DiT 在同一次前向中同时输出视频 flow 和动作 flow。未来视频相当于隐式的视觉规划,视频预测出错时动作往往跟随错误的视觉计划。此外,省略 RGB 解码并不能显著降低推理耗时,主要计算量在 DiT blocks 与 diffusion steps。
2. 核心组件详解
2.1 感知编码器
DreamZero 使用两个独立的预训练编码器提取视觉和语言条件。两者的原生输出维度并不相同,各自经过独立投影层后才对齐到 DiT 的 5120 维隐空间:
- 图像编码器:
open-clip-xlm-roberta-large-vit-huge-14(CLIP ViT-H/14),输出 1280 维 视觉 token,经 MLPProj(1280, 5120) 投影
- 文本编码器:
umt5-xxl(Google 多语言 T5),输出 4096 维 序列特征,经 text_embedding(Linear(4096, 5120) → GELU → Linear(5120, 5120))投影
⚠️ 配置文件 wan_flow_matching_action_tf.yaml 中的 input_embedding_dim: 1536 不是这两个编码器的实际维度。该字段在代码中仅被赋值给 self.input_embedding_dim(wan_flow_matching_action_tf.py:228)后再未被引用,是从上层框架沿用下来的失效字段。
graph LR
subgraph ImagePath["图像编码路径"]
VF["视频帧<br/>[B, T_img, C, H, W]"] --> CLIP_ENC["CLIP ViT-H/14 编码<br/>wan_video_image_encoder.py"]
CLIP_ENC --> CLIP_FEAT["图像特征<br/>[B, 257, 1280]"]
CLIP_FEAT --> CLIP_PROJ["MLPProj<br/>1280 → 5120"]
end
subgraph TextPath["文本编码路径"]
PROMPT["语言提示<br/>(string)"] --> T5_TOK["T5 分词器"]
T5_TOK --> T5_ENC["umt5-xxl 编码<br/>wan_video_text_encoder.py"]
T5_ENC --> TXT_FEAT["文本特征<br/>[B, T_txt, 4096]"]
TXT_FEAT --> TXT_PROJ["text_embedding<br/>4096 → 5120"]
end
subgraph Fusion["特征融合 → DiT Cross-Attention"]
CLIP_PROJ --> CAT["拼接为条件序列<br/>[B, ·, 5120]"]
TXT_PROJ --> CAT
CAT --> DIT_IN["输入 Causal WAN DiT<br/>(作为 K/V 供 cross-attn)"]
end
style ImagePath fill:#e3f2fd,stroke:#2196F3
style TextPath fill:#fff3e0,stroke:#FF9800
style Fusion fill:#e8f5e9,stroke:#4CAF50
2.2 视频 VAE
沿用 Wan2.1 的视频 VAE,作为压缩/解压工具(无生成能力),让 DiT 在低维 latent 空间工作以降低计算量。VAE 编码器和解码器服务于两个完全不同的目的,输入输出的 latent 含义不同。
graph LR
subgraph Encode["VAE 编码器(输入侧)"]
V_IN["观测帧像素<br/>[B, T, H, W, 3]"] --> VAE_E["VAE 编码器<br/>wan_video_vae.py"]
VAE_E --> LAT["观测 latent(条件)<br/>[B, 16, T, h, w]<br/>用途:填入 KV Cache,作为 DiT 生成的历史条件"]
end
subgraph Decode["VAE 解码器(输出侧)"]
LAT2["预测 latent(生成目标)<br/>[B, 16, T, h, w]<br/>DiT 从纯噪声去噪得到"] --> VAE_D["VAE 解码器"]
VAE_D --> V_OUT["预测视频帧像素<br/>[B, T, H, W, 3]<br/>(仅用于可视化,不影响机器人控制)"]
end
subgraph Params["关键参数"]
P1["latent 通道数: 16"]
P2["时域压缩: 4×"]
P3["空域压缩: 8×"]
end
style Encode fill:#e3f2fd,stroke:#2196F3
style Decode fill:#e8f5e9,stroke:#4CAF50
style Params fill:#fff9c4,stroke:#FFC107
2.3 Causal WAN DiT(核心扩散模型)
Causal WAN DiT 是 DreamZero 的核心,基于 Wan2.1 的 DiT 架构改造而来:40 层因果 Transformer,通过 Causal Masking 保证时序因果性,同时在 latent 序列中嵌入动作 token 与状态 token,实现视频与动作的联合生成。
序列排布:视频 token 在前,动作与状态 token 拼成 action_register 追加在后,构成 [video tokens | action tokens | state tokens] 一条序列。因此视频 token 与动作 token 在自注意力中可以相互访问——这是「视频分支影响控制」的机制来源,而非先生成视频、再用独立 IDM 从视频反推动作。
graph TB
subgraph InputSeq["输入序列构建"]
direction LR
VL["含噪预测 latent z_τ<br/>[B, T_vid, 16]<br/>(去噪目标,从纯噪声出发)"]
AC["含噪动作 a_τ<br/>[B, 24, action_dim]"]
ST["本体状态 q_t<br/>[B, ·, max_state_dim]"]
VIS["视觉条件<br/>[B, 257, 5120]"]
LNG["语言条件<br/>[B, T_txt, 5120]"]
TS["时间步 t (视频/动作各一路)"]
VL --> ROPE["RoPE 位置编码"]
AC --> ACT_PROJ["action_encoder 动作编码器<br/>(含时间步正弦编码)"]
ST --> ST_PROJ["state_encoder 状态编码器"]
TS --> TIME_EMB["时间步嵌入<br/>(正弦编码)"]
end
subgraph Block["单层 Causal DiT Block (×40)"]
direction TB
IN_B["输入 [B, T, 5120]"] --> LN1_B["RMSNorm"]
LN1_B --> MOD1["Modulation<br/>(shift, scale by timestep emb)"]
MOD1 --> CATTN["因果自注意力<br/>40 heads, Flash Attn 3<br/>Causal Mask"]
CATTN --> ADD1["+ 残差"]
IN_B --> ADD1
ADD1 --> LN2_B["RMSNorm"]
LN2_B --> XATTN["交叉注意力<br/>(视觉 + 语言条件)"]
XATTN --> ADD2["+ 残差"]
ADD1 --> ADD2
ADD2 --> LN3_B["RMSNorm"]
LN3_B --> FFN_B["FFN (SwiGLU)<br/>dim 5120 → 13824 → 5120"]
FFN_B --> ADD3["+ 残差"]
ADD2 --> ADD3
ADD3 --> OUT_B["输出 [B, T, 5120]"]
end
subgraph Output["输出解码"]
OUT_B --> SPLIT["按位置切分<br/>video / action / state token"]
SPLIT --> VP["视频 flow 预测 v_z<br/>(head + unpatchify)"]
SPLIT --> AP["动作 flow 预测 v_a<br/>action_decoder 动作解码器<br/>5120 → action_dim"]
end
ROPE --> Block
ACT_PROJ --> Block
ST_PROJ --> Block
TIME_EMB --> Block
VIS --> Block
LNG --> Block
style InputSeq fill:#e3f2fd,stroke:#2196F3
style Block fill:#e8f5e9,stroke:#4CAF50
style Output fill:#fce4ec,stroke:#E91E63
因果注意力掩码设计: 视频帧 token 只能 attend 到过去帧和当前帧,动作 token attend 到所有视频 token(全局条件),保证生成的时序一致性。
graph LR
subgraph Mask["Causal Attention Mask 示意"]
direction LR
F0["帧 0 token"] -->|"✓ attend"| F0S["自身"]
F1["帧 1 token"] -->|"✓ attend"| F0_1["帧 0"]
F1 -->|"✓ attend"| F1S["自身"]
F1 -->|"✗ 不 attend"| F2_BLK["帧 2+ (未来)"]
AT["动作 token"] -->|"✓ attend"| ALL["所有帧 token"]
end
style Mask fill:#fff9c4,stroke:#FFC107
2.4 本体适配机制(Embodiment Interface)
DreamZero 的本体适配可以概括为:
共享 Causal DiT + 本体专属 state/action 接口(embodiment-specific state/action interface)
而不是「共享 DiT + 一个 robot-ID prompt token」。
三个本体专属模块,而非一个
embodiment_id 不会被编码成向量再作为 token 喂给 Transformer。它的作用是从权重张量中索引出该本体专属的一组权重:
class CategorySpecificLinear(nn.Module):
def __init__(self, num_categories, input_dim, hidden_dim):
self.W = nn.Parameter(0.02 * torch.randn(num_categories, input_dim, hidden_dim))
self.b = nn.Parameter(torch.zeros(num_categories, hidden_dim))
def forward(self, x, cat_ids):
selected_W = self.W[cat_ids] # 按本体 ID 选权重
selected_b = self.b[cat_ids]
return torch.bmm(x, selected_W) + selected_b.unsqueeze(1)
即对本体 $e$ 执行 $y = x W_e + b_e$。
W_e / b_e 到底是什么、从哪来
一句话:它们就是普通 nn.Linear 的权重和偏置,只不过被"每个本体存一份"地堆成了一个多出一维的大张量,W_e 是从中切出的第 e 片。
对比一下就很清楚——普通线性层与这里的差别只在最前面多了一个「本体」维:
|
普通 nn.Linear |
CategorySpecificLinear |
| 权重形状 |
(input_dim, hidden_dim) |
(num_categories, input_dim, hidden_dim) |
| 偏置形状 |
(hidden_dim,) |
(num_categories, hidden_dim) |
| 前向 |
x @ W + b |
x @ W[e] + b[e],即 W_e = W[e] |
所以:
- 是什么:
W 和 b 是 nn.Parameter,即模型的可训练参数本身,和 DiT 里任何一个权重矩阵地位相同。W_e 不是运行时算出来的中间变量,只是对 W 沿第 0 维做一次索引 W[e],切出来的 (input_dim, hidden_dim) 就是一个再普通不过的权重矩阵。
- 从哪来(初始化):随机初始化,从零训起 ——
W = 0.02 * torch.randn(...),b = torch.zeros(...)。这三个模块在 Wan2.1 预训练权重里根本不存在(Wan2.1 是纯视频生成模型,没有机器人状态/动作的概念),是 DreamZero 新增的部件,因此没有预训练权重可继承。
- 怎么训出来:靠正常的反向传播。注意
W[cat_ids] 是一次索引操作,梯度是稀疏的——某个 batch 里只出现了本体 3,那么只有 W[3] 这一片拿到梯度,其他本体的切片原封不动。这正是"各本体互不干扰"的来源。
- LoRA 模式下也全量训练:
train_architecture=lora 时先把所有参数冻结,注入 LoRA 后又显式把这三个模块解冻(wan_flow_matching_action_tf.py:341-343),它们走的是完整梯度更新,不是低秩分解。原因很直接:新本体的接口层没有预训练权重可微调,低秩增量无从谈起。
torch.bmm 的意义:用 batch matmul 而非普通 matmul,是为了让同一个 batch 内不同样本可以用不同本体的权重,从而支持多本体数据混合训练。
一个直观的类比:把 W 想成一本按机器人型号编号的「适配器手册」,embodiment_id = e 就是翻到第 e 页。手册本身(W 整体)是训练出来的,翻页(索引)不消耗任何计算,也不产生 token。
这套机制分布在三个位置,而不只是动作输入端:
| 位置 |
模块 |
类型 |
作用 |
| 状态输入 |
state_encoder |
CategorySpecificMLP |
$q_t \rightarrow$ 状态 token(max_state_dim → 5120) |
| 动作输入 |
action_encoder |
MultiEmbodimentActionEncoder |
带噪动作 $a_\tau$ + 扩散时间步 $\tau \rightarrow$ 动作 token |
| 动作输出 |
action_decoder |
CategorySpecificMLP |
DiT 输出 $\rightarrow$ 本体动作空间的 flow(5120 → action_dim) |
MultiEmbodimentActionEncoder 内部由 3 个 CategorySpecificLinear(W1/W2/W3)加正弦位置编码构成:动作先经 W1 升维,与时间步编码拼接后经 W2 + swish,再经 W3 输出。
graph TB
subgraph InRobot["本体侧输入(每台机器人各不相同)"]
Q["本体状态 q_t<br/>关节角 / 夹爪开合等"]
A_NOISE["动作高斯噪声 ε_a"]
TAU["扩散时间步 τ"]
A_NOISE --> A_TAU["带噪动作 a_τ<br/>待去噪的动作序列"]
end
subgraph Select["embodiment_id 的真实作用:查表选权重"]
EID2["embodiment_id = e<br/>本体编号,一个整数"]
BANK["权重库 W, b (nn.Parameter)<br/>W 形状 (本体数, in_dim, out_dim)<br/>b 形状 (本体数, out_dim)<br/>随机初始化,反向传播训练得到"]
EID2 --> PICK["沿第 0 维索引出第 e 片<br/>W_e = W[e] → (in_dim, out_dim)<br/>b_e = b[e] → (out_dim,)<br/><b>即该本体专属的一套线性层权重</b><br/><b>不生成任何 token</b>"]
BANK --> PICK
end
subgraph IfaceIn["本体专属输入接口(把异构机器人对齐到统一维度)"]
Q --> SE["state_encoder<br/>类型 CategorySpecificMLP<br/>状态编码器:max_state_dim → 5120"]
A_TAU --> AE["action_encoder<br/>类型 MultiEmbodimentActionEncoder<br/>动作编码器:W1/W2/W3 + 正弦时间编码"]
TAU --> AE
end
subgraph SharedDiT["共享主干(与本体无关,所有机器人复用同一套权重)"]
SE --> SEQ["拼接为一条序列<br/>[视频 token | 动作 token | 状态 token]"]
AE --> SEQ
VID_TOK["视频 token z_τ"] --> SEQ
SEQ --> DIT_PROC["Causal WAN DiT<br/>40 层 · 14B 参数<br/>世界模型先验就在这里"]
end
subgraph IfaceOut["本体专属输出接口(映射回该机器人的动作空间)"]
DIT_PROC --> SLICE["从输出序列中切出动作 token"]
SLICE --> AD["action_decoder<br/>类型 CategorySpecificMLP<br/>动作解码器:5120 → action_dim"]
AD --> VA["动作 flow v_a<br/>交给 scheduler 迭代去噪"]
end
PICK -.->|"该层用 W_e, b_e 做 y = x·W_e + b_e"| SE
PICK -.->|"同上"| AE
PICK -.->|"同上"| AD
style InRobot fill:#e3f2fd,stroke:#2196F3
style Select fill:#fff8e1,stroke:#FBC02D
style IfaceIn fill:#ede7f6,stroke:#673AB7
style SharedDiT fill:#e8f5e9,stroke:#4CAF50
style IfaceOut fill:#fce4ec,stroke:#E91E63
⚠️ 当前公开实现并非「单 checkpoint 多本体自由切换」
尽管类名(MultiEmbodimentActionEncoder)和字段(max_num_embodiments)都暗示多本体,当前公开主分支实际只有一个 category:
# wan_video_dit_action_casual_chunk.py:1325,紧接在赋值 self.max_num_embodiments 之后
max_num_embodiments = 1 # 局部变量被硬覆盖为 1
self.state_encoder = CategorySpecificMLP(num_categories=max_num_embodiments, ...)
self.action_encoder = MultiEmbodimentActionEncoder(num_embodiments=max_num_embodiments, ...)
self.action_decoder = CategorySpecificMLP(num_categories=max_num_embodiments, ...)
并且两条前向路径都会把传入的 embodiment_id 覆盖为 0:
# :1715(常规前向)与 :2011(TensorRT 前向)
embodiment_id = torch.tensor([0], device=x.device).repeat(x.shape[0])
所以代码结构预留了多本体权重选择机制,但当前发布的权重并没有真正用一个 checkpoint 同时学习多个本体。论文的说法与之一致:
- AgiBot G1 与 Franka 分别预训练;
- multi-embodiment joint training 列为未来工作;
- AgiBot → YAM 是拿 AgiBot checkpoint 在 YAM 约 30 分钟数据上继续 post-training,不是推理时把
embodiment_id 从 AgiBot 改成 YAM 就能直接部署。
需要注意的是,数据侧的 dreamzero_cotrain.py 里确实存在完整的 embodiment 白名单(AGIBOT / OXE_DROID / GR1_UNIFIED / MECKA_HANDS / XDOF / YAM 等,见 transform/base.yaml 的 embodiment_tag_to_projector_index),但这些 ID 用于数据变换阶段的分发,模型侧最终仍被压到单一 category。
适配新本体实际需要做什么
不能简化为「换个 ID 加个 MLP」。实际流程包含:
- 训练新的 state / action interface(上表三个模块);
- 调整动作归一化与维度(
max_state_dim / max_action_dim,如 AgiBot 配置为 64 / 32);
- 对共享 DiT 做 post-training——公开脚本提供两套:
droid_training_lora.sh(train_architecture=lora)与 droid_training_full_finetune.sh(train_architecture=full);AgiBot / YAM 脚本默认走 LoRA;
- LoRA 模式下三个 state/action adapter 仍保持可训练:
# wan_flow_matching_action_tf.py:343-345
self.model.state_encoder.requires_grad_(True)
self.model.action_encoder.requires_grad_(True)
self.model.action_decoder.requires_grad_(True)
- 依赖视频世界模型已有的视觉动力学先验完成迁移。
相对动作计算: 各机器人体在数据加载时将绝对动作转为相对值(action - reference_state),使模型学到的动作表示更加泛化。
3. 训练流水线
DreamZero 基于 Flow Matching 框架进行训练,同时对视频 latent 和动作 token 施加扩散损失。训练使用 DeepSpeed ZeRO-2 分布式优化,并通过 LoRA 高效微调 14B 参数基础模型。
graph TB
subgraph DataLoad["数据加载 (LeRobot 格式)"]
direction LR
PARQ["LeRobot parquet 文件"] --> DS["lerobot.py / lerobot_sharded.py"]
DS --> MOD["模态变换<br/>· 视频 resize → 320×176<br/>· 状态/动作归一化<br/>· 相对动作计算<br/>· 语言编码"]
MOD --> BATCH_OUT["训练 Batch"]
end
subgraph Encode["编码阶段"]
BATCH_OUT --> IMG_E["CLIP 图像编码<br/>→ 视觉特征"]
BATCH_OUT --> TXT_E["T5 文本编码<br/>→ 语言特征"]
BATCH_OUT --> VAE_E2["VAE 编码<br/>视频 → latent"]
end
subgraph FlowMatch["Flow Matching 加噪(视频/动作各一路)"]
NOISE_V2["视频噪声 ε_z ~ N(0,I)<br/>torch.randn_like(latents)"]
NOISE_A2["动作噪声 ε_a ~ N(0,I)<br/>torch.randn_like(actions)"]
VID_GT["真实视频 latent z_1"]
ACT_GT["真实动作 a_1"]
NOISE_V2 --> INTERP_V["z_τ = (1-σ)·z_1 + σ·ε_z"]
VID_GT --> INTERP_V
NOISE_A2 --> INTERP_A["a_τ = (1-σ)·a_1 + σ·ε_a"]
ACT_GT --> INTERP_A
end
subgraph Forward["前向传播"]
IMG_E --> DIT_FWD["Causal WAN DiT<br/>[video | action | state] 同序列"]
TXT_E --> DIT_FWD
VAE_E2 --> DIT_FWD
INTERP_V --> DIT_FWD
INTERP_A --> DIT_FWD
DIT_FWD --> PRED_ALL["联合预测<br/>视频 flow v_z + 动作 flow v_a"]
end
subgraph Loss["损失计算"]
PRED_ALL --> L_VID["视频 Flow Matching Loss<br/>target = ε_z − z_1"]
PRED_ALL --> L_ACT["动作 Flow Matching Loss<br/>target = ε_a − a_1"]
L_VID --> L_TOTAL["总损失<br/>L = λ_vid · L_vid + λ_act · L_act<br/>(各自按 timestep 加权)"]
L_ACT --> L_TOTAL
end
subgraph Optim["优化"]
L_TOTAL --> BACK["反向传播"]
BACK --> LORA["LoRA rank=4, alpha=4<br/>或 train_architecture=full 全量微调<br/>(两套脚本均已开源)"]
LORA --> OPT["AdamW<br/>β₁=0.95, β₂=0.999<br/>Cosine LR + Warmup"]
OPT --> DS_ZERO["DeepSpeed ZeRO-2<br/>梯度累积"]
end
DataLoad --> Encode
Encode --> FlowMatch
FlowMatch --> Forward
Forward --> Loss
Loss --> Optim
style DataLoad fill:#e3f2fd,stroke:#2196F3
style Encode fill:#fff3e0,stroke:#FF9800
style FlowMatch fill:#f3e5f5,stroke:#9C27B0
style Forward fill:#e8f5e9,stroke:#4CAF50
style Loss fill:#fce4ec,stroke:#E91E63
style Optim fill:#e0f7fa,stroke:#00BCD4
训练模式: 公开仓库提供两套脚本。
- LoRA(
droid_training_lora.sh / agibot_training.sh / yam_training.sh,train_architecture=lora):基础模型(Wan2.1,约 14B 参数)冻结,在 DiT 的 q,k,v,o,ffn.0,ffn.2 上插入 LoRA 适配器(rank=4, alpha=4),同时保持 state_encoder / action_encoder / action_decoder 三个本体接口全量可训练。
- 全量微调(
droid_training_full_finetune.sh,train_architecture=full,save_lora_only=false):更新全部 DiT blocks,对应论文中的主预训练设置。
4. 推理流水线
推理时视频 latent 与动作都从各自的高斯噪声出发,在同一次 DiT 前向中得到两路 flow,再由两个独立的 scheduler 分别积分更新。动作不是 DiT 一次前向直接回归出来的干净值,而是多步去噪的结果。DreamZero 支持通过 WebSocket 的分布式多 GPU 推理服务。
graph TB
subgraph Init["初始化"]
OBS["当前观测图像<br/>[B, 1, C, H, W]"]
LANG_IN["语言指令"]
Q_IN["本体状态 q_t"]
NV["视频噪声 ε_z → z_τ"]
NA["动作噪声 ε_a → a_τ"]
end
subgraph EncoderOnce["编码(仅执行一次)"]
OBS --> CLIP_INF["CLIP 编码"]
LANG_IN --> T5_INF["T5 编码"]
CLIP_INF --> COND["条件特征缓存<br/>(+ KV Cache)"]
T5_INF --> COND
Q_IN --> SE_INF["state_encoder 状态编码器"]
end
subgraph DenoiseLoop["迭代去噪循环 (num_inference_timesteps 步)"]
direction TB
FWD["单次 DiT 前向 (cond + uncond)<br/>输入 [z_τ | a_τ | q]"]
FWD --> VZ["视频 flow v_z<br/>uncond + s·(cond − uncond)<br/>← CFG 只作用于视频"]
FWD --> ADEC["action_decoder 动作解码器<br/>→ 动作 flow v_a<br/>只取 cond 分支,不做 CFG"]
VZ --> STEP_V["sample_scheduler.step<br/>更新 z_τ"]
ADEC --> STEP_A["sample_scheduler_action.step<br/>更新 a_τ"]
STEP_V --> FWD
STEP_A --> FWD
end
subgraph PostProcess["后处理"]
STEP_A --> ACT_POST["动作后处理<br/>(反归一化, 相对→绝对)"]
ACT_POST --> ACT_EXEC["执行前 N 步动作<br/>(Action Chunking)"]
STEP_V --> VAE_DEC2["VAE 解码 → RGB 视频帧<br/>(可选,仅可视化需要)"]
end
Init --> EncoderOnce
COND --> DenoiseLoop
SE_INF --> DenoiseLoop
NV --> DenoiseLoop
NA --> DenoiseLoop
style Init fill:#e3f2fd,stroke:#2196F3
style EncoderOnce fill:#fff3e0,stroke:#FF9800
style DenoiseLoop fill:#e8f5e9,stroke:#4CAF50
style PostProcess fill:#fce4ec,stroke:#E91E63
两个 scheduler 均为 FlowUniPCMultistepScheduler,用同样的步数、同样的 sigma_shift,但各自持有独立的 timestep 序列。开启 decouple_inference_noise 时,视频 sigma 序列会被重映射到 [σ_max, video_inference_final_noise](默认 0.8,即视频故意不完全去噪),而动作始终按标准 1000→0 完整去噪——再次说明视频分支的价值在于它在 DiT 内部提供的世界模型先验,而不是最终那张图好不好看。
另外循环里还有一个 should_run_model() 的跳步缓存:某些步会直接复用上一次的 flow_pred,不重跑 DiT,所以「去噪步数」不等于「DiT 前向次数」。
分布式推理服务
graph LR
subgraph Client["客户端 (test_client_AR.py)"]
CLI["机器人控制程序"] --> WS_CLI["WebSocket 客户端<br/>发送观测 + 指令"]
end
subgraph Server["推理服务器 (socket_test_optimized_AR.py)"]
WS_SRV["WebSocket 服务器<br/>(Flask-SocketIO)"]
REDIS["Redis<br/>会话状态管理"]
RAY_W["Ray Worker Pool<br/>多 GPU 并行推理"]
MODEL["DreamZero 模型<br/>(GB200 / H100)"]
WS_SRV --> REDIS
WS_SRV --> RAY_W
RAY_W --> MODEL
end
WS_CLI -->|"观测数据"| WS_SRV
MODEL -->|"动作预测"| WS_SRV
WS_SRV -->|"动作序列"| WS_CLI
WS_CLI -->|"执行动作"| CLI
style Client fill:#e3f2fd,stroke:#2196F3
style Server fill:#e8f5e9,stroke:#4CAF50
5. 关键超参数表
5.1 架构常量(由预训练权重固定,不可改)
| 参数 |
值 |
说明 |
| 模型总参数 |
~14B |
基于 Wan2.1-I2V-14B |
| DiT 层数 |
40 |
Causal WAN DiT |
| DiT hidden dim |
5120 |
dim,全部 token 统一到该维度 |
| 注意力头数 |
40 |
head_dim = 128 |
| 文本编码维度 |
4096 → 5120 |
UMT5 输出 text_dim=4096,经 text_embedding 两层 MLP 升维 |
| 图像编码维度 |
1280 → 5120 |
CLIP ViT-H 输出 1280,经 MLPProj(1280, dim) 升维 |
| VAE latent 通道 |
16 |
视频压缩维度 |
| 训练 timestep 桶数 |
1000 |
num_timestep_buckets |
⚠️ 配置里的 input_embedding_dim: 1536 是失效字段,不是任何编码器的真实维度,详见 2.1 节。
5.2 配置项(随脚本/数据集变化)
下表左列是官方 scripts/train/*.sh(droid / agibot / yam / aloha_lekiwi / pen_right_arm 全部一致)里实际传入的值;右列是仓库内另一套数据配置 base_48_wan_fine_aug_relative.yaml 的取值。这些是配置而非架构常量——换数据集时会变。
| 参数 |
训练脚本 |
base_48_..._relative.yaml |
视频帧数 num_frames |
33 |
49 |
action_horizon |
24 |
48 |
| 图像分辨率 |
320×176 |
480×256 |
5.3 训练配置
| 参数 |
值 |
说明 |
| 训练模式 |
lora / full |
droid_training_lora.sh vs droid_training_full_finetune.sh |
| LoRA rank / alpha |
4 / 4 |
默认值,作用于 q,k,v,o,ffn.0,ffn.2 |
| LoRA 额外可训练模块 |
三个本体接口 |
state / action 编码器与 action 解码器始终 requires_grad_(True) |
| 分布式策略 |
DeepSpeed ZeRO-2 |
分布式训练 |
| 优化器 |
AdamW |
β₁=0.95, β₂=0.999 |
| 学习率调度 |
Cosine + Warmup |
分布式多卡训练 |
| 精度 |
bfloat16 |
混合精度训练 |
6. 关键源文件表
| 组件 |
文件路径(相对 /home/zhuyilong/dreamzero/) |
| 核心 VLA 模型 |
groot/vla/model/dreamzero/base_vla.py |
| Action Head (Flow Matching) |
groot/vla/model/dreamzero/action_head/wan_flow_matching_action_tf.py |
| Causal WAN DiT |
groot/vla/model/dreamzero/modules/wan_video_dit_action_casual_chunk.py |
| 视频 VAE |
groot/vla/model/dreamzero/modules/wan_video_vae.py |
| 文本编码器 |
groot/vla/model/dreamzero/modules/wan_video_text_encoder.py |
| 图像编码器 |
groot/vla/model/dreamzero/modules/wan_video_image_encoder.py |
| 数据变换 |
groot/vla/model/dreamzero/transform/ |
| 数据集加载 |
groot/vla/data/dataset/lerobot.py |
| 分片数据集 |
groot/vla/data/dataset/lerobot_sharded.py |
| 训练器 |
groot/vla/experiment/experiment.py |
| 训练基类 |
groot/vla/experiment/base.py |
| 分布式推理服务 |
socket_test_optimized_AR.py |
| 推理客户端 |
test_client_AR.py |
| Hydra 模型配置 |
groot/vla/configs/model/ |
| 训练脚本 |
scripts/train/ |
7. Attention 机制深入解析
7.1 Q / K / V 的本质分工
Attention 的计算分四步:
1. q = W_q(x) # x 投影成 Q
2. k = W_k(context) # context 投影成 K
3. v = W_v(context) # context 投影成 V(同一原材料,不同矩阵)
4. 权重 = softmax(q · kᵀ / √d)
5. 输出 = 权重 · v
K 和 V 的原材料相同(都来自 context),但投影矩阵不同,作用完全不同:
| 矩阵 |
参与哪一步 |
对输出的影响 |
学到的是 |
| W_q |
步骤 4(点积) |
间接(通过权重) |
我该去找什么样的 K |
| W_k |
步骤 4(点积) |
间接(通过权重) |
什么样的 Q 应该找到我 |
| W_v |
步骤 5(加权求和) |
直接(构成输出) |
被选中后我该提供什么内容 |
K 是抽象意义上的索引——它不直接出现在输出里,只决定权重分布,告诉模型"哪些位置比较重要"。V 才是真正被读取的内容。
graph LR
subgraph Input["输入"]
X["x(视频 token)"]
CTX["context(文本/图像特征)"]
end
subgraph Proj["投影"]
X --> WQ["W_q"] --> Q["Q\n寻址用"]
CTX --> WK["W_k"] --> K["K\n索引(被动等待匹配)"]
CTX --> WV["W_v"] --> V["V\n内容(直接构成输出)"]
end
subgraph Attn["Attention 计算"]
Q --> DOT["q · kᵀ / √d\n点积"]
K --> DOT
DOT --> SM["softmax\n权重分布"]
SM --> WV2["权重 · V\n加权求和"]
V --> WV2
WV2 --> OUT["输出\n[B, T_x, dim]"]
end
style Input fill:#e3f2fd,stroke:#2196F3
style Proj fill:#fff3e0,stroke:#FF9800
style Attn fill:#e8f5e9,stroke:#4CAF50
7.2 Self-Attention vs Cross-Attention
区别只有一个:Q、K、V 的原材料从哪来。
|
Q 来自 |
K、V 来自 |
用途 |
| Self-Attention |
序列自身 x |
序列自身 x |
序列内部 token 互相交流 |
| Cross-Attention |
序列自身 x |
另一个序列 context |
从外部条件(文本/图像)读取信息 |
DreamZero 的每个 DiT Block 里两者都有:先 Self-Attention 让视频 token 内部交流,再 Cross-Attention 读取文本和图像条件。
graph LR
subgraph ImagePath["图像编码路径"]
VF["视频帧<br/>[B, T_img, C, H, W]"] --> CLIP_ENC["CLIP ViT-H/14 编码<br/>wan_video_image_encoder.py"]
CLIP_ENC --> CLIP_FEAT["图像特征<br/>[B, 257, 1280]"]
CLIP_FEAT --> CLIP_PROJ["MLPProj<br/>1280 → 5120"]
end
subgraph TextPath["文本编码路径"]
PROMPT["语言提示<br/>(string)"] --> T5_TOK["T5 分词器"]
T5_TOK --> T5_ENC["umt5-xxl 编码<br/>wan_video_text_encoder.py"]
T5_ENC --> TXT_FEAT["文本特征<br/>[B, T_txt, 4096]"]
TXT_FEAT --> TXT_PROJ["text_embedding<br/>4096 → 5120"]
end
subgraph Fusion["特征融合 → DiT Cross-Attention"]
CLIP_PROJ --> CAT["拼接为条件序列<br/>[B, ·, 5120]"]
TXT_PROJ --> CAT
CAT --> DIT_IN["输入 Causal WAN DiT<br/>(作为 K/V 供 cross-attn)"]
end
style ImagePath fill:#e3f2fd,stroke:#2196F3
style TextPath fill:#fff3e0,stroke:#FF9800
style Fusion fill:#e8f5e9,stroke:#4CAF50
0
7.3 为什么 K 不能等于 Q
K=Q 意味着 W_k = W_q,K 和 Q 是同一个向量。此时:
权重 = softmax(q · qᵀ / √d) # x 和自身的相似度
权重退化成"找和自己像的",模型只能检索与自身相似的 token,无法找到"内容上互补但方向不同"的信息。
类比:图书馆里每本书的索引标签直接写的是"我想找机器人抓取的书"——所有书都在大喊需求,没有一本书在说"我能提供什么",检索系统完全失效。
K 和 Q 必须用不同的投影矩阵,让"被检索"和"去检索"解耦。
7.4 多头注意力(Multi-Head Attention)
多头就是把 dim 切成 h 份,并行跑 h 套独立的 Q/K/V:
head_i 输出 = softmax(Q_i · K_iᵀ / √d) · V_i
最终输出 = Concat(head_0, ..., head_h-1) · W_o
多头的价值来源于 Q 的多样性——每个头问不同的问题,从 context 里提取不同侧面的信息。
graph LR
subgraph ImagePath["图像编码路径"]
VF["视频帧<br/>[B, T_img, C, H, W]"] --> CLIP_ENC["CLIP ViT-H/14 编码<br/>wan_video_image_encoder.py"]
CLIP_ENC --> CLIP_FEAT["图像特征<br/>[B, 257, 1280]"]
CLIP_FEAT --> CLIP_PROJ["MLPProj<br/>1280 → 5120"]
end
subgraph TextPath["文本编码路径"]
PROMPT["语言提示<br/>(string)"] --> T5_TOK["T5 分词器"]
T5_TOK --> T5_ENC["umt5-xxl 编码<br/>wan_video_text_encoder.py"]
T5_ENC --> TXT_FEAT["文本特征<br/>[B, T_txt, 4096]"]
TXT_FEAT --> TXT_PROJ["text_embedding<br/>4096 → 5120"]
end
subgraph Fusion["特征融合 → DiT Cross-Attention"]
CLIP_PROJ --> CAT["拼接为条件序列<br/>[B, ·, 5120]"]
TXT_PROJ --> CAT
CAT --> DIT_IN["输入 Causal WAN DiT<br/>(作为 K/V 供 cross-attn)"]
end
style ImagePath fill:#e3f2fd,stroke:#2196F3
style TextPath fill:#fff3e0,stroke:#FF9800
style Fusion fill:#e8f5e9,stroke:#4CAF50
1
Q 不能共享,K/V 可以共享(MQA):
|
多K多V |
单K单V |
| 多Q(问题不同) |
标准多头 MHA |
MQA(现代大模型常用) |
| 单Q(问题相同) |
无意义(退化) |
普通单头 |
MQA(Multi-Query Attention)中所有 Q 头共享同一组 K/V,节省 KV Cache 显存,但不影响 Q 的多样性——40 个人用不同问题检索同一个书架,仍然能找出不同的重点。
7.5 W_q、W_k、W_v 各自如何被训练
三个矩阵的梯度路径不同,学到的东西也不同。
graph LR
subgraph ImagePath["图像编码路径"]
VF["视频帧<br/>[B, T_img, C, H, W]"] --> CLIP_ENC["CLIP ViT-H/14 编码<br/>wan_video_image_encoder.py"]
CLIP_ENC --> CLIP_FEAT["图像特征<br/>[B, 257, 1280]"]
CLIP_FEAT --> CLIP_PROJ["MLPProj<br/>1280 → 5120"]
end
subgraph TextPath["文本编码路径"]
PROMPT["语言提示<br/>(string)"] --> T5_TOK["T5 分词器"]
T5_TOK --> T5_ENC["umt5-xxl 编码<br/>wan_video_text_encoder.py"]
T5_ENC --> TXT_FEAT["文本特征<br/>[B, T_txt, 4096]"]
TXT_FEAT --> TXT_PROJ["text_embedding<br/>4096 → 5120"]
end
subgraph Fusion["特征融合 → DiT Cross-Attention"]
CLIP_PROJ --> CAT["拼接为条件序列<br/>[B, ·, 5120]"]
TXT_PROJ --> CAT
CAT --> DIT_IN["输入 Causal WAN DiT<br/>(作为 K/V 供 cross-attn)"]
end
style ImagePath fill:#e3f2fd,stroke:#2196F3
style TextPath fill:#fff3e0,stroke:#FF9800
style Fusion fill:#e8f5e9,stroke:#4CAF50
2
W_q 和 W_k 通过同一条路径(点积)互相塑造,W_v 的梯度路径完全独立,只由输出质量驱动。
7.6 QK Norm:为什么归一化能稳定训练
原始点积:
q · k = |q| × |k| × cos(θ)
模长和方向混在一起,W_q 和 W_k 需要同时控制两件事,容易震荡。
QK Norm 之后:
q_norm = q / |q|, k_norm = k / |k|
q_norm · k_norm = cos(θ) # 值域固定在 [-1, 1]
点积退化为纯余弦相似度,只剩方向这一个自由度,模长的干扰被彻底消除。W_q 和 W_k 只需要各自学"朝哪个方向",优化目标从三个因素(|q|、|k|、θ)变成一个(θ),训练更稳定。
DreamZero 代码中对应实现(wan2_1_submodule.py):
q = self.norm_q(self.q(x)) # QK Norm
k = self.norm_k(self.k(context)) # QK Norm