最近在调一个nlp分类模型,数据量不大(大概2万条),单卡3090训练。试了torch.compile,用默认模式跑,第一轮确实慢(编译开销),但后面几个epoch速度只提升了10%左右,和官方宣传的“训练提速30%-50%”差很多。而且用动态shape(比如变长序列padding到不同长度)直接报错,回退到eager模式。想问问大家:是不是小模型小数据量根本没必要上compile?还是我哪里设置不对?有一说一,模型里用了transformers的BertForSequenceClassification,加了个自定义分类头。有没有类似场景的老哥分享下实际收益?现在纠结要不要为了部署时的推理加速去折腾这个。
PyTorch 2.0的compile到底值不值得用?静态图加速不明显还报错
全部回复
共 88 条你这情况跟我之前调bert-base时挺像的,小模型加小batch确实吃不满compile的红利,官方那个30%-50%多半是大模型大batch下的benchmark。动态shape报错基本无解,要么padding到固定长度,要么干脆只在推理时开compile,训练阶段收益真不大。我后来试了inductor的mode=max-autotune,虽然编译更久但能多挤几个点,你可以试试。另外自定义分类头如果太简单,其实瓶颈都在transformer那块,compile优化不到啥,不如把精力放在数据加载和混合精度上。
2w条数据单卡3090,这个量级compile的编译开销摊不平收益,10%其实算正常了,官方那30%-50%多半是大模型大batch或者CV场景跑出来的。动态shape报错太常见了,transformers里很多op本来就不友好,我试过把padding改成固定长度再配合reduce-overhead能稍微好点,但提升也就那样。你要是主要纠结部署推理,不如直接上ONNX或者TensorRT,那个收益比compile实在多了。
说实话你这个场景我太熟了,之前用electra做类似的分类任务也是这德行,2万条数据在3090上根本喂不饱显卡,compute bound都没到,compile优化的是kernel launch和显存带宽,小batch下收益自然被稀释。动态shape那个报错确实无解,torch.compile对变长序列的support一直很迷,就算你padding到固定长度,内部如果有多分支或者python控制流,它还是会回退。我个人建议是小模型直接别折腾compile,把精力放在gradient accumulation和混合精度上,收益来得更直接,我试过fp16加上多步累计,整体训练时间能压掉将近一半。至于部署推理,如果你不是追求极致延迟,其实纯eager加torch.inference_mode就够了,真要上compile也得等模型结构完全冻结,而且得用fullgraph=True去逼它做整图优化,默认模式那点优化幅度确实不值得换那堆报错。另外你说transformers的模型,可以试试把自定义分类头单独提出来compile,backbone保持原样,有时候这样能避开不少坑,但说实话提升也就那样。反正这玩意儿现阶段更像是给大模型或者CNN那种固定shape场景准备的,NLP小模型拿它性价比确实不高。
小模型真没必要折腾compile,那点提升还不够调bug的时间,推理时直接上onnx或者TensorRT更香。
跟你情况差不多,之前试过在bert-base上开compile,小数据集下收益确实就那样,10%左右算正常,官方那个30%-50%估计得大模型加静态shape才跑得出来。动态padding那个坑我也踩过,后来干脆固定长度padding到512,虽然浪费点显存但至少不报错,速度还稳一点。你要是主要纠结部署推理,不如直接上onnx或者tensorrt,那个提升比compile明显多了,训练阶段真没必要折腾。
这场景我太熟了,跟你差不多配置,2万条数据上compile属实有点鸡肋。提速10%算正常,官方那30%-50%多半是CV大模型或者动态shape不严重的理想情况。动态shape报错基本无解,建议直接关掉或者用dynamic=False硬凑,不然编译缓存反复失效反而更慢。你这规模其实把batch size调大点、梯度累积搞上,收益比折腾compile实在。推理端倒是可以单独试试,毕竟部署时静态shape多,提速比训练明显。
你这情况太真实了,我拿差不多的模型试过,2万条数据加3090,compile收益基本就10%左右,官方那个30%-50%得看模型规模和batch大小,小模型根本吃不满。动态shape报错也是老毛病了,别硬上,把padding固定到最大长度或者用bucket分桶能缓解一点,但提升也就那样。如果只是训练,没必要折腾,部署推理时用onnx或者TensorRT反而更稳,收益也更明显。
你这情况我太熟了,2万条数据加bert这级别真没必要折腾compile,官方那30%-50%都是大模型大batch堆出来的,小模型光算子融合省下的那点时间还不够填编译和显存调度的坑。动态shape报错更是常态,transformers里一堆带条件的tensor操作,默认模式根本兜不住,我建议要么固定长度padding要么直接放弃。推理阶段倒是可以试试,毕竟部署不吃训练那套动态逻辑,但收益也就那样。