206 lines
11 KiB
Markdown
206 lines
11 KiB
Markdown
# 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 存档
|
||
```
|
||
|
||
关键点先记住三条:
|
||
|
||
1. **特征是离线算好的**:训练时不跑视觉塔,直接从 `.pt` 加载各模态特征(省显存/算力)。
|
||
2. **全参数微调**:不是"冻结编码器只训投影头",而是整个 LLM + 各投影层一起训。
|
||
3. **双目标损失**:语言建模损失(`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 会按其特征条数被复制**——
|
||
|
||
```python
|
||
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 的词向量替换成感官特征向量**:
|
||
|
||
1. 先正常 `embed_tokens(input_ids)` 得到文本 embedding;
|
||
2. 把各模态特征过投影层对齐到 LLM 隐藏维:
|
||
- `scene` / `visual` → 复用 LLaVA 的 **`mm_projector`**
|
||
- `tactile` → 新增的 **`tactile_projector`**(2 层 GELU MLP)
|
||
- `sound` → 新增的 **`sound_projector`**(2 层 GELU MLP)
|
||
3. `_insert_feature` 按 `*_insert_loc` **逐位把占位 token 的 embedding 覆盖成对应特征**,同时把这些位置的 label 设为 `-100`(不计入语言建模损失);
|
||
4. 之后就是标准 LLaMA 前向 + `lm_head`。
|
||
|
||
> 这与 LLaVA 处理图像 patch 的机制同源——多感官信息以"替换占位词向量"的方式进入 LLM。
|
||
|
||
---
|
||
|
||
## 4. 训练主循环(`fsdp_train.py`)
|
||
|
||
### 4.1 模型加载与可训设置(`main`)
|
||
|
||
```python
|
||
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`(FSDP `FULL_STATE_DICT`,rank0 存 `checkpoint_{epoch}.pt`)。
|
||
|
||
### 4.3 单步训练(`train_one_epoch`)—— 这是核心
|
||
|
||
```python
|
||
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 做贪心生成并与真值精确匹配:
|
||
|
||
```python
|
||
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 训练的三句话
|
||
|
||
1. **离线**把视/触/听等多感官信号各自编码成 1024 维特征 `.pt`;
|
||
2. **在线**在 LLaVA-1.5-7B 上做全参数指令微调,用"占位 token 替换成特征向量"的方式把多感官喂进 LLM,只在答案上算语言损失;
|
||
3. 叠加一个**物体选择注意力损失**(hidden state × 物体特征 的加权 BCE),让模型学会"挑出目标物体",与语言损失一起用 FSDP/AdamW/fp16 训练。
|