File size: 7,048 Bytes
f664f3f
 
d58698c
f664f3f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2460459
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f664f3f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2460459
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
# trainers/trainers.md · 训练脚本

`trainers/` 按模态分子包,每个脚本暴露 `main(default_config=None)`(通过 `python -m trainers.<mod>` 调用)。

## 文本(lm/)

| 脚本 | 任务 | 关键损失/算法 |
| --- | --- | --- |
| `pretrain.py` | 预训练 | 下一 token CE |
| `full_sft.py` | 全量 SFT | CE(loss mask 仅 assistant) |
| `lora.py` | LoRA 微调 | 低秩适配,仅训 A/B |
| `dpo.py` | DPO | 偏好对齐(参考比损失) |
| `distillation.py` | 知识蒸馏 | 师生 KL |
| `ppo.py` | PPO | Actor-Critic + 奖励 |
| `grpo.py` | GRPO | 分组相对策略优化 |
| `agent.py` | Agent RL | 工具调用强化学习 |
| `rollout_engine.py` | — | torch / sglang 推理引擎(被 ppo/grpo/agent 复用) |
| `train_tokenizer.py` | — | tokenizer 训练(学习用),结果保存到 `checkpoint/tokenizer/` |

## 视觉(vlm/)

- `pretrain.py`:视觉预训练
- `full_sft.py`:视觉 SFT(含 `vlm_collate_fn`## 全模态(vam/)

- `full_sft.py`:全模态 SFT(文本 + 视觉 + 音频,双 head 损失)

## VAM SFT 详解

### 模型初始化(`init_omni_model`)

```python
model, tokenizer = init_omni_model(omni_config,
    from_weight='omni-v',
    tokenizer_path='checkpoint/omni/native_hf',
    audio_encoder_path='checkpoint/sensevoice',
    vision_model_path='checkpoint/siglip',
    model_dir='checkpoint/omni-v')  # 权重源目录
```

- `from_weight` + `model_dir` 决定加载哪个 checkpoint
- 优先加载 `{model_dir}/sft_omni_{hidden_size}.pth`
- 若加载权重不含 talker 参数,自动从 thinker 后几层复制初始化
- 编码器(SenseVoice / SigLIP)从独立路径初始化,不在 checkpoint 中保存

### 训练模式(mode)

| mode | trainable params | 用途 |
| --- | --- | --- |
| `all` | 全部(113M) | 全参数 SFT |
| `audio_proj` | 仅 audio_proj(1.0M) | 音频特征对齐 |
| `vision_proj` | 仅 vision_proj(1.2M) | 视觉特征对齐 |

```python
if args.mode == 'audio_proj':
    for p in model.parameters(): p.requires_grad = False
    for p in model.audio_proj.parameters(): p.requires_grad = True
```

optimizer 使用 `filter(requires_grad)` 避免为冻结参数维护动量:

```python
optimizer = optim.AdamW(
    filter(lambda p: p.requires_grad, model.parameters()),
    lr=args.learning_rate
)
```

### 损失计算(双 head)

```python
# 文本损失
text_loss = CE(logits, labels, ignore_index=-100)

# 音频损失(对 8 层 Mimi code 逐层 CE,stop token 10× 加权)
audio_loss = 0
for i, al in enumerate(res.audio_logits):
    layer_loss = CE(al.view(-1, al.size(-1)), audio_labels[:, i, :].reshape(-1))
    stop_mask = (targets == audio_stop_token).float()  # 2050
    weighted = layer_loss * valid_mask * (1 + stop_mask * 9)
    audio_loss += weighted.sum() / valid_mask.sum()
audio_loss = audio_loss / 8

# 总损失(accumulation_steps 用于梯度累积)
loss = (text_loss + audio_loss + res.aux_loss) / args.accumulation_steps
```

### DataLoader 与 collate_fn

VAM 的 `omni_collate_fn` 处理**变长**的音频和视觉输入:

```python
def omni_collate_fn(batch):
    # batch 包含:input_ids, labels, audio_labels, audio_inputs, audio_lens, pixel_values, spk_emb

    # 1. 文本:直接 stack(已 padding)
    input_ids = torch.stack(input_ids)

    # 2. 音频:padding 到 batch 内最大长度
    valid_audios = [a for a in audio_inputs if a is not None]
    max_t = max(a.size(1) for a in valid_audios)
    padded = [pad(a, max_t) for a in valid_audios]
    audio_inputs = torch.cat(padded, dim=0)

    # 3. 视觉:SigLIP 返回 dict(pixel_values + attention_mask)
    valid_images = [p for p in pixel_values if p is not None]
    pixel_values = {k: torch.cat([d[k] for d in valid_images], dim=0) for k in keys}
```

### 音频处理流程

A2A 数据包含 `question_audios`(二进制音频文件),在 `__getitem__` 中按需解码:

```python
def load_audio_inputs(self, audio_bytes):
    wav, sr = sf.read(io.BytesIO(audio_bytes))
    if wav.ndim > 1: wav = wav.mean(axis=1)
    if sr != 16000:
        wav_t = torch.from_numpy(wav).unsqueeze(0)
        wav_t = AF.resample(wav_t, sr, 16000)  # torchaudio
        wav = wav_t.squeeze(0).numpy()
    inputs = self.audio_processor(wav, sampling_rate=16000, ...)
    return inputs.input_features, valid_len
```

- 使用 `soundfile` 解码音频 bytes
- `torchaudio.functional.resample` 重采样到 16kHz(SenseVoice 要求)
- `SenseVoiceAudioProcessor` 提取 fbank 特征

### 3 阶段训练配置

运行示例:

```bash
# Stage 1: T2A mode=all(从预训练初始化)
python -m trainers.vam.full_sft --config configs/vam/vam_t2a_all_mini_omni-v.yaml

# Stage 2: A2A audio_proj(从 Stage 1 初始化)
python -m trainers.vam.full_sft --config configs/vam/vam_a2a_audio_proj_mini.yaml

# Stage 3: A2A mode=all(从 Stage 2 初始化)
python -m trainers.vam.full_sft --config configs/vam/vam_a2a_all_mini.yaml
```

各配置文件的 `model_dir``from_weight` 构成训练链:
- Stage 1 → 从 `omni-v.pth` 初始化 → 输出到 `vam_t2a_all_mini_omni-v/`
- Stage 2 → 从 Stage 1 输出初始化 → 输出到 `vam_a2a_audio_proj_mini/`
- Stage 3 → 从 Stage 2 输出初始化 → 输出到 `vam_a2a_all_mini/`

### 检查点保存

```python
# 推理权重(仅 LLM 部分,fp16)
torch.save({k: v.half().cpu() for k, v in clean_state_dict.items()}, ckp)

# 续训检查点(含 optimizer + scaler 状态)
omni_checkpoint(omni_config, weight=..., model=..., optimizer=..., ...)
```

推理权重过滤掉 `audio_encoder.` 前缀(编码器需在各训练脚本中单独加载),确保 checkpoint 格式兼容。

## 通用训练循环(以 full_sft 为例)

```python
for epoch in range(epochs):
    loader = DataLoader(ds, batch_sampler=SkipBatchSampler(...))
    for step, (input_ids, labels) in enumerate(loader):
        loss = model(input_ids, labels=labels).loss + res.aux_loss
        loss = loss / accumulation_steps
        scaler.scale(loss).backward()
        if step % accumulation_steps == 0:
            clip_grad_norm_; scaler.step(optimizer); zero_grad()
        # 定期保存权重到 save_dir + 保存 optimizer/ckpt 到 checkpoint/
```

支持:分布式(`init_distributed_mode` + `DistributedDataParallel`)、混合精度(`autocast` + `GradScaler`)、梯度累积、断点续训(`from_resume`)、可选 wandb/swanlab。

## 要点(面试)

- **`SkipBatchSampler`**:分布式下跳过已训 step,配合 `from_resume` 实现精确续训。
- **`aux_loss`**:MoE 路由均衡损失,只在 `use_moe` 时非 0,需显式加到总损失。
- **RL trainer 复用 `rollout_engine`**:生成样本与训练解耦,可换 torch / sglang 后端。
- 保存分两份:`save_dir`(最终权重 `.pth`)+ `checkpoint/`(optimizer/scheduler 状态用于续训)。
- **VAM 特殊点**:双 head 损失、变长 collate_fn、audio_proj 模式只训 1% 参数、filter(requires_grad) 优化器。