最近把之前写的一个图像分类项目从PyTorch 1.13迁移到2.0,听说torch.compile能白嫖加速,就试着给模型加了个@torch.compile。结果发现训练速度反而慢了20%,报错还一堆,比如什么“dynamic shape not supported”。我用的就是标准ResNet50,输入尺寸固定,数据加载也没啥花活。是不是我姿势不对?还是说这玩意儿只对某些特定场景有效?求有经验的老哥指点一下,到底该怎么用compile才能不翻车?
PyTorch 2.0的torch.compile到底能不能直接加速老项目?踩坑了
全部回复
共 152 条小模型和CNN上compile收益真不大,反而有额外开销,试试大模型或transformer结构才划算。
说实话torch.compile在ResNet这种CNN上收益真的不大,它主要红利在Transformer和动态图结构上,我试过EfficientNet也是差不多甚至变慢。你那个dynamic shape报错大概率是dataloader里有个别tensor维度不是严格固定,比如最后一批batch size不同,建议把drop_last=True开起来或者给输入加个pad。另外编译本身有预热开销,建议先跑几十个step热身后再计时,小数据集上直接对比很容易被启动时间误导。想省心的话可以先试试torch.compile(mode="reduce-overhead"),对CNN老模型兼容性会好不少,实在不行就只在eval模式开inference模式编译,训练还是老实关掉吧。
我之前也踩过这个坑,torch.compile对动态shape特别敏感,你虽然觉得输入固定,但模型里如果有python控制流或者某些op会触发重编译,反而拖慢速度。建议先用torch.compile的mode="reduce-overhead"试试,再关掉cudagraphs看看,或者直接用torch._dynamo的日志看哪里被fallback了。另外老项目里那些自定义层和旧版API经常不兼容,最好先把模型改成纯nn.Module原生写法再说加速的事。反正我最后是只在推理阶段用compile,训练还是老老实实关掉,收益明显但坑太多。
另一个风格:这问题我也遇到过,ResNet50理论上是最适合compile的场景了,慢20%大概率是编译开销没摊薄,或者某些算子没融合。你试试把batch size调大点,或者用torch.compile(model, mode="max-autotune")多跑几个epoch看平均时间,别只看前几步。报错dynamic shape一般是因为有个别地方用了可变长操作,比如自适应池化或者view(-1)这种,建议先用静态shape的dummy input跑通再换真实数据。我自己的经验是,compile对GPU利用率高的任务效果明显,CPU瓶颈或者数据加载瓶颈的话确实会负优化,先profile一下再决定要不要开。
torch.compile对训练场景的加速确实不如推理来得明显,尤其是小batch或者GPU没跑满的时候,编译开销和graph break反而会拖后腿。你那个dynamic shape报错,八成是某些op在trace时被判定为动态了,比如数据加载里偶尔出现的tensor形状抖动,哪怕只有一次也会触发。我自己的经验是,先别急着全模型compile,只包住backbone或者耗时最长的block试试,同时把torch.compile的mode调成max-autotune跑一遍看真实耗时。另外建议看一眼编译后的统计报告,用torch._dynamo.explain定位具体哪里break了,很多问题其实是第三方库或自定义loss引起的。
说实话你这情况我真遇到过,当时也是信心满满加了compile结果直接给我整不会了。后来翻了下GitHub讨论才发现,torch.compile对静态shape和GPU利用率高的场景收益才明显,ResNet50这种CNN反而容易因为graph break和CUDA graph的额外开销拖慢速度,尤其在小batch下几乎必亏。我建议你先检查下是不是开了模式默认的reduce-overhead,这选项在Ampere以下架构上反而有负优化,试下mode="max-autotune"或者干脆把dynamic=True去掉固定死shape。还有个坑是数据加载的pin_memory和non_blocking跟compile的缓存策略偶尔会冲突,你可以先关掉pin_memory跑一轮对比下。要是还慢的话,用torch._dynamo.explain看下graph break都断在哪,大概率是你自定义的loss或者预处理里有python控制流,把那部分改成tensor操作或者挪到compile外面就行。最后说句实话,如果显存不吃紧、训练规模没到那种需要榨干算力的程度,老项目真没必要硬迁移,收益可能还不如把torch.backends.cudnn.benchmark开着。
说实话你这个情况我太熟了,刚升2.0那会儿我也在ResNet上踩过一模一样的坑。torch.compile不是无脑加个装饰器就完事的,它默认会做很多激进优化,但你要是没给它足够的静态信息,它反而会在图编译和回退上浪费大量时间。你训练变慢20%很可能是因为每次迭代都在重新编译或者频繁触发guard检查,尤其是如果你用了dataloader的默认行为,哪怕输入尺寸固定,tensor的shape或者device属性稍有变化它都会觉得是dynamic shape。我建议你先试试在compile函数里加上mode="reduce-overhead"或者"max-autotune",同时明确指定dynamic=False,如果还有问题就检查一下是不是有python控制流或者自定义autograd.Function里用了不支持的API。另外老项目里如果混着很多原地操作或者非张量返回值,也可能导致图捕获失败,这时候可以先只compile模型的forward部分而不是整个训练step。说白了这玩意儿对计算密集、图结构稳定的模型收益最明显,像你这种标准ResNet其实应该能提速的,但前提是代码得符合它那套图编译的“洁癖”,如果项目里历史包袱多,那真不如先用torch.jit.script或者干脆别折腾,等后面版本兼容性更好了再上。
torch.compile对训练场景确实不友好,尤其是batch size小或者GPU没吃满的时候,图编译和优化开销反而会拖慢速度。我之前在检测模型上试过,只有开cudagraphs加max-autotune才有点提升,但调参过程比换分布式还折腾。你这情况建议先试试inference模式,或者检查下是不是dataloader的num_workers太低导致编译线程抢资源。另外2.0的compile对动态shape特别敏感,建议把模型输入用torch.zeros固定个shape跑一遍warmup再正式训练。反正我现在是只在推理部署才用它,训练还是老实用原生autocast+混合精度。
ResNet50这种静态图确实不是compile的主场,试试大模型或者动态shape场景,收益才明显。
老项目别无脑上,先关掉cudnn benchmark再试,小模型compile反而容易因图优化开销得不偿失。
我这边也遇到过类似情况,torch.compile对ResNet这种结构其实挺挑的。你试试加mode="max-autotune"看看,默认模式有时候反而会拖后腿。另外dynamic shape那个报错大概率是某个op触发的,可以先用torch._dynamo.explain跑一下定位问题。我第一次也是直接套上去就翻车,后来发现得配合warmup和固定batch size才有效果。
我前阵子也踩过同样的坑,ResNet50加compile后单卡训练反而慢了一截,后来发现是编译本身的开销没摊平。torch.compile默认走inductor后端,第一次跑会花大量时间做图捕获和kernel生成,如果你只跑几个epoch就停,那整体时间肯定被拖垮。而且它默认开dynamic shape,哪怕你输入固定,如果dataloader里最后的batch尺寸不一样,或者有随机resize之类的操作,它就会反复重编译。可以试试加mode="reduce-overhead"或者显式设dynamic=False,再把warmup拉长到几十个step看steady state。另外老项目里如果混用了numpy或者自定义cuda extension,图会断掉,加速就没了。真要白嫖的话建议先用torch.compile(model, fullgraph=False)跑一遍看graph break的日志,确认瓶颈到底在哪儿。
torch.compile这玩意儿确实不是无脑加个装饰器就能提速的,我一开始也踩过类似的坑。ResNet50这种经典结构按理说应该支持得不错,但关键得看你有没有触发graph break,比如数据加载里如果混了numpy操作或者自己写的python逻辑,编译图就会断成好几段,反而比eager还慢。dynamic shape那个报错一般是输入尺寸在某个地方变了,或者你用了mask之类的动态维度,固定尺寸的话可以试试在compile里指定dynamic=False。另外首次编译有开销,如果只跑几个epoch可能还没热身完就结束了,建议先跑个几十步看看稳定后的速度。我后来是把模型forward里的条件分支和print都清掉,再把optimizer那步也纳进去,才看到实际收益。你可以先用torch._dynamo.explain看看哪里断了图,比瞎猜管用。
我拿ResNet也试过,第一次同样翻车,后来发现是mode没设对,默认模式对小模型反而不友好。你试试torch.compile(model, mode="reduce-overhead"),再把dynamic=False显式加上,固定输入尺寸下效果会好很多。另外第一次编译那几十秒别算进速度里,跑几个epoch再看才有意义。