omni / docs /interview /training-systems.md
chenbhao's picture
refactor(docs): restructure docs with modular organization
f664f3f
|
Raw
History Blame Contribute Delete
13.4 kB
# 面试:训练系统深度
> 本仓库 `src/trainers/` 的训练逻辑,覆盖 pretrain、SFT、DPO、PPO、GRPO、蒸馏
## 0. 训练流程概览
```
数据集 (dataset/)
SkipBatchSampler → DataLoader → 批次拼接
Trainer (trainers/)
├── 模型初始化(支持 checkpoint 恢复)
├── 混合精度训练(bfloat16/float16)
├── 梯度累积 + 梯度裁剪
├── 学习率调度(cosine)
└── 定期保存 checkpoint
```
---
## Q1. 损失里 label 为什么要平移?
### 标准 Next-Token Prediction
```python
# src/models/lm/model.py:69-70
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
loss = F.cross_entropy(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))
```
### 为什么平移?
- `logits[:, :-1]` 预测的是 `labels[:, 1:]` 的内容
- 标准 next-token 预测:给定前缀,预测下一个 token
- 不平移的话,模型会学到"预测自己",没有意义
> 面试点:如果 label 不平移会怎样?→ 模型会学到恒等映射,训练 loss 很低但生成质量很差
---
## Q2. SFT 为什么只标 assistant 段?
### Loss Mask 机制
```python
# src/dataset/sft.py
def generate_labels(self, input_ids):
labels = input_ids.clone()
# prompt/system 段 label 置 -100
labels[:prompt_end_pos] = -100
# 只对 assistant 回复计算损失
return labels
```
### 为什么只标 assistant?
1. **避免学用户说话**:模型应该学习生成回复,不拟合用户输入
2. **提高效率**:只计算有意义部分的损失
3. **符合实际使用**:推理时只生成 assistant 回复
> 面试点:如果标全部 token 会怎样?→ 模型会学用户输入的模式,生成时可能重复用户说话风格
---
## Q3. 配置系统优先级?
### 三级优先级
```
CLI 参数 > YAML 默认值 > 代码默认
```
### 实现(`src/utils/training.py:25-38`)
```python
def apply_config(parser, default_config=None):
# 1. 先加载 YAML 默认值
if default_config:
with open(default_config) as f:
config = yaml.safe_load(f)
# 2. 扁平化注入 argparse 默认值
for key, value in config.items():
if key in [a.dest for a in parser._actions]:
parser.set_defaults(**{key: value})
# 3. CLI 参数覆盖
return parser.parse_args()
```
### 为什么这样设计?
1. **灵活性**:可以用 YAML 配置常用参数,CLI 临时覆盖
2. **可复现**:YAML 文件可以版本控制
3. **向后兼容**:代码默认值保证基本功能
---
## Q4. 续训怎么保证精确?
### SkipBatchSampler(`src/utils/training.py:177-200`)
```python
class SkipBatchSampler:
def __init__(self, dataset, batch_size, step):
self.step = step # 已完成的 step 数
def __iter__(self):
# 跳过前 step 个 batch
indices = list(range(len(self.dataset)))
indices = indices[self.step * self.batch_size:]
...
```
### Checkpoint 恢复
```python
# src/utils/training.py:105-158
def lm_checkpoint(model, optimizer, scheduler, step, path):
# 原子保存
tmp_path = path + ".tmp"
torch.save({
'step': step,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'scheduler_state_dict': scheduler.state_dict(),
}, tmp_path)
os.replace(tmp_path, path)
```
### 世界规模自适应(`src/utils/training.py:153-156`)
```python
# GPU 数量变化时自动调整 step
if world_size != saved_world_size:
step = int(step * saved_world_size / world_size)
```
> 面试点:为什么需要世界规模自适应?→ 多机训练时 GPU 数量可能变化,需要按比例调整 step 保证训练进度一致
---
## Q5. DPO Loss 实现
### 标准 DPO Loss(`src/trainers/lm/dpo.py:32-48`)
```python
def dpo_loss(policy_logratios, reference_logratios, beta=0.15):
# DPO Loss = -logsigmoid(β * (π_logratios - ref_logratios))
loss = -F.logsigmoid(beta * (policy_logratios - reference_logratios))
return loss.mean()
```
### 关键参数
- `beta=0.15`:控制策略偏离参考模型的程度
- `lr=4e-8`:极小学习率防止遗忘
### 参考模型(`src/trainers/lm/dpo.py:187-189`)
```python
# 冻结的 ref_model
self.ref_model = init_model(args, config)
for param in self.ref_model.parameters():
param.requires_grad = False
```
> 面试点:为什么 DPO 学习率这么小?→ DPO 直接优化策略,太大的学习率会导致策略偏离参考模型太远,生成质量下降
---
## Q6. GRPO 实现(`src/trainers/lm/grpo.py:119-142`)
### 核心思想
Group Relative Policy Optimization:每个 prompt 生成多个候选,组内标准化优势。
### 实现细节
```python
def grpo_loss(self, logprobs, old_logprobs, advantages, clip_epsilon=0.2):
# 每个 prompt 生成 num_generations=6 个候选
# 优势估计:组内标准化
advantages = (reward - mean) / (std + 1e-4)
# PPO Clip
ratio = torch.exp(logprobs - old_logprobs)
clipped = torch.clamp(ratio, 1 - clip_epsilon, 1 + clip_epsilon)
loss = -torch.min(ratio * advantages, clipped * advantages)
# KL 惩罚
kl = logprobs - old_logprobs
kl_penalty = torch.exp(kl) - kl - 1
return loss.mean() + self.kl_coef * kl_penalty.mean()
```
### 两种 Loss 模式
1. **"cispo"**:高端 clamped ratio(`epsilon_high=5.0`
2. **"grpo"**:标准 PPO clip(`epsilon=0.2`
### 奖励函数(`src/trainers/lm/grpo.py:35-66`)
```python
def reward_fn(self, text):
reward = 0
# 长度奖励:20-800 字符 +0.5
if 20 < len(text) < 800:
reward += 0.5
# thinking 奖励:20-300 字符 +1.0
if '<think>' in text and 20 < think_len < 300:
reward += 1.0
# thinking 次数奖励:恰好 1 次 +0.25
if text.count('<think>') == 1:
reward += 0.25
# 重复惩罚:3-gram 重复度
reward -= repetition_penalty(text)
# Reward Model 分数
rm_score = self.reward_model(text)
reward += rm_score
return reward
```
> 面试点:GRPO 和 PPO 的区别?→ GRPO 不需要 critic model,直接用组内标准化估计优势,更简单高效
---
## Q7. PPO 实现(`src/trainers/lm/ppo.py`)
### CriticModel(`src/trainers/lm/ppo.py:35-47`)
```python
class CriticModel(LMForCausalLM):
def __init__(self, config):
super().__init__(config)
# lm_head 替换为 value_head (hidden_size -> 1)
self.value_head = nn.Linear(config.hidden_size, 1, bias=False)
del self.lm_head
```
### GAE(`src/trainers/lm/ppo.py:138-145`)
```python
def compute_gae(rewards, values, gamma=1.0, lam=0.95):
advantages = []
gae = 0
for t in reversed(range(len(rewards))):
delta = rewards[t] + gamma * values[t + 1] - values[t]
gae = delta + gamma * lam * gae
advantages.insert(0, gae)
return advantages
```
### 早停机制(`src/trainers/lm/ppo.py:181-188`)
```python
def ppo_update(self, ...):
approx_kl = (logprobs - old_logprobs).mean()
if approx_kl > 0.25:
# 早停,但保持 DDP 通信闭环
loss = loss * 0.0
return loss
```
> 面试点:为什么早停时要 `loss * 0.0`?→ DDP 要求所有 rank 都参与前向/反向传播,`loss * 0.0` 保持计算图连通,避免死锁
---
## Q8. 蒸馏实现(`src/trainers/lm/distillation.py`)
### KL 散度蒸馏(`src/trainers/lm/distillation.py:23-34`)
```python
def kl_divergence(teacher_logits, student_logits, temperature=1.5):
# KL(teacher || student) * temperature^2
teacher_probs = F.softmax(teacher_logits / temperature, dim=-1)
student_log_probs = F.log_softmax(student_logits / temperature, dim=-1)
kl = F.kl_div(student_log_probs, teacher_probs, reduction='batchmean')
return kl * temperature ** 2
```
### 混合损失(`src/trainers/lm/distillation.py:91`)
```python
loss = alpha * ce_loss + (1 - alpha) * kl_loss
# 默认 alpha=0.5, temperature=1.5
```
### Teacher/Student 独立配置
```python
# Teacher 和 Student 可以有不同的配置
teacher_config = LMConfig(hidden_size=768, num_hidden_layers=8)
student_config = LMConfig(hidden_size=384, num_hidden_layers=4)
```
> 面试点:为什么蒸馏要乘以 temperature²?→ 保持梯度尺度一致,避免温度变化影响学习率
---
## Q9. Rollout Engine 设计(`src/trainers/lm/rollout_engine.py`)
### 策略模式
```python
class RolloutEngine(ABC):
@abstractmethod
def generate(self, prompts, **kwargs):
pass
class TorchRolloutEngine(RolloutEngine):
def generate(self, prompts, **kwargs):
# 原生 PyTorch 推理
...
class SGLangRolloutEngine(RolloutEngine):
def generate(self, prompts, **kwargs):
# 通过 HTTP API 调用 SGLang 服务
...
```
### 权重同步(`src/trainers/lm/rollout_engine.py:165-188`)
```python
def sync_weights(self, model):
# 仅 rank 0 执行保存
if self.rank == 0:
# 保存到磁盘
torch.save(model.state_dict(), 'tmp_model.pt')
# HTTP 请求 SGLang 热加载
requests.post('http://localhost:8000/load_model', ...)
# 广播成功标志
dist.broadcast(success_flag, src=0)
```
> 面试点:为什么要权重同步?→ RL 训练需要最新的策略生成 rollout,必须确保推理服务使用最新权重
---
## Q10. 混合精度训练
### 本仓库的混合精度策略
```python
# src/trainers/lm/pretrain.py
scaler = GradScaler(enabled=(args.dtype == 'float16'))
with autocast(device_type='cuda', dtype=dtype):
logits, _, aux_loss = model(input_ids, labels=labels)
loss = criterion(logits, labels) + aux_loss
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
```
### 为什么用 bfloat16 而不是 float16?
- **bfloat16**:8 位指数,范围大,不易溢出
- **float16**:5 位指数,范围小,容易溢出
> 面试点:GradScaler 的作用?→ 防止 float16 梯度下溢,动态调整 loss 缩放因子
---
## Q11. 梯度累积
### 为什么需要梯度累积?
显存有限时,可以用小 batch 大 accumulation 模拟大 batch。
### 实现(`src/trainers/lm/pretrain.py`)
```python
accumulation_steps = 8 # 默认
for i, batch in enumerate(dataloader):
with autocast(device_type='cuda', dtype=dtype):
logits, _, aux_loss = model(batch)
loss = criterion(logits, labels) + aux_loss
loss = loss / accumulation_steps # 梯度累积需要除以步数
scaler.scale(loss).backward()
if (i + 1) % accumulation_steps == 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
```
> 面试点:为什么 loss 要除以 accumulation_steps?→ 保持梯度尺度一致,避免累积步数影响学习率
---
## Q12. 学习率调度(`src/utils/training.py:82-83`)
### Cosine 调度
```python
def get_lr(step, lr, total):
# 最低衰减到 10%
return lr * (0.1 + 0.45 * (1 + math.cos(math.pi * step / total)))
```
### 为什么用 Cosine?
1. **平滑衰减**:避免学习率骤降
2. **实验效果好**:比 step decay 和 linear decay 更稳定
3. **标准做法**:几乎所有大模型都用 cosine
---
## Q13. Checkpoint 原子写入(`src/utils/checkpoint.py:7-10`)
### 问题
训练过程中保存 checkpoint 时,如果中途崩溃,可能导致 checkpoint 损坏。
### 解决方案
```python
def save_checkpoint(model, optimizer, scheduler, path):
# 先保存到临时文件
tmp_path = path + ".tmp"
torch.save({...}, tmp_path)
# 原子替换
os.replace(tmp_path, path)
```
### 为什么用 os.replace?
`os.replace` 是原子操作,要么完成要么不发生,不会出现部分写入的情况。
---
## Q14. DDP 初始化(`src/utils/distributed.py`)
### 标准 DDP 初始化
```python
def init_distributed():
dist.init_process_group(backend='nccl')
local_rank = int(os.environ['LOCAL_RANK'])
torch.cuda.set_device(local_rank)
return local_rank
```
### DDP 死锁防护
```python
# PPO 早停时保持通信闭环
if early_stop:
loss = loss * 0.0 # 保持计算图连通
```
> 面试点:为什么 DDP 会死锁?→ 如果某个 rank 不参与前向/反向传播,其他 rank 会等待它,导致死锁
---
## Q15. 世界规模自适应(`src/utils/training.py:153-156`)
### 问题
多机训练时 GPU 数量可能变化,需要按比例调整 step。
### 实现
```python
def lm_checkpoint(model, optimizer, scheduler, step, path):
# 保存世界规模
torch.save({
'step': step,
'world_size': dist.get_world_size(),
}, path)
# 恢复时检查世界规模
saved_world_size = checkpoint['world_size']
if world_size != saved_world_size:
step = int(step * saved_world_size / world_size)
```
> 面试点:为什么需要这个?→ 保证训练进度一致,避免某些 rank 重复训练或跳过训练