epoch、batch、loss、val_loss:深度学习训练四要素
先说结论调模型调到最后还能天天盯着看的指标其实就那么四个——epoch、batch、loss、val_loss。几乎所有人在入门深度学习的第一天就听过它们但真正能把四者之间的关系讲清楚、并在实际训练里用对的人没那么多。我见过太多人把 epoch 直接拉到几百轮才开始调参也见过有人 batch size 随手填个 32 就再也没动过最后抱怨显存不够、收敛太慢、验证集一动不动。这篇文章我打算把这四个概念从定义、换算关系、参数选择逻辑到训练曲线怎么读、出问题怎么排查完整捋一遍中间夹带一些我自己踩过的坑和实际项目里的经验数值希望能让刚上手的人少走点弯路也让已经跑过几个模型的人重新校准一下自己的直觉。1. 先把四个概念摆到桌面上一次训练到底在干什么1.1 从原始数据到一次参数更新要理解 epoch 和 batch得先接受一个事实神经网络没法一次性吃下你所有的数据。假设你有 10 万张图片要训练一个分类模型如果每一轮都拿全部 10 万张算一次梯度再更新一次参数会发生两件糟糕的事。第一显存扛不住10 万张图的前向激活值放在显存里普通卡直接爆掉。第二梯度方向会很稳但更新次数太少一个 epoch 只更新一次参数收敛速度慢到没法用。所以从业者的做法是把数据切开一块一块喂给模型。每次喂进去的这一小块就叫一个batch。模型对这一个 batch 做一次前向传播算出损失也就是loss做一次反向传播算出梯度然后优化器更新一次参数。这个取一个 batch 到更新一次参数的完整动作叫做一个iteration也有人叫 step。而当全部数据都被完整地过了一遍不管中间切了多少个 batch这一整轮就叫做一个epoch。换句话说epoch 是数据过一遍的宏观单位iteration 是参数更新一次的微观单位batch 是连接它们两个的中间量。这里有个容易混淆的点很多人把 batch 和 iteration 当成一回事其实严格说 batch 是数据块的容量名词属性更强iteration 是动作的次数。20 万条样本、batch size 取 200那需要 1000 个 iteration 才能完成一个 epoch。这个换算关系建议刻进脑子里后面调参全靠它。提示很多人问epoch 到底要设多少这个问题的答案永远是看验证集但前提是你得先把 iteration 和 epoch 的换算搞清楚否则你连模型这辈子更新过多少次参数都不知道调参就是盲猜。1.2 batch、iteration、epoch 三者的换算关系把换算关系写成公式比记文字靠谱得多每个 epoch 的 iteration 数 ceil(样本总数 / batch_size)总 iteration 数 每个 epoch 的 iteration 数 × epoch 数有效 batch size 单卡 batch_size × 梯度累积步数 × 数据并行卡数拿一个真实场景举例。手头数据集 12.8 万条样本单卡 batch size 设 64那一个 epoch 就是 128000 / 64 2000 个 iteration。如果训练 30 个 epoch总迭代次数就是 6 万次也就意味着参数被更新了 6 万次。参数含义典型取值范围主要影响batch_size一次前向反传的样本数8 ~ 512显存占用、训练速度、梯度噪声iteration一次参数更新的动作由数据集和 batch 决定决定总训练步数epoch完整过一遍数据集5 ~ 200训练总时长、过拟合风险loss训练集上的损失值越小越好但要看趋势判断模型是否在学val_loss验证集上的损失值越小越好是选模型依据判断泛化能力上表这几个数值不是随便写的是我在图像分类和文本分类任务里反复试出来比较稳的区间。当然具体项目差异很大检测、分割这类任务的 batch size 常常只有 2 到 8因为单张图的激活值太大了。这个后面会展开。2. epoch 和 batch 的选择逻辑为什么我不建议一上来就调它们2.1 epoch 太多会怎样过拟合的现场新手最常见的误区是epoch 拉满总没错。理论上训练集上的 loss 会随着 epoch 增加一直往下走看起来很美。但模型的目的是在没见过的数据上表现好而 val_loss 往往在某个 epoch 之后开始掉头往上走。这个拐点就是过拟合开始的地方。我印象很深的一次做一个工业质检的二分类任务训练集只有 4000 多张图我图省事设了 200 个 epoch。结果前 30 个 epoch val_loss 一路降到 0.18很漂亮从第 40 个 epoch 开始train loss 还在降掉到 0.02但 val_loss 已经悄悄爬回 0.35 了。模型把训练集里的每张图的噪声都背下来了换到新图上直接歇菜。正确的思路是epoch 设一个偏大的上限比如 100 或 200然后靠**早停Early Stopping**来截断。早停的逻辑很朴素——盯着 val_loss如果连续 N 个 epoch 都不再创新低就停下来。这个 N 通常叫 patience取 5 到 10 比较常见太小容易被正常波动误伤太大浪费时间。注意早停必须配合保存最佳模型一起用。如果你只早停不保存最后拿到的可能是已经过拟合的那一版参数。保存逻辑一定挂在 val_loss 最低的那个 epoch 上而不是最后一个 epoch。还有一个细节如果你是做微调fine-tune预训练模型epoch 通常只要 3 到 10 就够了甚至 2 到 3 个 epoch 就能出很好的效果。因为预训练权重本身已经很强你只是把它往你的任务上掰一点用力过猛反而会把原有的知识覆盖掉这个现象在圈子里叫灾难性遗忘。我在做文本分类微调时10 万条数据一般也就跑 3 个 epoch再多就开始伤模型了。2.2 batch size 与显存、速度、泛化的三角关系batch size 是这四个参数里最需要认真对待的一个因为它同时牵扯三件事显存、速度、泛化。显存角度batch size 和显存占用基本是线性关系。batch 翻倍激活值占用翻倍权重本身不变。所以显存不够时的第一反应就是降 batch size。但要注意降低 batch size 会让某些依赖批统计的层出问题比如 BatchNorm。当 batch size 小到 1 或 2 时BatchNorm 算出的均值方差极不稳定训练会抖得很厉害。这时候的正确做法是把 BatchNorm 换成 GroupNorm或者冻结 BN 的统计量。速度角度这里有个反直觉的地方。很多人以为 batch size 越小越快其实在 GPU 上恰恰相反。小 batch 会导致 GPU 的并行计算单元吃不满每个 iteration 虽然算得快但利用率低单位时间内处理的样本数反而下降。我实测过 ResNet50 在单张卡上的表现batch size 从 16 提到 128每秒处理的图片数大概能翻一倍。泛化角度这个更微妙。小 batch 因为每次梯度噪声大反而有轻微的正则化效果有时最终精度比大 batch 还好一点。大 batch 梯度方向准收敛平滑但容易陷进 sharp minimum测试集上稍差。这也是为什么有些论文用超大 batch 训练时要额外加 warmup 或者改用 LARS、LAMB 这类优化器。综合下来我的经验是在显存允许范围内把 batch size 设到能跑满 GPU 的档位通常是 32 到 256 之间。如果碰到显存瓶颈再考虑梯度累积。2.3 梯度累积小显存做大 batch 的折中梯度累积Gradient Accumulation是个非常实用的技巧值得单独讲。它的思路是我不一次性塞 256 个样本而是每次只塞 32 个连续算 8 次但前 7 次不更新参数只把梯度累加起来第 8 次才真正更新一次。这样参数更新的效果和 batch size 256 几乎一样但显存只需要 batch size 32 的量。伪代码大概长这样accum_steps 8 optimizer.zero_grad() for i, (x, y) in enumerate(loader): out model(x) loss criterion(out, y) / accum_steps # 关键除以累积步数 loss.backward() if (i 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()这里有个新手特别容易踩的坑loss 一定要除以 accum_steps。因为 PyTorch 的 backward 默认是把梯度累加的如果不除累积 8 次之后梯度就是原来的 8 倍等效于把学习率放大了 8 倍训练直接炸掉。我就见过有人这么写loss 上来就是 nan排查了一整天。提示梯度累积只是模拟大 batch 的更新频率它并不改变 BatchNorm 的统计量。如果你的模型里有 BatchNorm累积出来的统计量仍然是按小 batch 算的。这点要清楚。3. loss 与 val_loss训练中唯一值得你信任的两个数字3.1 loss 的数学本质loss 这个词被用得太随意了其实它是一个非常明确的数学对象衡量模型预测和真实标签之间差多远的一个标量函数。训练的本质就是不断调整参数让这个标量变小。回归任务里最常见的是均方误差MSE把预测值和真实值的差平方后求平均。分类任务里最常见的是交叉熵Cross Entropy它衡量的是模型输出的概率分布和真实分布之间的差异。交叉熵的公式是L -Σ y_i * log(p_i)其中 y_i 是真实标签的独热编码p_i 是模型给出的第 i 类概率。直观理解就是模型对真实类别给出的概率越高log 出来的负值越小loss 就越低如果模型把真实类别的概率给得很低loss 就会很大逼着模型改。这里我要顺带讲一个高频热词——focal loss。它是在标准交叉熵上加了一个调节因子FL -α * (1 - p_t)^γ * log(p_t)其中 p_t 是模型对真实类别的预测概率γ 是聚焦参数通常取 2α 是类别平衡权重。这个 (1 - p_t)^γ 的作用是对已经分得很准的样本p_t 接近 1自动降低权重把注意力集中到那些还没分对的难样本上。它为什么火因为在目标检测、欺诈识别这类正负样本极度不均衡的场景里标准交叉熵会被大量的简单负样本淹没——模型只要把所有东西都预测成负类loss 就已经很低了梯度就没什么动力去学分得少但重要的正样本。focal loss 就是来解决这个的。我在做电商异常订单识别时正样本占比只有 0.3%换成 focal loss 后召回率从 0.61 提到 0.78效果非常明显。3.2 val_loss 为什么有时候比 loss 还小这是新手最困惑的问题之一明明验证集是模型没见过的数据为什么 val_loss 反而比 train loss 还低第一个原因是训练时的 loss 是一个滑动平均值。你看到的是最近若干个 iteration 的均值这里面包含了训练早期 loss 很高的那一段所以整体被拉高了。而 val_loss 是在当前这个已经训练得不错的模型上一次性算出来的自然更低。第二个原因是训练时的额外扰动。训练阶段通常会开 dropout、数据增强、随机裁剪等操作这些都会让 train loss 偏高。而验证阶段这些全部关闭模型是在最干净的状态下评估loss 自然更低。第三个原因是正则项。如果你的 loss 里加了 L2 正则weight decay那 train loss 本身就含了正则项而 val_loss 通常只算纯粹的预测误差。如果排除了以上原因val_loss 依然显著低于 train loss那就要警惕是不是验证集和训练集有重叠或者数据划分时没打乱导致分布不一致。这个坑我踩过一次做时间序列预测时用随机划分替代了按时间划分结果验证集里有一部分样本和训练集高度相关val_loss 低得离谱上线后一塌糊涂。3.3 早停与模型保存盯哪个指标前面说了早停要看 val_loss这里展开一下具体怎么做。best_val_loss float(inf) patience 7 counter 0 for epoch in range(num_epochs): train_loss train_one_epoch(model, train_loader) val_loss evaluate(model, val_loader) if val_loss best_val_loss: best_val_loss val_loss counter 0 torch.save(model.state_dict(), best.pt) else: counter 1 if counter patience: print(f早停于 epoch {epoch}) break有两个细节值得说。第一保存的模型应该是 val_loss 最低的那一版不是最后一版。很多人训练完直接拿最后一步的权重去推理白白浪费了早停的意义。第二patience 的设置跟学习率调度有关。如果你用了 ReduceLROnPlateau 这种val_loss 不降就降学习率的策略patience 要设得比它大否则学习率还没降下去你就早停了白白错过一波提升。另外提一个容易忽略的点val_loss 和下游指标不一定完全正相关。做分类任务时val_loss 最低的模型准确率或 F1 不一定最高因为在决策边界附近的样本概率的微小变化对 loss 影响大但对最终分类结果没影响。所以如果你的项目对某个具体指标特别敏感比如医疗场景的召回率那就应该直接拿那个指标来做模型选择而不是死盯 val_loss。4. 损失函数选型与训练循环实操4.1 常见 loss 选型对照表选 loss 这件事很多人是抄来的看到别人用交叉熵就用交叉熵看到别人用 Dice 就用 Dice其实每种 loss 都有它适合的场景。我把常用的整理成一张表任务类型推荐 loss适用场景注意事项单标签分类CrossEntropyLoss类别互斥标签是类别索引别自己转 one-hot多标签分类BCEWithLogitsLoss一个样本多个标签内部自带 sigmoid别再手动加类别不均衡Focal Loss正负样本悬殊γ 从 2 起调α 按类别频率倒数设语义分割Dice Loss CE前景占比小单用 Dice 早期梯度不稳建议联合目标检测框回归Smooth L1 / GIoU坐标回归直接 MSE 对离群框太敏感回归预测MSELoss / HuberLoss连续值有异常值时用 Huber 更稳度量学习TripletLoss人脸、检索难样本挖掘策略比 loss 本身更关键这张表里的每一条我基本都在项目里用过。这里特别说一下类别不均衡用 Focal Loss这一行很多人一遇到不均衡就上 focal loss其实如果只是轻微不均衡比如 3:1直接给 CrossEntropy 加个 class_weight 就够了focal loss 调参反而更麻烦。真正需要 focal loss 的是那种 100:1 甚至更极端的场景。还有一点多任务学习里各 loss 的加权是门手艺。我一般会先让每个任务的 loss 单独跑一段观察它们大致的数值量级然后让各项在总 loss 里的贡献比例大致相当。有些团队会专门定义一个 loss ratio 指标来跟踪每个损失项占总损失的比重比如总 loss 是 1.0其中分类项占 0.6、回归项占 0.4如果发现某一项长期占 0.95那基本说明其他任务被压制了得手动调权重。这个做法在多任务项目里非常实用比拍脑袋设权重靠谱得多。4.2 学习率与 batch size 的联动这是个必须讲的点因为它直接决定了你的 loss 能不能正常下降。有一个经验规律叫线性缩放规则如果把 batch size 扩大 k 倍学习率也应该相应扩大 k 倍或者取平方根视情况而定。原因是 batch 变大后梯度的噪声变小方向更准你可以用更大的步子走。但这里有个配套技巧叫warmup。因为训练刚开始时参数是随机的梯度方向很乱如果一上来就用大学习率loss 很容易直接炸到 nan。warmup 的做法是前几百个 iteration 让学习率从 0 线性升到目标值然后再按正常策略衰减。def lr_lambda(step): if step warmup_steps: return step / warmup_steps return 0.5 * (1 math.cos(math.pi * (step - warmup_steps) / (total_steps - warmup_steps))) scheduler torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)我实测下来的经验是batch size 在 64 以下时学习率 1e-3 到 3e-4 比较稳batch size 到了 512 以上学习率可以上到 1e-2但必须配 warmup。如果是微调预训练模型学习率要再降一个数量级通常用 2e-5 到 5e-5这个范围和 BERT 系列论文里给的建议基本一致。4.3 手写一个可复现的训练循环把前面这些东西串起来一个完整、可复现、带早停和梯度累积的训练循环大概是这样import torch import math model build_model() optimizer torch.optim.AdamW(model.parameters(), lr3e-4, weight_decay0.01) criterion torch.nn.CrossEntropyLoss() accum_steps 4 best_val_loss float(inf) patience, counter 7, 0 for epoch in range(50): model.train() optimizer.zero_grad() running_loss 0.0 for i, (x, y) in enumerate(train_loader): x, y x.cuda(), y.cuda() out model(x) loss criterion(out, y) / accum_steps loss.backward() running_loss loss.item() * accum_steps if (i 1) % accum_steps 0: torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() optimizer.zero_grad() train_loss running_loss / len(train_loader) model.eval() val_loss, correct, total 0.0, 0, 0 with torch.no_grad(): for x, y in val_loader: x, y x.cuda(), y.cuda() out model(x) val_loss criterion(out, y).item() correct (out.argmax(1) y).sum().item() total y.size(0) val_loss / len(val_loader) acc correct / total print(fEpoch {epoch}: train{train_loss:.4f} val{val_loss:.4f} acc{acc:.4f}) if val_loss best_val_loss: best_val_loss val_loss counter 0 torch.save(model.state_dict(), best.pt) else: counter 1 if counter patience: print(触发早停) break这段代码里几个点值得单独强调。梯度裁剪clip_grad_norm_在 RNN、Transformer 类模型里几乎是标配能有效防止梯度爆炸每轮结束手动切换 model.train() 和 model.eval()忘了切换会导致 dropout 和 BatchNorm 行为不对验证结果完全不可信val_loss 累加后要除以 len(val_loader)不然得到的数值跟 batch 数量挂钩没有可比性。注意如果你用的是 CrossEntropyLoss标签必须是类别索引长整型不用做 one-hot。而 BCEWithLogitsLoss 的标签必须是浮点型这个类型不匹配是最常见的报错来源之一报错信息往往还很难看懂。5. 训练曲线诊断与问题排查实录5.1 四种典型曲线走势训练曲线是 loss 和 val_loss 随 epoch 变化的图读图的能力比调参技巧更重要。我总结下来常见的走势有四种第一种是理想型train loss 和 val loss 同步下降最后都趋于平缓两条线之间保持一个稳定的间距。这说明模型容量、正则化强度、数据量都比较匹配是最省心的状态。第二种是过拟合型train loss 继续降val loss 在某个点后拐头上升两条线开始分叉。这时候的动作是加正则dropout、weight decay、做数据增强、或者干脆减小模型。数据这种从根上解决问题的办法永远优先于调参。第三种是欠拟合型两条线都在高位平缓降不下去。说明模型太简单或者学习率太小。这种时候换个更大的模型、调大学习率、去掉过强的正则通常马上见效。第四种是发散型loss 直接往上跑或者变成 nan。这基本是学习率太大、梯度爆炸或者数据里有脏样本标签错、数值 inf/nan。先降学习率再加梯度裁剪最后检查数据。5.2 常见问题速查表我把这些年遇到过的问题和对应处理整理成一张表方便对着查现象可能原因排查方向处理办法loss 一直是 nan学习率过大、数据含 inf打印第一批数据降 lr、加 warmup、清洗数据loss 完全不降lr 太小、模型没接对检查参数是否在更新调大 lr、打印梯度范数val_loss 剧烈震荡batch 太小、BN 不稳看 batch size增大 batch 或换 GroupNormval_loss 远低于 train数据泄漏检查划分逻辑按时间/主体重新划分训练到一半突然变差学习率调度不合适看 lr 曲线改余弦退火或加 warmup显存突然爆掉序列长度不均、缓存未释放看输入 shape加 padding 上限、清 cache某轮起 loss 卡住不变梯度消失、数据加载卡死打印梯度范数换激活函数、检查 dataloader这张表里梯度范数这个排查手段我要重点推荐。很多人调 loss 不降的时候只会反复改学习率其实在 backward 之后打印一下各层的梯度范数一眼就能看出问题如果梯度范数接近 0是梯度消失如果是几十上百是梯度爆炸如果正好是 0那可能是某个环节 detach 了或者这层根本没参与计算。5.3 踩坑记录说几个我印象比较深的真实坑。第一个坑是Dataloader 的 num_workers 设太大。我一度以为 num_workers 越大加载越快设成了 32结果训练速度反而下降因为进程切换开销加上内存拷贝把收益吃掉了。一般的经验是 num_workers 设成 CPU 核数的 1/4 到 1/2配合 pin_memoryTrue 就够了。第二个坑是验证集的预处理跟训练集不一致。训练时我用了随机裁剪加归一化验证时顺手复制了同一套 transform结果验证集也做了随机裁剪每次评估的结果都不一样val_loss 抖得跟心电图似的。正确的做法是验证集只保留确定性的 resize 和归一化绝对不能有随机操作。第三个坑是滑动平均导致的假下降。有段时间我用了指数移动平均EMA来平滑曲线结果 loss 看起来很稳实际上是平滑掩盖了剧烈波动真正的训练问题一直没被发现。现在我更倾向于原始曲线和滑动平均曲线都画出来两条一起看。提示训练初期一定要做一次小样本过拟合测试。随便取 16 条数据把模型跑上几十个 iteration如果 loss 不能降到接近 0说明模型结构、loss 或者标签对接有问题这时候不用浪费时间跑全量。这个习惯能帮你省下大量调试时间。6. 命名撞车batch 和 loss 在其他领域指什么6.1 批处理的几种常见语境batch 这个词本身是一批、一炉的意思所以在很多行业里都有它自己的含义跨领域沟通时经常闹笑话。比如在 3D 建模和游戏美术流程里batch FBX export指的是把多个模型文件一次性批量导出成 FBX 格式。这跟深度学习的 batch 没有半点关系只是借用了批量这个语义。你在搜索引擎里搜 batch搜出来一堆 3D 软件的操作教程就是这个原因。工业扫描领域有batch scan wizard指的是扫描仪软件里的批量扫描向导用来连续扫描多张纸、多个工件。化工流程模拟里有aspen batch process指的是间歇式批处理生产工艺的建模也就是一炉一炉投料的生产方式跟连续生产工艺相对。这些 batch 都是批量处理的意思语义上是一致的只是应用场景完全不同。深度学习里的 batch 其实也是这个语义的延伸——把一批样本一起处理。所以你在跟非算法同事交流时直接说我一次喂给模型 64 张图比说我的 batch size 是 64要清楚得多。6.2 名字相同但含义完全不同的坑更有意思的是 loss 和 power loss 这种撞车。BIOS 里有一项 AC Power Loss 的设置指的是服务器或工控机在突然断电、之后又来电时主板应该怎么反应——是自动开机、保持关机还是恢复断电前的状态。运维同事排查掉电后服务器不自动开机时会去改这个 BIOS 选项跟神经网络里的损失函数一点关系都没有。类似地电力行业的线损、金融行业的损失率翻译成英文都可能是 loss但那是完全不同的概念。我曾经在一次跨部门会议上听到有人在讨论loss ratio一边理解为模型各项损失的占比一边理解为业务上的赔付率聊了半天才发现双方说的不是一件事。所以我的建议是在跨领域协作时第一次提到这些词一定要带上限定语。说训练损失而不是loss说批量大小而不是batch说一轮训练而不是epoch。多花几秒钟把话说全能省掉后面几十分钟的互相误解。关于 epoch 和 batch 的选择我自己最后沉淀下来的默认配置是这样的图像分类任务batch size 从 64 起试显存允许就往上提到 128 或 256epoch 上限设 100 并配 patience 为 7 的早停学习率 3e-4 配余弦退火和 500 步 warmuploss 先跑标准交叉熵看基线如果发现类别不均衡再换 focal loss。这套配置在我手上大部分中等规模数据集里都能跑出可用的结果剩下的就是根据 val_loss 曲线的形状微调。真要说有什么忠告就是别急着调参先把数据切分、标签对齐、预处理一致性这几个基础环节确认清楚我见过的问题里有七成以上最后都溯源到数据环节而不是 epoch 或者 batch 本身设错了。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →