在第 8.1 节完成框架与环境配置、第 8.2 节完成模型初始化之后,预训练的核心阶段正式启动。所谓“训练”,本质上是一个被精密编排的循环:数据加载 → 前向计算 → 反向传播 → 参数更新。这四个步骤每迭代一次,模型参数就向着“更擅长预测下一个 Token”的方向微调一次。
对于百亿以上参数的模型,这个循环不是在单卡上悠哉跑完的,而是在千卡集群上顶着显存上限、通信带宽和稳定性压力进行的极限操作。本节按工程落地视角,把每一步拆开讲透,告诉你显存在哪里被吃掉、时间在哪里被浪费、以及哪些参数调错会直接让训练崩掉。
一、数据加载(Data Loading):别让 GPU 等 CPU
训练开始前,原始文本需要经过 7.4 节提到的 Tokenizer 被切分为整数 ID 序列。数据加载阶段的任务,就是把这些 ID 高效地喂到 GPU 显存中。
1. 样本构造与 Batch 组装
预训练通常采用因果语言建模(CLM,详见 4.1 节),因此每个样本是一个固定长度的 Token 序列(如 2048、4096 或 8192)。数据加载器会:
- 将语料流式拼接成一个长序列;
- 按
seq_length切分,得到(batch_size, seq_length)的输入张量X; - 标签
Y通常是X左移一位(预测下一个 Token)。
2. 工程配置要点
- num_workers:PyTorch DataLoader 的多进程加载数。太小则 CPU 预处理跟不上 GPU,太大则进程间抢内存。单机多卡场景通常设为 4–8。
- pin_memory:必须开启。将数据页锁定在内存中,实现 CPU→GPU 的异步零拷贝传输。
- prefetch_factor:预取倍数。在 GPU 计算当前 batch 时,CPU 提前准备下一个 batch,掩盖数据搬运延迟。
- 分布式 sampler:配合数据并行(第 9.2 节),确保每个 epoch 中各卡拿到的数据子集不重复、不遗漏。
3. 长上下文训练的加载差异
当 seq_length 从 4K 扩展到 128K 或 1M 时,单个样本的显存占用暴涨。此时往往不再使用标准 DataLoader 预加载全量样本,而是采用内存映射(mmap)+ 流式读取,或配合 Megatron-LM/DeepSpeed 的自定义数据管道,避免主机内存(RAM)被压垮。
实用陷阱:如果训练时发现 GPU 利用率(
nvidia-smi中的 Volatile GPU-Util)周期性掉到 0%,大概率是数据加载瓶颈。先检查num_workers和存储 I/O,别急着加卡。
二、前向计算(Forward Pass):从 Embedding 到 Logits
数据到达 GPU 显存后,模型开始执行前向传播。以标准的 Decoder-only Transformer(如 Llama、GPT)为例,流程如下:
1. 输入嵌入(Input Embedding)
整数 Token ID 先经过词表嵌入矩阵 W_e,维度从 (batch, seq_len) 映射到 (batch, seq_len, hidden_size)。同时加入位置编码(RoPE 或绝对位置编码,详见 3.5 节)。
2. Transformer Block 堆叠
依次通过 N 个解码器层,每层包含:
- RMSNorm / LayerNorm:稳定数值分布;
- 自注意力(Self-Attention):计算 Q、K、V,执行
Softmax(QK^T/√d)V,得到上下文向量; - 残差连接:Attention 输出与输入相加;
- FFN(前馈网络):通常采用 SwiGLU 或 GELU 激活,升维再降维(如 4× hidden_size);
- 再次残差连接。
3. 语言模型头(LM Head)
最后经过一次 Norm,输入输出嵌入矩阵(部分模型与输入共享权重),投影到词表维度 (batch, seq_len, vocab_size),得到每个位置对所有 Token 的Logits。
4. 显存占用的第一杀手:激活值(Activations)
前向过程中,每一层的输入、Attention 矩阵、FFN 中间结果都需要保存在显存中,供反向传播使用。以 7B 模型、batch_size=4、seq_len=4096、FP16 训练为例,激活值可能占到 20–40 GB 显存。
工程救星:Activation Checkpointing(激活重计算)
通过设置 gradient_checkpointing=True,模型不再保存中间激活值,而是在反向传播时重新计算一次前向。代价是约 20–30% 的额外计算时间,但可节省 50–70% 的激活显存。对于大模型预训练,这是必选项,而非可选项。
5. 混合精度前向
现代训练几乎默认使用混合精度(详见 4.5 节):
- BF16/FP16:矩阵乘法用低精度,加速且省显存;
- FP32 主权重:维护一份高精度权重副本,防止低精度累积误差;
- Loss Scaling:FP16 训练时,对损失乘以一个缩放因子,避免梯度下溢为零。
三、反向传播(Backward Pass):梯度从何而来
前向得到了 Logits,接下来要计算损失并回传梯度。
1. 损失计算(Loss Computation)
将 Logits 与标签 Y 送入交叉熵损失函数(Cross-Entropy Loss,详见 4.3 节):
loss = cross_entropy(logits.view(-1, vocab_size), labels.view(-1))
注意:在 CLM 中,每个 Token 位置都参与损失计算(除了 padding 部分被 mask)。最终损失是整个序列的平均。
2. 梯度回传(Backpropagation)
PyTorch 的 loss.backward() 自动执行链式法则,从 LM Head 逐层回传到 Embedding 层,计算出每个可训练参数对损失的敏感度——即梯度(Gradient)。
3. 梯度累积(Gradient Accumulation)
由于单卡显存放不下理想的 global batch_size(如 4M Token),工程上采用梯度累积:
- 将 global batch 拆分为
N个 micro-batch; - 每个 micro-batch 独立前向+反向,但不立即更新参数;
- 累积
N个 micro-batch 的梯度后,执行一次参数更新。
公式上,这等价于增大了 batch_size,却不需要一次性加载全量数据到显存。例如,单卡 micro_batch=1,梯度累积 32 步,则单卡等效 batch=32。
4. 分布式下的梯度同步
在数据并行(DP/DDP)场景下,各卡计算的是不同数据子集的梯度。反向传播结束后,需要通过 AllReduce(第 9.7 节) 操作,将多卡梯度求平均,确保各卡参数更新方向一致。DeepSpeed 的 ZeRO 系列(第 9.5 节)还会在此阶段对优化器状态和梯度进行分片,进一步节省显存。
5. 梯度裁剪(Gradient Clipping)
大模型训练极易因个别异常样本导致梯度爆炸。通常设置 max_grad_norm=1.0,在参数更新前将梯度范数裁剪到阈值内,这是训练稳定性的标配(详见 4.5 节)。
四、参数更新(Parameter Update):优化器真正干活
梯度就绪后,进入一次训练迭代的最后一步。
1. 优化器计算(以 AdamW 为例,详见 4.4 节)
AdamW 维护两套动量状态:
- 一阶动量(m):梯度的指数移动平均(类似惯性);
- 二阶动量(v):梯度平方的指数移动平均(自适应学习率)。
对于每个参数 θ,更新规则为:
m_t = β1 * m_{t-1} + (1-β1) * g_t
v_t = β2 * v_{t-1} + (1-β2) * g_t^2
m_hat = m_t / (1-β1^t)
v_hat = v_t / (1-β2^t)
θ_t = θ_{t-1} - lr * m_hat / (sqrt(v_hat) + ε) - lr * weight_decay * θ_{t-1}
工程现实:AdamW 的显存开销是惊人的。除了 FP32 主权重,还需保存 m 和 v,导致每参数占用约 12–16 字节(FP32 权重 4B + 梯度 4B + m 4B + v 4B,混合精度下略有差异)。一个 7B 模型,仅优化器状态就可能吃掉 28 GB+ 显存,这也是为什么必须引入 ZeRO-Offload 或 ZeRO-3 进行分片卸载。
2. 学习率调度(Learning Rate Schedule)
预训练极少使用固定学习率,标准配方是:
- Warmup:前 0.1%–1% 的步数内,LR 从 0 线性增至峰值。防止初期参数不稳定直接飞掉;
- Cosine Decay:Warmup 后按余弦曲线衰减至最低值(通常是最小 LR 的 10%);
- Decay Ratio:峰值 LR 与模型大小、数据量相关,7B 模型常见峰值在
1e-4到3e-4量级,万亿参数模型可能低至1e-5量级。
3. 权重衰减(Weight Decay)
在 AdamW 中,权重衰减与 L2 正则解耦,直接对参数进行衰减,防止过拟合。默认值通常设为 0.1,但部分代码模型训练会调到 0.01。
4. 混合精度更新
若使用 BF16/FP16,更新发生在 FP32 主权重上:
- 将计算得到的 FP32 梯度更新到 FP32 主权重;
- 将更新后的 FP32 权重拷贝(cast)回 BF16/FP16,供下一轮前向使用。
5. 一次迭代收尾
执行完 optimizer.step() 和 lr_scheduler.step() 后,必须清空梯度:
optimizer.zero_grad()
否则下一轮累积的梯度会包含上一轮残余,导致训练震荡。
五、完整循环的工程速览
把以上四步串起来,一个标准的预训练 step 如下:
for batch in dataloader:
# 1. 数据加载:GPU 显存中已有 (batch, seq_len) 的张量
input_ids = batch["input_ids"].cuda()
labels = batch["labels"].cuda()
# 2. 前向计算(混合精度 + 激活重计算)
with autocast(dtype=torch.bfloat16):
logits = model(input_ids)
loss = cross_entropy(logits, labels)
loss = loss / gradient_accumulation_steps # 累积归一化
# 3. 反向传播
loss.backward()
# 若达到累积步数,执行更新
if (step + 1) % grad_accum_steps == 0:
# 梯度裁剪
clip_grad_norm_(model.parameters(), max_norm=1.0)
# 4. 参数更新 + 学习率调整
optimizer.step()
lr_scheduler.step()
optimizer.zero_grad()
六、常见问题与调优备忘
| 现象 | 可能原因 | 解决方向 |
|------|----------|----------|
| 显存 OOM( out of memory) | 激活值过大或 batch_size 过高 | 开启 Activation Checkpointing,减小 micro_batch,启用 ZeRO-2/3 |
| Loss 在初期飙升后 NaN | 学习率过大或 FP16 下溢 | 降低峰值 LR,加长 Warmup,换 BF16(比 FP16 更稳定) |
| 吞吐量(tokens/s/GPU)低 | 通信瓶颈或数据加载跟不上 | 检查 IB/NCCL 配置,增加 num_workers,使用 FlashAttention |
| 梯度范数持续>10 | 存在异常数据或初始化问题 | 检查数据清洗,加梯度裁剪,排查 abnormal loss spike |
七、小结
训练执行的四步循环看似简单,但在大模型工程中是显存、计算、通信三重约束下的精密平衡:
- 数据加载决定 GPU 是否有米下锅;
- 前向计算是显存占用的主战场,激活重计算是必开的阀门;
- 反向传播是分布式训练的同步枢纽,梯度累积决定了等效 batch_size;
- 参数更新背后站着 AdamW 和它的巨额状态开销,学习率调度则是防止训练翻车的方向盘。
这四个步骤跑通一轮,称为一个 step;跑完所有数据一轮,称为一个 epoch。由于预训练数据量极大(动辄数万亿 Token),预训练通常不以 epoch 计数,而是直接按 step 或消耗的总 Token 数(如 3×10^12 tokens)来规划。
在 8.4 节中,我们将进入训练过程的“驾驶舱”——如何监控损失曲线、梯度范数、显存占用和吞吐量,以及哪些指标异常是训练崩溃的前兆。