本节摘要:训练一整个周期(epoch)的内部动作,以及"批处理"这个不起眼却影响深远的设定。本节解剖一次完整的前向/反向/更新,讲解批大小如何影响梯度噪声、内存占用和泛化,并给出划分训练跑批的落地流程。
阅读完本节,你应当能够:
上一节装好了目标和优化器,这一节回答问题:训练脚本在每一个 epoch 里内部到底做了什么。很多人把训练当按钮,一按就跑,跑完看分数——这没错,但一旦出了问题,你就失去诊断能力。把流程剖开,后面排错才有着手点。
这里的关键是把"训练步"和"训练轮"两个单位分清。一次"步"(迭代)是模型在一个批量上完成一次前向+反向+更新;一个"epoch"是把整个训练集完整过一遍。一个 epoch 里有多少步,取决于训练集大小除以批大小。譬如一万条样本、批大小 100,那一个 epoch 恰好是 100 步。这个换算关系,后面读日志、配置步数时几乎每天都要用到。
一个 epoch = 把整个训练集过一遍上面的循环,通常包含若干次"一步"(迭代)。每次迭代执行四步:
这里有一个数字先讲清:批大小等于每次更新用多少样本。批大小直接落在三步棋上:梯度质量、显存压力、泛化性格。
理解"反向"这一步时,一个常见的认知误区是误以为每一步都只算一次梯度。实际上,批大小为 B 的一步,B 越大,这一个更新的梯度就越接近"在全量数据上的真实梯度",代价是每步的计算量也越大。这也是为什么"1 个 epoch = 样本数 ÷ 批大小 步"这个换算关系很关键——你把批从 32 翻到 128,一个 epoch 的步数就变成原来的四分之一,而每次更新用的样本数变成四倍,两者的乘积(横穿的样本总量)保持不变。
用蒙特卡洛的直觉理解最生动:一次更新见的样本越少,梯度越像"用一点样本来估全局"——噪声大、direction 乱,但它便宜、且自带一定正则化效果(噪声能帮忙逃离局部坑)。批越大,梯度估得越稳,但内存/显存线性涨、且可能收敛到更尖的坑、泛化有时反而略差。
给一张速查:
| 批大小 | 梯度噪声 | 显存 | 泛化倾向 | 适用 |
|---|---|---|---|---|
| 小(8-64) | 大 | 低 | 偏正则、好泛化 | 大模型、内存紧 |
| 中(64-512) | 中 | 中 | 均衡 | 多数日常 |
| 大(512+) | 小 | 高 | 偶尔略差、可配学习率缩放 | 有大规模算力 |
批大小不是孤立旋钮:实践常按"线性缩放规则"——批变大时把学习率成比例调大,才能保持收敛步长相当。这提示批大小与学习率在超参调优里要一起看,第六章会再回这个耦合。
举一个具体操作:基线在批大小 64、学习率 1e-3 下的曲线还算理想。你想试试批 256 只求更快,就先把学习率按比例提到 4e-3(64 的 4 倍),这样每次更新的"步长"量与基线大致相当。如果你只把批放大却不动学习率,收敛就会明显变慢,这常常是"换了大批反而更慢"这一困惑的真相。
连带要养成的一个好习惯是把"步/轮/样本吞吐"这几个读数盯着看。很多排错的起点其实就藏在这里:当你发现"明明批大小没变,一个 epoch 却比以前慢了很多",多半不是模型变笨了,而是数据加载或预处理路径出了岔——某一步不小心退化成并行预取被关掉、或者重复清洗了同一个数据。把"每轮耗时、每轮吞吐、验证抖动"三行数字记进实验记录,你就能先于损失曲线发现这类"看得见却说不清"的回退。这也正是第一章建立的"用验证信号判据"纪律在生产环境里的延伸——不只盯分数,还要盯那些能解释分数为什么变化的中间读数。
显存是硬道理。模型大、批太大 → 显存爆掉,最常见的处理是减小批,若还爆就用梯度累积:连续跑若干个微批、累加梯度,攒够了再更新一次,等价于一个大批但内存省下来。
这里要提醒一个方向的误区:并不是批越小越好。批太小(比如只有 1-2 个样本)会让梯度噪声大到几乎每步都在乱走,收敛变得异常缓慢,甚至不收敛;批太大则吞显存又可能收敛到更尖的坑。所以现实里"恰好装得下、又不太吵"的那个区间,往往需要你亲手量一下显存余量再定——这也正是它成为超参数选项卡里常客的原因。
# 示意:梯度累积把大档案拆成若干微批 accum_steps = 4 optimizer.zero_grad() for i, (xb, yb) in enumerate(loader): loss = loss_fn(model(xb), yb) / accum_steps loss.backward() if (i + 1) % accum_steps == 0: optimizer.step() optimizer.zero_grad()
⚠️ 常见坑:随机切批时每批混进测试集样本。数据 loader 里若有测试逃逸,验证分数必然失真。保证 loader 只读训练子集,测试独立装 loader。
💡 关键直觉:训练速度瓶颈常常不在 GPU 计算而在"喂数据太慢"。打好数据加载的并行与预取,常比调超参更直接地缩短迭代周期。
说到梯度怎么走,下一节专门钻进梯度下降及其变种。