Llama 模型训练配方
Llama 系列模型是开源大模型的事实标准。本文总结 Llama 模型从预训练到微调的完整训练配方,包括超参数选择、训练配置和工程实践。
预训练配方
Llama 2 训练配置
| 参数 | 7B | 13B | 70B |
|---|---|---|---|
| 层数 | 32 | 40 | 80 |
| 隐藏维度 | 4096 | 5120 | 8192 |
| 注意力头数 | 32 | 40 | 64 |
| KV 头数 | 32 | 40 | 8 (GQA) |
| 中间维度 | 11008 | 13824 | 28672 |
| 学习率 | 3e-4 | 3e-4 | 1.5e-4 |
| Batch Size | 4M tokens | 4M tokens | 4M tokens |
| 训练 Tokens | 2T | 2T | 2T |
| LR Schedule | Cosine | Cosine | Cosine |
| Warmup | 2000 steps | 2000 steps | 2000 steps |
| Weight Decay | 0.1 | 0.1 | 0.1 |
| Grad Clip | 1.0 | 1.0 | 1.0 |
预训练启动脚本
bash
#!/bin/bash
# Llama-2-7B 预训练
deepspeed --num_gpus=8 pretrain.py \
--model_name_or_path meta-llama/Llama-2-7b-hf \
--data_path /data/processed/redpajama \
--bf16 \
--learning_rate 3e-4 \
--lr_scheduler_type cosine \
--warmup_steps 2000 \
--max_steps 500000 \
--per_device_train_batch_size 4 \
--gradient_accumulation_steps 16 \
--weight_decay 0.1 \
--max_grad_norm 1.0 \
--save_steps 2000 \
--save_total_limit 5 \
--logging_steps 10 \
--deepspeed ds_config_zero2.json \
--use_flash_attn \
--seed 42SFT 微调配方
python
# SFT 超参数
sft_config = {
"learning_rate": 2e-5,
"lr_scheduler_type": "cosine",
"num_train_epochs": 3,
"per_device_train_batch_size": 4,
"gradient_accumulation_steps": 4,
"max_seq_length": 4096,
"weight_decay": 0.01,
"warmup_ratio": 0.03,
"bf16": True,
}LoRA 高效微调
python
from peft import LoraConfig, get_peft_model
lora_config = LoraConfig(
r=16, # LoRA 秩
lora_alpha=32, # 缩放因子
target_modules=["q_proj", "v_proj", "k_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
)
model = get_peft_model(base_model, lora_config)
print(f"可训练参数: {model.print_trainable_parameters()}")
# 可训练参数: ~0.8% of totalLoRA 秩选择
- 7B 模型:r=8-16 通常足够
- 70B 模型:r=16-64 效果更好
- 目标模块:所有线性层 > 仅 Q/V > 仅 Q
训练监控要点
| 指标 | 正常范围 | 异常信号 |
|---|---|---|
| Loss | 稳步下降 | Spike > 2× 正常值 |
| Grad Norm | 稳定 | 突然增大 |
| LR | 按计划变化 | 未正确衰减 |
| Throughput | 稳定 | 逐步下降(数据瓶颈) |
| GPU Util | >80% | <60%(检查数据加载) |
Loss Spike 处理
Loss Spike 通常由坏数据或学习率过大引起。处理方案:1) 回滚到上一个检查点;2) 降低学习率;3) 检查数据质量。