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 维隐空间:

⚠️ 配置文件 wan_flow_matching_action_tf.yaml 中的 input_embedding_dim: 1536 不是这两个编码器的实际维度。该字段在代码中仅被赋值给 self.input_embedding_dimwan_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 想成一本按机器人型号编号的「适配器手册」,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 同时学习多个本体。论文的说法与之一致:

需要注意的是,数据侧的 dreamzero_cotrain.py 里确实存在完整的 embodiment 白名单(AGIBOT / OXE_DROID / GR1_UNIFIED / MECKA_HANDS / XDOF / YAM 等,见 transform/base.yamlembodiment_tag_to_projector_index),但这些 ID 用于数据变换阶段的分发,模型侧最终仍被压到单一 category。

适配新本体实际需要做什么#

不能简化为「换个 ID 加个 MLP」。实际流程包含:

  1. 训练新的 state / action interface(上表三个模块);
  2. 调整动作归一化与维度(max_state_dim / max_action_dim,如 AgiBot 配置为 64 / 32);
  3. 对共享 DiT 做 post-training——公开脚本提供两套droid_training_lora.shtrain_architecture=lora)与 droid_training_full_finetune.shtrain_architecture=full);AgiBot / YAM 脚本默认走 LoRA;
  4. 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)
  1. 依赖视频世界模型已有的视觉动力学先验完成迁移。

相对动作计算: 各机器人体在数据加载时将绝对动作转为相对值(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

训练模式: 公开仓库提供两套脚本。


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