1. 引言:自动驾驶中的未来预测困境
在自动驾驶系统的决策过程中,存在一个本质性的矛盾:未来场景的演化依赖于当前考虑的驾驶动作,而合理的驾驶动作又依赖于对未来场景的准确预测。这种相互依赖的关系在传统的World Action Model(世界动作模型,简称WAM)中并未得到充分体现。现有的WAM要么采用并行分支独立预测世界状态和驾驶动作,要么采用严格的预测-规划流水线,先固定未来场景预测,再基于此进行动作规划。
这些设计存在根本性局限:并行设计虽然允许世界和动作产生关联,但无法让二者在生成过程中相互重塑;顺序设计虽然让动作依赖于预测的未来,但这个未来是固定的假设,并非随着动作假设演化的动态未来。在复杂的交互式驾驶场景中,这种结构性解耦导致模型无法充分理解动作依赖性的未来演化特性。
DAWN(Denoising Actions and World iNteractive model)正是在这样的背景下诞生的。它提出了World-Action Interactive Model(世界-动作交互模型,简称WAIM)的新范式,将未来世界状态和驾驶动作视为需要在推理过程中共同推断的耦合变量,而非独立生成或单向依赖的输出。这一理念的核心在于:在决策相关的未来本质上依赖于所考虑的动作时,模型不应先预测一个被动的世界未来再在其中行动,而应该联合推断一个世界演化与决策制定保持相互对齐的未来。

2. 从WAM到WAIM:范式转变的本质
2.1 传统WAM的结构性局限
传统的自动驾驶策略直接建模条件概率分布 p ( a 1 : H ∣ o , l ) p(a_{1:H}\mid o,l) p(a1:H∣o,l),其中 o o o 表示当前观测, l l l 表示任务指令, a 1 : H a_{1:H} a1:H 表示未来 H H H 步的动作序列。这种直接映射忽略了未来世界演化对决策的影响。
World Action Model通过引入未来世界表示 v 1 : T v_{1:T} v1:T,将问题扩展为联合建模 p ( v 1 : T , a 1 : H ∣ o , l ) p(v_{1:T}, a_{1:H}\mid o,l) p(v1:T,a1:H∣o,l)。理论上,动作分布可以通过对所有可能未来的边缘化得到:
p ( a 1 : H ∣ o , l ) = ∫ p ( v 1 : T , a 1 : H ∣ o , l ) d v 1 : T p(a_{1:H}\mid o,l)=\int p(v_{1:T},a_{1:H}\mid o,l)\,dv_{1:T} p(a1:H∣o,l)=∫p(v1:T,a1:H∣o,l)dv1:T
然而,现有WAM的实现在推理时往往采取两种极端策略:
-
并行生成:使用独立的分支从共享的视觉上下文同时预测未来场景和动作,二者在训练时相关但在生成时各自独立。
-
顺序流水线:先预测未来观测、占用栅格或潜在场景状态,然后基于这些预测的未来进行动作规划。
这两种策略的共同问题在于:在生成时刻,一方相对于另一方是固定的。并行设计允许相关性但不允许迭代重塑;顺序设计让动作依赖于一个冻结的未来假设,而非随动作假设共同演化的未来。

2.2 WAIM的核心思想
WAIM将World-Action Interactive Model定义为一类特殊的WAM,其中未来世界和未来动作被推断为耦合变量,而非独立生成或按固定单向顺序生成。形式化地,WAIM寻求一个自洽的对 ( v ^ 1 : T , a ^ 1 : H ) (\hat{v}{1:T}, \hat{a}{1:H}) (v^1:T,a^1:H),使得:
v ^ 1 : T = F θ ( o , l , a ^ 1 : H ) , a ^ 1 : H = G ϕ ( o , l , v ^ 1 : T ) \hat{v}{1:T}=F{\theta}(o,l,\hat{a}{1:H}),\qquad \hat{a}{1:H}=G_{\phi}(o,l,\hat{v}_{1:T}) v^1:T=Fθ(o,l,a^1:H),a^1:H=Gϕ(o,l,v^1:T)
在实践中,这可以通过迭代交互实现:
( v 1 : T ( k + 1 ) , a 1 : H ( k + 1 ) ) = I Θ ( v 1 : T ( k ) , a 1 : H ( k ) ; o , l ) (v_{1:T}^{(k+1)},a_{1:H}^{(k+1)}) =\mathcal{I}{\Theta}(v{1:T}^{(k)},a_{1:H}^{(k)};o,l) (v1:T(k+1),a1:H(k+1))=IΘ(v1:T(k),a1:H(k);o,l)
其中 I Θ \mathcal{I}_{\Theta} IΘ 表示一个交互算子,它同时更新世界假设和动作假设。
关键区别在于:WAM联合建模未来世界和动作,而WAIM通过交互联合推断它们。这种区别在决策相关的未来依赖于所考虑的动作,而非仅依赖于场景动力学时尤为重要。
在本地代码里,WAIM的交互思想主要落在验证入口 DAWN/scripts/infer_nuscenes_val.py:先构建 encoder/predictor/planner/token_ae,再把这些模块交给 run_validation 执行预测和规划闭环。
python
metrics = run_validation(
encoder=encoder,
predictor=predictor,
planner=planner,
val_loader=val_loader,
val_sampler=val_sampler,
config=config,
epoch=args.epoch,
rank=rank,
world_size=world_size,
use_tubelet_repeat=config.data.use_tubelet_repeat,
vis_output_dir=vis_output_dir,
token_ae=token_ae,
)
3. DAWN架构设计:紧凑潜在空间中的交互式推理
3.1 整体架构概览
DAWN在紧凑的语义潜在空间中实例化WAIM,避免了昂贵的像素级未来渲染。其核心思想是使用短期显式潜在推演来支持复杂交互场景中的长期动作生成。架构包含以下关键组件:
- Student Vision-Encoder:从当前观测中提取密集视觉特征
- Teacher Vision-Encoder:仅在训练时使用,为未来观测提供监督信号
- Auto-Encoder Resampler:将密集编码器特征压缩为紧凑的潜在世界表示
- World Predictor:在潜在空间中预测未来世界状态
- World-Conditioned Action Denoiser:基于预测的未来世界去噪动作假设
- Action Head:将去噪后的动作状态解码为最终轨迹
这种设计的独特之处在于,它不需要在整个动作规划时域内推演世界状态,也不需要在像素空间中进行未来渲染。相反,它在一个紧凑的潜在空间中执行短期世界推演,这足以支持长期轨迹生成。

、
| 论文模块 | 本地代码位置 | 作用 |
|---|---|---|
| Runtime入口 | DAWN/scripts/infer_nuscenes_val.py |
构建encoder、predictor、planner、dataloader,并加载checkpoint |
| 模型初始化 | DAWN/app/vjepa_cowa_world_model/training/models.py |
初始化V-JEPA encoder、TokenAE、World Predictor、Diffusion Planner |
| Auto-Encoder Resampler | DAWN/app/vjepa_cowa_world_model/models/token_ae.py |
将密集ViT tokens压缩为少量latent world tokens |
| World Predictor | DAWN/src/models/ac_predictor.py |
VisionTransformerPredictorAC,动作条件的未来latent预测器 |
| Action Denoiser | DAWN/app/vjepa_cowa_world_model/models/diffusion_planner.py |
DiffusionPlanner / TrajectoryDiT,基于latent world tokens生成轨迹 |
| 验证闭环 | DAWN/app/vjepa_cowa_world_model/val_command.py |
执行encoder -> predictor rollout -> planner -> metrics |

3.2 视觉编码与潜在压缩
给定当前观测 o o o,Student Vision-Encoder提取密集视觉tokens:
u = E s t u ( o ) u=E_{\mathrm{stu}}(o) u=Estu(o)
在实现中,Student和Teacher分支都使用V-JEPA 2 Large作为视觉骨干网络。由于密集编码器tokens的直接推演成本过高,DAWN引入了Auto-Encoder Resampler,这是一个在token空间中操作的学习瓶颈自编码器:
z = R s t u ( u ) z=R_{\mathrm{stu}}(u) z=Rstu(u)
这产生了一个紧凑的潜在世界表示,用于下游交互。在训练过程中,未来观测 o + o^{+} o+ 通过Teacher Vision-Encoder及其对应的resampler处理,生成目标未来潜在表示:
z t a r g e t = R t e a ( E t e a ( o + ) ) z_{\mathrm{target}}=R_{\mathrm{tea}}(E_{\mathrm{tea}}(o^{+})) ztarget=Rtea(Etea(o+))
这种设计的关键创新在于:它不是在高维像素空间或密集特征空间中进行世界建模,而是学习一个极其紧凑的潜在表示(例如,仅16个tokens),这个表示保留了动作规划所需的关键场景信息,同时大幅降低了推演的计算成本。
实现细节显示,Resampler使用16个注意力头,4层编码器和2层解码器架构,MLP比率为4.0,并采用可学习的位置嵌入。这种配置在表达能力和计算效率之间取得了良好平衡。
本地实现里,TokenAE的推理压缩接口在 DAWN/app/vjepa_cowa_world_model/models/token_ae.py,验证流程里的实际调用在 DAWN/app/vjepa_cowa_world_model/val_command.py。
来源:DAWN/app/vjepa_cowa_world_model/models/token_ae.py
python
def compress(self, x: torch.Tensor, num_frames: int) -> torch.Tensor:
"""Compress tokens for inference (no reconstruction, no loss).
Parameters
----------
x : [B, T * tokens_per_frame, D]
num_frames : int
Returns
-------
z : [B, T * num_latent_tokens, D]
"""
return self.encode(x, num_frames)
来源:DAWN/app/vjepa_cowa_world_model/val_command.py
python
def _compress_tokens_with_token_ae(token_ae, tokens: torch.Tensor, num_frames: int) -> torch.Tensor:
"""Apply frozen Token AE compression on concatenated frame tokens."""
ae_tokens_per_frame = int(getattr(token_ae, "tokens_per_frame"))
expected_tokens = int(num_frames) * ae_tokens_per_frame
if tokens.size(1) != expected_tokens:
if tokens.size(1) % ae_tokens_per_frame != 0:
raise ValueError(
"Cannot infer TokenAE frame count: "
f"tokens={tokens.size(1)}, num_frames={num_frames}, ae_tokens_per_frame={ae_tokens_per_frame}"
)
num_frames = tokens.size(1) // ae_tokens_per_frame
return token_ae.encode(tokens, num_frames=num_frames)
3.3 递归交互机制:世界预测与动作去噪的耦合
DAWN的核心创新在于World Predictor和World-Conditioned Action Denoiser之间的递归交互。World Predictor实现为一个因果Transformer,从当前潜在上下文和当前动作假设预测未来潜在世界tokens。World-Conditioned Action Denoiser实现为DiT(Diffusion Transformer),它基于潜在上下文和预测的未来世界对动作tokens进行去噪。
设条件tokens c包含自车状态和高级动作或路线tokens。Action Denoiser额外接收角色特定查询,指示它是生成初始提议还是使用预测器推演来细化动作。DAWN执行以下递归过程:
a 1 : H ( 0 ) = G ϕ ( q p r o p , c , z ) a_{1:H}^{(0)}=G_{\phi}(q_{\mathrm{prop}},c,z) a1:H(0)=Gϕ(qprop,c,z)
z f u t u r e ( r ) = P θ ( z , c , a 1 : H ( r ) ) z_{\mathrm{future}}^{(r)}=P_{\theta}(z,c,a_{1:H}^{(r)}) zfuture(r)=Pθ(z,c,a1:H(r))
a 1 : H ( r + 1 ) = G ϕ ( q r e f ( r ) , c , z f u t u r e ( r ) , a 1 : H ( r ) ) a_{1:H}^{(r+1)} =G_{\phi}(q_{\mathrm{ref}}^{(r)},c,z_{\mathrm{future}}^{(r)},a_{1:H}^{(r)}) a1:H(r+1)=Gϕ(qref(r),c,zfuture(r),a1:H(r))
这里的关键设计是:
- q p r o p q_{\mathrm{prop}} qprop 和 q r e f ( r ) q_{\mathrm{ref}}^{(r)} qref(r) 是角色特定的查询嵌入,用于提议生成和细化
- 去噪器权重在两个角色之间共享,仅输入源和查询嵌入不同
- 预测的未来世界 z f u t u r e ( r ) z_{\mathrm{future}}^{(r)} zfuture(r) 条件动作去噪
- 去噪的动作假设 a 1 : H ( r ) a_{1:H}^{(r)} a1:H(r) 被反馈以更新世界预测
这种设计的深刻之处在于:它在训练和推理时都保持相同的交互模式。在训练阶段,模型学习两个角色:首先从resampler潜在上下文生成初始提议,然后基于预测器推演细化动作。不同的查询和源嵌入指定去噪器当前操作的是提议还是交互式细化角色。
代码实现中,World Predictor使用12层Transformer,嵌入维度384,12个注意力头,RoPE位置编码,以及激活检查点以节省内存。Action Denoiser采用DiT-style骨干,隐藏维度384,12层,12个注意力头,MLP比率4.0。与初始扩散规划器块相比,DiT块更紧密地遵循原始adaLN-Zero设计:时间步/状态条件向量不仅调制自注意力和MLP分支,还调制交叉注意力分支到潜在世界tokens。
VisionTransformerPredictorAC在forward里显式接收视觉tokens、动作、状态和可选外参;这对应论文中的 P θ ( z , c , a ) P_{\theta}(z,c,a) Pθ(z,c,a)。DiffusionPlanner.forward接收预测器输出的 z_ar 和自车状态 status_feature,这对应世界条件动作去噪器。
来源:DAWN/src/models/ac_predictor.py
python
# Fwd prop
for i, blk in enumerate(self.predictor_blocks):
if self.use_activation_checkpointing:
x = torch.utils.checkpoint.checkpoint(
blk,
x,
mask=None,
attn_mask=attn_mask,
T=T,
H=self.grid_height,
W=self.grid_width,
action_tokens=cond_tokens,
use_reentrant=False,
)
else:
x = blk(
x,
mask=None,
attn_mask=attn_mask,
T=T,
H=self.grid_height,
W=self.grid_width,
action_tokens=cond_tokens,
)
# Split out action and frame tokens
x = x.view(B, T, cond_tokens + self.grid_height * self.grid_width, D) # [B, T, K+H*W, D]
x = x[:, :, cond_tokens:, :].flatten(1, 2)
x = self.predictor_norm(x)
x = self.predictor_proj(x)
return x
来源:DAWN/app/vjepa_cowa_world_model/models/diffusion_planner.py
python
def forward(
self,
z_ar: torch.Tensor,
status_feature: torch.Tensor,
z_context: Optional[torch.Tensor] = None,
z_observed: Optional[torch.Tensor] = None,
action_history: Optional[torch.Tensor] = None,
gt_trajectory: Optional[torch.Tensor] = None,
anchor_state: Optional[torch.Tensor] = None,
) -> Dict[str, torch.Tensor]:
"""
Forward pass.
Args:
z_ar: [B, T*P, encoder_dim] --- predictor output tokens
status_feature: [B, status_dim] --- ego status
z_context: [B, P, encoder_dim] --- optional first-frame encoder tokens
z_observed: [B, T_obs*P, encoder_dim] --- optional observed-frame encoder tokens
action_history: [B, T_obs, action_history_dim] --- optional observed history tokens
gt_trajectory: [B, num_poses, 6] --- ground truth in 6-dim format
(x, y, vx, vy, cos_yaw, sin_yaw). Only needed for training.
anchor_state: [B, 6] or [B, 1, 6] --- current ego state for anchor frame.
Only used when ``use_anchor_frame=True``. If *None*, a
default zero-origin anchor is used.
Returns:
Training: {"loss", "reg_loss", "conf_loss", "cover_loss"}
Inference: {"trajectories": [B, K, num_poses, 3], "confidences": [B, K]}
"""
# Prepare conditioning
context_tokens = self._prepare_context(z_ar, z_context, z_observed, action_history) # [B, N, hidden_dim]
status_emb = self._prepare_status(status_feature) # [B, hidden_dim]
if gt_trajectory is not None and self.training:
return self._training_forward(context_tokens, status_emb, gt_trajectory, anchor_state)
else:
return self._inference_forward(context_tokens, status_emb, anchor_state)

4. 训练流程:四阶段渐进式学习策略
4.1 阶段一:大规模视觉预训练
DAWN采用四阶段训练策略,每个阶段都有明确的学习目标。第一阶段在大规模驾驶视频数据上预训练Student Vision-Encoder,包括OpenScene、DrivingDojo和CoVLA等数据集。所有数据集被转换为统一的视频格式,使用滑动窗口采样。
预训练在256×512分辨率和2Hz帧率下进行。这一阶段的关键是建立强大的视觉先验,使编码器能够从原始像素中提取驾驶场景的语义信息。预训练目标通常是视频级别的自监督任务,如V-JEPA中的masked prediction,使模型学习场景的时空结构和动态特性。
这一阶段可以概括为:
u = E s t u ( o ) , L s s l = ℓ s s l ( u , m a s k e d t a r g e t s ) u=E_{\mathrm{stu}}(o),\qquad \mathcal{L}{\mathrm{ssl}}=\ell{\mathrm{ssl}}(u,\mathrm{masked\ targets}) u=Estu(o),Lssl=ℓssl(u,masked targets)
Student Encoder通过自监督损失更新,Teacher Encoder通过EMA同步:
E t e a ← E M A ( E s t u , E t e a ) E_{\mathrm{tea}}\leftarrow \mathrm{EMA}(E_{\mathrm{stu}},E_{\mathrm{tea}}) Etea←EMA(Estu,Etea)
这一阶段的训练数据规模通常达到数百万甚至上千万视频片段,确保编码器能够泛化到各种驾驶场景、天气条件和光照变化。
DAWN/app/vjepa_cowa_world_model/training/models.py
python
def init_encoder(
config: TrainingConfig,
device: torch.device,
) -> Tuple[nn.Module, nn.Module]:
"""
初始化 encoder 和 target_encoder
根据 config.model.backbone 选择 V-JEPA 2 或 V-JEPA 2.1 编码器。
Args:
config: 训练配置
device: 设备
Returns:
Tuple[nn.Module, nn.Module]: (encoder, target_encoder)
"""
encoder = init_context_encoder(config, device)
target_encoder = copy.deepcopy(encoder)
logger.info("end init encoder")
# 打印参数量(仅主进程)
encoder_params = sum(p.numel() for p in encoder.parameters())
target_encoder_params = sum(p.numel() for p in target_encoder.parameters())
if _is_main_process():
logger.info(f"init encoder_params: {encoder_params / 1e6:>8.2f}M")
logger.info(f"init target_encoder_params: {target_encoder_params / 1e6:>8.2f}M")
return encoder, target_encoder
DAWN/app/vjepa_cowa_world_model/training/ema.py
python
@torch.no_grad()
def update_ema(m: float) -> None:
"""
使用动态 momentum 更新 target encoder 的参数
采用高效的原地操作方式
Args:
m: 当前的 momentum 值 (从 momentum_scheduler 获取)
"""
# 收集所有参数
params_k = []
params_q = []
for param_q, param_k in zip(encoder.parameters(), target_encoder.parameters()):
params_k.append(param_k)
params_q.append(param_q)
# 高效的原地操作
torch._foreach_mul_(params_k, m)
torch._foreach_add_(params_k, params_q, alpha=1 - m)
4.2 阶段二:Token空间瓶颈学习
第二阶段在相同的预训练语料上训练Auto-Encoder Resampler。这个阶段从预训练的编码器开始,学习一个紧凑的token空间瓶颈,将密集视觉tokens压缩为潜在世界tokens,同时保留未来预测和动作生成所需的信息。
关键创新是同时训练一个辅助扩散规划器头,与Action Denoiser使用相同的DiT-style配置。这个辅助头鼓励压缩tokens保留动作相关信息,确保瓶颈表示不仅适合世界建模,更适合下游规划任务。
这一阶段可以概括为:
u = E s t u ( o ) , z = R s t u ( u ) u=E_{\mathrm{stu}}(o),\qquad z=R_{\mathrm{stu}}(u) u=Estu(o),z=Rstu(u)
u ^ = R s t u d e c ( z ) , L r e c = ℓ r e c ( u ^ , u ) \hat{u}=R_{\mathrm{stu}}^{\mathrm{dec}}(z),\qquad \mathcal{L}{\mathrm{rec}}=\ell{\mathrm{rec}}(\hat{u},u) u^=Rstudec(z),Lrec=ℓrec(u^,u)
为了让压缩后的tokens保留动作相关信息,还引入辅助规划损失:
L A E = L r e c + L a u x _ a c t \mathcal{L}{\mathrm{AE}}=\mathcal{L}{\mathrm{rec}}+\mathcal{L}_{\mathrm{aux\_act}} LAE=Lrec+Laux_act
这一设计的深刻之处在于:它不是简单的降维,而是学习一个任务特定的压缩,使得紧凑表示既能支持未来世界推演,又能支持动作规划。
DAWN/app/vjepa_cowa_world_model/training/models.py
python
def prepare_runtime_tokens(
tokens: torch.Tensor,
num_frames: int,
normalize_reps: bool,
token_ae: Optional[nn.Module] = None,
) -> torch.Tensor:
"""Compress per-frame tokens with TokenAE if present, then apply runtime normalization."""
input_dtype = tokens.dtype
if token_ae is not None:
ae_tokens_per_frame = int(getattr(token_ae, "tokens_per_frame"))
expected_tokens = int(num_frames) * ae_tokens_per_frame
if tokens.size(1) != expected_tokens:
if tokens.size(1) % ae_tokens_per_frame != 0:
raise ValueError(
"Cannot infer TokenAE frame count: "
f"tokens={tokens.size(1)}, num_frames={num_frames}, "
f"ae_tokens_per_frame={ae_tokens_per_frame}"
)
num_frames = tokens.size(1) // ae_tokens_per_frame
token_ae_parameters = getattr(token_ae, "parameters", None)
ae_param = next(token_ae_parameters(), None) if callable(token_ae_parameters) else None
if ae_param is not None and tokens.is_floating_point() and tokens.dtype != ae_param.dtype:
tokens = tokens.to(dtype=ae_param.dtype)
tokens = token_ae.encode(tokens, num_frames=num_frames)
if tokens.is_floating_point() and tokens.dtype != input_dtype:
tokens = tokens.to(dtype=input_dtype)
if normalize_reps:
4.3 阶段三:World Predictor训练
第三阶段在下游任务数据集如nuScenes和NAVSIM上训练World Predictor。此阶段,预测器学习从预训练编码器和resampler产生的紧凑潜在上下文中推演任务相关的未来潜在世界状态。
World Predictor被设计为因果Transformer,能够自回归地生成未来的潜在tokens序列。训练目标是最小化预测的未来潜在表示与Teacher分支生成的目标之间的距离:
这一阶段可以概括为:
z = R s t u ( E s t u ( o ) ) , z t a r g e t = R t e a ( E t e a ( o f u t u r e ) ) z=R_{\mathrm{stu}}(E_{\mathrm{stu}}(o)),\qquad z_{\mathrm{target}}=R_{\mathrm{tea}}(E_{\mathrm{tea}}(o_{\mathrm{future}})) z=Rstu(Estu(o)),ztarget=Rtea(Etea(ofuture))
z ^ f u t u r e = P θ ( z , c ) , L W M = d ( z ^ f u t u r e , z t a r g e t ) \hat{z}{\mathrm{future}}=P{\theta}(z,c),\qquad \mathcal{L}{\mathrm{WM}}=d(\hat{z}{\mathrm{future}},z_{\mathrm{target}}) z^future=Pθ(z,c),LWM=d(z^future,ztarget)
这一阶段的关键是使预测器学习驾驶场景的动态演化规律,包括其他车辆的运动模式、交通规则的隐式约束,以及场景几何的时空变化。
DAWN/app/vjepa_cowa_world_model/training/predictor_loss.py
python
observed_steps = int(num_observed_steps) if num_observed_steps is not None else config.train.num_observed_frames
if bool(getattr(config.train, "use_parallel_predictor", False)):
future_start_step = observed_steps if config.train.predictor_inference_consistent else 1
future_offset = future_start_step * tokens_per_frame
jloss = loss_fn(
z_tf[:, future_offset:],
h_target,
offset=future_offset,
)
if config.train.predictor_use_z_ar_supervision:
if z_ar is None:
raise ValueError("z_ar must be provided when predictor_use_z_ar_supervision=True")
sloss = loss_fn(z_ar, h_target, offset=future_offset)
else:
sloss = jloss * 0.0
jepa_loss = jloss + sloss
return jepa_loss, jloss, sloss
4.4 阶段四:联合世界-动作训练
最后阶段从阶段三初始化World Predictor,附加World-Conditioned Action Denoiser和Action Head,在目标数据集上联合训练世界和动作分支。这是WAIM范式实现的关键阶段。
在此阶段,预测器和动作去噪器同时优化。Action Denoiser在两个角色下训练,权重共享:首先从resampler潜在上下文生成初始提议,然后基于预测器推演细化动作。不同的查询和源嵌入指定去噪器当前操作的角色。
这一阶段可以概括为:
z = R s t u ( E s t u ( o ) ) , z t a r g e t = R t e a ( E t e a ( o f u t u r e ) ) z=R_{\mathrm{stu}}(E_{\mathrm{stu}}(o)),\qquad z_{\mathrm{target}}=R_{\mathrm{tea}}(E_{\mathrm{tea}}(o_{\mathrm{future}})) z=Rstu(Estu(o)),ztarget=Rtea(Etea(ofuture))
a 1 : H ( 0 ) = G ϕ ( q p r o p , c , z ) a_{1:H}^{(0)}=G_{\phi}(q_{\mathrm{prop}},c,z) a1:H(0)=Gϕ(qprop,c,z)
第 r r r 轮递归交互为:
z f u t u r e ( r ) = P θ ( z , c , a 1 : H ( r ) ) z_{\mathrm{future}}^{(r)}=P_{\theta}(z,c,a_{1:H}^{(r)}) zfuture(r)=Pθ(z,c,a1:H(r))
a 1 : H ( r + 1 ) = G ϕ ( q r e f ( r ) , c , z f u t u r e ( r ) , a 1 : H ( r ) ) a_{1:H}^{(r+1)} =G_{\phi}(q_{\mathrm{ref}}^{(r)},c,z_{\mathrm{future}}^{(r)},a_{1:H}^{(r)}) a1:H(r+1)=Gϕ(qref(r),c,zfuture(r),a1:H(r))
最终轨迹和联合损失为:
τ ^ = H a c t ( a 1 : H ( R ) ) \hat{\tau}=H_{\mathrm{act}}(a_{1:H}^{(R)}) τ^=Hact(a1:H(R))
L t o t a l = L W M + L p l a n \mathcal{L}{\mathrm{total}} =\mathcal{L}{\mathrm{WM}}+\mathcal{L}_{\mathrm{plan}} Ltotal=LWM+Lplan
训练方案的核心在于:模型不仅学习预测未来世界和生成动作,更重要的是学习如何让这两个过程通过递归交互相互增强。
DAWN/app/vjepa_cowa_world_model/training/models.py
python
token_ae_enabled = bool(getattr(config, "token_ae", None) and config.token_ae.enabled)
if token_ae_enabled:
effective_tokens_per_frame = resolve_effective_tokens_per_frame(config)
else:
effective_tokens_per_frame = (
int(raw_tokens_per_frame_override)
if raw_tokens_per_frame_override is not None
else resolve_effective_tokens_per_frame(config)
)
runtime_normalize_reps = resolve_predictor_runtime_normalize_reps(config)
token_ae = None
if token_ae_enabled:
token_ae, runtime_normalize_reps = load_frozen_token_ae(
config,
device=device,
encoder_embed_dim=encoder_embed_dim,
tokens_per_frame=raw_tokens_per_frame,
normalize_reps=runtime_normalize_reps,
dtype=config.dtype,
)
predictor = init_predictor_for_ae(
config,
device=device,
encoder_embed_dim=encoder_embed_dim,
num_latent_tokens=effective_tokens_per_frame,
latent_grid_size=getattr(token_ae, "latent_grid_size", config.token_ae.latent_grid_size),
)
else:
predictor = init_predictor(
config,
device,
encoder_embed_dim,
predictor_img_size_override=predictor_img_size_override,
)
return predictor, token_ae, effective_tokens_per_frame, runtime_normalize_reps