Skip to content

训练工程

写出一个模型 forward 只是起点——真正把它”训起来”且训得稳、训得快、训得省显存,靠的是训练工程(Training Engineering)。本页把训练循环的每个工程组件讲透:Dataset/DataLoader、完整训练循环、梯度累积、超参调优、混合精度与梯度检查点、分布式训练(DDP/FSDP/DeepSpeed ZeRO)。前置阅读:PyTorch 入门、梯度下降与优化器、混合精度训练、分布式训练、学习率调度策略。

把训练工程想象成”开一家工厂流水线”:

  • Dataset/DataLoader= 原料仓库 + 传送带。仓库负责按索引取货(Dataset),传送带负责打包、排序、按批次送货到车间(DataLoader)。
  • 训练循环= 车间的标准作业程序(SOP):取料→加工→检验→调整机器→下一批。每个 epoch 把全部原料跑一遍。
  • 梯度累积= 显存太小装不下大订单,就分几次小批量加工,攒够再统一调机——模拟大 batch 效果。
  • 超参调优= 找最佳工艺参数(温度、压力、速度)。盲目试错慢,用贝叶斯优化”聪明地猜”。
  • 混合精度= 粗加工用半精度(快、省料),关键测量用全精度——在精度与速度间找平衡。
  • 梯度检查点= 车间仓库太小放不下所有半成品,就边做边扔、用到时再重算——用算力换显存。
  • 分布式训练= 开多家分厂同步生产。数据并行是各做各的零件、最后对账;模型并行是把一台大机器拆开分给多家。

PyTorch 的数据加载分两层:

  • Dataset:定义”如何取一条数据”。你继承它,实现 __len__(数据总量)和 __getitem__(按下标取一条样本)。
  • DataLoader:定义”如何把数据组成 batch 送往 GPU”。它包装 Dataset,负责批处理、打乱、多进程加载、自动批拼装(collate)。
from torch.utils.data import Dataset, DataLoader
import torch
# 1. Map-style Dataset(最常用,支持按下标随机访问)
class TextDataset(Dataset):
def __init__(self, path):
self.samples = [line.strip() for line in open(path)]
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
text = self.samples[idx]
input_ids = self.tokenize(text) # 自定义分词
label = self.get_label(text)
return {"input_ids": input_ids, "label": label}
# 2. IterableDataset(流式,数据太大放不下内存/硬盘时用)
class StreamDataset(Dataset):
def __init__(self, fileobj):
self.fileobj = fileobj
def __iter__(self):
for line in self.fileobj:
yield self.tokenize(line)
# 3. 直接用 TensorDataset(数据已在内存里)
ds = torch.utils.data.TensorDataset(X_tensor, y_tensor)

大模型场景:预训练语料动辄几百 GB 到几 TB,必须用流式读取(IterableDataset + memory-mapped 文件如 .bin/.npy/webdataset),绝不能一次性 load 进内存。

loader = DataLoader(
ds,
batch_size=32,
shuffle=True, # 训练集打乱,验证集不打乱
num_workers=4, # 多进程预取,加速 I/O
pin_memory=True, # 锁页内存,加速 CPU→GPU 传输
drop_last=True, # 丢弃不完整的最后一个 batch(训练时常见,保证 batch 大小一致)
collate_fn=custom_fn, # 自定义如何把多条样本拼成 batch
)

几个容易踩坑的点:

  • num_workers:设大能并行预取,但太多会耗尽内存、增加进程切换开销。经验值是每 GPU 2–8 个。Windows/Mac 上多进程可能不稳定,调试时先设 0。
  • pin_memory=True:把数据放在锁页内存(Page-Locked Memory)里,GPU 可以用 DMA 直接读取,省一次 CPU 拷贝。几乎总是该开。
  • collate_fn:当样本是变长序列(如文本)时,默认 collate 无法 stack。需自定义:对变长部分 pad 到等长,并返回 attention mask。
# 变长文本的 collate:pad 到 batch 内最大长度
def collate_fn(batch):
input_ids = [torch.tensor(x["input_ids"]) for x in batch]
labels = torch.tensor([x["label"] for x in batch])
# pad_sequence 自动补 0 到等长
input_ids = torch.nn.utils.rnn.pad_sequence(input_ids, batch_first=True, padding_value=0)
return input_ids, labels
  • sampler:控制采样顺序。类别不平衡时用 WeightedRandomSampler 给少数类更高采样权重;分布式训练用 DistributedSampler 保证各卡数据不重叠。

一个生产级的训练循环远不止 loss.backward(); optimizer.step()。下面是一个包含梯度裁剪、EMA、梯度累积、checkpoint 的完整骨架:

import torch, os, math
def train_one_epoch(model, loader, optimizer, scheduler, scaler, ema, args):
model.train()
optimizer.zero_grad(set_to_none=True) # set_to_none=True 比 zero_() 更快
accum_steps = args.grad_accum_steps # 梯度累积步数
for step, batch in enumerate(loader):
batch = {k: v.to(args.device, non_blocking=True) for k, v in batch.items()}
# --- 混合精度 forward ---
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
outputs = model(**batch)
loss = outputs.loss / accum_steps # loss 缩放(梯度累积)
# --- 混合精度 backward ---
scaler.scale(loss).backward()
# --- 每 accum_steps 步做一次参数更新 ---
if (step + 1) % accum_steps == 0:
# 梯度裁剪(防爆炸)
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), args.max_grad_norm)
# 参数更新
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad(set_to_none=True)
# 学习率调度(按 optimizer step 而非 micro-step)
scheduler.step()
# EMA 更新(指数滑动平均)
if ema is not None:
ema.update(model)
# --- 定期保存 checkpoint ---
if (step + 1) % args.save_steps == 0:
save_checkpoint(model, optimizer, scheduler, step, args)
def save_checkpoint(model, optimizer, scheduler, step, args):
ckpt = {
"model": (ema.get_model() if ema else model).state_dict(),
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
"step": step,
}
path = os.path.join(args.output_dir, f"ckpt-{step}.pt")
torch.save(ckpt, path)

梯度爆炸(Gradient Explosion)是深层网络和 RNN/Transformer 训练的头号杀手。梯度裁剪把梯度的全局范数限制在一个上限内:

if ∥∇L∥2>gmax⁡,∇L←gmax⁡∥∇L∥2⋅∇L\text{if } \|\nabla \mathcal{L}\|_2 > g_{\max}, \quad \nabla \mathcal{L} \leftarrow \frac{g_{\max}}{\|\nabla \mathcal{L}\|_2} \cdot \nabla \mathcal{L}

这叫按范数裁剪(clip by norm,clip_grad_norm_),比”按值裁剪”(clip by value,逐元素截断)更常用,因为它保留了梯度方向。gmax⁡g_{\max} 通常取 1.0。大模型训练几乎必开梯度裁剪。

EMA 维护一份参数的滑动平均副本,推理时用这份平均参数(往往比直接用最后一刻的参数更稳、泛化更好):

θˉt=α⋅θˉt−1+(1−α)⋅θt\bar{\theta}_{t} = \alpha \cdot \bar{\theta}_{t-1} + (1-\alpha) \cdot \theta_t

其中 α∈[0.99,0.9999]\alpha \in [0.99, 0.9999] 是衰减率。EMA 在扩散模型(Diffusion Models)训练中几乎是标配,在 LLM 训练中也常用。它的本质是对参数轨迹做低通滤波,平滑掉训练后期的抖动。

class EMA:
def __init__(self, model, decay=0.9999):
self.decay = decay
self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}
@torch.no_grad()
def update(self, model):
for k, v in model.state_dict().items():
if v.dtype.is_floating_point:
self.shadow[k].mul_(self.decay).add_(v, alpha=1 - self.decay)
def get_model(self):
# 返回用 EMA 参数的模型副本(用于推理/保存)
...
  • 保存内容:模型权重、优化器状态、学习率调度器状态、当前 step/epoch、随机数种子。这样断点可精确续训。
  • 保存频率:按步数(如每 1000 步)或按 epoch。大模型训练崩溃常见,务必勤存。
  • 保留策略:只保留最近 N 个 + 性能最好的几个,避免磁盘爆满。
  • 分布式:DDP/FSDP 下只需 rank 0 保存,避免多卡重复写盘。

3. Gradient Accumulation(梯度累积)

Section titled “3. Gradient Accumulation(梯度累积)”

显存不够装不下大 batch?梯度累积(Gradient Accumulation)让你用小显存模拟大 batch 的训练效果。

PyTorch 的梯度是累加的——backward() 不会清零梯度,而是叠加到 .grad 上(直到你调用 zero_grad())。利用这一点,可以把一个大 batch 拆成 KK 个 micro-batch,依次前向+反向,梯度自然累加,第 KK 次再调 optimizer.step()。这等价于用 batch size =K×micro_batch= K \times \text{micro\_batch} 训练。

数学上,设 micro-batch 上的损失为 ℓi\ell_i,则累积后的总损失为:

Leffective=1K∑i=1Kℓi\mathcal{L}_{\text{effective}} = \frac{1}{K}\sum_{i=1}^{K} \ell_i

这正是大 batch 的平均损失。因此梯度累积在数学上精确等价于大 batch(假设 BatchNorm 不跨 micro-batch 统计——这是唯一的细微差异)。

梯度累积示意图:4 个 micro-batch 累积成一次参数更新

accum_steps = 4
optimizer.zero_grad()
for i, batch in enumerate(loader):
loss = model(batch) / accum_steps # ① loss 除以累积步数
loss.backward() # ② 梯度自然累加
if (i + 1) % accum_steps == 0:
clip_grad_norm_(model.parameters(), 1.0)
optimizer.step() # ③ 等价于大 batch 的一步
optimizer.zero_grad()

三个关键细节:

  1. loss 要除以 accum_steps——否则梯度被放大 KK 倍,等效 batch 的损失应该是平均而非求和。
  2. 学习率调度按 optimizer step 计——不是按 micro-batch 计,否则调度会被稀释 KK 倍。
  3. BatchNorm 警告——BatchNorm 的统计量是按 micro-batch 算的,micro-batch 太小(如 1–2)会让统计量方差极大。大 batch 模拟失效。解决:用 GroupNorm/LayerNorm,或冻结 BN 统计量。

大模型实践:LLaMA/Qwen 训练时,单卡 micro-batch 可能只有 1–4 条序列,靠 grad_accum=16~64 拼出有效 batch 数千。这是消费级硬件训大模型的核心技巧。

超参数(Hyperparameter)是训练前设定、不参与梯度更新的参数(如学习率、batch size、正则系数、网络层数)。找最佳超参组合的过程叫超参调优(Hyperparameter Optimization, HPO)。

策略思想优点缺点
Grid Search(网格搜索)枚举所有组合简单、可并行维度灾难,组合数指数爆炸
Random Search(随机搜索)随机采样组合比网格更高效(Bergstra 2012 证明)无学习,纯靠运气
Bayesian Optimization(贝叶斯优化)用代理模型拟合”超参→性能”函数,选最可能更好的点试样本高效,少试几次串行性强,难大规模并行
Population / 进化算法维护一群超参,交叉变异淘汰适合非凸、多目标收敛慢

Random Search 为什么常胜 Grid Search?因为多数情况下只有少数超参重要(低有效维度)。网格搜索在无效维度上浪费大量试验,而随机搜索在每个维度上都均匀探索,更可能命中重要维度的好的取值。

贝叶斯优化用高斯过程(Gaussian Process, GP)或TPE(Tree-structured Parzen Estimator)建模目标函数 f(超参)=验证集性能f(\text{超参}) = \text{验证集性能},用采集函数(Acquisition Function)决定下一个试验点——在”探索未知区域”与”利用已知好区域”间权衡。

  • EI(Expected Improvement):选最可能改进当前最优的点。
  • UCB(Upper Confidence Bound):f(x)+κ⋅σ(x)f(x) + \kappa \cdot \sigma(x),平衡均值与不确定性。

数学上,给定已观测数据 D\mathcal{D},GP 给出每点的后验均值 μ(x)\mu(x) 和方差 σ2(x)\sigma^2(x),EI 定义为:

EI(x)=E[max⁡(f(x)−f∗,0)]=∫f∗∞(f−f∗) p(f∣x,D) df\text{EI}(x) = \mathbb{E}\left[\max(f(x) - f^*, 0)\right] = \int_{f^*}^{\infty} (f - f^*)\, p(f \mid x, \mathcal{D})\, df

每次试验后更新 GP,迭代推进。贝叶斯优化适合”单次试验昂贵”的场景(如训练一个大模型)。

Optuna 是当下最流行的 HPO 框架,默认用 TPE 采样器,API 极简:

import optuna
def objective(trial: optuna.Trial):
# 定义搜索空间
lr = trial.suggest_float("lr", 1e-5, 1e-2, log=True)
hidden = trial.suggest_categorical("hidden", [64, 128, 256])
dropout = trial.suggest_float("dropout", 0.0, 0.5)
n_layers = trial.suggest_int("n_layers", 1, 4)
model = build_model(hidden, dropout, n_layers)
val_score = train_and_eval(model, lr) # 训练并返回验证集分数
return val_score # 最大化
study = optuna.create_study(direction="maximize",
sampler=optuna.samplers.TPESampler(seed=42))
study.optimize(objective, n_trials=50, n_jobs=1) # n_jobs 控制并行试验数
print(study.best_params, study.best_value)

Optuna 还支持:剪枝(Pruning,提前终止差的试验)、可视化(参数重要性、优化历史)、分布式试验。

大模型单次训练动辄几百 GPU 小时,跑几十组超参不现实。实践做法:

  • 先在小模型/小数据上调,再把超参迁移到大模型(学习率通常按 sqrt(scaling) 调整)。
  • 只调最关键的几个:学习率、warmup 步数、权重衰减、dropout。其余用社区默认值。
  • Population-Based Training(PBT):把进化算法嵌入训练,多组并行训、差的向好的”复制变异”。Ray Tune 支持。

5. 混合精度训练回顾 + Gradient Checkpointing

Section titled “5. 混合精度训练回顾 + Gradient Checkpointing”

混合精度用 FP16/BF16 做 forward/backward(快、省显存),用 FP32 累积梯度和更新参数(精度)。核心机制是损失缩放(Loss Scaling)防止 FP16 小梯度下溢。详见 混合精度训练 的完整推导。

2024–2025 年的关键趋势:BF16(bfloat16)逐渐取代 FP16 成为默认。BF16 与 FP32 同为 8 位指数,动态范围相同,无需损失缩放,大幅降低训练不稳定性。现代 GPU(A100/H100/Blackwell)和 TPU 都原生支持 BF16。PyTorch 里只需:

with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
outputs = model(**batch)
loss = outputs.loss
loss.backward()

优先用 BF16;只在老硬件(V100 及更早,不支持 BF16)才用 FP16 + 损失缩放。

Gradient Checkpointing(梯度检查点)

Section titled “Gradient Checkpointing(梯度检查点)”

反向传播需要保留前向所有中间激活值(用于链式法则求导)。深层网络激活值显存占用极大——这是大模型训练的主要显存瓶颈。Gradient Checkpointing(又叫 Activation Checkpointing / Recomputation)的思路:前向时不保存部分中间激活,反向需要时重新前向计算一次——用一次额外前向的算力换大幅显存节省。

数学上,标准反向的显存占用 O(L)O(L)(LL 为层数),检查点后降至 O(L)O(\sqrt{L})(按层分块,只保存检查点层的激活),计算开销增加约 33%(一次额外前向)。

from torch.utils.checkpoint import checkpoint
class TransformerBlock(torch.nn.Module):
def forward(self, x):
return self.layer(x)
class MyModel(torch.nn.Module):
def forward(self, x):
for block in self.blocks:
# use_reentrant=False 是 PyTorch 新推荐(更稳、支持更多算子)
x = checkpoint(block, x, use_reentrant=False)
return x

显存三件套:对大模型训练,几乎总是组合使用 混合精度 + 梯度检查点 + 梯度累积。三者叠加,单卡能训的模型规模可放大数倍。

单卡训不动大模型时,必须分布式。三大范式:数据并行(每卡完整模型、不同数据)、模型并行(模型拆开分到多卡)、混合(FSDP/ZeRO/3D 并行)。详见 分布式训练 与 DeepSpeed/FSDP 指南 的系统论述,这里给出上手代码。

DDP(DistributedDataParallel)——数据并行入门

Section titled “DDP(DistributedDataParallel)——数据并行入门”

DDP 是最简单、最常用的分布式方式:每张卡持有完整模型,各自处理不同数据,反向时用 All-Reduce 同步梯度。

import torch, torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP # 仅示意
from torch.nn.parallel import DistributedDataParallel as DDP
dist.init_process_group(backend="nccl")
torch.cuda.set_device(dist.get_rank())
model = MyModel().cuda()
model = DDP(model, device_ids=[dist.get_rank()])
# 注意:DataLoader 要用 DistributedSampler,保证各卡数据不重叠
sampler = torch.utils.data.distributed.DistributedSampler(dataset, shuffle=True)
loader = DataLoader(dataset, batch_size=per_gpu_bs, sampler=sampler)
for epoch in range(epochs):
sampler.set_epoch(epoch) # 关键!保证每 epoch 数据打乱方式不同
for batch in loader:
loss = model(batch)
loss.backward()
optimizer.step()
optimizer.zero_grad()

DDP 的瓶颈:每卡都要装下完整模型 + 优化器状态 + 梯度 + 激活。对 7B+ 的模型,单卡 80GB 显存也吃不消。

FSDP 把模型参数、梯度、优化器状态分片(shard)到各卡,用到时再临时 All-Gather 拼回。它本质上就是 ZeRO-3 的 PyTorch 原生实现,是大模型训练的主流方案。

from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import ShardingStrategy, MixedPrecision
mp_policy = MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
buffer_dtype=torch.bfloat16,
)
model = MyModel().cuda()
model = FSDP(
model,
sharding_strategy=ShardingStrategy.FULL_SHARD, # = ZeRO-3
mixed_precision=mp_policy,
use_orig_params=True, # 兼容梯度检查点
activation_checkpointing=True,
)
# 后续用法与普通 model 一致

FSDP 的分片层级对应 ZeRO 的三阶段:

  • ZeRO-1:分片优化器状态(省 4x 显存)。
  • ZeRO-2:再分片梯度(省更多)。
  • ZeRO-3 / FULL_SHARD:再分片参数(最省显存,但通信最多)。

DeepSpeed 通过 JSON 配置驱动,与 PyTorch 训练循环解耦:

{
"train_micro_batch_size_per_gpu": 2,
"gradient_accumulation_steps": 16,
"optimizer": {
"type": "AdamW",
"params": { "lr": 2e-5, "weight_decay": 0.01 }
},
"fp16": { "enabled": "auto" },
"bf16": { "enabled": true },
"zero_optimization": {
"stage": 2,
"offload_optimizer": { "device": "cpu" }
},
"gradient_clipping": 1.0,
"activation_checkpointing": { "partition_checkpointing": true }
}
import deepspeed
model_engine, optimizer, _, _ = deepspeed.initialize(
model=model, model_parameters=model.parameters(), config="ds_config.json"
)
for batch in loader:
loss = model_engine(batch)
model_engine.backward(loss)
model_engine.step() # DeepSpeed 内部处理梯度累积、裁剪、step、zero_grad

DeepSpeed 的优势:配置驱动、支持 CPU/NVMe Offload(把优化器状态卸载到 CPU 内存或 SSD,进一步突破显存墙)、集成 1F1B 流水线并行。劣势:抽象较重,调试不如原生 PyTorch 直观。

场景推荐方案
单卡能放下模型(<7B 训练、微调)单卡 + AMP + 梯度累积
多卡、模型放得下、要加速DDP(最简单)
多卡、模型放不下(7B+ 全量训练)FSDP 或 DeepSpeed ZeRO-2/3
超大模型(70B+)多机FSDP + 流水线并行 + 张量并行(3D 并行),或 Megatron-LM/torchtitan
消费级显卡微调LoRA + 8-bit 优化器 + 梯度检查点(详见 量化)

训练工程的各个组件不是孤立的——学习率调度(详见 学习率调度策略)与 batch size、梯度累积强耦合。下图对比四种主流调度策略(均带 warmup)的曲线形态:

四种学习率调度策略对比(含 warmup)

工程经验法则:

  • 增大 batch size 时,学习率通常要相应增大。经验公式(线性缩放规则):ηnew=ηbase×batchnewbatchbase\eta_{\text{new}} = \eta_{\text{base}} \times \frac{\text{batch}_{\text{new}}}{\text{batch}_{\text{base}}}(大 batch 时要 cap,且 warmup 要更长)。
  • WSD 调度(Warmup-Stable-Decay)在 2024–2025 年大模型中流行:长时间稳定学习率训练,末期快速衰减。优点是”训到一半决定要不要继续”很灵活——stable 段可任意延长。
  • 梯度累积不影响有效学习率,但影响 optimizer step 数,调度器的 total_steps 要按 optimizer step 算。
  1. 先在小规模跑通,再 scale up:小数据小模型验证循环正确性(loss 下降、梯度不爆),再上大数据大模型。
  2. 永远监控梯度范数:训练日志里记 grad_norm、loss、lr、tokens/s。梯度范数突增是爆炸的前兆,loss spike 是不稳定信号。
  3. 混合精度优先 BF16,省去损失缩放的麻烦。
  4. 显存三件套常备:混合精度 + 梯度检查点 + 梯度累积。三者叠加才玩得转大模型。
  5. 勤存 checkpoint,分布式训练崩溃是常态,从断点续训能省巨额成本。
  6. 调参优先级:学习率 > warmup > batch size > 权重衰减 > dropout。别在不重要的超参上浪费时间。
  7. DDP 用 DistributedSampler 并 set_epoch,否则各卡数据打乱方式固定,模型会过拟合到固定顺序。
  8. zero_grad(set_to_none=True) 比 zero_() 快,因为省去内存写入。
  • BF16 成为绝对主流,FP16 + 损失缩放基本退出大模型训练;新一代 Blackwell GPU 原生支持 FP8(见 混合精度训练 的 FP8 章节)。
  • FSDP2 重写:PyTorch 用基于 per-parameter 的新 FSDP 实现替代旧版,内存效率更高、与 torch.compile 兼容性更好、调试更直观,2025–2026 年逐步成为默认。
  • torchtitan:PyTorch 官方的轻量预训练参考实现,提供 Llama/GPT 类模型的完整 FSDP2 + 张量并行 + FP8 训练配方,是学习工业级训练工程的最佳起点。
  • FP8 训练普及:H100/Blackwell 上 FP8 训练吞吐较 BF16 再提 ~2x,但需精细的 E4M3/E5M2 分配与张量级缩放因子管理。
  • 3D 并行成熟:数据并行 × 流水线并行 × 张量并行的组合成为千亿模型标配;Megatron-LM、torchtitan、DeepSpeed 提供开箱即用配方。
  • 离线/异步并行:如 DeepSpeed 的 ZenFlow、异步卸载引擎,让通信与计算 overlap,逼近线性扩展。
  • 超参搜索与训练耦合:Ray Tune + Optuna 支持大规模异步 HPO;PBT 及其变体(如 ASHA 早停调度)让大模型 HPO 在合理预算内可行。
  • 长序列训练:百万级 token 上下文训练(如 ALST、Ring/USB Attention)对显存与通信提出新挑战,催生新的激活重计算与序列并行(Sequence Parallelism)技术。