torch.compile结合flash-attention报错:无法追踪trampoline_autograd_fwd为单图
解决方案
问题根源是flash-attention的自定义Autograd Function无法被TorchDynamo正确解析,导致编译失败或回退到eager模式,以下是可行的解决办法:
1. 升级依赖版本
PyTorch 2.0早期版本与flash-attention的兼容性存在缺陷,升级到最新稳定版可解决大部分TorchDynamo追踪问题:
- 升级PyTorch到2.1及以上版本
- 升级flash-attention到2.4及以上版本
执行命令:
pip install --upgrade torch torchvision torchaudio pip install --upgrade flash-attn --no-build-isolation
2. 标记自定义Autograd Function为可追踪
将flash-attention中无法被追踪的函数标记为允许纳入Dynamo图中,在代码开头添加:
import flash_attn from torch._dynamo import allow_in_graph allow_in_graph(flash_attn.ops.triton.attention.trampoline_autograd_fwd) allow_in_graph(flash_attn.ops.triton.attention.trampoline_autograd_bwd)
若找不到上述路径,可直接包裹你调用flash-attention的自定义函数:
from torch._dynamo import allow_in_graph @allow_in_graph def my_flash_attention(qkv, **kwargs): return flash_attn.flash_attn_qkvpacked_func(qkv, **kwargs)
3. 调整torch.compile参数
尝试使用对自定义Autograd Function兼容性更好的模式:
model = torch.compile(model, mode="reduce-overhead")
若模型存在动态shape,添加dynamic=True参数:
model = torch.compile(model, dynamic=True)
4. 局部禁用编译
如果上述方法无效,可在调用flash-attention的部分临时禁用Dynamo,确保该部分走eager模式,不影响模型其他编译部分:
from torch._dynamo import disable def forward(self, x): # 其他编译友好的模型操作 ... # 禁用编译处理flash-attention with disable(): attn_output = flash_attn.flash_attn_qkvpacked_func(qkv, dropout_p=self.dropout) ... return output
内容的提问来源于stack exchange,提问作者Vaio
相关产品推荐
相关产品推荐

