预训练——Scaling Laws、数据工程与训练稳定性 / Pretraining with Scaling Laws, Data Engineering, and Stability
📅 创建时间:2026-06-02 🏷️ 标签:#预训练 #Scaling-Laws #数据工程 #学习率调度 #Loss-Spike #断点续训 📚 前置知识:[[01-gpu-hardware]](GPU 硬件基础) [[02-distributed-training]](分布式训练) [[03-memory-optimization]](显存优化) [[04-mixed-precision]](混合精度) 📚 相关知识:[[06-posttraining-sft]](SFT) [[09-training-engineering]](训练工程)
场景:Llama-3 用 15T token 训练,怎么保证训练稳定
┌─────────────────────────────────────────────────────────────┐
│ │
│ 你决定训练一个 Llama-3-70B 级别的模型。 │
│ │
│ 核心数据: │
│ • 参数:70B │
│ • 训练数据:15T(15 万亿)token │
│ • 硬件:4096 张 H100,跑 90 天 │
│ • 成本:约 5000 万美元 │
│ │
│ 你的担忧: │
│ │
│ 问题 1:这个数据量够吗? │
│ → 15T token 是怎么算出来的?Scaling Laws 怎么用? │
│ │
│ 问题 2:训练到一半 Loss spike 了怎么办? │
│ → 90 天的训练,中途坏了意味着数百万美元打水漂 │
│ │
│ 问题 3:数据从哪里来,怎么清洗? │
│ → 互联网上那么多文本,哪些可以用于训练? │
│ │
│ 问题 4:训练完了怎么验证效果? │
│ → 只有 Loss 不够,需要建立能力评估体系 │
│ │
│ 这就是预训练要回答的核心问题。 │
│ │
└─────────────────────────────────────────────────────────────┘第1节:Scaling Laws——训练需要多少数据和算力
Chinchilla 定律——数据量和模型参数量同样重要
┌─────────────────────────────────────────────────────────────┐
│ Scaling Laws:Kaplan vs Chinchilla │
├─────────────────────────────────────────────────────────────┤
│ │
│ Kaplan et al. (2020, GPT-3 论文): │
│ → 模型性能主要随参数量扩展,数据量相对不那么重要 │
│ → 推荐:10B 参数模型,训练 200B token │
│ → 模型越大,数据效率越高(每个 token 的价值更高) │
│ │
│ Chinchilla (Hoffmann et al., 2022, DeepMind): │
│ → 重新定义 Scaling Laws │
│ → 模型大小和训练 token 数应该同比例扩展 │
│ → 推荐:10B 参数模型,训练 200B token(相同!) │
│ → 但更准确的公式:Training Token ≈ 20 × Parameters │
│ │
│ 核心公式(Chinchilla): │
│ │
│ Loss ≈ (a × N^α + b × C^β + c)^γ │
│ │
│ 其中: │
│ • N = 模型参数量 │
│ • C = 计算量(FLOPs) │
│ • α ≈ 0.73, β ≈ 0.28, γ ≈ -0.34(拟合参数) │
│ │
│ 最优配置(给定计算预算 B FLOPs): │
│ • N* ∝ B^0.5(参数量) │
│ • C* ∝ B^0.5(计算量) │
│ • T* ∝ B^0.5(token 数) │
│ │
│ 结论:参数量翻倍时,token 数也应该翻倍 │
│ │
└─────────────────────────────────────────────────────────────┘LLM 的 Scaling Laws 实践
┌─────────────────────────────────────────────────────────────┐
│ 各模型的 Scaling 配置对比 │
├─────────────────────────────────────────────────────────────┤
│ │
│ │ 模型 │ 参数 │ Token 数 │ 比例 │
│ ├───────────────┼─────────┼───────────┼───────────────┤
│ │ GPT-3 │ 175B │ 300B │ 1.7x │
│ ├───────────────┼─────────┼───────────┼───────────────┤
│ │ Chinchilla │ 70B │ 1.4T │ 20x ✓ │
│ ├───────────────┼─────────┼───────────┼───────────────┤
│ │ PaLM │ 540B │ 780B │ 1.4x │
│ ├───────────────┼─────────┼───────────┼───────────────┤
│ │ LLaMA 1 │ 65B │ 1.4T │ 22x ✓ │
│ ├───────────────┼─────────┼───────────┼───────────────┤
│ │ LLaMA 2 │ 70B │ 2.0T │ 29x ✓ │
│ ├───────────────┼─────────┼───────────┼───────────────┤
│ │ LLaMA 3 │ 70B │ 15.0T │ 214x │
│ ├───────────────┼─────────┼───────────┼───────────────┤
│ │ LLaMA 3 │ 405B │ 15.0T │ 37x │
│ │
│ 注意:LLaMA 3 远超过 Chinchilla 最优比例 │
│ → 可能的解释:数据质量大幅提升(超过滤后的数据) │
│ → 数据质量 ↑ → 可以用更多 token 训练而不饱和 │
│ │
│ 新理解(数据质量调整后的 Scaling): │
│ • 低质量数据:1B token ≈ 1 token(饱和快) │
│ • 高质量数据:1B token ≈ 5-10 token(饱和慢) │
│ → Scaling Laws 需要根据数据质量调整 │
│ │
└─────────────────────────────────────────────────────────────┘计算量的估算方法
┌─────────────────────────────────────────────────────────────┐
│ LLM 训练的计算量估算 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 标准公式(每个 token 的 FLOPs): │
│ │
│ FLOPs/token ≈ 2 × N │
│ │
│ 推导: │
│ • Forward:每个参数参与 2 次乘加(乘 + 加) │
│ • Backward:约 2x Forward(梯度计算) │
│ • 总计:约 6N FLOPs/token │
│ │
│ 实际中考虑 Activation 重计算: │
│ FLOPs/token ≈ 6N × (1 + checkpoint_ratio) │
│ │
│ 完整训练的计算量(总 FLOPs): │
│ │
│ Total FLOPs = 6 × N × T │
│ │
│ 其中: │
│ • N = 参数量 │
│ • T = 训练 token 数 │
│ • 6 = 前向 2 + 反向 4 的经验系数 │
│ │
│ 例子:Llama-3-70B,训练 15T token │
│ Total FLOPs = 6 × 70B × 15T = 6300 PFLOPS-days │
│ 在 4096 张 H100(989 TFLOPS FP8)上: │
│ 时间 = 6300 × 10^15 / (4096 × 989 × 10^12) │
│ ≈ 1.57 天(纯计算,无效率损失) │
│ 考虑 MFU=50%:约 3.14 天 │
│ 考虑实际各种 overhead:约 90 天 │
│ │
└─────────────────────────────────────────────────────────────┘第2节:数据工程——训练语料从哪来,怎么处理
预训练语料的类型和来源
┌─────────────────────────────────────────────────────────────┐
│ LLM 预训练语料来源与比例 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 主流预训练语料构成(以 LLaMA 3 为例): │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Common Crawl(网页爬取): 约 46% │ │
│ │ → 来源最广,但噪声最多,需要大量清洗 │ │
│ ├─────────────────────────────────────────────────────┤ │
│ │ C4(Colossal Clean Crawled Corpus):约 15% │ │
│ │ → Google 的清洗版本,质量较好 │ │
│ ├─────────────────────────────────────────────────────┤ │
│ │ GitHub(代码): 约 5% │ │
│ │ → 代码能力的关键来源,GitHub 协议允许使用 │ │
│ ├─────────────────────────────────────────────────────┤ │
│ │ Wikipedia / Books: 约 5% │ │
│ │ → 知识密集,语言规范,但数据量有限 │ │
│ ├─────────────────────────────────────────────────────┤ │
│ │ arXiv(学术论文): 约 2% │ │
│ │ → 数学、科学的知识来源 │ │
│ ├─────────────────────────────────────────────────────┤ │
│ │ StackExchange(问答): 约 3% │ │
│ │ → 知识问答格式,数据质量高 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 不同模型侧重的语料差异: │
│ • Code Llama:代码比例提升到 85%+ │
│ • 数学模型:提升 arXiv + 数学教科书比例 │
│ • 对话模型:提升 Reddit / 社交媒体比例 │
│ │
└─────────────────────────────────────────────────────────────┘数据清洗流水线
┌─────────────────────────────────────────────────────────────┐
│ 数据清洗的完整流水线 │
├─────────────────────────────────────────────────────────────┤
│ │
│ Step 1:去重(Deduplication) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ • URL 去重:相同 URL 只保留一个 │ │
│ │ • Exact 去重:完全相同的文本只保留一份 │ │
│ │ • Near-Duplicate:MinHash/LSH 找近似重复 │ │
│ │ │ │
│ │ 效果:Common Crawl 去重后减少 40-60% │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 2:质量过滤(Quality Filtering) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ • 语言识别:只保留目标语言(英语/中文等) │ │
│ │ • 长度过滤:过滤过短(<100 chars)或过长(>100KB) │ │
│ │ • 噪声过滤: │ │
│ │ - 包含大量特殊字符/乱码 │ │
│ │ - 重复内容过多("the the the the...") │ │
│ │ - 包含不良内容标记 │ │
│ │ • 模型打分:用小模型(如 BERT)预测质量分数 │ │
│ │ - 高质量文档打分 > 阈值,保留 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 3:安全过滤(Safety Filtering) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ • NSFW 内容识别 │ │
│ │ • 恶意软件/钓鱼内容 │ │
│ │ • 个人信息(PII)识别:姓名、电话、邮箱等 │ │
│ │ • 版权内容(可选,取决于法律考量) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 4:格式标准化(Normalization) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ • Unicode 规范化(NFKC) │ │
│ │ • HTML/Markdown 解析为纯文本 │ │
│ │ • 统一换行符、空格处理 │ │
│ │ • 分句、分段处理 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────┘数据格式:Arrow 和 Parquet
┌─────────────────────────────────────────────────────────────┐
│ 预训练数据格式:Arrow vs JSON vs Raw Text │
├─────────────────────────────────────────────────────────────┤
│ │
│ 传统方式(Raw Text): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ text.txt: │ │
│ │ The quick brown fox jumps... │ │
│ │ Another document starts here... │ │
│ │ │ │
│ │ 问题: │ │
│ │ • 无法并行读取(需要全文扫描找边界) │ │
│ │ • 元数据缺失(来源、语言、分数等) │ │
│ │ • 不支持随机访问 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 推荐方式(Apache Arrow / Parquet): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ dataset/: │ │
│ │ ├── train-00000.parquet # 100K 文档/文件 │ │
│ │ ├── train-00001.parquet │ │
│ │ └── ... │ │
│ │ │ │
│ │ Parquet schema: │ │
│ │ { │ │
│ │ "text": "string", # 文档内容 │ │
│ │ "source": "string", # 来源(cc/gutenberg等) │ │
│ │ "language": "string", # 语言(en/zh) │ │
│ │ "quality_score": "float", # 质量分数 │ │
│ │ "num_tokens": "int", # token 数 │ │
│ │ "url": "string", # 原始 URL │ │
│ │ } │ │
│ │ │ │
│ │ 优势: │ │
│ │ • 列式存储,只读取需要的列 │ │
│ │ • 支持过滤器下推(WHERE language='en') │ │
│ │ • 支持多进程并行读取 │ │
│ │ • 压缩率高(列内重复数据) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ PyTorch DataLoader 使用 Arrow/Parquet: │
```python
from datasets import load_dataset
ds = load_dataset("parquet", data_files="dataset/train-*.parquet")
# 支持流式加载,不需要把所有数据加载到内存
def tokenize(examples):
return tokenizer(examples["text"], truncation=True, max_length=seq_len)
tokenized_ds = ds.map(
tokenize,
batched=True,
num_proc=64, # 多进程并行处理
remove_columns=["text"] # 删除原文节省内存
)│ │ └─────────────────────────────────────────────────────────────┘
### 数据配比——课程学习的重要性┌─────────────────────────────────────────────────────────────┐ │ 数据配比与课程学习策略 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 问题:不同来源的数据,质量差异巨大 │ │ → 直接混合训练,可能被高质量数据利用不足 │ │ → 被低质量数据带偏 │ │ │ │ 策略 1:质量加权采样(Quality-Weighted Sampling) │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ • 每个文档有质量分数 q ∈ [0, 1] │ │ │ │ • 采样概率 P(doc) ∝ q^α │ │ │ │ • α 控制采样倾向: │ │ │ │ - α=0:均匀采样 │ │ │ │ - α=1:按质量加权 │ │ │ │ - α 过高:可能过拟合高质量数据 │ │ │ │ │ │ │ │ LLaMA 3 的经验:α ≈ 0.5-1.0 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 策略 2:课程学习(Curriculum Learning) │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 早期: │ │ │ │ → 多用 Wikipedia/Books(语言规范、结构清晰) │ │ │ │ → 少用网页(噪声多) │ │ │ │ → 帮助模型建立基础语言能力 │ │ │ │ │ │ │ │ 中期: │ │ │ │ → 增加代码比例 │ │ │ │ → 增加 Common Crawl(提升知识覆盖面) │ │ │ │ │ │ │ │ 后期: │ │ │ │ → 增加高质量对话数据 │ │ │ │ → 增加数学/科学数据 │ │ │ │ → 培养复杂推理能力 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ └─────────────────────────────────────────────────────────────┘
---
## 第3节:训练稳定性——Loss Spike 和学习率调度
### Loss Spike——预训练的最大噩梦┌─────────────────────────────────────────────────────────────┐ │ Loss Spike:原因与处理 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 什么是 Loss Spike: │ │ → 正常情况:Loss 从 4.0 稳步下降到 2.5 │ │ → Spike:Loss 突然跳到 10、50、甚至 100,然后恢复或不恢复 │ │ │ │ 原因 1:坏数据(Bad Data Batch) │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ • Batch 内混入了噪声极大的文档 │ │ │ │ • 编码特殊字符导致 tokenizer 产生异常 token │ │ │ │ • 梯度过大,更新后参数进入不稳定区域 │ │ │ │ │ │ │ │ 特征:Spike 后快速恢复(1-2 步) │ │ │ │ 处理:跳过坏 batch,继续训练 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 原因 2:学习率过高(Learning Rate too High) │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ • 参数更新幅度过大,跳出局部最优 │ │ │ │ • 某些层进入饱和区(激活饱和、logits爆炸) │ │ │ │ │ │ │ │ 特征:Spike 后缓慢恢复或持续不稳定 │ │ │ │ 处理:降低学习率,从 checkpoint 恢复 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 原因 3:数值溢出(Numerical Overflow) │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ • BF16 仍然溢出了极少数极端值 │ │ │ │ • Softmax 前 logits 过大(e^1000 → inf) │ │ │ │ │ │ │ │ 特征:Loss 直接跳到 nan │ │ │ │ 处理:检查溢出位置,添加数值裁剪 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 原因 4:架构问题(Architecture Issues) │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ • 某些层初始化不当 │ │ │ │ • 残差连接有问题 │ │ │ │ • Attention 缩放因子 1/sqrt(d) 有误 │ │ │ │ │ │ │ │ 特征:规律性的小幅 Spike,模型能力受限 │ │ │ │ 处理:修复架构,从头训练 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ └─────────────────────────────────────────────────────────────┘
### 学习率调度——训练稳定性的关键┌─────────────────────────────────────────────────────────────┐ │ 学习率调度:Warmup + Cosine Decay │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 标准调度(GPT-3 / LLaMA 等主流模型采用): │ │ │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ lr │ │ │ │ │╲ │ │ │ │ │ ╲ │ │ │ │ │ ╲ │ │ │ │ │ ╲ │ │ │ │ │ ╲ │ │ │ │ │ ╲____ │ │ │ │ │ ‾‾‾‾‾‾‾‾‾ │ │ │ │ └────────────────────────────────────────────── │ │ │ │ 0 warmup peak decay │ │ │ │ steps │ │ │ │ │ │ │ │ 1. Warmup: 线性从 0 增到 peak_lr(2-5% 总步数) │ │ │ │ 2. Peak: 保持 peak_lr(0%-5% 总步数,可选) │ │ │ │ 3. Cosine Decay: 余弦曲线衰减到 min_lr │ │ │ │ 4. Linear Decay(可选): 最后线性衰减到 0 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 关键参数: │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 参数 │ 推荐值 │ │ │ │ ──────────────────────┼────────────────────────────│ │ │ │ peak_lr │ 1e-4 ~ 3e-4(GPT-3: 1.2e-4)│ │ │ │ min_lr(cosine 终点)│ peak_lr / 100 │ │ │ │ warmup_steps │ 2000 ~ 20000 │ │ │ │ total_steps │ 由训练 token 数决定 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 为什么需要 Warmup: │ │ → 初始参数随机,梯度方向可能不稳定 │ │ → 冷启动时,过大的学习率会导致参数跳变 │ │ → Warmup 让优化器逐步"热身",找到稳定方向 │ │ │ └─────────────────────────────────────────────────────────────┘
### 训练稳定性的其他关键实践┌─────────────────────────────────────────────────────────────┐ │ 训练稳定性的其他关键实践 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 1. 权重初始化(Weight Initialization) │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ GPT-2 / LLaMA 采用: │ │ │ │ W ~ Normal(0, sqrt(2/n_in)) # RMSNorm 友好 │ │ │ │ │ │ │ │ 残差分支(Attention + FFN)的初始化: │ │ │ │ → 输出投影 W_o 乘以缩放因子 1/sqrt(2n) │ │ │ │ → 确保残差连接足够强,不会被旁路掩盖 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 2. 梯度裁剪(Gradient Clipping) │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ grad_norm = torch.nn.utils.clip_grad_norm_( │ │ │ │ model.parameters(), max_norm=1.0 │ │ │ │ ) │ │ │ │ │ │ │ │ 作用:防止梯度爆炸,限制参数更新幅度 │ │ │ │ 值:通常 0.5 ~ 1.0,预训练用 1.0 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 3. 激活函数选择 │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ SwiGLU(Llama 2/3 采用): │ │ │ │ f(x) = x * silu(W_g(x)) │ │ │ │ → 比 ReLU 更平滑,梯度流动更好 │ │ │ │ → 比 GeLU 计算更快(sigmoid 近似) │ │ │ │ │ │ │ │ GeLU(LLaMA 1 / GPT-4 采用): │ │ │ │ f(x) = x * Phi(x) │ │ │ │ → 精度更高,但计算稍慢 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ └─────────────────────────────────────────────────────────────┘
---
## 第4节:断点续训——90 天训练的中途管理
### Checkpoint 的保存策略┌─────────────────────────────────────────────────────────────┐ │ Checkpoint 保存:频率 vs 存储成本 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 问题:保存太频繁 → 存储压力大、影响训练 │ │ 保存太少 → 中途故障损失大 │ │ │ │ 常见策略: │ │ │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 策略 1:固定间隔保存 │ │ │ │ → 每 1000 步保存一个 checkpoint │ │ │ │ → 最简单,但可能浪费存储 │ │ │ │ │ │ │ │ 策略 2:指数间隔保存(推荐) │ │ │ │ → 1, 2, 4, 8, 16, 32, 64, 128, 256, ... 步 │ │ │ │ → 训练早期保存频繁(参数不稳定) │ │ │ │ → 训练后期保存稀疏(参数稳定) │ │ │ │ │ │ │ │ 策略 3:基于时间保存 │ │ │ │ → 每 30 分钟保存一次(与步数无关) │ │ │ │ → 适合长时间训练 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 70B 模型 Checkpoint 大小: │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 模型权重(FP32):280 GB │ │ │ │ 优化器状态(FP32):560 GB │ │ │ │ 梯度(FP32):280 GB │ │ │ │ 总计(完整保存):~1.1 TB │ │ │ │ │ │ │ │ 优化方案(只保存权重 + 优化器状态): │ │ │ │ → ~840 GB │ │ │ │ │ │ │ │ 极致优化(ZeRO-3,只保存权重切片): │ │ │ │ → ~280 GB(需要 rank 0 聚合) │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ └─────────────────────────────────────────────────────────────┘
### 断点续训时的数据一致性┌─────────────────────────────────────────────────────────────┐ │ 断点续训的最大陷阱:数据顺序 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 问题:恢复训练时,数据从哪开始继续? │ │ │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 假设: │ │ │ │ • 训练集有 1T token │ │ │ │ • 当前步:500,000 │ │ │ │ • 训练中断,要从 checkpoint 恢复 │ │ │ │ │ │ │ │ 错误做法: │ │ │ │ → 直接从 shard 0 重新开始数据加载 │ │ │ │ → 会重复训练前 500K 步 的数据 │ │ │ │ → 模型过拟合,数据污染 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 正确做法: │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 方案 1:保存数据加载器状态 │ │ │ │ → 保存当前 shard index + position in shard │ │ │ │ → 恢复时精确恢复到同一位置 │ │ │ │ → 最精确,但 checkpoint 更复杂 │ │ │ │ │ │ │ │ 方案 2:Epoch-based 训练 + 全局 shuffle │ │ │ │ → 训练前对全部数据做一次全局 shuffle(随机种子) │ │ │ │ → 分成固定数量的 epoch │ │ │ │ → 恢复时,只要记录当前 epoch + step in epoch │ │ │ │ → 简单,但需要足够大的 epoch 保证随机性 │ │ │ │ │ │ │ │ 方案 3:数据状态保存到 metadata │ │ │ │ → Checkpoint 目录保存: │ │ │ │ { │ │ │ │ "global_step": 500000, │ │ │ │ "data_shard": "webtext-0042", │ │ │ │ "position_in_shard": 1234567, │ │ │ │ "rng_state": {...} │ │ │ │ } │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ └─────────────────────────────────────────────────────────────┘
### 故障检测和自动恢复┌─────────────────────────────────────────────────────────────┐ │ 训练故障检测与自动恢复 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 故障类型: │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 硬件故障(最常见): │ │ │ │ • GPU ECC Error / Xid Error │ │ │ │ • NCCL Timeout(某张卡无响应) │ │ │ │ • 网络 IB 断开 │ │ │ │ • NVMe 写入失败 │ │ │ │ │ │ │ │ 软件故障: │ │ │ │ • Python 进程崩溃(OOM, Segmentation Fault) │ │ │ │ • PyTorch CUDA 错误(illegal memory access) │ │ │ │ • NCCL 内部错误 │ │ │ │ │ │ │ │ 数据故障(隐蔽): │ │ │ │ • 某个 shard 损坏,读到乱码 │ │ │ │ • 数据管道卡住(死锁) │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 监控指标(需要实时追踪): │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ • Loss 趋势(是否有 spike) │ │ │ │ • GPU 利用率(是否有 GPU 掉队) │ │ │ │ • 梯度范数(是否有梯度爆炸) │ │ │ │ • Learning Rate(调度是否正常) │ │ │ │ • 数据吞吐量(samples/sec 是否稳定) │ │ │ │ • NCCL 通信时间(是否有通信瓶颈) │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ │ 自动恢复流程(用 Ray / SkyPilot 等调度器): │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 1. 检测到故障(GPU 掉队 / NCCL Timeout) │ │ │ │ 2. 保存当前 progress 到持久化存储 │ │ │ │ 3. 终止所有进程 │ │ │ │ 4. 请求新 GPU 资源 │ │ │ │ 5. 从最近 checkpoint 恢复 │ │ │ │ 6. 重新启动训练 │ │ │ │ 7. 记录故障日志(用于后续分析) │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ └─────────────────────────────────────────────────────────────┘
---
## 升华:预训练的工程哲学┌─────────────────────────────────────────────────────────────┐ │ 预训练的核心工程哲学 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 1. Scaling Laws是规划工具,不是金科玉律 │ │ → Chinchilla 最优比例是理论值,实际中数据质量差异巨大 │ │ → 高质量数据可以训练更多 token(Chinchilla 也认) │ │ → 重要的是测量,而不是盲目遵循公式 │ │ │ │ 2. 数据质量比数据数量更重要 │ │ → Common Crawl 有 100T+ token,但可用不到 10% │ │ → 质量过滤的价值往往被低估 │ │ → 一个高质量的 1T token 数据集 > 低质量的 10T 数据集 │ │ │ │ 3. 训练稳定性是一切的前提 │ │ → 90 天的训练,一次 Loss Spike 可能浪费数百万美元 │ │ → 学习率调度、权重初始化、数值稳定性必须做对 │ │ → 不要过早优化,先让训练稳定跑起来 │ │ │ │ 4. 故障是必然,不是偶然 │ │ → 10000 张卡跑 30 天,任何一张卡故障概率 ≈ 100% │ │ → 必须在设计阶段就把故障恢复考虑进去 │ │ → 定期 checkpoint、自动化监控、快速恢复流程是标配 │ │ │ │ 一句话总结: │ │ 预训练是 Scaling Laws、数据工程、训练稳定性的三位一体。 │ │ 任何一个短板,都会成为整个系统的瓶颈。 │ │ │ └─────────────────────────────────────────────────────────────┘
---
## "AI 可查 vs 必须理解"清单AI 可查: ✅ 不同模型的 Scaling Laws 系数(Kaplan vs Chinchilla 的精确拟合参数) ✅ 具体的数据过滤规则(长度阈值、质量分数阈值) ✅ PyArrow / datasets 库的具体 API
必须理解: 🔴 Chinchilla Scaling Laws:参数量和 token 数的最优比例 ≈ 20-30x 🔴 为什么 LLaMA 3 训练了 15T token(远超 Chinchilla 最优比例) 🔴 数据清洗的三个阶段:去重 → 质量过滤 → 安全过滤 🔴 为什么需要 Warmup + Cosine 学习率调度,以及各自的典型参数 🔴 Loss Spike 的四种原因,以及对应的处理方法 🔴 断点续训时,如何保证数据不重复(shard + position vs epoch-based) 🔴 70B 模型完整 checkpoint 约 1TB,其中优化器状态占大头
---
**学习状态**:🟡 开始学习