训练一启动就报 CUDA out of memory,先不要同时改批次大小、序列长度和精度。比较稳的顺序是:先用训练日志确认爆显存发生在哪个阶段,用显存采样留下一个峰值基线;然后固定序列长度不动,先减小批次大小,用梯度累积补回等效批量;批次大小降到 1 仍然不够,再考虑缩短序列长度,或者开梯度检查点、换低精度。这样安排的原因不是批次大小更“重要”,而是它改动代价小、接近线性、好回退,出现问题时也容易归因。
显存不足时通常先动批次大小,而不是序列长度:批次大小对显存的影响近似线性,对训练语义的影响相对小,配合梯度累积还能保住等效批量。序列长度会同时改变注意力计算量、位置编码范围和全量激活大小,属于第二杠杆,改之前要确认数据切分和位置编码仍然成立。这一切的前提是先定位报错阶段并记录峰值显存;没有基线数值,任何改动都无法判断是参数起效还是环境抖动。
在训练日志里确认显存报错发生在哪个阶段
PyTorch 的显存报错文本通常长这样:torch.cuda.OutOfMemoryError: CUDA out of memory. Tried to allocate X GiB ... GPU 0 has a total capacity of ...。光看这行不够,关键是看报错前最后几行属于哪个环节。读法是:报错前最后一行还在读数据集、做 tokenize 或拼 batch,说明压力来自数据加载与预处理;最后一行停在模型 forward 或某个层名,就是前向阶段;已经出现过 backward 相关栈,则是反向阶段连同优化器状态一起占满了显存;训练循环结束、进入 eval 或 generate 才报错,多半是采样阶段的 KV cache 随生成长度增长。
再看报错时对应的步数。第 0 步或前十步内就爆,一般是静态配置问题,批次大小、序列长度、精度三者之一超了;训练到中途(比如几百步之后)才爆,要考虑序列长度是否被动态拉长、缓存是否随步数累积、是否同时跑了验证或定期保存。
用显存监控命令记录峰值占用
先把手感换成数值。终端侧按固定间隔采样,推荐把输出直接落盘,方便和训练日志对齐:
nvidia-smi `--query-gpu`=timestamp,memory.used,utilization.gpu \
`--format`=csv,noheader,nounits -l 5 | tee logs/run003_mem.csv
脚本侧更适合记录单卡进程的真实分配量,因为它不受其他进程干扰。在训练步里按 log_every 打印一次,并在打印后重置峰值统计,否则峰值会一直停留在历史最大值:
import torch
if step % log_every == 0:
alloc = torch.cuda.max_memory_allocated() / 1024 ** 3
reserved = torch.cuda.max_memory_reserved() / 1024 ** 3
print(f'step={step} peak_alloc={alloc:.2f}GB peak_reserved={reserved:.2f}GB')
torch.cuda.reset_peak_memory_stats()
max_memory_allocated 是张量实际占用,max_memory_reserved 是缓存分配器向驱动申请的池子,两者差距大通常意味着碎片或预留过多。对齐方式很简单:让脚本打印和 nvidia-smi 采样使用同一段训练过程,以脚本的 step 为主轴,把采样时间戳落在两次打印之间即可。启动采样命令的时刻建议写进日志文件头部,避免事后对不上时间轴。
按批次大小、序列长度的顺序做单变量对照
每轮只改一个参数,其他全部冻结。第一轮固定序列长度、精度和数据,把 device_batch_size 依次减半,直到能稳定跑过若干步;如果目标是复现原来的训练动态,减批次的同时按比例增加 grad_accum_steps,让每步的等效批量大致不变。第二轮固定已经跑通的批次大小,再尝试把 max_seq_len 缩短一档,确认注意力部分的激活是否才是主要占用;注意缩短序列长度会改变样本切分与位置编码覆盖范围,要同步检查数据侧配置。
每轮记录这些字段,缺一项后面就无法对比:run_id、device_batch_size、max_seq_len、dtype、grad_accum_steps、是否开梯度检查点、峰值显存(alloc / reserved)、是否溢出、报错阶段、以及前若干步的 loss 走向。对比方式是横向看同一 run_id 里只动过的那一列,纵向看峰值显存是否真的降低、以及损失是否还在正常下降。如果峰值没降而只是没报错,说明只是踩在边缘上,换一个数据批次就可能再爆。
用梯度累积或低精度设置换取能跑通的配置
不换硬件的前提下,先让训练跑起来比追求最快更重要。下面是一份通用配置骨架,参数名以你实际训练脚本为准,不要照抄字段名:
device_batch_size: 4 # 先降到能跑通的值
max_seq_len: 1024 # 第一轮保持不动
grad_accum_steps: 8 # 让 grad_accum * device_batch 接近原等效批量
dtype: bfloat16 # 或改用 autocast,而不是整模型转半精度
gradient_checkpointing: true
改动之后必须验证损失仍在正常下降,而不是只看它没崩:打印每步或每若干步的 loss,和之前的基线在相同 token 数口径下比走势;确认没有出现 NaN 或长期不降。梯度检查点通常以增加计算时间换取激活显存下降,属于止损手段,不是性能优化;bfloat16 与 float16 的动态范围不同,如果开启 fp16 后出现溢出,优先检查 loss scale 相关设置是否配套。梯度累积只改变每步参与前向的样本数,等效批量变大时优化器步数变少,学习率是否需要同步调整,需要结合环境确认。
把成功与失败配置整理成对照表
把每次尝试收敛成一张表,下次复现时直接查,而不用重新试错。可以用最朴素的 CSV 保存,字段固定下来:
run_id,device_batch_size,max_seq_len,dtype,grad_accum,grad_ckpt,peak_alloc_GB,peak_reserved_GB,oom,oom_stage,loss_trend
run001,8,2048,bfloat16,1,false,,,yes,forward,-
run002,4,2048,bfloat16,2,false,___,___,no,-,平稳下降
run003,4,1024,bfloat16,2,false,___,___,no,-,平稳下降
run004,4,2048,float16,2,true,___,___,no,-,需确认 loss scale
配套的保存方式是把配置单独写成文件,并把训练输出和显存采样按同一个 run_id 落盘,例如把配置存为 configs/run003.yaml,再用 python train.py `--config` configs/run003.yaml 2>&1 | tee logs/run003.log 执行,同时把 nvidia-smi 采样写入 logs/run003_mem.csv。这样“参数取值、峰值显存、是否溢出、损失趋势”四项能在同一目录下互相印证,也便于把某个失败配置整体复现或整体废弃。