Skip to content
Gains Summary
Main Navigation 首页 / Home
C++ 编程 / C++ Programming
系统与高性能 / Systems & Performance
Web 开发 / Web Development
人工智能 / Artificial Intelligence
工业软件 / Industrial Software
其他内容 / Other Topics
C++ 编程 / C++系统与性能 / SystemsWeb 开发 / Web人工智能 / AI工业软件 / Industrial

外观

Sidebar Navigation

← 人工智能 / Artificial Intelligence

AI 基础设施 / AI Infrastructure

集群基础设施 / Cluster Infrastructure

1. GPU 集群基础设施全景——训练框架之下、硬件之上的那一层 / GPU Cluster Infrastructure Between Training Frameworks and Hardware

2. GPU 集群硬件架构——从 NVLink 到 InfiniBand / GPU Cluster Hardware from NVLink to InfiniBand

3. 异构硬件生态——CPU/DPU/NPU 的集群角色 / Roles of CPUs, DPUs, and NPUs in Heterogeneous Clusters

4. GPU 虚拟化与资源隔离——一张卡多人用 / GPU Virtualization and Resource Isolation

5. 作业调度系统——Kubernetes 和 Slurm / Job Scheduling with Kubernetes and Slurm

6. 多作业与多租户管理——让集群被所有人高效使用 / Multi-Job and Multi-Tenant Cluster Management

7. 网络架构与 RDMA——让 GPU 之间的通信更快 / Network Architecture and RDMA for Faster GPU Communication

8. NCCL 集群组网——大规模集合通信调优 / NCCL Cluster Networking and Collective Communication Tuning

9. 分布式存储——让数据跑得比 GPU 快 / Distributed Storage That Keeps GPUs Fed with Data

10. 集群运营与故障处理——让万卡集群稳定运行 / Operations and Failure Recovery for Large GPU Clusters

训练系统 / Training Systems

1. AI Infra 训练侧全景——让千亿参数模型跑起来需要什么 / Training-Side AI Infrastructure for Hundred-Billion-Parameter Models

2. GPU 硬件基础——为什么 GPU 比 CPU 快,显存为什么总是不够 / GPU Hardware, Parallel Throughput, and Memory Capacity

3. 分布式训练——如何把大模型分到多张卡上 / Distributing Large-Model Training Across Multiple GPUs

4. 显存优化——让 70B 模型在有限显存中跑起来 / Memory Optimization for Running 70B Models

5. 混合精度与通信——BF16 为什么是 LLM 训练的主流选择 / Mixed Precision and Communication with BF16

6. 预训练——Scaling Laws、数据工程与训练稳定性 / Pretraining with Scaling Laws, Data Engineering, and Stability

7. 后训练 SFT——从预训练模型到助手模型 / Supervised Fine-Tuning from Pretrained Model to Assistant

8. 后训练 RLHF/DPO——从助手模型到对齐模型 / RLHF and DPO from Assistant Model to Aligned Model

9. 高效微调——LoRA 和 QLoRA 让大模型走进消费级 GPU / Efficient Fine-Tuning with LoRA and QLoRA on Consumer GPUs

10. 训练工程——千卡集群的管理与故障恢复 / Training Engineering for Thousand-GPU Cluster Operations and Recovery

本页目录

预训练——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
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

第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 数也应该翻倍                   │
│                                                             │
└─────────────────────────────────────────────────────────────┘
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

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 需要根据数据质量调整                       │
│                                                             │
└─────────────────────────────────────────────────────────────┘
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

计算量的估算方法 ​

┌─────────────────────────────────────────────────────────────┐
│                 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 天                          │
│                                                             │
└─────────────────────────────────────────────────────────────┘
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

第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 / 社交媒体比例                    │
│                                                             │
└─────────────────────────────────────────────────────────────┘
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

数据清洗流水线 ​

┌─────────────────────────────────────────────────────────────┐
│                 数据清洗的完整流水线                            │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  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 解析为纯文本                      │  │
│  │  • 统一换行符、空格处理                             │  │
│  │  • 分句、分段处理                                   │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
└─────────────────────────────────────────────────────────────┘
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

数据格式: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
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

│ │ └─────────────────────────────────────────────────────────────┘


### 数据配比——课程学习的重要性
1
2

┌─────────────────────────────────────────────────────────────┐ │ 数据配比与课程学习策略 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 问题:不同来源的数据,质量差异巨大 │ │ → 直接混合训练,可能被高质量数据利用不足 │ │ → 被低质量数据带偏 │ │ │ │ 策略 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——预训练的最大噩梦
1
2
3
4
5
6

┌─────────────────────────────────────────────────────────────┐ │ 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,模型能力受限 │ │ │ │ 处理:修复架构,从头训练 │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ └─────────────────────────────────────────────────────────────┘


### 学习率调度——训练稳定性的关键
1
2

┌─────────────────────────────────────────────────────────────┐ │ 学习率调度: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
2

┌─────────────────────────────────────────────────────────────┐ │ 训练稳定性的其他关键实践 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 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 的保存策略
1
2
3
4
5
6

┌─────────────────────────────────────────────────────────────┐ │ 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 聚合) │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ └─────────────────────────────────────────────────────────────┘


### 断点续训时的数据一致性
1
2

┌─────────────────────────────────────────────────────────────┐ │ 断点续训的最大陷阱:数据顺序 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 问题:恢复训练时,数据从哪开始继续? │ │ │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 假设: │ │ │ │ • 训练集有 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": {...} │ │ │ │ } │ │ │ └─────────────────────────────────────────────────────┘ │ │ │ └─────────────────────────────────────────────────────────────┘


### 故障检测和自动恢复
1
2

┌─────────────────────────────────────────────────────────────┐ │ 训练故障检测与自动恢复 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 故障类型: │ │ ┌─────────────────────────────────────────────────────┐ │ │ │ 硬件故障(最常见): │ │ │ │ • 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
2
3
4

┌─────────────────────────────────────────────────────────────┐ │ 预训练的核心工程哲学 │ ├─────────────────────────────────────────────────────────────┤ │ │ │ 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 必须理解"清单
1
2
3
4

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,其中优化器状态占大头


---

**学习状态**:🟡 开始学习
1
2
3
4

最后更新于:

Pager
上一篇5. 混合精度与通信——BF16 为什么是 LLM 训练的主流选择 / Mixed Precision and Communication with BF16
下一篇7. 后训练 SFT——从预训练模型到助手模型 / Supervised Fine-Tuning from Pretrained Model to Assistant

持续记录,持续成长

Copyright © Tidenflow