Audio-Visual Flamingo 这类多模态模型出现内存溢出时,先别急着换小模型,先把峰值来源找出来。它同时处理文本、图像和音频,输入 token 会拼接成长序列,显存大头通常落在中间激活和注意力缓存上,而不是模型权重本身。排查时先确定是训练还是推理,再按“缩小输入、减小 batch、开启精度、检查点、清缓存”的顺序处理。
多模态模型的内存溢出,常见原因不是模型参数太多,而是长序列输入和中间激活占用叠加。建议按“定位峰值来源 → 缩短输入/批次 → 开启混合精度与梯度检查点 → 清理碎片和残留进程”四步处理。先区分训练与推理,再动手改配置,避免一上来就换小模型。
先判断是训练还是推理,再定位显存峰值
训练时的 OOM 多发生在 backward,优化器状态和梯度会额外占用一片显存;推理时的 OOM 则基本来自 forward 过程中的激活缓存和长序列注意力分数。用 nvidia-smi -l 1 实时看显存曲线,或者直接在代码里打印峰值:
torch.cuda.reset_peak_memory_stats()
# 跑一个 step
print(torch.cuda.max_memory_allocated() / (1024**3))
- 如果 batch size 设为 1 后不再 OOM,说明梯度累积或批量叠加是主因,先减 batch 并配合梯度累积。
- 如果 batch size 为 1 仍 OOM,说明输入序列长度或模型前向本身占用过高,直接检查输入分辨率和帧率/采样点。
- 如果只在加载数据时 OOM,可能是 DataLoader 的 num_workers 开得过高或内存碎片,和模型关系不大。
从输入长度和 batch size 入手,这是最直接的调整
图像、音频和文本会被编码成 token 后拼到一个序列里。视频的帧数、音频的采样长度、文本的 prompt 长度,任何一项拉长,注意力矩阵都会按平方上涨。先把单段输入长度切掉一半,或把视频帧率降下来,通常比减小 batch 更明显。如果 batch size 必须保持,可以使用梯度累积:
# 原 batch size 为 16,显存只能跑 4,则 accumulation_steps = 4
for step, batch in enumerate(dataloader):
loss = model(batch) / 4
loss.backward()
if (step + 1) % 4 == 0:
optimizer.step()
optimizer.zero_grad()
注意:这种写法会保持“有效 batch”不变,但中间计算峰值仍是单次 batch 的峰值。验证方式是看完整训练曲线是否与原先一致。
开启混合精度和梯度检查点
如果你的环境支持 CUDA 且 PyTorch 版本较新,可以将大部分计算切到 bf16 或 fp16,显存峰值会明显下降。梯度检查点则是把 forward 中的激活丢弃,到 backward 时重新计算,直接降低同时存储的激活量。以 PyTorch 为例:
model.gradient_checkpointing_enable()
scaler = torch.amp.GradScaler('cuda', enabled=torch.cuda.is_available())
for step, batch in enumerate(dataloader):
with torch.amp.autocast('cuda', dtype=torch.bfloat16):
loss = model(**batch).loss
scaler.scale(loss / accumulation_steps).backward()
if (step + 1) % accumulation_steps == 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
torch.cuda.empty_cache()
这里的 torch.cuda.empty_cache() 只在显存确实需要回收时使用,不要每步都调。开启梯度检查点后训练时间会变长,但显存峰值通常明显下降;如果你用的是多机多卡,需要先确认模型是否支持 gradient checkpointing,不是所有自定义实现都有这个接口。
检查显存碎片、残留进程和缓存分配器
有时候 OOM 并不是模型占用太大,而是显存被碎片化或旧进程占用。运行 nvidia-smi 看有哪些进程还在占用,如果之前崩过,会有遗留 Python 进程,需要 kill 掉。PyTorch 环境变量 PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True 可以降低碎片率,但需要你的 CUDA 版本支持;设置后重启进程,再重新跑同一个 batch。
验证方式:用同一段数据跑 50 步,观察 nvidia-smi 的显存曲线是否逐渐升高。如果曲线只增不降,检查代码里是否在每个 batch 都保存了不需要的中间张量。
最后才考虑 offload 或换小模型
如果以上步骤都试过,仍然不够,就只能把部分计算挪到 CPU 或换更小的模型。例如 Hugging Face transformers 的 device_map='auto' 可以自动把部分层放到 CPU,但代价是每步的前向/反向都会增加数据传输时间,多机多卡场景会更明显。这不是内存溢出的正解,而是一种兜底。
同样,换小模型不是首选,因为精度和任务表现都需要重新验证。先明确当前任务的指标基线,再决定要不要降级。