11 KiB
MultiPLY 训练过程讲解
本文基于 MultiPLY 官方仓库(
model_release/下的fsdp_train.py、dataset.py、llava/llava/model/llava_arch.py、builder.py)的真实代码,逐步讲清它"怎么训"。 配套阅读:同目录《MultiPLY技术讲解.md》(含完整架构与代码级剖析)。 一句话定位:MultiPLY = 在 LLaVA-1.5-7B 上、用预提取的多感官特征做指令微调,并叠加一个物体选择(grounding)注意力损失形成双目标训练。
0. 全局概览
训练可以拆成"一个离线阶段 + 一个在线阶段":
阶段零(离线,仿真器里做) 阶段一(在线,fsdp_train.py 做)
仿真采集 → 各模态特征预提取 .pt 加载 LLaVA-1.5-7B → 注入多感官 token
视觉/点云 (CLIP) → 删除视觉塔、全参数可训
触觉 (DiffTactile) → FSDP 分布式
撞击声 (ObjectFolder/CLAP) → 双损失:语言建模 + 物体选择
逐物体目标标签 prediction → AdamW / fp16 / 每 epoch 存档
关键点先记住三条:
- 特征是离线算好的:训练时不跑视觉塔,直接从
.pt加载各模态特征(省显存/算力)。 - 全参数微调:不是"冻结编码器只训投影头",而是整个 LLM + 各投影层一起训。
- 双目标损失:语言建模损失(
loss1)+ 物体选择损失(loss2),两者相加。
1. 阶段零:离线特征预提取(数据准备)
训练读的不是原始图像/声音,而是预先算好的特征张量。每条样本对应一个 JSON 字典(来自 all_questions.json),可能含这些字段:
| 字段 | 含义 | 训练时如何用 |
|---|---|---|
question / answer |
指令与目标回答 | 拼成文本序列 |
scene |
场景 id | 去 dataset/feature_dict/<scene>/<obj_id>.pt 取每个物体的 1024 维特征(CLIP 风格,多视角平均;默认场景特征形状 (256,1024),即最多 256 个物体) |
visual |
观察到的物体 | 取该物体的视觉/点云特征 |
tactile_reading |
触觉读数路径 | 取 data5/<...>/marker4.pt 并按 marker 维取均值(对应 DiffTactile) |
impact_sound |
撞击声 | 取 impact_sound_*/0.pt;或场景环境音 audioset embedding |
temperature |
温度 | 代码中部分实现(投影器在公开版未完整释出) |
prediction |
逐物体的目标性二值标签 | 用于 loss2 物体选择损失 |
这些特征由
simulator/里的网格扫描 + 各模态仿真器离线生成(详见《MultiPLY技术讲解》第 6.5 节)。训练阶段只管"加载 + 喂进 LLM"。
2. 训练数据怎么组织(dataset.py)
2.1 文本模板与"占位符复制"技巧
MultisensoryDataset.__getitem__ 把样本拼成固定模板:
Question: {question} Answer: {answer} {eos}
其中文本里嵌有特殊占位 token(<scene> <visual> <tactile> <sound> 等)。关键技巧是:每个占位 token 会按其特征条数被复制——
text = text.replace(self.scene_token, self.scene_token * len(scene_feature))
.replace(self.tactile_token, self.tactile_token * len(tactile_feature))
.replace(self.sound_token, self.sound_token * len(sound_feature))
.replace(self.visual_token, self.visual_token * len(visual_feature))
这样"占位 token 的数量 == 要插入的特征向量数量",保证后面逐位替换严格对齐(见 §3)。
2.2 tokenize 与插入位置
文本经 tokenizer(max_length=2048,padding 到最大长度)后:
- 记录每类占位 token 在序列里的下标
*_insert_loc(scene_insert_loc/visual_insert_loc/tactile_insert_loc/sound_insert_loc); - 各模态特征按插入位置数量截断对齐;
prediction转成 0/1 浮点张量(>0的置 1),作为物体选择监督。
2.3 collate(成 batch)
collate_wrapper 把一个 batch 拼起来:场景特征 zero-pad 到 batch 内最大物体数,记录 max_scene_length,并把各 insert_loc 改写成 [batch_idx, 位置] 的形式,便于跨 batch 定位。batch_size=2。
3. 模型怎么"吃"多感官特征(llava_arch.py)
前向第一步是 prepare_inputs_labels_for_multimodal,核心是把占位 token 的词向量替换成感官特征向量:
- 先正常
embed_tokens(input_ids)得到文本 embedding; - 把各模态特征过投影层对齐到 LLM 隐藏维:
scene/visual→ 复用 LLaVA 的mm_projectortactile→ 新增的tactile_projector(2 层 GELU MLP)sound→ 新增的sound_projector(2 层 GELU MLP)
_insert_feature按*_insert_loc逐位把占位 token 的 embedding 覆盖成对应特征,同时把这些位置的 label 设为-100(不计入语言建模损失);- 之后就是标准 LLaMA 前向 +
lm_head。
这与 LLaVA 处理图像 patch 的机制同源——多感官信息以"替换占位词向量"的方式进入 LLM。
4. 训练主循环(fsdp_train.py)
4.1 模型加载与可训设置(main)
model_path = "liuhaotian/llava-v1.5-7b"
tokenizer, model, image_processor, context_len = load_pretrained_model(
model_path, None, model_name, device_map=None, add_multisensory_token=True)
add_multisensory_token=True:在builder.py里把<scene>/<visual>/<tactile>/<sound>/<observe>/<touch>/<hit>等加进 tokenizer 并resize_token_embeddings(扩词表)。model.requires_grad_(True):全参数可训(整 LLM + 投影层)。del model.model.vision_tower:删掉 LLaVA 自带的在线 CLIP 视觉塔(特征已离线预提取,省显存)。
4.2 分布式与优化器
- FSDP(Fully Sharded Data Parallel):
auto_wrap_policy按LlamaDecoderLayer自动 wrap;MixedPrecision全 fp16(param/reduce/buffer);ShardingStrategy.SHARD_GRAD_OP(分片梯度与优化器状态)。
- 优化器:AdamW,lr = 1e-6。
- 数据:
DistributedSampler+DataLoader(batch_size=2, num_workers=4)。 - 训练
num_epochs轮,每轮train_one_epoch后save_checkpoint(FSDPFULL_STATE_DICT,rank0 存checkpoint_{epoch}.pt)。
4.3 单步训练(train_one_epoch)—— 这是核心
labels = input_ids.clone()
answer_indices = torch.where(labels == 22550)[1] # 22550 = "Answer" 分隔符
for j, answer_idx in enumerate(answer_indices):
labels[j, :answer_idx+2] = -100 # 只对"答案"部分算语言损失
labels[labels == tokenizer.pad_token_id] = -100 # padding 不算损失
with torch.autocast(device_type="cuda"):
outputs = llava_model(input_ids, attention_mask, labels=labels,
feature_dict=feature_dict, output_hidden_states=True)
# —— loss2:物体选择(grounding)——
hidden_state = outputs['hidden_states'][-1][:, -1, :].unsqueeze(1) # 最后层、最后位置
scene_feature = llava_model.model.mm_projector(sample.scene_feature) # 投影后的物体特征
attention = torch.einsum("abf,acf->abc", scene_feature, hidden_state).squeeze(-1)
# 加权 BCE:正样本权重1、负样本0.2、忽略0;pos_weight=5
loss2 = F.binary_cross_entropy_with_logits(attention, prediction,
weight=weights, pos_weight=pos_weight)
loss = outputs.loss + loss2 # loss1 + loss2
loss.backward(); optimizer.step()
要点拆解:
- 标签掩码:用 token id
22550(LLaMA 分词下的 "Answer")定位答案起点,把它之前的 token 全置-100——即只在"答案"部分计算语言建模损失(问题、占位符、padding 都不算)。 - loss1(语言建模):标准下一 token 交叉熵(在
llava_llama.py的 forward 里算好,移位后CrossEntropyLoss)。 - loss2(物体选择 / grounding):用 LLM 最后一层、最后一个位置的 hidden state 与投影后的场景物体特征做点积(
einsum),得到每个物体的"被选中分数",再对二值标签prediction做加权 BCE:- 权重
weights:目标物体(pred==1)权重 1,非目标(pred==0)权重 0.2,忽略(pred==-1)权重 0;超过max_scene_length的 padding 物体权重 0; - 额外
pos_weight=5缓解正负样本不平衡。
- 权重
- 总损失:
loss = loss1 + loss2,一起反传。
这个
loss2把"物体检索/指代消解"显式变成可监督的注意力对齐任务——正是 MultiPLY 在物体检索基准上大幅领先(56.7% vs 次优 48.9%)的工程原因之一。
5. 推理 / 评测(eval)
评测对短答案 QA 做贪心生成并与真值精确匹配:
output_ids = model.generate(input_ids, feature_dict=feature_dict,
do_sample=False, max_new_tokens=10)
# 解码后与 ground-truth 字符串 .lower().strip() 精确匹配,统计准确率
推理时同样先把多感官特征注入占位 token,再自回归生成;在完整的具身设定下,模型还会生成动作 token(<observe>/<touch>/<hit> 等)驱动智能体交互、把新观测作为状态 token 回填——形成"动作→反馈→再生成"的闭环(详见架构文档)。
6. 关键超参数 / 配置一览
| 项 | 取值 | 出处 |
|---|---|---|
| 主干 | liuhaotian/llava-v1.5-7b(Vicuna-7B / LLaMA-2 系) |
fsdp_train.py |
| 可训范围 | 全参数(requires_grad_(True)),删除视觉塔 |
fsdp_train.py |
| 优化器 / 学习率 | AdamW / 1e-6 | fsdp_train.py |
| batch size | 2 | fsdp_train.py |
| 序列长度 | 2048 | dataset.py |
| 分布式 | FSDP,SHARD_GRAD_OP,fp16 混合精度 |
fsdp_train.py |
| 损失 | loss1(LM CE) + loss2(物体选择加权 BCE, pos_weight=5) |
train_one_epoch |
| 答案分隔符 | token id 22550("Answer") |
train_one_epoch |
| 特征维度 | 各模态 1024 维 → 投影到 LLM 隐藏维 | dataset.py / llava_arch.py |
| 投影层 | scene/visual→mm_projector;tactile/sound→各自 2 层 GELU MLP |
llava_arch.py |
| 存档 | 每 epoch,FSDP FULL_STATE_DICT,rank0 保存 | save_checkpoint |
7. 与论文叙述的差异 / 注意点
- "冻结"误区:直觉上 LLaVA 式训练常冻结编码器,但 MultiPLY released 脚本是全参数微调(视觉塔已删、特征离线)。
- 双损失少被强调:论文正文偏重"动作/状态 token 闭环",但代码里
loss2(物体选择)对检索类任务很关键。 - 温度模态未完整:
<temperature>在词表/数据里出现,但llava_arch.py没有对应投影器、公开训练循环也未真正喂入——属于部分释出。 - 研究级代码:
all_questions.json、预提取特征、requirements等需自行准备;__getitem__用try/except容错,数据完整性需使用者保证。
8. 小结:MultiPLY 训练的三句话
- 离线把视/触/听等多感官信号各自编码成 1024 维特征
.pt; - 在线在 LLaVA-1.5-7B 上做全参数指令微调,用"占位 token 替换成特征向量"的方式把多感官喂进 LLM,只在答案上算语言损失;
- 叠加一个物体选择注意力损失(hidden state × 物体特征 的加权 BCE),让模型学会"挑出目标物体",与语言损失一起用 FSDP/AdamW/fp16 训练。