去年年底我在一台 8 卡 A100 的机器上训一个 ViT-B/16,跑了两天,loss 降得比预期慢一半。第一反应是模型有问题——把 attention 换成 FlashAttention-2,重新跑,没区别。后来开 nvtop 盯着看,8 张卡的利用率在 30%~40% 之间来回跳,跟心电图似的。也就是说我花了两天钱,其中六成时间 GPU 在等 CPU 喂数据。
这事我不止踩过一次。2021 年做 OCR 的时候也一样,当时以为是 batch size 太小,从 32 调到 128,结果显存爆了,利用率还降了。所以现在我的习惯是:训练脚本写完,第一件事不是看 loss,是看 GPU 利用率。低于 80% 就别急着调模型,先查数据管道。
先算一笔账,再决定优化什么
判断是不是数据侧瓶颈,其实不用什么高级工具,做个除法就行。假设单张图在 CPU 上完成 decode + resize + 增强要 6ms(这个数我很保守了,500x500 的 JPEG 用 PIL 解码大约 1.5ms,配合 torchvision 的 RandomResizedCrop 到 224 再加 ColorJitter,实测能到 5~8ms)。8 个 DataLoader worker 并行,理论上每秒能产出 8 / 0.006 ≈ 1333 张图。
再看模型这边:ViT-B/16 在 A100 上,batch size 64,fp16 + 单步前反向大约 45ms(这个数会随序列长度和具体实现浮动,但量级差不多)。也就是每秒需要 64 / 0.045 ≈ 1422 张图才能喂饱。
1333 < 1422。差这不到 100 张,就是利用率上不去的原因。
我是用 torch.profiler 确认的,profile 里 dataloader 那段占比 58%,模型 forward 只占 19%。py-spy dump 到 worker 进程上,栈里全是 PIL 的 C 函数。这两个工具装起来都不麻烦,pip install py-spy,然后 py-spy dump --pid <worker_pid> 就行。
我试过的三种做法,和它们的代价
第一种:把 num_workers 从 0 调上去。我的经验是 8 到 12 是甜点区,具体看你 CPU 核数。有一次我调到 32,利用率反而从 78% 掉到 61%,原因是每个 worker 都要复制一份 Python 解释器和数据集索引,内存拷贝和锁竞争把收益吃掉了。另外记得配 persistent_workers=True,否则每个 epoch 结束 worker 全销毁重建,如果你的 Dataset.init 里有扫目录的操作,100 万张图扫一次要 30 秒,10 个 epoch 就是 5 分钟白扔。还有个小坑:pin_memory=True 在 num_workers=0 的时候基本没意义,反而可能卡住主线程。
第二种:把增强搬到 GPU 上。kornia 或者 NVIDIA DALI 都行。代价是显存,DALI 的 pipeline 大概多吃 1~2GB,另外调试起来是真难受,报错信息经常只给你一句 pipeline 内部错误。我后来只在 decode 这一层用 DALI,增强还是留在 CPU。
第三种,也是我现在最常用的:预处理。训练前把所有图统一 resize 到 256x256,重新存成 JPEG(quality 85),打包成 WebDataset 的 tar 分片,每片 1000 张左右。存储涨了一倍多,但 epoch 时间从 42 分钟降到 17 分钟,利用率稳定在 94%。代价是你没法再做那种依赖原图的增强,比如随机裁剪后再看重细节,损失了一定精度,我这边大概掉了 0.3 个点的 top-1。这个取舍值不值,看你数据量和算力预算。
一个不太受欢迎的观点
我不太赞成“单卡利用率没到 80% 就去上 DDP”这个做法。DDP 解决的是多卡通信和梯度同步,它不会帮你解决 CPU 侧喂不满的问题。我在两台机器上见过同一个错误:单卡利用率 40%,上了 4 卡 DDP,每张卡利用率变成 15%,总吞吐只涨了 50%,比线性加速差远了——因为数据管道还是那一个,4 张卡抢一个瓶颈。
顺序应该是:先单卡把利用率干到 85% 以上,再上多卡。诊断工具也不用多高级,nvtop 看利用率,htop 看 CPU 有没有打满,py-spy 看 worker 卡在哪,三个够了。torch.compile 在 PyTorch 2.0 之后确实能提速,但它主要作用在模型图上,对 DataLoader 那一段基本没影响,别指望它。