DreamVLA 方法原理说明
DreamVLA 方法原理通俗详解
1. 先看全流程

DreamVLA 的一句话核心是:
让机器人一边学习动作,一边学习预测动作之后的未来世界。
普通 VLA 策略通常做:
图像 + 语言 + 状态 -> 动作
DreamVLA 做的是:
图像 + 语言 + 状态
-> 动作
-> 未来 RGB
-> 未来深度
-> 未来轨迹
-> 未来 DINO/SAM 视觉特征
推理时真正执行的还是动作。未来 RGB、深度、轨迹、DINO/SAM 特征主要是训练时的辅助监督,用来逼模型学到更好的世界动态理解。
先给一段总伪代码,后面每一节会拆开讲:
def dreamvla_train_step(batch):
# 1. 从一段机器人演示里取历史窗口
images_primary = batch.primary_rgb[:, :S]
images_wrist = batch.wrist_rgb[:, :S]
states = batch.robot_state[:, :S]
language = batch.language
# 2. 把输入变成 token
text_tokens = encode_text(language)
state_tokens = encode_state(states)
visual_tokens = encode_and_resample_images(images_primary, images_wrist)
# 3. 加入“问题卡片”:动作查询 + 未来世界查询
tokens = build_tokens(
text_tokens,
state_tokens,
visual_tokens,
action_queries,
future_world_queries,
)
# 4. 时序 Transformer 负责融合历史,但不能偷看未来
hidden = causal_transformer(tokens, attention_mask=dreamvla_mask)
# 5. 不同 query 的 hidden state 解码成不同答案
pred_action = decode_action(hidden[action_query_positions])
pred_rgb = decode_future_rgb(hidden[rgb_query_positions])
pred_depth = decode_future_depth(hidden[depth_query_positions])
pred_traj = decode_future_trajectory(hidden[trajectory_query_positions])
pred_dino = decode_future_dino(hidden[dino_query_positions])
pred_sam = decode_future_sam(hidden[sam_query_positions])
# 6. 动作 loss 是主任务;未来世界 loss 是辅助任务
loss = action_loss(pred_action, batch.actions)
loss += weighted_future_world_losses(
pred_rgb, pred_depth, pred_traj, pred_dino, pred_sam, batch
)
return loss
把它说成人话:
1. 先看最近几帧机器人经历了什么。
2. 把语言、图像、机器人状态都翻译成 token。
3. 放入一些查询 token,让模型回答“下一步怎么动”和“未来会怎样”。
4. 用时序 Transformer 做推理。
5. 训练时同时监督动作和未来世界。
6. 推理时只拿动作去控制机器人。
2. 输入数据是什么
每个训练样本是一段连续机器人轨迹。DreamVLA 关心四类基础输入:
1. 主视角 RGB 图像
2. 腕部/夹爪相机 RGB 图像
3. 语言指令
4. 机器人本体状态
训练时还可能有这些标签:
1. 专家动作
2. 未来深度
3. 未来 DINO 特征
4. 未来 SAM 特征
5. 未来轨迹/光流式位移
可以把一个 batch 粗略写成:
batch = {
"image_primary": Tensor[B, window_size, 3, H, W],
"image_wrist": Tensor[B, window_size, 3, H, W],
"text_token": Tensor[B, text_len],
"state": Tensor[B, window_size, state_dim],
"action": Tensor[B, window_size, 7],
"depth_primary": Optional[Tensor],
"depth_wrist": Optional[Tensor],
"dino_primary": Optional[Tensor],
"dino_wrist": Optional[Tensor],
"sam_primary": Optional[Tensor],
"sam_wrist": Optional[Tensor],
"track_infos": Optional[Dict],
}
这里最重要的几个符号:
B = batch size
S = sequence_length,模型真正输入的历史长度
H = action_pred_steps,每个时间步预测多少步未来动作
F = future_steps,未来观测标签相对当前向后偏移多少步
P = pred_num,每个时间步预测多少个未来观测目标
Q = num_resampler_query,每个相机保留多少个视觉摘要 token
D = hidden_dim,Transformer 内部隐藏维度
训练循环一开始会从窗口里取前 S 步作为模型输入:
input_primary = image_primary[:, :S]
input_wrist = image_wrist[:, :S]
input_state = state[:, :S]
input_text = repeat(text_token, times=S)
动作是 7 维:
前 6 维:机械臂位姿/运动控制
最后 1 维:夹爪开合
训练动作标签通常是:
label_actions[t] = actions[t : t + H]
也就是说,在时间步 t,模型可以预测未来 H 步动作,而不是只能预测一步。
3. 语言和状态怎么变成 token
3.1 语言 token
语言指令先由 CLIP 文本编码器编码,再投影到 DreamVLA 的 hidden 维度。
def encode_text(text_token):
# CLIP 输出 512 维文本特征
with no_grad():
text_feature = CLIP.encode_text(text_token) # [B*S, 512]
# 投影到 Transformer 使用的 hidden_dim
text_embedding = text_projector(text_feature) # [B*S, D]
text_embedding = reshape(text_embedding, [B, S, 1, D])
return text_embedding
源码确认:CLIP 文本编码器被冻结,主要训练后面的投影层和策略网络。
为什么要这样做?因为 CLIP 已经有较强语言理解能力,DreamVLA 不需要从零学语言,而是把 CLIP 语言特征接到机器人策略里。
3.2 状态 token
机器人状态被拆成机械臂状态和夹爪状态。
arm_state = state 前 6 维
gripper_state = state 夹爪维度
如果不是 gripper width 模式,夹爪状态会转成开/关 one-hot;如果是 gripper width 模式,就直接用连续夹爪宽度。
伪代码:
def encode_state(state):
arm = state[..., :6]
gripper = state[..., 6:]
arm_feature = linear_arm_state(arm)
if not gripper_width:
gripper = to_open_close_onehot(gripper)
gripper_feature = linear_gripper_state(gripper)
state_embedding = state_projector(concat(arm_feature, gripper_feature))
state_embedding = reshape(state_embedding, [B, S, 1, D])
return state_embedding
为什么要单独给状态?因为只看图像不一定能准确知道机械臂的位姿、夹爪开合和控制状态。机器人策略必须知道“自己在哪里”。
4. 图像怎么变成 token
DreamVLA 使用两个相机:
主视角:看整体桌面、物体、任务场景
腕部视角:从夹爪附近看局部接触和细节
图像编码分两步:
图像 -> 冻结视觉编码器 -> 大量 patch token
大量 patch token -> Perceiver Resampler -> 少量视觉摘要 token
先看总体伪代码:
def encode_and_resample_images(primary, wrist):
# primary, wrist: [B, S, 3, 224, 224]
with no_grad():
primary_tokens = vision_encoder(primary.reshape(B*S, 3, 224, 224))
wrist_tokens = vision_encoder(wrist.reshape(B*S, 3, 224, 224))
primary_cls = primary_tokens[:, :1]
wrist_cls = wrist_tokens[:, :1]
primary_patch = primary_tokens[:, 1:]
wrist_patch = wrist_tokens[:, 1:]
primary_resampled = perceiver_resampler(primary_patch)
wrist_resampled = perceiver_resampler(wrist_patch)
primary_resampled = image_primary_projector(primary_resampled)
wrist_resampled = image_wrist_projector(wrist_resampled)
primary_cls = cls_token_primary_projector(primary_cls)
wrist_cls = cls_token_wrist_projector(wrist_cls)
return concat(primary_resampled, wrist_resampled, primary_cls, wrist_cls)
4.1 默认 MAE 路径下的 patch 数
默认 MAE 视觉编码器使用 224 x 224 输入和 16 x 16 patch。
224 / 16 = 14
14 x 14 = 196
所以一张图会得到:
1 个 cls/global token
196 个 patch token
源码实现里会把 cls token 单独拿出来,剩下 196 个 patch token 交给 Perceiver Resampler。
image_cls = image_feature[:, :1, :] # [B*S, 1, 768]
image_patch = image_feature[:, 1:, :] # [B*S, 196, 768]
如果启用 DINO + SigLIP 路径,patch 数和维度会不同。当前实现中 DINO/SigLIP 路径按 256 个 patch token 处理,并把 DINO 和 SigLIP patch 特征拼接成更高维视觉特征。核心思想不变:先得到很多 patch token,再压缩成固定数量视觉 token。
5. 图像压缩不是随机采样
这一节单独讲清楚“采样”怎么做。DreamVLA 不是随机挑几个 patch,也不是固定取图像中心或角落。它用 Perceiver Resampler 做软注意力汇聚。
5.1 Perceiver Resampler 的输入输出
默认 MAE 路径下,单个相机单帧图像进入 Resampler 前是:
196 个 patch token,每个 token 是 768 维
如果 num_resampler_query = 16,输出就是:
16 个视觉摘要 token
形状变化:
[196, 768] -> [16, 768] -> 线性投影 -> [16, hidden_dim]
两个相机都做一遍:
主视角:196 patch -> 16 摘要 token
腕部视角:196 patch -> 16 摘要 token
再加上两个 cls/global token:
主视角:16 resampler token + 1 cls token
腕部视角:16 resampler token + 1 cls token
视觉相关 token 总数 = 34
5.2 Resampler 具体怎么汇聚
Perceiver Resampler 内部有 Q 个可学习 latent token。可以把它们理解成 Q 个观察员。
latents = Parameter(shape=[Q, vision_dim])
每个 latent 都会去看全部图像 patch。注意力过程大概是:
def perceiver_resampler(patch_tokens):
# patch_tokens: [196, vision_dim]
latents = learned_latents # [Q, vision_dim]
for layer in layers:
# latent 当 query
q = linear_q(norm(latents))
# patch token 和 latent 一起当 key/value
kv_input = concat(norm(patch_tokens), norm(latents))
k, v = linear_kv(kv_input)
# 每个 latent 对所有 patch/latent 算注意力
weights = softmax(q @ k.T)
# 加权汇聚,更新 latent
latents = latents + weights @ v
latents = latents + feed_forward(latents)
return norm(latents)
重点是:
每个输出 token 不是某一个固定 patch。
每个输出 token 是对整张图所有 patch 的加权融合。
这些权重由 attention 学出来。
所以它更像“软采样”:
不是选中某几个点,而是学习每个摘要 token 应该关注哪些区域。
例如 16 个 latent 可以学成不同的关注模式:
latent 1:更关注夹爪附近
latent 2:更关注目标物体
latent 3:更关注桌面和容器
latent 4:更关注背景里的任务线索
...
这些角色不是人工指定的,而是训练过程中由损失推动形成的。
5.3 为什么不用硬采样
硬采样需要提前规定保留哪些 patch,比如中心区域或规则网格。但机器人任务里的关键区域是动态变化的:
杯子可能在左边,也可能在右边。
夹爪会移动。
接触区域每一步都不同。
任务相关物体也会变。
Perceiver Resampler 的优点是:
模型可以根据当前图像内容,自适应决定每个摘要 token 关注哪里。
这比固定采样更适合机器人操作。
6. 每个时间步的 token 怎么拼起来
经过语言、状态、图像编码后,每个时间步先有一组真实输入 token:
语言 token: 1 个
状态 token: 1 个
主视角 resampler token: Q 个
腕部 resampler token: Q 个
主视角 cls/global token: 1 个
腕部 cls/global token: 1 个
所以真实输入 token 数是:
num_A = 1 + 1 + Q + Q + 1 + 1 = 2Q + 4
如果 Q = 16:
num_A = 36
伪代码:
def build_real_input_tokens(text, state, primary_vis, wrist_vis, primary_cls, wrist_cls):
tokens = concat([
text, # [B, S, 1, D]
state, # [B, S, 1, D]
primary_vis, # [B, S, Q, D]
wrist_vis, # [B, S, Q, D]
primary_cls, # [B, S, 1, D]
wrist_cls, # [B, S, 1, D]
], dim="token")
return tokens # [B, S, 2Q+4, D]
这部分 token 描述的是“模型真实看到的信息”。
7. 查询 token 是什么
DreamVLA 不只放真实输入 token,还会放一批可学习查询 token。它们像模型内部的问题卡片。
动作查询 token:问“下一步或未来几步怎么动?”
RGB 查询 token:问“未来画面是什么样?”
深度查询 token:问“未来深度是什么样?”
轨迹查询 token:问“未来哪些点怎么动?”
DINO 查询 token:问“未来语义特征是什么?”
SAM 查询 token:问“未来区域/边界特征是什么?”
这些 token 不是数据集提供的输入,而是模型参数。
action_pred_token = Parameter([1, 1, action_pred_steps, D])
obs_tokens = Parameter([1, 1, num_rgb_query, D])
depth_tokens = Parameter([1, 1, num_depth_query, D])
dino_tokens = Parameter([1, 1, num_dino_query, D])
sam_tokens = Parameter([1, 1, num_sam_query, D])
trajectory_tokens = Parameter([1, 1, num_traj_query, D])
每个时间步都会复制一份这些查询 token:
query_tokens = []
if obs_pred:
query_tokens.append(repeat(obs_tokens, B, S))
if depth_pred:
query_tokens.append(repeat(depth_tokens, B, S))
if dino_feat_pred:
query_tokens.append(repeat(dino_tokens, B, S))
if sam_feat_pred:
query_tokens.append(repeat(sam_tokens, B, S))
if trajectory_pred:
query_tokens.append(repeat(trajectory_tokens, B, S))
if action_pred_steps > 0:
query_tokens.append(repeat(action_pred_token, B, S))
完整的单时间步 token 是:
tokens_t = concat(real_input_tokens_t, query_tokens_t)
经过 Transformer 后:
动作查询 token 的 hidden state -> 解码动作
RGB 查询 token 的 hidden state -> 解码未来 RGB
深度查询 token 的 hidden state -> 解码未来深度
轨迹查询 token 的 hidden state -> 解码未来轨迹
DINO/SAM 查询 token 的 hidden state -> 解码未来视觉特征
这就是 DreamVLA 的核心机制:同一个时序 Transformer 内部同时回答多个问题。
8. 时序 Transformer 和 attention mask
DreamVLA 使用 GPT2 风格 Transformer,但它不输入文字 token id,而是直接输入多模态 embedding。
8.1 token 序列怎么展平
每个时间步有 N 个 token:
N = 真实输入 token 数 + 查询 token 数
历史窗口有 S 个时间步,所以送进 Transformer 前会展平成:
[B, S, N, D] -> [B, S*N, D]
伪代码:
tokens = concat(real_input_tokens, query_tokens, dim="token")
tokens = tokens + timestep_position_embedding
tokens = tokens.reshape(B, S * N, D)
tokens = layer_norm(tokens)
hidden = GPT2_backbone(inputs_embeds=tokens, attention_mask=mask)
hidden = hidden.reshape(B, S, N, D)
8.2 为什么要 mask
训练时,模型不能偷看未来真实观测。比如第 3 帧预测动作时,不能看到第 5 帧的真实图像。
DreamVLA 的 mask 大概做三件事:
1. 当前时间步不能看未来时间步。
2. 预测查询 token 不能被普通 token 当作真实输入去读取。
3. 动作查询 token 可以读取观测预测查询 token,利用模型内部“想象的未来”。
简化伪代码:
def build_attention_mask(S, num_A, num_B):
# num_A: 每步真实输入 token 数
# num_B: 每步查询 token 数
N = num_A + num_B
mask = zeros([S*N, S*N])
for i in range(S):
cur_start = i * N
cur_end = (i + 1) * N
# 1. 第 i 步不能看 i 之后的时间步
mask[cur_start:cur_end, cur_end:] = -inf
# 2. 查询 token 是答案槽,不让所有 token 随便读它
query_start = cur_start + num_A
mask[:, query_start:cur_end] = -inf
# 3. 但动作 query 可以读观测 query
allow_action_to_read_obs_queries(mask, i)
return mask
8.3 atten_only_obs 的直觉
atten_only_obs 会进一步限制动作 token 的可见范围,让它更集中地读:
当前视觉 token
观测预测 query token
可选机器人本体状态
通俗理解:让动作决策更依赖模型自己形成的“未来世界表示”,而不是随意从所有 token 里找捷径。
8.4 atten_goal 的直觉
atten_goal 有两个作用:
1. 训练时窗口尾部没有足够未来标签,所以最后 atten_goal 个时间步不算某些 loss。
2. 配合 atten_goal_state 时,观测预测 token 可以读取未来目标状态 token。
直观上,它让模型学习:
从当前状态到未来目标状态之间,世界应该怎么变化。
9. 未来世界预测怎么做
未来预测不是一个单独模块,而是一组可选任务。每个任务都有自己的 query token、decoder 和 loss。
9.1 未来 RGB 预测
RGB 预测不是直接输出整张图片,而是输出图像 patch。
流程:
rgb_query_hidden = hidden[rgb_query_positions]
rgb_query_hidden = image_decoder_projector(rgb_query_hidden)
decoder_input = concat(rgb_query_hidden, mask_tokens)
decoder_input = decoder_input + image_2d_position_embedding
decoded = image_decoder(decoder_input)
future_patch_pred = linear(decoded[mask_token_positions])
它预测的是:
每个 patch 的 RGB 像素值
标签处理:
future_images = images[:, F : F + S - atten_goal + P - 1]
future_image_patches = patchify(future_images, patch_size)
future_image_patches = normalize_each_patch(future_image_patches)
损失:
loss_image = mse(pred_primary_rgb, label_primary_rgb)
loss_image += mse(pred_wrist_rgb, label_wrist_rgb)
loss_image *= 0.5
如果开启 flow_as_mask,RGB loss 只重点计算轨迹显示会动的区域:
motion_mask = build_mask_from_tracks(track_infos)
loss_image = mse(pred_rgb * motion_mask, label_rgb * motion_mask)
直觉:未来 RGB 让模型学习“动作之后画面会怎么变”。
9.2 未来深度预测
深度预测结构和 RGB 类似,但输出是深度。
depth_query_hidden = hidden[depth_query_positions]
depth_decoder_input = concat(project(depth_query_hidden), depth_mask_tokens)
depth_pred = depth_decoder(depth_decoder_input)
标签:
future_depth = depths[:, F : F + S - atten_goal + P - 1]
损失使用 SiLogLoss:
loss_depth = silog(pred_depth_primary, label_depth_primary)
loss_depth += silog(pred_depth_wrist, label_depth_wrist)
loss_depth *= 0.5
直觉:深度监督让模型学习物体和夹爪的空间关系。
9.3 未来轨迹预测
轨迹来自 CoTracker。它会追踪图像网格点未来移动到哪里。
标签大概是:
每个网格点的二维位移 dx, dy
训练时:
traj_query_hidden = hidden[trajectory_query_positions]
traj_decoder_input = concat(project(traj_query_hidden), traj_mask_tokens)
traj_pred = trajectory_decoder(traj_decoder_input)
如果不开 no_unshuffle,轨迹标签会经过 pixel unshuffle 降采样到 patch 网格:
label_tracks = rearrange_to_grid(label_tracks)
label_tracks = pixel_unshuffle(label_tracks)
label_tracks = flatten_back_to_patch_tokens(label_tracks)
损失:
loss_traj = mse(pred_primary_tracks, label_primary_tracks)
loss_traj += mse(pred_wrist_tracks, label_wrist_tracks)
loss_traj *= 0.1
直觉:轨迹监督让模型知道“哪里会动、往哪动”。
9.4 未来 DINO 特征预测
DINO 特征是提前离线提取的视觉语义特征。模型不预测像素,而是预测未来图像经过 DINO 后的 patch feature。
dino_query_hidden = hidden[dino_query_positions]
dino_pred = dino_decoder(dino_query_hidden)
损失是余弦距离:
loss_dino = mean(1 - cosine_similarity(dino_pred, dino_label))
直觉:DINO 监督让模型学习物体和场景的高层语义,而不只是像素。
9.5 未来 SAM 特征预测
SAM 特征更偏物体区域、边界和可分割结构。
sam_query_hidden = hidden[sam_query_positions]
sam_pred = sam_decoder(sam_query_hidden)
损失也是余弦距离:
loss_sam = mean(1 - cosine_similarity(sam_pred, sam_label))
直觉:SAM 监督帮助模型理解“哪里是物体、哪里是边界、哪些区域可操作”。
9.6 多个未来任务怎么共享 query
默认情况下,不同任务可以有不同 query token:
RGB query
Depth query
DINO query
SAM query
Trajectory query
如果启用 share_query,多个任务共享一组观测 query,然后按 hidden 维度切分:
shared = hidden[obs_query_positions]
rgb_feature = shared[..., :D//4]
depth_feature = shared[..., D//4:D//2]
dino_feature = shared[..., D//2:3*D//4]
sam_feature = shared[..., 3*D//4:]
直觉:共享 query 可以减少 token 数,但要求同一组 query 同时承载多个预测任务的信息。
10. 动作怎么预测
DreamVLA 有两种动作头。
10.1 MLP 动作头
普通模式下,动作查询 token 的输出直接进 MLP。
action_hidden = hidden[action_query_positions]
action_hidden = action_decoder_mlp(action_hidden)
arm_action = tanh(arm_action_head(action_hidden)) # 6 维
gripper = sigmoid(gripper_action_head(action_hidden)) # 1 维
训练损失:
loss_arm = smooth_l1(arm_action, label_action[..., :6])
loss_gripper = bce(gripper, label_action[..., 6:])
推理时:
gripper_open_close = gripper > 0.5
gripper_env_value = (gripper_open_close - 0.5) * 2 # 转成 -1 / 1
action = concat(arm_action, gripper_env_value)
10.2 DiT 扩散动作头
如果开启 use_dit_head,动作不是 MLP 直接回归,而是扩散模型生成。
训练时:
condition = hidden[action_query_positions]
x0 = label_actions
noise = randn_like(x0)
t = random_diffusion_timestep()
xt = add_noise(x0, noise, t)
pred_noise = DiT(xt, t, condition)
loss_action = mse(pred_noise, noise)
推理时:
condition = hidden[action_query_positions]
x = randn([B, action_pred_steps, 7])
for step in ddim_steps:
pred_noise = DiT(x, step, condition)
x = denoise_one_step(x, pred_noise, step)
action = x
直觉:普通 MLP 更像“直接算答案”;扩散动作头更像“从随机动作开始,一步步修成合理动作”。当一个状态下有多种可行动作时,扩散头更有表达能力。
11. 总损失怎么组合
当前实现的总损失是:
loss =
loss_arm_action_ratio * loss_arm_action
+ loss_gripper_action_ratio * loss_gripper_action
+ 0.1 * loss_image
+ 0.001 * loss_depth
+ 0.1 * loss_trajectory
+ 0.01 * loss_dino_feat
+ 0.01 * loss_sam_feat
可以写成伪代码:
loss = 0
if loss_action:
loss += loss_arm_action_ratio * loss_arm_action
loss += loss_gripper_action_ratio * loss_gripper_action
if loss_image:
loss += 0.1 * loss_image
if loss_depth:
loss += 0.001 * loss_depth
if loss_trajectory:
loss += 0.1 * loss_trajectory
if loss_dino_feat:
loss += 0.01 * loss_dino_feat
if loss_sam_feat:
loss += 0.01 * loss_sam_feat
重点理解:
动作 loss 是主任务。
未来世界预测 loss 是辅助任务。
辅助任务的作用是塑造 Transformer 内部表征,让动作预测站在更好的世界理解上。
12. 完整训练流程
把前面所有模块合起来,训练流程是:
for batch in dataloader:
# 1. 取历史观测
primary = batch.image_primary[:, :S]
wrist = batch.image_wrist[:, :S]
state = preprocess_state(batch.state[:, :S])
text = repeat_language(batch.text_token, S)
# 2. 构造动作标签
label_action = []
for j in range(action_pred_steps):
label_action.append(batch.actions[:, j:S-atten_goal+j])
label_action = stack(label_action, dim="future_step")
# 3. 编码输入
text_tokens = encode_text(text)
state_tokens = encode_state(state)
visual_tokens = encode_and_resample_images(primary, wrist)
# 4. 拼真实输入 token
real_tokens = concat(text_tokens, state_tokens, visual_tokens)
# 5. 拼查询 token
query_tokens = build_query_tokens(enabled_tasks)
tokens = concat(real_tokens, query_tokens, dim="token")
# 6. 时序 Transformer
tokens = add_timestep_position_embedding(tokens)
hidden = causal_transformer(flatten_time(tokens), attention_mask)
hidden = unflatten_time(hidden)
# 7. 解码动作和未来世界
outputs = decode_all_enabled_heads(hidden)
# 8. 构造未来标签并计算 loss
loss = compute_action_loss(outputs, label_action)
loss += compute_enabled_future_losses(outputs, batch)
# 9. 更新参数
loss.backward()
clip_grad_norm(model, 0.1)
optimizer.step()
scheduler.step()
optimizer.zero_grad()
注意:CLIP 文本编码器和视觉骨干通常被冻结;训练的重点是投影层、Perceiver Resampler、时序 Transformer、查询 token、各个 decoder/action head。
13. 完整推理流程
推理时没有未来标签,也不需要计算未来预测 loss。
评估 wrapper 会维护历史队列:
history_primary = deque(maxlen=S)
history_wrist = deque(maxlen=S)
history_state = deque(maxlen=S)
history_text = deque(maxlen=S)
每个环境 step:
def policy_step(obs, instruction, timestep):
primary = preprocess_primary_image(obs)
wrist = preprocess_wrist_image(obs)
state = preprocess_robot_state(obs)
text = tokenize(instruction)
history_primary.append(primary)
history_wrist.append(wrist)
history_state.append(state)
if history_text is empty:
history_text.extend([text] * S)
# 前几步历史不足时,用最后一帧补齐
model_primary = pad_to_length(history_primary, S)
model_wrist = pad_to_length(history_wrist, S)
model_state = pad_to_length(history_state, S)
outputs = DreamVLA(
image_primary=model_primary,
image_wrist=model_wrist,
state=model_state,
text_token=history_text,
mode="test",
)
action_seq = outputs.action
action = select_current_action(action_seq, current_history_length)
action = convert_gripper_to_env_format(action)
return action
如果模型一次预测多步动作,LIBERO wrapper 还可以做 action ensembling:同一个当前时刻可能被多个过去时间步预测到,于是对这些候选动作做加权平均,让动作更平滑。
14. 离线预处理标签从哪里来
未来世界监督里,有些标签不是原始数据直接给的,需要预处理。
14.1 CoTracker 轨迹
CoTracker 在图像上放规则网格点,然后跟踪这些点未来的位置。
grid_points = make_grid_points(image_size=224, patch_size=8)
tracks, visibility = CoTracker(video, grid_points)
save_npz(tracks, visibility)
这些标签用于:
1. 训练 trajectory_pred
2. 可选地生成 flow_as_mask,让 RGB loss 关注动态区域
14.2 DINO 特征
DINO 脚本对每一帧图像提取 patch token:
dino_feature = DINOv2(image)["x_norm_patchtokens"]
save_pt(dino_feature)
训练时模型预测未来 DINO feature,用余弦损失和标签对齐。
14.3 SAM 特征
SAM 脚本用 SAM image encoder 提取空间特征,并池化到较粗网格:
sam_feature = SAM.image_encoder(image)
sam_feature = avg_pool(sam_feature, kernel_size=4)
save_pt(sam_feature)
训练时模型预测未来 SAM feature,用来学习物体区域和边界结构。
15. 预训练和微调
15.1 预训练学什么
预训练主要学习通用的视觉-语言-动作时序表征。
CALVIN 预训练脚本中,常见设置是:
sequence_length = 14
action_pred_steps = 3
future_steps = 3
obs_pred = true
loss_image = true
loss_action = true
atten_goal = 4
对应意思:
看 14 步历史。
每步预测 3 步动作。
预测 3 步之后的未来图像。
同时训练动作和未来 RGB。
最后 4 步因为未来标签不够,不参与部分监督。
15.2 微调学什么
微调阶段从预训练 checkpoint 加载权重,再加入更具体的任务监督。
例如 CALVIN 微调脚本展示了:
depth_pred + loss_depth
sam_feat_pred + loss_sam_feat
use_dit_head
flow_as_mask
load_track_labels
通俗理解:
预训练:先学通用的机器人世界理解。
微调:再适应具体 benchmark 的动作分布和更细的视觉/几何监督。
16. 关键配置开关
--obs_pred
开启未来 RGB 查询 token 和 RGB 解码器。
--depth_pred
开启未来深度查询 token 和深度解码器。
--trajectory_pred
开启未来轨迹查询 token 和轨迹解码器。
--dino_feat_pred
开启未来 DINO 特征预测。
--sam_feat_pred
开启未来 SAM 特征预测。
--loss_action
计算动作监督损失。
--loss_image
计算未来 RGB 损失。
--loss_depth
计算未来深度损失。
--loss_trajectory
计算未来轨迹损失。
--loss_dino_feat
计算未来 DINO 特征损失。
--loss_sam_feat
计算未来 SAM 特征损失。
--use_dit_head
用 DiT 扩散动作头,而不是 MLP 直接回归动作。
--use_fm
在动作头里使用 flow matching 风格训练。
--num_resampler_query
每个相机视角用多少个 Perceiver latent 汇聚图像 patch。
--num_obs_token_per_image
每个视角用于未来观测预测的查询 token 数。
--action_pred_steps
每个时间步预测多少步未来动作。
--future_steps
未来观测标签相对当前往后偏移多少步。
--pred_num
每个时间步预测多少个未来观测目标。
--atten_only_obs
限制动作 query 的可见范围,使其更依赖视觉和观测 query。
--atten_goal
控制未来目标跨度,也避免窗口尾部缺少未来标签。
--atten_goal_state
允许观测预测 query 读取未来目标状态 token。
--attn_robot_proprio_state
在 atten_only_obs 下允许动作 query 读取机器人本体状态。
--mask_l_obs_ratio
随机屏蔽一部分观测 query 到动作 query 的连接,作为正则化。
--flow_as_mask
用轨迹产生动态区域 mask,让 RGB loss 更关注会动的区域。
--share_query
多个未来预测任务共享一组 query,并按 hidden 维度切分给不同任务。
17. 用一个例子完整串起来
任务:
把桌上的杯子拿起来
训练时,模型看到最近几帧:
主视角:杯子、桌面、机械臂整体位置
腕部视角:夹爪附近细节
状态:机械臂末端位姿、夹爪开合
语言:拿起杯子
模型内部流程:
text = encode_text("拿起杯子")
state = encode_state(robot_state)
primary_patch = vision_encoder(primary_images) # 每帧约 196 patch
wrist_patch = vision_encoder(wrist_images)
primary_vis = perceiver_resampler(primary_patch) # 196 -> 16
wrist_vis = perceiver_resampler(wrist_patch) # 196 -> 16
tokens = concat(
text,
state,
primary_vis,
wrist_vis,
action_query,
future_rgb_query,
future_depth_query,
future_traj_query,
future_sam_query,
)
hidden = causal_transformer(tokens, mask)
action = action_decoder(hidden[action_query])
future_rgb = rgb_decoder(hidden[future_rgb_query])
future_depth = depth_decoder(hidden[future_depth_query])
future_traj = traj_decoder(hidden[future_traj_query])
future_sam = sam_decoder(hidden[future_sam_query])
训练目标会推动模型同时学会:
夹爪下一步往哪里移动
夹爪什么时候闭合
几步之后杯子在图像里会怎么变
夹爪和杯子的深度关系怎么变化
哪些图像点会随着杯子或夹爪移动
杯子的区域/边界特征怎么变化
如果这些都学得好,模型就不只是记住“这张图对应这个动作”,而是更接近理解:
杯子在哪
夹爪在哪
怎么靠近杯子
靠近后视觉和几何会怎么变化
什么时机该闭合夹爪
18. 最容易误解的点
18.1 推理时会不会真的生成未来图像
通常不会。未来图像、深度、轨迹、DINO、SAM 主要是训练监督。推理时通常只取动作输出。
18.2 图像压缩是不是随机采样
不是。它是 Perceiver Resampler 的 attention 汇聚。输出 token 是对全部 patch 的软加权结果,不是固定挑几个 patch。
18.3 查询 token 是不是输入数据
不是。查询 token 是模型参数,是可学习的问题槽位。它们进入 Transformer 后,从上下文里读信息,再被不同 decoder 解码。
18.4 为什么动作 query 能读未来观测 query
未来观测 query 不是未来真实图像,而是模型根据当前和过去自己推理出的内部表示。动作 query 读取它们,相当于基于“模型想象的未来”做决策。
18.5 未来预测越多越好吗
不一定。更多监督可能带来更强表征,但也增加数据预处理、显存和训练难度。实际要根据任务和资源开启。
19. 最后总结
DreamVLA 的方法可以压缩成五句话:
1. 把语言、双视角图像、机器人状态变成统一 hidden_dim 的 token。
2. 图像不是硬采样,而是用 Perceiver Resampler 从所有 patch 中软汇聚出少量视觉摘要 token。
3. 每个时间步加入动作 query 和未来世界 query,让同一个 Transformer 同时回答“怎么动”和“未来会怎样”。
4. attention mask 保证模型不能偷看未来真实观测,同时让动作 query 可以利用模型内部形成的未来世界表示。
5. 训练时动作 loss 是主任务,RGB/深度/轨迹/DINO/SAM 预测是辅助世界模型监督;推理时只执行动作。
最关键的理解是:
DreamVLA 不是单纯的图文到动作回归模型。
它是用未来世界预测来训练机器人策略的 VLA 模型。
本文来自博客园,作者:S-X-Q,转载请注明原文链接:https://www.cnblogs.com/sxq-blog/p/21468069

浙公网安备 33010602011771号