人人都会AI编程

8.3 训练执行全流程:数据加载、前向计算、反向传播、参数更新

更新时间:2026-07-09

在第 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 主权重,还需保存 mv,导致每参数占用约 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-43e-4 量级,万亿参数模型可能低至 1e-5 量级。

3. 权重衰减(Weight Decay)
在 AdamW 中,权重衰减与 L2 正则解耦,直接对参数进行衰减,防止过拟合。默认值通常设为 0.1,但部分代码模型训练会调到 0.01

4. 混合精度更新
若使用 BF16/FP16,更新发生在 FP32 主权重上:

  1. 将计算得到的 FP32 梯度更新到 FP32 主权重;
  2. 将更新后的 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 节中,我们将进入训练过程的“驾驶舱”——如何监控损失曲线、梯度范数、显存占用和吞吐量,以及哪些指标异常是训练崩溃的前兆。