单张图能跑、换成批量就撑不住,先别急着把 batch_size 调小或调大。显存瓶颈和数据加载瓶颈的表现不一样:显存瓶颈通常是批量涨到某个值后直接 OOM,或者因为显存回收不干净、碎片增多,让单步耗时突然变陡;数据加载瓶颈则是显存还剩不少,单步耗时却随批量上升得不成比例,GPU 利用率上不去。要分清这两类,靠的是在同一时间轴上采样的显存峰值与单步耗时,而不是先改参数再看手感。
先把 batch=1 的显存峰值和端到端耗时固定成基线,再按 1、2、4、8、16 逐档递增,每档记录峰值显存与单步耗时,直到报 OOM,记下最后成功的批大小。若显存先封顶、耗时随批量近似线性增长,优先考虑降批大小或降输入尺寸;若显存仍有余量而耗时异常放大,先查读图与预处理是否落在主线程、是否在关键路径上等待。判断要结合本机显卡、驱动和框架版本确认,换页、缓存这类止血手段不当作提速方案。
固定单图基线:记录显存峰值与单张耗时
基线的作用是让后面每一档批量都有可比对的参照。采样时机上,建议先跑若干张把 CUDA 上下文初始化、模型加载和算子自动调优的噪声排掉,再开始正式计时;如果第一次计时就记录,耗时往往偏大,会让后面的批处理曲线看起来“变快了”。耗时口径要写清楚:端到端包含读图、预处理、前向和后处理,另外单独记一列纯前向时间,否则读图慢会被误算成模型慢。用 GPU 异步执行时,计时前必须显式同步,不然测到的是任务入队时间。
import time
import torch
def measure(step_fn, iters=10, warmup=3):
for _ in range(warmup):
step_fn()
torch.cuda.synchronize()
torch.cuda.reset_peak_memory_stats()
t0 = time.perf_counter()
for _ in range(iters):
step_fn()
torch.cuda.synchronize()
dt = (time.perf_counter() - t0) / iters
peak = torch.cuda.max_memory_allocated() / 1024 ** 3
return dt, peak
把 step_fn 分别传成“读图+预处理+前向+后处理”和“纯前向”,就能拆出读图占了多少。外部还可以用 nvidia-smi `--query-gpu`=memory.used,utilization.gpu `--format`=csv -l 1 轮询观察显存和利用率随时间的变化。基线记录表建议至少包含这些字段:
| 字段 | 填写内容 | 说明 |
|---|---|---|
| 模型与权重版本 | 如 lingbot-vision / 权重文件名 | 换版本后要重测 |
| 输入尺寸与精度 | 如 1024x1024 / fp16 | 分辨率直接决定显存 |
| 批大小 | 1 | 基线档 |
| 读图 / 预处理 / 前向 / 后处理耗时 | 毫秒 | 四段分开记 |
| 端到端单张耗时 | 毫秒 | 预热后取中位数 |
| 峰值显存 | GB | allocated 与 reserved 都记 |
批量递增扫描直到报显存不足
递增方式建议按 1、2、4、8、16 逐档走,不要跳档,跳档会定位不到报错边界。每个档位重复跑几轮,取中位数或最大值,同时记录:批大小、峰值显存、整个 batch 的单步墙钟耗时、折算后的每样本耗时、是否 OOM。这样得到的是一条显存随批量增长的曲线,以及一条耗时随批量变化的曲线。
有一处容易踩坑:OOM 之后显存缓存未必立刻归还,如果在同一进程里继续往上加批量,测到的是被污染的状态。更稳妥的做法是每档起一个独立进程,让测量干净。可以用下面的方式批量循环:
import subprocess
import sys
for bs in [1, 2, 4, 8, 16, 32]:
r = subprocess.run(
[sys.executable, 'bench.py', '`--batch-size`', str(bs)],
capture_output=True, text=True)
print(bs, r.returncode, r.stdout.strip()[-200:])
返回码非零基本就是这一档没跑过,把它前面的成功档位记成边界。如果边界附近的档位是显存刚好顶到上限,说明是容量问题;如果边界远高于实际使用的批量,而耗时已经明显变差,问题多半不在显存。
区分显存瓶颈与数据加载瓶颈
两条曲线放在一起看,判断方向比较清楚。显存随批量近似线性上升并很快封顶,同时每样本耗时基本持平或缓慢下降(批处理把固定开销摊薄了),这是容量限制,动批大小或输入尺寸更直接。反过来,显存还有明显余量,单步耗时却不随批量下降甚至上升,每样本耗时反而变差,通常说明批量里的数据没有及时送上来,GPU 在等 CPU。
一个简单的对照实验:在读图环节人为加长一段延迟,比如在读图函数里插入几十毫秒 sleep,再跑同一档批量。如果总耗时随注入的延迟等量增加,说明读图处在关键路径上,批处理掩盖不了它;如果总耗时几乎没变,说明读图已经被并发预取盖住了,瓶颈更可能在计算侧。这个 sleep 只是用来做对照,不是优化手段,验证完要删掉。
def load_one(path):
img = read_image(path)
time.sleep(0.05) # 对照实验:验证读图是否在关键路径
return img
调整读取与预处理位置的观察
如果对照实验指向数据管线,通用做法是把解码和预处理从主线程挪出去,让它们和 GPU 计算重叠。骨架大致是:多进程加载器负责读图和 resize,主进程只做张量搬运和前向;开启固定内存和预取,减少每次迭代的等待;worker 数量先按 CPU 核数的一半试,再上下调,观察每样本耗时的变化。注意把预处理放在 GPU 上并不一定更快,它会占显存并与推理抢算力,需要结合显存余量判断。
loader = DataLoader(
dataset,
batch_size=bs,
num_workers=4,
pin_memory=True,
persistent_workers=True,
prefetch_factor=2,
)
改动前后要在同一批大小、同一输入尺寸下对比,记录每样本耗时和 GPU 利用率。常见的情况是:读完图这一步的耗时下降,每样本耗时随之下降,说明之前确实卡在数据侧;如果每样本耗时几乎不变,说明数据管线不是主要矛盾,可以回到批量与显存那一侧继续调。
给出可复现的吞吐记录表
结论要能被别人复算,记录格式就得把变量写全。建议按下面的字段整理,每一行是一档配置:
| GPU 型号 / 显存 | 驱动与框架版本 | 输入尺寸 | 精度 | 批大小 | 读图配置 | 峰值显存 | 单步耗时 | 样本每秒 |
|---|---|---|---|---|---|---|---|---|
| 按实际填写 | 按实际填写 | 如 1024x1024 | 如 fp16 | 1 / 2 / 4 / 8 | 如 workers=0 或 workers=4 | GB | 毫秒 | 由批量 ÷ 单步耗时折算 |
复算时固定同一份数据子集、先预热、跑固定轮数取中位数,并注明读图是串行还是多进程预取。同一张表里如果 workers=0 的行耗时明显偏高,而显存曲线还很宽松,那答案就是数据管线,而不是显存不够。