最近在调一个7B的LLM微调,显存卡在24G边缘,听群里大佬说torch.compile能省显存还提速,就试了一下。结果compile之后显存直接飙到30G+,batch size被迫减半,速度也没快多少。我用的是默认模式,没加任何自定义backend。代码大概是这样:model = torch.compile(model),然后照常forward/backward。想问问大家,是不是compile默认会保留更多中间张量用于反向?还是说需要配合reduce-overhead或者max-autotune这类模式调优?另外,有没有可能是我模型里有动态shape操作(比如attention_mask导致的padding)导致graph break,反而增加了内存开销?求有实际部署经验的哥们儿指点一下,感激不尽。
楼主
6小时前
PyTorch模型用torch.compile后显存暴涨,是用法不对还是正常现象?
请 登录 后发表回复
全部回复
共 0 条暂无回复,快来抢沙发吧