☰
深度学习训练调参:batch size、epoch、iteration 口径
2026/10/1 1:23:50 网站建设 项目流程

调参这件事,真正让人翻车的往往不是网络结构写错,而是 batch、batch size、epoch、iteration 这四个词没对齐口径。我见过太多人在群里发问:"我训练了 10 个 epoch 为什么 loss 还在 2.3?"结果一看,他的 batch size 是 4,数据集有 6 万张图,10 个 epoch 也就是 15 万次 iteration,而别人 10 个 epoch 已经是 600 万次参数更新了。同样叫"10 个 epoch",工作量差了 40 倍,这根本不是模型的问题,是概念没理清。这四个词是模型训练的地基,搞明白它们之间的关系,你才能算得清训练时间、估得准显存占用、调得动学习率。不管你是拿 ResNet 做四分类花卉识别,还是用 YOLO 训自己的检测数据集,或者跑 LoRA 微调一个中文大模型,这套口径都是通用的。下面我按自己实际带项目和踩坑的顺序,把这四个词拆开讲透。

1. 四个词其实站在不同尺度上:先把 iteration、batch、epoch 的分工摆清楚

很多教程一上来就丢公式,但我觉得更有效的做法是先想清楚"谁在计量什么"。训练过程本质上是一个循环套一个循环:外层循环一遍一遍地扫数据,内层循环一批一批地送数据、一次一次地更新参数。epoch、batch、iteration 分别描述这三个层次上的"一件事",混淆它们就像把"秒""分""小时"混着用,账一定算不平。

1.1 参数每更新一次,就是一个 iteration

一次 iteration 的完整动作是:取一小批样本送进网络,前向算出预测值和损失,反向算出梯度,优化器根据梯度把权重改一次。注意关键词是"改一次权重"。只要优化器的step()被调用了一次,就算一个 iteration 走完了。这里有个容易忽略的细节:如果你用了梯度累积,那step()可能每 N 个小批才被调用一次,这时候"前向+反向"发生了 N 次,但参数只更新了一次。严格来说 iteration 指的是参数更新次数,而"前向次数"是另一个量。这个区别在你读大模型训练日志时会特别明显——日志里写的global_step通常就是 iteration 数,而不是前向次数。

1.2 数据走完一遍,才算一个 epoch

epoch 的定义非常朴素:训练集里的每一条样本都被模型"看过"一次,就叫一个 epoch 结束。注意是"看过",不是"记住",也不是"训好"。所以 epoch 是一个关于数据覆盖范围的量,和 batch size 没有直接关系——不管你把 batch size 设成 8 还是 256,把所有数据过一遍都叫一个 epoch,只是走的步数不同。这也是为什么"训了几个 epoch"永远不能单独作为训练充分与否的指标,必须配合数据集规模一起看。

1.3 batch 是连接两者的中间单位

batch(也叫 mini-batch)是我们从数据集里一次捞出来送进网络的那一小撮样本。它既是显存占用的直接决定因素,也是 epoch 换算成 iteration 的除数。之所以不用全量数据算一次梯度(那叫 full-batch 梯度下降),是因为全量梯度虽然方向准,但一次更新要扫完所有数据,收敛慢且内存扛不住;而一次只用一个样本(随机梯度下降)噪声又太大,梯度方向抖得厉害。mini-batch 是在这两者之间取的折中,既让梯度有一定的统计意义,又能让显存装得下、GPU 跑得满。

概念计量对象一句话定义谁在决定它
batch size单次送入的样本数一次前向吃多少条数据显存、任务类型、经验值
iteration参数更新次数优化器 step 一次数据量 ÷ batch size
epoch数据覆盖轮数全量数据被扫一遍训练策略、早停条件

把这张表记住,后面所有的换算和调参都从它推。

2. epoch 与 iteration 的换算,以及进度条背后那笔账

理解了三个尺度之后,换算就是纯算术。但恰恰是这步算术,能解释你 90% 的"训练怎么这么慢"的困惑。

2.1 换算公式和一段可验证的代码

核心关系只有两条:

# 每个 epoch 的迭代次数 iters_per_epoch = ceil(len(dataset) / batch_size) # 不舍弃最后一个残缺 batch iters_per_epoch = len(dataset) // batch_size # drop_last=True 时 # 总迭代次数 total_iters = iters_per_epoch * num_epochs

拿一个具体场景算:数据集 60000 张,batch size = 64,drop_last=True,那么iters_per_epoch = 937。跑 50 个 epoch 就是 46850 次参数更新。如果单步耗时 0.18 秒,总训练时间约 2.34 小时。你可以用这个反推:当你听说别人训 YOLO 用了一天,可以先问清他的 batch size、图片数量和轮数,很多"一天"和"三小时"的差距,就是这几个数不同,而不是硬件碾压。

计算的时候还有两个隐藏项要留意。一是数据加载时间:如果num_workers设得太小、硬盘又是机械盘,GPU 会空转等数据,实际单步耗时可能比纯计算时间高一倍。二是验证集评估:很多框架在每个 epoch 结束后跑一次 validation,如果验证集有 10000 张而 batch size 只有 8,光验证就要 1250 步,这部分时间经常被新手忽略,结果发现"训得比预期慢很多"。

2.2 最后一个不满的 batch:drop_last 到底丢不丢

当数据集大小不是 batch size 的整数倍时,最后一批会少于 batch size。这时候有两个选择:丢掉它(drop_last=True)或者保留它(drop_last=False)。很多人的直觉是"数据别浪费,当然保留",但实际经验恰恰相反。

保留最后一个残缺 batch 会带来一个问题:如果 BatchNorm 这类对 batch 内统计量敏感的层存在,一个只有 3 条样本的 batch 算出来的均值和方差噪声极大,会让这一步的梯度特别脏,模型反而被带偏。而且如果剩余样本数正好是 1,很多框架直接报错。所以图像分类、检测这类任务的实践里我基本都开drop_last=True,特别是做小批量微调时。只有当数据集本来就很小(比如几千条以内),丢掉一批意味着损失掉可观的样本量时,我才会保留,并配合把 BatchNorm 关掉或换成 GroupNorm。

数据集大小batch sizedrop_last每 epoch 迭代数备注
6000064True937标准图像分类常见配置
6000064False938最后一批仅 32 条
500016True312小数据集,丢掉 8 条可接受
500016False313最后一批仅 8 条,BN 会抖

3. batch size 的三方拉扯:显存、梯度方差、吞吐量

batch size 是这四个词里唯一需要你"人为拍板"的量,也是最容易调错的。它同时被三个因素绑架:显存上限、梯度质量、训练吞吐。这三个方向经常互相打架,所以没有放之四海皆准的"最佳 batch size"。

3.1 显存到底被谁吃掉了

先算清楚显存账,你才知道 batch size 的物理上限在哪。以全量微调一个 7B 参数的大模型为例,用 Adam 加混合精度训练时,显存主要花在四块:

  • 模型参数本身:fp16 存储约 14GB
  • 梯度:fp16 约 14GB
  • 优化器状态:Adam 需要一阶、二阶动量各一份 fp32,合计约 56GB
  • 激活值:与 batch size、序列长度、层数成正比,浮动最大

前三项加起来已经 84GB 左右,这就是为什么单张消费级显卡根本放不下 7B 的全量微调。激活值那一项正是 batch size 直接影响的:batch 翻倍,激活大致翻倍。所以当你把 batch size 从 4 调到 8 时显存爆了,不是模型变大了,是激活这块顶到了天花板。这时候正确的思路是砍激活——用梯度检查点(gradient checkpointing)把一部分中间结果丢掉重算,代价是训练慢 20%~30%,但能换来 2~3 倍的实际 batch size 空间。这个交易在显存紧张时几乎总是划算的。

3.2 梯度噪声与泛化:小 batch 未必吃亏

有个流传很广的说法是"batch size 越大越好,能训得更快更准"。这话只对了一半。大 batch 确实让梯度估计更接近真实梯度,优化路径更平滑,但代价是泛化能力有时反而下降。原因在于小 batch 带来的梯度噪声本身就起到了一种正则化效果——噪声把模型从尖锐的极小值里推出来,逼它去找更"宽"的解。这也是为什么很多论文发现,在相同 epoch 数下,超大 batch 训出来的模型在测试集上未必赢过中等 batch。

反过来说,小 batch 也有它的问题:梯度方差大,损失曲线会抖,学习率不能设太高,否则容易震荡甚至发散。所以实践中的阶梯是这样的:显存允许的前提下,先试一个"中等偏大"的 batch(比如单卡 32~64),跑几百步看 loss 是否平稳下降;如果抖得厉害,往下调;如果显存没吃满又有余力,往上加并同步放大学习率。我个人的经验是,分类任务里 batch size 对最终精度的影响往往在 1 个百分点以内,但对训练稳定性和时间的影响可能是翻倍的,所以别为了那 0.5 个点死磕,先把时间和复现性拿稳。

3.3 吞吐量和 GPU 利用率:小 batch 的隐性浪费

还有一个纯工程视角:batch 太小,GPU 根本喂不饱。现代 GPU 的算力是按"大批量矩阵运算"设计的,当 batch size 只有 2 或 4 时,很多算子在等待和数据搬运上花的时间超过了实际计算时间,算力利用率可能只有 30%~50%。这时候你把 batch size 翻倍,单步耗时几乎不变,等于白捡一倍的训练速度。

判断方法很直接,看nvidia-smi里的利用率和显存占用:如果显存只用了 40% 而利用率长期在 50% 以下,八成是 batch 太小或者数据加载成了瓶颈。前者靠加 batch 或梯度累积解决,后者靠加num_workers、把数据预取到内存、用更快的存储解决。我在做 OCR 类任务(比如 EasyOCR 这类自带训练配置的框架)时,把 batch 从 8 提到 32、num_workers从 2 提到 8,整体训练时间直接砍掉近一半,模型效果几乎没变。

4. 显存不够时的两套补救方案:梯度累积与学习率缩放

不是每个人都有 A100,更多时候你手上只有一张 8GB 或 12GB 的卡,想把等效 batch size 做大只能走别的路。

4.1 梯度累积的原理和最容易写错的两行代码

梯度累积的思路很朴素:既然一次装不下 64 条,那就分 4 次各装 16 条,把这 4 次算出来的梯度加起来再更新一次。这样等效 batch size 就是 16 × 4 = 64,显存按 16 条算,梯度质量按 64 条算,堪称性价比之王。

accum_steps = 4 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,否则梯度会被放大 4 倍,等效于学习率翻了 4 倍,很容易训飞。第二,zero_grad()的位置必须在step()之后,而不是每个循环开头,否则梯度一累积就被清空了。第三,如果模型里有 BatchNorm,累积并不能解决小 batch 导致的统计噪声问题——它只解决梯度层面的等效,不解决层内的批次统计。这点特别容易混淆,很多人以为用了梯度累积就能安全地用小 batch 配 BN,结果效果还是差。

4.2 学习率怎么跟着 batch size 走

等效 batch size 变大之后,学习率通常也要跟着调整,否则收敛会变慢。最常用的是线性缩放规则:batch 扩大 k 倍,学习率也乘 k。

base_batch = 256 base_lr = 1e-3 lr = base_lr * (global_batch / base_batch)

但这个规则在大 batch 时会失效——缩放过头会导致训练初期震荡。所以实践中会加一个平方根缩放的备选(lr * sqrt(k)),以及一段 warmup:前几百到几千步把学习率从接近 0 线性升到目标值,让优化器先"热身"再全速跑。大模型微调日志里常见的warmup_steps: 100说的就是这个。

缩放策略公式适用场景风险
线性缩放lr × kbatch 扩大在 4 倍以内大 batch 下初期震荡
平方根缩放lr × √k极大 batch 或 Transformer 类收敛偏慢
线性 + warmup先升后恒定大模型微调、检测模型需要调 warmup 步数

5. 落到具体任务:分类、检测、语音、大模型微调的起点值

概念讲完,得有能直接抄的起点。下面这些值是我在不同任务上反复用过的起步配置,具体还要按硬件微调。

5.1 图像分类与 ResNet 系列

拿 ResNet 系列做四分类花卉这类中小规模任务,我的起点是 batch size = 32、初始学习率 1e-3(SGD 时用 0.01)、训练 30~50 个 epoch。数据集如果有几万张,这个配置基本两三小时内能跑完。特别提醒:ResNet 里全连接层之前是全局平均池化,对 batch size 并不敏感,真正敏感的是每个 Block 里的 BatchNorm。所以一旦你把 batch 降到 8 以下,分类精度掉得会比想象中快,这时候要么改用 GroupNorm 版本,要么老老实实开梯度累积但把 BN 换成 SyncBN(多卡)或干脆冻结 BN 的统计量。

5.2 目标检测与 YOLO 系列

检测任务的 batch size 选择比分类更受图片分辨率牵制。YOLO 系列在 640×640 输入下,单卡 batch = 16 是个很稳的起点,显存 8GB 卡上把分辨率降到 416 也能跑 batch = 16。很多 YOLO 实现支持batch=-1自动探测显存能装下的最大值,这个功能在换卡、换分辨率时非常好用,但注意它探出来的是"能装下的最大 batch",未必是"训练效果最好的 batch",探完之后你可能还要往下收一档。检测任务里有个额外变量是正负样本比例,batch 太小时一张图里可能一个目标都没有,导致某一步全是背景负样本,梯度方向偏得厉害。所以检测我一般不会把 batch 压到 8 以下。

5.3 大模型微调与 LoRA

到了大模型这块,口径要换一下:通常不叫 batch size,而叫 global batch size(全局批大小),计算公式是per_device_batch × 梯度累积 × 卡数。7B 模型全量微调时代价太高,现在主流是 LoRA 之类的参数高效方法,可训练参数只有原来的千分之几,显存一下就宽松了。LoRA 微调的起点我一般用per_device_batch = 1~4、累积 8~16 步,从而让 global batch 落在 32~64 这个区间。原因是大模型优化对梯度估计质量要求更高,global batch 太小会让 loss 抖得看不清趋势,太大又会让单轮迭代变少、总步数不足。序列长度也是个大变量:同样的 batch size,序列从 512 拉到 2048,显存大概涨 3~4 倍,所以长文本微调往往要先把 batch 压到 1 再靠累积补回来。

6. 只在实跑中才会暴露的坑

前面都是可以提前算的账,真正折磨人的是那些算不出来、只能跑出来才发现的坑。

6.1 BatchNorm 和 batch size 的强绑定关系

这是我最想强调的一条。BatchNorm 在训练时用当前 batch 的均值和方差做归一化,所以 batch size 变小,这两个统计量的估计误差就变大,模型看到的输入分布和推理时(用滑动平均统计量)的分布偏差也随之变大。表现出来就是:训练 loss 看着还行,验证集指标却明显掉。判断方法很简单,把 batch size 从 32 改成 4 再跑一次,如果验证指标掉得超过两三个点,基本可以确定是 BN 的问题。解决办法要么换 GroupNorm、LayerNorm,要么把 BN 层设成 eval 模式并冻结(前提是预训练模型已经学好了统计量),要么把 batch 加回来。

6.2 DataLoader 里那几个容易被忽视的参数

num_workers、shuffle、pin_memory这三个参数看着不起眼,影响却很大。

  • num_workers决定用几个进程去读数据。设 0 就是在主进程里读,GPU 会频繁空等;设太大又可能把内存吃爆。经验值是从 CPU 核数的一半起,边跑边看 GPU 利用率上调。
  • shuffle=True在训练集上是必须的,否则模型会记住样本顺序,尤其当数据按类别排序时,前半个 epoch 全是同一类,训练直接崩。验证集则要shuffle=False,保证指标可复现。
  • pin_memory=True会把数据预先放到锁页内存,加快 CPU 到 GPU 的拷贝,几乎无脑开。

还有个小坑:shuffle和drop_last一起用时,每个 epoch 丢掉的样本不是固定的,这会让不同 epoch 的有效数据量有微小差异。数据量小的时候可能影响指标的可比性,我一般在做严谨对比实验时会固定随机种子,让 batch 划分在每个 epoch 都一致。

6.3 断点续训时 epoch、step 与调度器的对齐

训练中断后续训是最容易出错的场景。这里要对齐三样东西:模型参数、优化器状态、以及学习率调度器的内部步数。只保存模型参数是远远不够的——如果你用的是 cosine 或 step 调度器,续训时它得知道"我已经走到第几步了",否则学习率会从初始值重新开始衰减,等于人为制造了一次学习率跳变,损失曲线会出现一个明显的台阶。

state = { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict(), "epoch": epoch, "global_step": global_step, }

另外,epoch 与 iteration 在日志里的记录要一致。有些框架按 epoch 存 checkpoint,有些按 step 存,混用时会算错"还剩多久",也会让早停条件判断失准。我现在的习惯是统一以global_step为准,epoch 只用来对外汇报,这样即使中途改了 batch size 或数据集,续训逻辑也不会乱。

最后分享一个我自己用了很久的小习惯:开训前先在纸上或注释里写下这五个数——数据集大小、batch size、每 epoch 迭代数、总 epoch、预期总步数。跑起来之后跟日志对一遍,对不上就说明有地方(数据加载、梯度累积、drop_last)和你以为的不一样。这两年我在本地跑各种微调任务,凡是训到一半发现"怎么和预期差这么多"的情况,九成都是这五个数里有一个没对齐。把账先算清楚,比事后调参省下的时间多得多。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询