前言
过去半年在 Agent 方向投入了大量精力,从 Agent 记忆架构设计到 Skill 工程化落地,积累了不少实践经验。但回头审视,发现自己一直停留在应用层------知道怎么调、怎么排、怎么蒸馏,但对底层为什么是这样缺乏系统理解。
于是决定给自己开一门"补课",从 Transformer 架构出发,系统学习大模型底层知识。不追求面面俱到,但求每个概念都能知道"为什么",本次来了解一下【预训练】的大致流程,后续再来逐步拆解,文章会持续更新。
核心训练目标:Next Token Prediction
预训练的核心任务极其简单:看到前面的词,预测下一个词。

为什么这个如此简单的任务能学到语言的一切能力?因为要准确预测下一个词,模型必须理解:
- 语法:英语的SVO结构、中文的SOV结构
- 事实知识:「巴黎是___的首都」→ 法国
- 推理逻辑:「如果A>B且B>C,那么A___C」→ >
- 代码语法:「def fibonacci(n):」→ 「if」
- 常识:「太阳从___升起」→ 东
预测下一个词是一个极其信息密集的任务------它迫使模型学习语言的方方面面。
训练数据工程------从原始数据到训练语料
数据预处理管道
原始数据不能直接喂给模型,必须经过严格的清洗流程:
第一步:去重(MinHash/LSH)
去除完全相同或高度相似的文档。使用 MinHash + LSH(局部敏感哈希)算法快速检测文档相似度,避免精确比较带来的计算开销。
原因:重复数据会让模型「背」而不是「学」------模型记住了重复文本的统计特征,但不会学到新的语言规律。
了解MinHash + LSH可以看这篇文章:blog.milvus.io/docs/zh/min...
第二步:质量过滤
通过三层过滤策略挑选高质量文本:
- 语言模型打分:用小模型给文本打分,保留高质量文本,剔除低质量段落。
- 规则过滤:去除太短(低于阈值)、太长(超出模型上下文)、乱码、广告、垃圾内容。
- 安全过滤:去除有害、色情、暴力等不适合训练的内容。
第三步:格式统一
将不同来源的数据统一为标准格式:统一编码(UTF-8)、统一换行符(LF)、去除 HTML 标签和特殊标记。
第四步:分词(Tokenization)
使用 BPE(Byte Pair Encoding)算法将文本切分成 token 序列。模型不认识「字」,只认识「token」------分词是模型理解语言的翻译器。
第五步:配比混合
按预定比例混合不同来源的数据,生成最终的训练语料。
分词器 + 词表构建
大模型不直接处理字符,而是处理token------token是介于字符和词之间的单位。
原始文本是字符串(中文汉字、英文单词、符号、数字、换行),神经网络只能接收数字数组输入,无法直接读文字。
分词器(Tokenizer)两件核心工作:
- 编码 encode:字符串 → 离散整数 ID(送入模型 Embedding 层)
- 解码 decode:模型输出整数 ID → 还原可读文本
目前主流的分词算法:
-
BPE(Byte Pair Encoding,字节对编码):OpenAI GPT1/2/3、LLaMA 原版使用
-
SentencePiece(内部也是 BPE,谷歌封装):LLaMA2、Qwen、Llama3、大多国产大模型标配,最大优势:原生支持中文、不需要预先分词、可直接对原始 Unicode 字符串训练词表
-
WordPiece(BERT、RoBERTa):带概率的 BPE 变体,优先合并高频且损失最小的片段,适合理解类 Encoder 模型,生成式预训练基本不用
这里大概介绍一下 BPE 算法
分步迭代示例(标准写法)
原始单词集合:low, lower, newest, widest
- 初始拆分(按单个字符分割,结尾添加
</w>标记单词边界)l o w </w>、l o w e r </w>、n e w e s t </w>、w i d e s t </w>基础单元集合:{l, o, w, e, r, n, i, d, s, t, </w>} - 第一轮统计所有相邻字符对频次(l,o):2、(o,w):2、(w,):1、(w,e):2、(e,r):1、(n,e):1、(e,w):1、(e,s):2、(s,t):2、(t,):2、(w,i):1、(i,d):1、(d,e):1最高频配对:
(l,o)合并:l o → lo文本更新:lo w </w>、lo w e r </w>、n e w e s t </w>、w i d e s t </w> - 第二轮重新统计配对频次最高频配对:
(lo,w)合并:lo w → low文本更新:low </w>、low e r </w>、n e w e s t </w>、w i d e s t </w> - 持续循环不断统计全局相邻单元频率,每次合并频次最高的单元对,新子词写入合并列表,全文同步替换;循环终止:词表总量达到预设数值(32k/64k),训练结束。
- 推理分词逻辑陌生词汇先拆成最基础字符,严格依照训练时的合并顺序从高优先级到低优先级做拼接,切分成词表里存在的子词 Token。
优势:
- 解决未登录词 OOV:陌生词汇会自动拆解为基础字节 / 基础子词,不会出现无法编码的文本
- 折中权衡:固定词表上限,高频词组合并为单个 Token 缩短序列长度,降低模型计算开销
- 通用性极强:底层基于 UTF-8 字节,英文、中文、符号、程序代码、多语言文本无需额外分词预处理均可适配
- 贪心统计驱动:完全依靠语料真实分布学习子词,无人工分词规则偏见,适配海量互联网预训练语料
词表(Vocab):就是「子串 <=> 唯一数字 ID」的映射字典,Embedding 矩阵的维度 = 词表总大小。
预训练全部流程
初始化环境
训练开始前,必须先搭建好运行环境。这一步做几个件事:
-
CUDA 底层全局配置配置显存动态分配策略、计算精度、cuDNN 加速等系统参数。作用:优化显卡显存占用行为,提升矩阵计算速度,避免训练时显存异常占用、TF32 计算无法启用等问题。
-
分布式初始化读取 WORLD_SIZE、RANK、LOCAL_RANK 分布式环境变量,依靠 NCCL 后端搭建所有显卡的互通通信组。作用:打通多卡之间的数据、梯度传输链路,所有分布式相关操作的前置必要步骤,必须最先执行。
-
分配显卡设备读取 LOCAL_RANK 编号,将当前进程绑定到对应编号的 GPU 硬件。作用:限定当前进程只在指定显卡执行运算,多进程之间硬件资源隔离,不会出现进程抢占显卡的问题。
-
配置全局随机种子设置基础固定种子,多卡环境中每个进程的种子叠加自身全局 rank 偏移。作用:基础种子固定保证实验结果可复现;种子错开可以让各个显卡采样到不同的数据序列,维持分布式训练的数据多样性。
-
项目目录初始化由主进程统一创建模型保存目录、断点目录、日志目录。作用:统一管理训练产出文件,多进程仅主进程读写文件,防止多进程并发写入造成文件损坏、日志错乱。
-
模型初始化与设备迁移实例化模型结构,可选加载历史断点权重,将整个模型参数搬运至绑定的 GPU 中。作用:生成可用于训练的模型实体,把模型数据从内存载入显存,满足 GPU 运算条件。
-
训练配套组件初始化创建混合精度缩放器、AdamW 优化器、学习率调度器,配置梯度裁剪、梯度累积等规则。作用:混合精度降低显存消耗;优化器负责参数更新;学习率调度控制训练过程的学习率变化;梯度相关配置防止梯度爆炸。
-
DDP 封装模型给模型外层包装 DDP,关闭 RoPE 内 freqs_cos、freqs_sin 固定常量参数的跨卡同步。作用:多卡前向各自独立计算,反向传播自动 AllReduce 聚合平均梯度,保证所有卡模型参数完全一致;屏蔽固定参数同步,减少网卡通信开销,提升训练速度。
配置混合精度训练
混合精度训练(Automatic Mixed Precision)是工业级训练的标准做法。核心思想是:部分计算用半精度(bfloat16),部分用全精度(float32) 。
为什么需要混合精度?现代 GPU 的 Tensor Core 对半精度计算有专门优化。使用 bfloat16 可以让显存占用减半、计算速度提升 1.5-2 倍,而精度损失几乎可忽略。
bfloat16 和 float16 的区别:
| 类型 | 尾数位 | 指数位 | 数值范围 | 精度 |
|---|---|---|---|---|
| float16 | 10位 | 5位 | 6.1e-5 ~ 65504 | 低,容易下溢 |
| bfloat16 | 7位 | 8位 | 1.2e-38 ~ 3.4e+38 | 中,下溢风险小 |
| float32 | 23位 | 8位 | 1.2e-38 ~ 3.4e+38 | 高 |
bfloat16 的指数位和 float32 一样,数值范围几乎相同,不容易溢出。尾数位比 float32 少但够用,因此推荐在支持的 GPU 上使用 bfloat16。
实际使用中,PyTorch 的 autocast 上下文管理器会自动完成精度切换:进入上下文后,矩阵乘法、卷积等操作自动转为半精度执行;softmax、LayerNorm 等对精度敏感的操作保持 float32;模型参数本身始终维护在 float32。
需要特别注意的是,bfloat16 和 float16 在梯度处理上有重要区别:bfloat16 的数值范围与 float32 几乎相同,梯度通常不会下溢为零,因此 bfloat16 训练不需要使用梯度缩放器(GradScaler);而 float16 的数值范围小得多,梯度很容易下溢,必须配合 GradScaler 使用------GradScaler 在前向传播时放大损失值,反向传播后反缩放梯度,防止半精度下的梯度丢失。CPU 训练时则完全不需要混合精度,使用全精度即可。
初始化模型与分词器
模型初始化包含几个关键操作:
创建模型实例:根据模型配置定义网络结构------隐藏层维度、Transformer 层数、注意力头数等。这些参数决定了模型的容量:层数越多、宽度越大,模型能学到的知识越丰富,但训练成本也越高。
创建分词器:分词器负责将文本转为 token id 序列。它有一个词汇表(vocabulary),每个 token 对应一个整数 ID。后续训练时,每段文本都会被分词器转为 token ID 序列喂给模型。
加载预训练权重(可选) :如果指定了 from_weight 参数,模型会从已有权重文件加载参数。这在迁移学习场景下非常有用------先用海量数据预训练一个基础模型,再用领域数据继续训练。
torch.compile 加速(可选) :PyTorch 2.0+ 的编译优化功能。原理是将模型的计算图编译为更高效的底层代码,通常提速 10%-30%。代价是首次运行有编译开销(几秒到几十秒)。
MoE(混合专家)架构(可选) :MoE 是一种稀疏激活的模型架构,核心思想是「模型很大,但每次只用一小部分」。MoE 模型包含多个专家网络(Expert),每一层通过一个门控网络(Gate/路由器)决定当前输入应该激活哪些专家。例如,一个 MoE 模型可能有 32 个专家,但每次只激活
加载训练数据
数据加载是预训练中工程量最大的环节之一。典型的预训练框架使用 DataLoader 体系,包含三层组件:
Dataset(数据集) :定义如何从磁盘读取数据并转为模型输入。预训练的数据格式通常是 jsonl------每行一个 JSON 对象,包含一段文本。Dataset 读取后会调用分词器转为 token id 序列,截断到最大长度,构造 (input_ids, labels) 对。
Sampler(采样器) :定义数据怎么打乱和分配。单卡训练时用 RandomSampler 随机打乱;多卡训练时用 DistributedSampler,确保每个 GPU 拿到不同的数据子集,避免重复计算。
DataLoader(数据加载器) :整合 Dataset 和 Sampler,提供批量加载、多线程预取、内存锁定等功能。pin_memory=True 将数据锁定在固定内存(pinned memory)中,加速 GPU 传输。num_workers 控制预取线程数,通常设为 CPU 核心数的 1/4 到 1/2。
断点续训的数据跳过机制:当从检查点恢复训练时,需要跳过已经训练过的 batch。这通过 SkipBatchSampler 实现------恢复时传入起始步数,Sampler 会跳过前面已完成的 batch,直接从断点处继续。每个 epoch 训练时还会用不同的随机种子重新打乱数据顺序,确保不同轮次的数据排列不同,增加数据多样性。
配置优化器与断点续训
优化器选择:预训练几乎无一例外地使用 AdamW 优化器。它结合了 Adam 的自适应学习率(每个参数有自己的学习率)和权重衰减正则化(防止过拟合)。
断点续训:训练大模型可能持续数天甚至数周。如果中途断电或崩溃,需要能从检查点恢复。恢复时需要加载四个状态:
- 模型权重:模型参数的当前值
- 优化器状态:AdamW 内部维护的一阶矩(m)和二阶矩(v),这些状态记录了每个参数的更新历史
- 梯度缩放器状态:混合精度训练中 GradScaler 的缩放因子
- 训练步数:当前训练到第几步
只恢复模型权重而不恢复优化器状态,会导致恢复后训练不稳定------因为优化器的动量信息和权重衰减状态被重置了。
DDP 对特殊参数的忽略:DDP 在同步梯度时,默认会对所有参与训练的参数进行 AllReduce。但有些缓冲区(如 RoPE 位置编码中的 freqs_cos 和 freqs_sin)是固定的查找表,各卡上的值完全相同,不需要同步。通过将这些缓冲区标记为不参与 AllReduce,可以节省不必要的通信开销,提升训练效率。
训练循环------核心执行流程
训练循环是预训练的灵魂。每一轮(epoch)中,模型遍历整个数据集一次,逐步更新参数。
数据搬运
从 DataLoader 拿到一个 batch 的 (input_ids, labels) 后,第一件事是把它搬到 GPU 上------调用 to(device) 将张量从 CPU 内存复制到 GPU 显存。这一步看起来简单,但实际上数据搬运的速度直接影响训练效率。如果 GPU 在计算时 CPU 在搬运数据,GPU 就会空闲等待。这就是 DataLoader 的 pin_memory 和多线程预取要解决的问题。
学习率调度
学习率不是固定不变的,而是随着训练步数动态调整:
- Warmup:训练初期从 0 线性增长到峰值学习率。原因是训练刚开始时参数是随机的,模型输出不稳定,如果一开始就用大学习率,梯度方向可能完全错误,导致训练崩溃。
- Cosine Decay:从峰值学习率按余弦曲线逐渐衰减到接近 0。训练后期模型已经学到了大部分知识,需要更小的步长做精细打磨。
前向传播
这是模型「看」数据的过程。input_ids 经过 Embedding 层转为向量表示,然后依次通过 Transformer 的每一层(Self-Attention + FFN),最终输出每个位置对下一个词的概率分布。
在混合精度模式下,前向传播的大部分计算(矩阵乘法等)自动用 bfloat16 执行,但 LayerNorm、softmax 等精度敏感操作保持 float32。
计算损失
模型输出的概率分布和真实的 labels 对比,用交叉熵损失(CrossEntropyLoss)计算一个标量损失值。损失值衡量「模型猜得多准」------越低越好。
对于 MoE(混合专家)模型,除了主损失外还有辅助损失(aux_loss),用于平衡不同专家模块的负载,避免某些专家被过度使用而其他专家闲置。具体来说,总损失由两部分组成:logits_loss(主损失,即交叉熵损失,衡量模型预测的准确性)和 aux_loss(辅助损失,衡量专家负载的均衡程度)。两者相加得到最终的总损失值,共同参与反向传播。logits_loss 直接反映模型的预测质量,而 aux_loss 只影响专家路由的优化方向,不直接影响模型的语言建模能力。
梯度累积
梯度累积是一种显存优化技巧。核心思想是:不立刻用梯度更新参数,而是把多步的梯度累加起来,最后一次性更新。
等效 batch size = 单步 batch size × 累积步数
比如单卡 batch_size=32、累积8步,等效 batch size 就是 256。这样可以模拟大 batch 训练的稳定性,同时不增加显存占用(因为每步只用处理 32 个样本的显存)。
反向传播
从损失值出发,通过链式法则逐层计算每个参数的梯度。在混合精度模式下,调用 scaler.scale(loss).backward()------先对损失值做缩放(防止半精度下的梯度下溢为零),再执行反向传播。
在多卡训练中,反向传播完成后,各卡的梯度会在 NCCL 通信组中自动同步(取平均)。这就是 DDP(DistributedDataParallel)的核心机制------各卡独立计算梯度,然后同步得到一致的参数更新方向。
参数更新
梯度计算完成后,需要三步操作更新参数:
- 反缩放梯度:scaler.unscale_(optimizer) 将梯度恢复到原始尺度(因为之前做过缩放)
- 梯度裁剪:clip_grad_norm_(model.parameters(), max_norm=1.0) 限制梯度的最大范数。梯度太大意味着参数更新步长过大,可能导致模型「一脚踏空」。裁剪是为了安全,不让模型「摔死」
- 更新参数:scaler.step(optimizer) 根据梯度更新参数。AdamW 会用每个参数的自适应学习率决定更新步长
最后更新优化器状态:optimizer.zero_grad(set_to_none=True)。set_to_none=True 比默认的 set_to_zero 更高效------它把梯度张量直接设为 None 而不是用零填充,节省显存。这是因为将梯度设为 None 后,JVM 的垃圾回收器可以立即回收对应的显存空间;而用零填充只是覆写了数据,显存仍然被占用。在大模型训练中,梯度张量往往非常大,set_to_none=True 可以显著降低峰值显存占用。
日志输出
每训练一定步数,打印关键指标:当前步数、损失值、Perplexity(困惑度)、实际学习率。这些指标帮助判断训练是否正常进行------损失持续下降说明学习有效,损失剧烈波动可能说明学习率太大。
模型保存
每训练一定步数,保存两个文件:
- 权重文件:只保存模型参数(float16 压缩),文件小,用于推理部署
- 完整检查点:保存模型参数 + 优化器状态 + 梯度缩放器状态 + 当前步数,用于断点续训
多卡训练时,只有主进程(rank 0)保存模型。其他进程的 DDP 包装器会自动同步参数,所以各卡权重完全一致,不需要重复保存。
保存模型时需要注意两个细节:(1) 保存前先将模型切换到评估模式(model.eval()),确保 Dropout 等随机操作被关闭,保存的是确定性的参数;保存后再切换回训练模式(model.train()),不影响后续训练。(2) 保存的应该是原始模型而非 DDP 包装后的模型。DDP 包装器在模型外面加了一层封装(.module),而如果还启用了 torch.compile,会再加一层(_orig_mod),因此需要通过 .module._orig_mod 或 .module 来获取原始模型参数,避免保存包含冗余包装的文件。
显存管理
大模型训练中,显存管理是一个容易被忽视但至关重要的环节。每个训练步结束后,需要主动释放不再需要的张量引用------将当前 batch 的 input_ids、labels、模型输出等变量用 del 显式删除,并手动触发 Python 垃圾回收(gc.collect())。这是因为 Python 的引用计数机制有时不能立即回收显存(特别是 PyTorch 的 CUDA 缓存),如果不主动清理,显存会持续累积,最终导致 OOM(显存溢出)崩溃。通常在模型保存或日志输出触发时执行一次显存清理,就足以维持稳定的训练。
训练完成后的清理
训练循环结束后,必须销毁分布式进程组(dist.destroy_process_group()),释放 NCCL 通信资源。这一步是分布式训练的标准收尾操作------如果不调用 destroy,NCCL 的通信上下文会一直占用系统资源,可能导致后续任务无法正常启动,甚至影响其他用户的分布式作业。
实际案例
LLaMA-3的三阶段预训练策略
LLaMA-3采用了一种创新的三阶段预训练策略:
| 阶段 | 上下文长度 | 数据量 | 目标 |
|---|---|---|---|
| 第一阶段(基础能力) | 4096 tokens | 15T tokens | 学习基本语言能力和知识 |
| 第二阶段(中等上下文) | 8192 tokens | 5T tokens | 逐步适应更长的上下文 |
| 第三阶段(长上下文) | 128K tokens | 5T tokens | 掌握超长文本理解能力 |
为什么分阶段? 如果一开始就用128K上下文训练,计算成本会非常高(O(n²)复杂度)。先在短上下文上训练基础能力,再逐步扩展到长上下文,既节省成本又能达到好的效果。
Chinchilla的最优训练策略
DeepMind 的 Chinchilla 论文给出了一个重要结论:
在给定算力预算下,模型参数和训练数据量应该等比例增加。
经验法则:Token数 ≈ 参数量 × 20
| 模型 | 参数量 | 最优tokens | 实际tokens | 是否最优 |
|---|---|---|---|---|
| GPT-3 | 1750亿 | 35万亿 | 3000亿 | ❌ 严重不足 |
| Chinchilla | 700亿 | 1.4万亿 | 1.4万亿 | ✅ 最优 |
| LLaMA-2 70B | 700亿 | 1.4万亿 | 2万亿 | ✅ 超额 |
| LLaMA-3 405B | 4050亿 | 8万亿 | 15万亿 | ✅ 大幅超额 |
GPT-3 的教训:参数太多但数据不够,相当于一个聪明人读的书太少。LLaMA 系列通过增加数据量(而非一味增大参数),实现了更好的性价比。
LLaMA-2 70B 的典型配置
| 参数 | 值 | 说明 |
|---|---|---|
| 序列长度 | 4096 tokens | 每个训练样本的长度 |
| 批次大小 | 4M tokens/step | 每步处理的token数 |
| 训练步数 | 1,400,000 steps | 总训练步数 |
| 训练数据量 | 2万亿 tokens | 总数据量 |
| 优化器 | AdamW | β1=0.9, β2=0.95, weight_decay=0.1 |
| 学习率 | 峰值 1.5e-4 | warmup 2000步 |
| 硬件 | 2048张A100 80G GPU | |
| 训练时间 | 约12天 | |
| 总FLOPs | 约840 ZetaFLOPs |
批次大小为什么这么大(4M tokens)? 大 batch size 有 3 个好处:(1) 梯度估计更准确(更多样本的平均);(2) GPU 利用率更高;(3) 训练更稳定。但太大会导致泛化能力下降,4M tokens 是经过大量实验得出的最优值。
评估指标:Perplexity(困惑度)
Perplexity 是衡量语言模型质量的核心指标:
Perplexity=exp(−N1∑i=1NlogP(wi∣w1,...,wi−1))
直观理解:Perplexity 表示模型在每个位置「犹豫」的程度。如果 Perplexity=10,意味着模型在每个位置像是在 10 个词中选一个。Perplexity 越低,模型越「确定」。
| 模型 | Perplexity(越低越好) |
|---|---|
| GPT-2 | ~20 |
| GPT-3 | ~15 |
| LLaMA-2 70B | ~5-7 |
| LLaMA-3 405B | ~3-5 |
最后
感谢你能看到这里,本文简单介绍了【预训练(Pre-training)】的流程概览,希望对你有用,文章不全,后续会跟着笔者的学习进度持续更新,更多 Agent、前端、Node、性能相关的技术文章和实践总结,可以查看我的代码花园: