ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

模型训练调参:batch、batch size、epoch、iteration

模型训练调参:batch、batch size、epoch、iteration 刚接触模型训练的人十个里有八个会被 batch、batch size、epoch、iteration 这四个词绕晕。它们看着像是同一类东西实际各自描述的是训练节拍里完全不同的刻度一个说的是把数据切成了多大一块一个说的是这块有多大另外两个则是两种不同层级的计数器。我见过不少朋友代码能跑起来loss 也在降但被问到你现在总共跑了多少个 iteration时答不上来也见过有人把 batch size 从 32 改到 256结果学习率没动训练直接不收敛回头怀疑是模型结构有问题。这些坑我在早期做图像分类和后来的大模型微调时都踩过所以想把这套东西系统地捋一遍。这篇内容面向的是所有正在或准备训练模型的人刚入门想搞清楚概念的新手、跑通了 demo 但调参靠感觉的进阶者、以及正在用微调框架做大模型适配的工程同学。我会从最基础的概念拆起把四个词之间的换算关系、显存账本怎么算、batch size 和学习率怎么联动、梯度累积为什么不完全等于大 batch、以及一大批实际故障的排查思路都讲清楚。里面所有代码和配置都可以直接抄走改参数计算过程我会把公式和推导摊开写而不是丢一堆结论让你自己猜。顺便说一句batch 这个词在别的软件里还有批处理的意思比如批量导出、批量扫描向导之类跟模型训练里的 batch 是两码事别被搜索结果里的同名词带偏。1. 把四个词摆到桌面上它们各自描述训练的哪个刻度1.1 batch一次参数更新用到的那一撮样本先把最容易被含糊过去的定义钉死。batch批次指的是在一次参数更新中同时送入模型计算的那一组样本。假设你手上有 50000 张图片作为训练集不可能一次性全部塞进显卡也不可能一张一张算——前者显存装不下后者没法利用 GPU 的并行能力而且单样本的梯度噪声大到几乎没法收敛。所以工程上的做法是把数据集切成很多小块每一小块就是一个 batch模型吃一块、算一次损失、反传一次梯度、更新一次参数然后再吃下一块。这里有个容易混淆的点batch 是一个集合是这一撮样本的整体而 batch size 是一个数值描述这一撮里有多少个样本。日常口语里大家经常混用batch 设成 64其实是batch size 设成 64的省略说法写代码的时候batch_size64才是准确的参数名。那为什么要切块而不是别的方式核心原因有三个。一是显存约束模型的前向激活值、梯度缓存都跟一次处理的样本量成正比切块是唯一能把大模型塞进有限显存的手段。二是梯度质量全量数据算出来的梯度最准但太慢单样本最快但太抖batch 是在算得准和跑得快之间找的平衡点。三是硬件效率GPU 的并行单元需要足够的计算量才能跑满batch 太小的时候计算单元大量空转吞吐量会断崖式下跌。1.2 batch size一口到底吃多少样本batch size 就是一个 batch 里包含的样本数量通常记作 B。这个数字是你训练配置里最需要反复权衡的超参数之一因为它同时牵动三件事显存占用、训练速度、以及最终模型的泛化表现。显存方面batch size 和激活值显存基本是线性关系——B 翻倍中间层激活占的显存大致也翻倍在不开梯度检查点的前提下。速度方面从 B1 涨到 B32 时吞吐量提升非常明显因为 GPU 终于有活干了但从 B512 涨到 B1024提升就平缓了因为计算单元已经接近饱和。泛化方面则更微妙大量实验观察到一个现象过大的 batch size 会让模型倾向于收敛到尖锐的极小值测试集表现反而不如中等 batch这个现象通常被叫做泛化间隙。实践中我会把 batch size 分成两种来管理per-device batch size单卡批量和global/effective batch size全局有效批量。前者是每张卡上每次实际吃多少后者是一次参数更新实际上等效用了多少样本也就是单卡批量乘以卡数、再乘以梯度累积步数。多卡训练时真正决定优化行为的是全局有效批量单卡批量只决定显存能不能扛住。这个区分在后面讲梯度累积时会变得非常重要。1.3 iteration 和 epoch两个层级的计数器iteration迭代也叫 step指的是参数更新了一次。一个完整的 iteration 包含四步取一个 batch 的数据、前向传播算出预测、反向传播算出梯度、优化器根据梯度更新参数。所以跑了多少个 iteration就等于参数被更新了多少次这是最细粒度的进度刻度。epoch轮次指的是整个训练集被完整遍历了一遍。注意这里的关键词是完整遍历——不是参数更新了几次而是所有样本都被模型看过至少一次。这是比 iteration 高一层的刻度反映的是数据被消化了几轮。为什么需要两个刻度因为它们回答的是不同的问题。iteration 回答训练跑了多久、还要跑多久用来控制日志频率、学习率调度、checkpoint 保存节奏epoch 回答这套数据我反复学了几遍用来判断是否过拟合。你可以只跑 500 个 iteration 而不关心跑了几轮也可以跑 10 个 epoch 而不关心中间更新了多少次两种描述方式各有适用场景。举个实际例子帮你建立体感。小数据集比如几千条做微调一个 epoch 可能只有几十个 iteration这时候大家习惯按 epoch 说事而大模型预训练动辄几百万条数据、几十张卡一个 epoch 要跑几万甚至几十万个 iteration这时候大家一律按 step 说事训练到 10000 步比训练到 0.3 个 epoch直观得多。所以不要把 epoch 当成唯一标准看场景选刻度。1.4 四者的换算关系与实例推演把关系写成公式就非常清楚了。设训练集样本总数为 Nbatch size 为 B训练轮数为 E则每个 epoch 的 iteration 数steps_per_epoch N / B取整方式取决于是否使用 drop_last总 iteration 数total_steps steps_per_epoch × E关于取整这里有个细节值得展开。如果 N50000、B64那么 50000/64 781.25不是整数。此时有两种处理方式向上取整得到 782 个 iteration最后一个 batch 只有 16 个样本这叫drop_lastFalse向下取整得到 781 个 iteration最后那 16 个样本直接丢掉这叫drop_lastTrue。两种都能用但我一般建议在图像任务里用 drop_lastTrue理由后面讲 BatchNorm 的时候会说。再补一个多卡场景的公式。假设用 4 张卡单卡 batch size 为 8梯度累积步数为 4那么全局有效 batch size 是8 × 4 × 4 128。此时每个 epoch 的优化器更新次数是N / 128而不是N / 8——因为每 4 个 micro-batch 才凑成一次真正的参数更新。很多人在多卡训练时算错学习率调度的总步数就是漏了梯度累积这一项。我可以再给一个真实的排查场景。有位朋友报告说他配置了num_train_epochs3但训练日志显示总步数只有 600 多他感觉跑得太少。一查发现数据 24000 条单卡 batch size 是 4梯度累积 88 张卡。全局有效批量 4×8×8 256每个 epoch 的更新次数 24000/256 ≈ 94三个 epoch 约 282 步。日志里的 600 多是因为包含了验证和日志额外的计数。这个例子说明光看 epoch 数是判断不出训练量的必须把 batch size 和累积步数一起算进去。术语英文描述对象常用记法作用批次batch一组样本的集合-切分数据的基本单位批量大小batch size数值B控制显存、速度、梯度质量迭代iteration / step计数-参数更新的次数轮次epoch计数E数据集被完整遍历的次数2. batch size 怎么定显存、速度与收敛的三角平衡2.1 显存账本参数、梯度、优化器状态、激活值要选 batch size先得知道显存都被谁吃了。训练时的显存开销可以拆成四块我把它们按是否随 batch size 变化分成两类来讲。第一类是与 batch size 无关的固定开销包括模型参数、参数梯度、优化器状态。以一个参数量为 P 的模型为例fp32 训练时参数占 4P 字节梯度占 4P 字节如果优化器是 Adam还需要保存一阶动量 m 和二阶动量 v各占 4P 字节。合起来是 16P 字节。听起来不多换算一下一个 7B 参数的模型P 7×10就是 112 GB——单张消费级卡基本没戏这就是为什么大模型训练离不开分片优化器状态、LoRA 这类手段。如果是 SGD 且不开动量就只有 12P 字节。第二类是随 batch size 线性增长的激活值包括每一层前向计算的中间结果因为反向传播时要用到它们算梯度。这部分没法用一个精确公式概括跟网络深度、序列长度、隐藏维度、算子实现都有关系。工程上有个粗略的估算方式激活显存 ≈ B × L × S × H × k其中 B 是 batch sizeL 是层数S 是序列长度图像任务可以理解成特征图的空间尺寸H 是隐藏维度k 是一个跟具体实现相关的系数经验范围大概在 8 到 20 之间取决于是否使用混合精度、算子融合程度等。这个式子只能用来做量级判断真要精确值还是得实测。用这个思路做个对比就清楚多了。同一个 ResNet-50约 25M 参数输入 224×224B32 时参数梯度SGD 动量约 300 MB激活值大约几百 MB总占用 1 GB 出头B128 时固定部分不变激活值涨到约 4 倍总占用可能到 3-4 GB这个差别解释了为什么在小模型上把 batch size 拉大几乎无感而在大模型上每加一点都要精打细算。注意激活值显存是峰值概念它出现在反向传播开始的那一刻。所以看显存不能看训练稳定后的平均值要用torch.cuda.max_memory_allocated()这类接口看峰值否则很容易在某个特定 step 突然 OOM。2.2 小 batch 与大 batch 的取舍不是越大越好batch size 的选择本质上是三方的博弈我把每一方的诉求列清楚。小 batch比如 8、16的优势梯度里带的噪声大这个噪声其实是有益的它相当于给优化过程加了一点随机扰动能帮模型跳出尖锐的局部极小值最终泛化往往更好显存压力小能上更大的模型或者更长的序列对数据分布的适应性更强因为每次看到的样本组合都不同。代价训练慢同样的 epoch 数需要更多 wall-clock 时间梯度抖动大loss 曲线看起来毛毛躁躁BatchNorm 的统计量估计不准小 batch 时均值和方差的噪声很大可能拖累收敛。大 batch比如 512、1024的优势GPU 利用率高单位时间的样本吞吐量大梯度估计更准loss 曲线平滑分布式训练时通信效率更高因为每次通信传输的信息量更大。代价显存占用高需要配合更大的学习率调参难度上升泛化间隙问题对数据加载速度的要求更高很容易出现GPU 等数据的瓶颈。我的经验做法是先找到显存能扛住的最大 batch size然后往下调一到两档。比如实测 B256 刚好不 OOM那就从 128 开始试而不是顶着上限用 256。留出的这档余量一是应对训练过程中激活值波动某些 step 的数据可能更长、更复杂二是给泛化留点余地。2.3 学习率与 batch size 的联动线性缩放和 warmup这是最容易出事的地方。改动 batch size 而不改学习率是训练不收敛的头号原因之一。背后的逻辑是batch size 变大意味着梯度估计的方差变小梯度的方向更可信所以你可以放心地朝这个方向迈更大的步子。反过来batch size 变小梯度噪声大步子迈大了容易直接跨过最优解。最常用的调整规则是线性缩放linear scalinglr_new lr_base × (B_new / B_base)比如你在 B32 时用 lr0.1 效果不错现在改成 B256那么 lr 大致应该设成 0.1 × 8 0.8。但这个规则在 batch size 特别大时会失效通常超过几千之后此时更稳妥的是用平方根缩放lr_new lr_base × sqrt(B_new / B_base)B256 时就是 0.1 × 2.83 ≈ 0.283。两种规则我都用过实践中我的判断是batch size 在 8 倍以内变化时线性缩放更准超过之后往平方根缩放靠拢。还有一个配套动作是warmup学习率预热。当你用大 batch 大学习率时训练刚开始的那几百个 step 最容易震荡甚至发散因为模型参数是随机的此时的大梯度配合大学习率很容易把参数推到很糟糕的区域。warmup 的做法是让学习率从一个小值比如峰值的 1% 或者直接 0在前若干个 step 内线性升到目标值。常见的配置是 warmup 占总步数的 3%-10%大 batch 场景下取上限。# 线性 warmup 余弦退火的调度示例 import math def get_lr(step, total_steps, base_lr, warmup_ratio0.05, min_lr_ratio0.1): warmup_steps int(total_steps * warmup_ratio) if step warmup_steps: # 从 0 线性升到 base_lr return base_lr * (step 1) / warmup_steps # 余弦退火从 base_lr 降到 base_lr * min_lr_ratio progress (step - warmup_steps) / max(1, total_steps - warmup_steps) cos 0.5 * (1 math.cos(math.pi * progress)) return base_lr * (min_lr_ratio (1 - min_lr_ratio) * cos)这个函数可以直接嵌进你的训练循环每个 step 调用一次然后写回optimizer.param_groups[0][lr]。2.4 一套可复用的 batch size 选型流程把上面的东西串成一套流程我平时就是按这个顺序走的。第一步确定显存预算。用nvidia-smi看总显存扣掉 CUDA 上下文、cuDNN 工作空间这些固定开销一般留 1-2 GB 余量剩下的才是可用空间。第二步用小 batch 跑通一个完整的 iteration确认前向、反向、优化器更新都没有报错此时记录一下基础显存占用。第三步逐步翻倍试探。从 B8 开始8 → 16 → 32 → 64 → 128每加一档跑 20 个 iteration 并记录峰值显存。当峰值超过预算的 85% 时停下取上一档作为上限。第四步按规则调整学习率先线性缩放跑 200 个 step 看 loss 曲线是否稳定下降。如果出现震荡或者直接 nan就换成平方根缩放或者把学习率再乘 0.5。第五步确认数据管道跟得上。用nvidia-smi观察 GPU 利用率如果长期低于 70%说明瓶颈在数据加载而不是计算这时候加 batch size 意义不大先把num_workers加上去。这一步我吃过亏曾经为了追吞吐把 batch size 从 64 拉到 256结果吞吐只涨了 5%因为瓶颈一直在 CPU 侧的图像解码上。场景建议起始 batch size学习率策略备注小模型图像分类单卡32 - 1280.1 起步 warmup 5%BatchNorm 对 batch 敏感目标检测单卡8 - 160.01 起步 warmup输入分辨率大激活值吃显存Transformer 文本分类16 - 642e-5 到 5e-5预训练权重微调学习率要小大模型 LoRA 微调单卡 1 - 4 梯度累积1e-4 到 3e-4全局有效批量凑到 64-128大模型全参微调单卡 1 - 2 梯度检查点1e-5 到 2e-5固定 max_steps 而非 epoch3. 实操从零搭一个可观测的训练循环3.1 数据管道DataLoader 里那几个容易被忽略的参数概念讲完了落到代码上。数据管道这块DataLoader的几个参数直接影响 batch 相关行为的正确性我逐个说。batch_size就是本章的主角不用多解释。shuffleTrue在每个 epoch 开始前打乱样本顺序这个必须开——不开的话每个 epoch 的 batch 组合完全一样梯度噪声的分布会失去随机性收敛质量明显下降。drop_lastTrue丢掉最后一个不满员的 batch作用是保证每个 batch 的大小一致这对 BatchNorm 特别重要下面单独讲。num_workers是子进程数量负责在后台预取数据经验值是 CPU 核心数的一半到全部我一般从 4 开始试不够再加。pin_memoryTrue让数据先放进锁页内存从 CPU 拷到 GPU 时能走更快的通道配non_blockingTrue一起用效果才完整。persistent_workersTrue让 worker 进程在 epoch 之间不销毁重建数据量大、epoch 多的时候能省下不少启动开销。prefetch_factor控制每个 worker 预取多少个 batch默认 2如果 GPU 经常饿着可以调到 4。from torch.utils.data import DataLoader train_loader DataLoader( train_dataset, batch_size64, shuffleTrue, num_workers8, pin_memoryTrue, drop_lastTrue, persistent_workersTrue, prefetch_factor4, ) steps_per_epoch len(train_loader) # drop_lastTrue 时已经是向下取整的结果 print(fsteps_per_epoch {steps_per_epoch})顺手提一下drop_last和 BatchNorm 的关系因为这是个真实且高频的坑。BatchNorm 在训练时用当前 batch 的均值和方差做归一化这个统计量只有在 batch 样本数足够时才靠谱。默认配置下 BatchNorm 要求每个通道至少有一定数量的样本如果最后一个 batch 只剩 2 个样本比如 N50002、B64BN 的统计量会极其不准梯度方向也会失真表现为训练后期偶尔出现的 loss 尖刺。drop_lastTrue直接把这个隐患消掉代价只是丢掉几个样本非常划算。注意如果你的模型全是 LayerNorm 而没有 BatchNorm现在的大模型基本都是这样drop_last的必要性就下降了因为 LayerNorm 是在单个样本内部做的归一化跟 batch 里有多少样本无关。这也是大模型在梯度累积时不受 batch 统计量影响的原因。3.2 手写训练循环把梯度累积写对下面是一个带梯度累积、混合精度、梯度裁剪的完整训练循环。这段代码我用了很久改改就能套到各种任务上。import torch import torch.nn as nn from torch.cuda.amp import autocast, GradScaler device torch.device(cuda) model model.to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr3e-4, weight_decay0.01) scaler GradScaler() accum_steps 4 # 梯度累积步数 global_batch 64 * accum_steps # 全局有效批量 epochs 20 for epoch in range(epochs): model.train() running_loss 0.0 optimizer.zero_grad(set_to_noneTrue) for i, (x, y) in enumerate(train_loader): x x.to(device, non_blockingTrue) y y.to(device, non_blockingTrue) with autocast(dtypetorch.float16): logits model(x) # 关键损失除以累积步数这样累积后的梯度量级才和真大 batch 一致 loss criterion(logits, y) / accum_steps scaler.scale(loss).backward() running_loss loss.item() * accum_steps if (i 1) % accum_steps 0: scaler.unscale_(optimizer) # 先还原梯度才能裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_noneTrue) if (i 1) % (accum_steps * 50) 0: print(fepoch {epoch} iter {i1} loss {running_loss / (accum_steps * 50):.4f}) running_loss 0.0 torch.save({model: model.state_dict(), epoch: epoch}, fckpt_epoch{epoch}.pt)这段代码里有三个点值得单独解释都是我曾经写错过的。第一损失为什么要除以 accum_steps。因为梯度是可加的。累加 N 个 micro-batch 的梯度等价于把它们的损失加在一起求导。如果你不做除法累加后的梯度就是真实大 batch 梯度的 accum_steps 倍相当于偷偷把学习率放大了四倍训练很容易发散。第二梯度裁剪必须在scaler.unscale_之后。混合精度训练时scaler.scale(loss).backward()会把梯度放大一个缩放因子来避免 fp16 下溢此时梯度值不是真实的。必须先 unscale 还原再裁剪再scaler.step()。顺序错了两件事都会失效。第三末尾残留的梯度怎么处理。如果steps_per_epoch不是accum_steps的整数倍循环结束时会有几个 micro-batch 的梯度累积着但没被使用下一个 epoch 开始时会跟新梯度混在一起相当于凭空多了一个奇怪的 batch。干净的做法是在每个 epoch 结束时判断一下# 处理 epoch 末尾未满 accum_steps 的残留梯度 if len(train_loader) % accum_steps ! 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_noneTrue)或者更简单直接让steps_per_epoch除以accum_steps取整多余的 micro-batch 跳过。3.3 日志与 checkpoint把 iteration 和 epoch 变成可读的信号训练跑起来之后你需要把抽象的数字变成能看懂的信息。我习惯记录这几个量。每个 step 记录当前学习率、当前 loss、梯度范数clip_grad_norm_的返回值、当前 GPU 显存峰值。梯度范数是诊断训练健康度的利器正常情况下它应该在一个相对稳定的范围内波动如果突然涨到几百甚至几千说明马上就要发散了可以据此把裁剪阈值调低。每 N 个 step 记录平滑后的平均 loss。原始 loss 抖动很大看单个值没意义至少要 50-100 个 step 的滑动平均。我一般用logging_steps10每 10 步打一条打印时对最近 10 条取平均。每个 epoch 记录验证集 loss 和任务指标准确率、F1、BLEU 等以及本次 epoch 的耗时。验证集 loss 的走势是判断过拟合的核心依据——训练 loss 还在降但验证 loss 开始上升就是过拟合的信号。checkpoint 保存有两条策略。按 epoch 存适合小数据集每轮存一个方便回溯按 step 存适合大模型训练因为一个 epoch 可能跑几天中途出事损失太大。保存内容上除了model.state_dict()我强烈建议把optimizer.state_dict()和当前 step 数一起存进去否则断点续训时优化器的动量状态全丢了loss 会出现一次明显的反弹。ckpt { model: model.state_dict(), optimizer: optimizer.state_dict(), scaler: scaler.state_dict(), epoch: epoch, global_step: global_step, best_val_loss: best_val_loss, } torch.save(ckpt, ckpt_last.pt)3.4 loss 曲线怎么读三种典型形态对应的病因训练日志拿到手怎么判断正不正常我把常见的曲线形态和对应病因整理一下。平稳下降最理想的情况训练 loss 和验证 loss 同步下降中间有小的波动但整体趋势清晰。说明 batch size、学习率、数据管道基本都配对了。剧烈震荡但整体下降曲线像心电图每个 step 的 loss 能差出好几倍。通常是 batch size 太小加上学习率偏大。可以先试把 batch size 翻倍或者把学习率减半看抖动是否收窄。另外要确认shuffleTrue开了没有——不开 shuffle 的话每个 batch 的样本组合固定如果数据本身按类别排序loss 会呈现出周期性的大起大落。先降后平再升这是典型的过拟合三段式。验证 loss 降到最低点之后开始回升而训练 loss 还在继续降两者分道扬镳。处理手段有三减少 epoch 数早停、加正则weight decay、dropout、加数据。这里要强调的是不是所有任务都需要训练到 loss 收敛验证集指标最好的那个 checkpoint 才是你要的后面的训练都是负收益。从头到尾在 2.3 附近横盘2.3 是ln(10)十分类任务的随机猜测水平。loss 一开始就卡在这个位置且完全不降八成是标签和输出没对齐、或者模型压根没在学。这时候先把学习率调大十倍试一次如果还是不动就要去查数据了——我遇到过的最离谱的一次是标签文件的行数和图片数量对不上导致所有样本都被配了错误的标签。4. 大模型微调场景batch 策略为什么完全不一样4.1 微调为什么通常只跑 1 到 3 个 epoch如果你从图像分类转到大模型微调第一个冲击就是epoch 数怎么这么少图像任务跑 50 甚至 100 个 epoch 都很正常而微调一个 7B 模型配置里写的经常是num_train_epochs: 2甚至1。原因有三层。第一层是数据量差异。图像分类的数据集可能只有几千张Mini-batch 走一遍也就几十步而微调用的指令数据集动辄几万到几十万条一个 epoch 就是几千到几万步模型见到的样本数其实并不少。第二层是预训练权重的起点。微调的模型已经在海量数据上预训练过了它需要的是适配而不是从零学习。模型参数的更新幅度很小学习率通常在 1e-5 到 3e-4 之间LoRA 可以高一些几轮下来就足够把风格和任务格式调整到位。第三层是过拟合风险。微调数据集相对于模型的参数量来说太小了尤其是几千条的高质量指令数据配 7B 模型跑三个 epoch 以上模型就开始死记硬背训练样本的措辞泛化能力下降。表现是训练 loss 很低但生成的内容开始重复、僵化。我的实操习惯是能固定步数就不用 epoch。把max_steps直接设成一个具体值比如 800、1500然后在中途每隔 200 步存一个 checkpoint最后拿几个 checkpoint 分别跑评测选最好的那个。这比训练 3 个 epoch 然后取最后一个科学得多因为最优解往往不在最后一站。4.2 单卡 batch size 只能设 1 或 2 时怎么办大模型微调的显存现实是这样的7B 模型做 LoRA序列长度 1024单卡 batch size 往往只能设到 4做全参微调的话单卡 1 都未必塞得下。但有效批量又需要足够大才能稳定收敛这个矛盾就靠梯度累积来解。在微调框架里通常有三个参数共同决定有效批量effective_batch per_device_train_batch_size × gradient_accumulation_steps × num_devices举例单卡 batch size 2累积步数 8用 4 张卡那么有效批量 2×8×4 64。这个量级对于大多数微调任务已经足够了。这里有个认知上的重要区别我在前面埋了伏笔现在展开梯度累积在数学上不完全等价于真实的大 batch。差异出在归一化层。如果模型里有 BatchNorm它在前向时用的是 micro-batch 的统计量累积再多步也改变不了这一点所以 BN 的统计量始终是按 micro-batch 算的跟真大 batch 的行为不同。而现在的 Transformer 架构全用 LayerNorm 或 RMSNorm是在单个样本内部做归一化跟 batch 里有多少样本无关所以梯度累积和真大 batch 的行为几乎完全一致。这也是为什么大模型场景下梯度累积可以用得这么放心。提示还有一个差异点是 dropout。如果模型有 dropoutmicro-batch 每次前向的随机 mask 不同累积后的梯度里包含了多组不同的 dropout 噪声这跟一次大 batch 单次前向的效果有细微差别。实践中这点差异可以忽略但如果你的 loss 曲线特别抖可以把 dropout 关掉对比一下。4.3 显存不够用的三板斧累积、检查点、混合精度除了梯度累积还有两个手段必须掌握三个配合起来基本能让你在单张消费级卡上微调 7B 模型。梯度检查点gradient checkpointing的思路是前向时只保存每一层的输入不保存中间激活值反向传播需要用到某层激活时临时重新算一遍。代价是计算量增加约 30%收益是激活值显存下降 60%-70%。开了它之后单卡 batch size 通常能从 1 提到 4。这是性价比最高的一个开关几乎所有微调框架里都有对应的配置项。混合精度bf16 / fp16让前向和反向计算用 16 位浮点做显存占用和计算时间都减半。优先用 bf16因为它的数值范围和 fp32 一致不需要损失缩放训练稳定性更好只有显卡不支持 bf16比如一些较老的架构时才退回 fp16此时必须配 GradScaler 防止梯度下溢。需要注意的是优化器里保存的 master 权重和动量状态仍然是 fp32所以模型状态那部分的显存并不会减半减的主要是激活值。三个手段叠加的效果可以粗略估一下假设不优化时单卡 batch size 只能设 1开了 bf16 能到 2再开梯度检查点能到 4-6再加梯度累积累积步数 8有效批量就能到 32-48配两张卡就是 64-96。这个水平足够跑大多数微调任务了。4.4 一份可以直接改的微调配置下面这份配置是我在微调平台上常用的模板参数组合是实测过的你可以直接拿来改。# 数据相关 cutoff_len: 1024 # 截断长度越长越吃显存 train_on_inputs: false # 只对回答部分算损失通常关掉输入 # batch 相关核心 per_device_train_batch_size: 2 # 单卡 micro batch gradient_accumulation_steps: 8 # 累积步数 # 有效批量 2 × 8 × 卡数 # 训练轮次 num_train_epochs: 3.0 max_steps: -1 # 设为正数则覆盖 epoch 配置 learning_rate: 2.0e-4 lr_scheduler_type: cosine warmup_ratio: 0.03 # 前 3% 的步数做 warmup # 显存优化 bf16: true gradient_checkpointing: true optim: adamw_torch weight_decay: 0.0 # 日志与保存 logging_steps: 10 # 每 10 个 step 打一条日志 save_steps: 200 # 每 200 步存一次 save_total_limit: 3 # 只留最近 3 个防止磁盘爆掉 eval_steps: 100几个参数之间的关系我梳理一下。cutoff_len直接决定激活值大小从 1024 提到 2048 大概会让激活显存翻倍如果 OOM 了就先把这里降下来优先级高于降 batch size。warmup_ratio配合有效批量一起看有效批量越大warmup 应该越长。save_total_limit一定要设我有一次忘了设训练跑了 8 个小时把磁盘写满进程直接崩了前功尽弃。如果你要换成固定步数训练把max_steps设成一个具体值比如 1000num_train_epochs就自动失效了。这种模式下建议把save_steps设小一点比如 100多留几个中间 checkpoint 做对比。5. 常见问题与排查技巧实录5.1 现象速查表看到症状直接对号入座这张表是我这些年攒下来的遇到问题先扫一遍能省掉大量瞎试的时间。现象高概率原因优先排查动作loss 第一个 step 就是 nan学习率过大 / fp16 溢出 / 数据含脏样本先降 lr 十倍fp16 换 bf16抽查数据loss 剧烈抖动不收敛batch size 太小 / lr 太大 / 没开 shufflebatch 翻倍或 lr 减半确认 shuffleTrueloss 长期横盘不降lr 太小 / 标签错位 / 梯度被截断没回传lr 放大十倍试用 32 条样本做过拟合测试训练 loss 降但验证 loss 升过拟合减少 epoch加正则扩充数据每经过 epoch 边界 loss 尖刺末尾残留梯度未更新 / 数据未 shuffle处理累积残留检查 shuffleGPU 利用率长期低于 60%数据管道是瓶颈增加 num_workers开 pin_memory训练中途 OOM某些样本序列过长 / 激活峰值波动看 max_memory_allocated限制最大长度断点续训后 loss 反弹优化器状态没保存checkpoint 里加 optimizer.state_dict()多卡训练效果不如单卡有效批量算错 / 学习率没缩放重新核算 effective_batch 并调 lr5.2 三个真实故障的排查全过程表格给的是结论但排查过程才是真正能学到东西的部分。我挑三个印象最深的案例展开。案例一换到四卡之后 loss 完全不动了。一个文本分类任务单卡 B32、lr2e-5 跑得很好改成四卡 DDP 之后前 500 步 loss 一动不动。第一反应是通信问题但检查梯度同步日志发现一切正常。后来把有效批量算了一遍四卡各 32总有效批量变成 128是原来的四倍而学习率还停在 2e-5。按线性缩放规则应该调到 8e-5。改成 8e-5 之后 loss 立刻开始下降。这个坑的本质是多卡训练时全局批量变了但很多人只改了并行配置忘了调学习率。案例二训练到第 8 个 epoch 突然 OOM。前 7 个 epoch 都好好的第 8 个 epoch 刚开始就爆显存。查了半天模型和配置都没变最后打印了每个 batch 的输入长度发现数据是按长度排序的——前 7 个 epoch 顺序固定所以每个 epoch 的显存曲线完全一样第八个 epoch 我打开了 shuffle一批长序列被随机凑到了一起激活值直接顶破上限。解决办法有两个把最大长度限制从 2048 降到 1536或者保证每个 batch 的长度分布相对均匀。我用的是前者简单直接。这个案例的教训是OOM 不一定是配置问题很可能是数据分布问题尤其是变长输入的任务。案例三LoRA 微调出来效果极差怀疑是框架有问题。一位同学用 3 张卡微调配置写的是per_device_train_batch_size: 4、gradient_accumulation_steps: 4但他以为有效批量就是 16。实际上乘以卡数之后是 48而他按 16 的量级把学习率设成了 5e-4对 48 的有效批量来说偏小模型学得很慢跑完三个 epoch 效果自然差。把学习率提到 8e-4 重新跑效果立刻正常。有效批量到底是多少这个问题一定要亲手算一遍不要凭印象。5.3 我踩过的坑和一些不讲道理但管用的技巧最后分享几条经验都是常规文档里不会写的。先做过拟合测试再开长跑。正式训练之前从训练集里挑 32 条样本让模型在上面反复训练到 loss 接近 0。这个测试能在几分钟内验完一整条链路数据加载对不对、标签对齐没有、梯度能不能回传、学习率量级合不合适。如果 32 条样本都过拟合不了那问题一定在代码或数据里不用浪费几小时去跑全量。我现在的习惯是每换一次数据集或者改一次数据管道这个测试都要重跑一遍。日志里一定要有 step 号不能只有 epoch。只有 epoch 的日志在排查问题时几乎是废的。看到第 3 个 epoch 出问题你根本不知道是 epoch 刚开始还是快结束时。加上 step 号之后第 3 个 epoch 的第 812 步开始 loss 飙升这种信息才有诊断价值。batch size 和学习率一起调永远不要只动一个。这是本文出现频率最高的一句话因为它是最高频的故障源。改配置的时候把这两个参数当成一个整体改完立刻跑 200 步看曲线。学习率调度器的总步数要算准。用余弦退火的时候total_steps如果算错整个学习率曲线就错了。如果用了梯度累积total_steps是有效总样本数 / 有效批量而不是总样本数 / 单卡批量。我第一次用余弦调度时就是这里算错了导致学习率提前衰减到最低值模型后半程基本没在学。保存 checkpoint 时顺手记一下当时的验证指标。文件名里直接带上指标比如ckpt_step1200_f1-0.863.pt后期挑选的时候一目了然不用一个个加载回来重新评测。关于现在本地模型还需不需要自己训练微调这个最近被反复问到的问题。我的看法是如果通用模型在你的任务上已经够用那就不需要但只要你有垂直领域的术语、特定的输出格式要求、或者希望模型说人话的风格贴合你的业务微调仍然是性价比最高的路径。而微调绕不开的就是本文这一套 batch 相关的配置你不需要成为专家但至少得知道有效批量怎么算、学习率怎么跟着走、显存不够时候先动哪个开关。这几个问题搞清楚大部分微调任务都能自己跑下来。
返回列表