PyTorch中ViT模型融合注意力禁用及梯度/注意力图异常问题
ViT模型注意力图与梯度相关问题解答
背景
我在PyTorch中研究不同Vision Transformer(ViT)模型的差异,相关代码如下:
retrained_vit_weights = torchvision.models.ViT_B_16_Weights.DEFAULT # requires torchvision >= 0.13, "DEFAULT" means best available pretrained_vit_1 = torchvision.models.vit_b_16(weights=retrained_vit_weights).to(device) pretrained_vit_2 = torch.hub.load('facebookresearch/deit:main', 'deit_tiny_patch16_224', pretrained=True).to(device) for block in pretrained_vit_2.blocks: block.attn.fused_attn = False
我需要运行模型处理图像后返回各ViT层注意力图,以及自注意力后dropout层的梯度(参考vit-explain实现),但遇到两个问题:
- 注释DeiT模型的
fused_attn禁用代码后,注意力图和梯度为空; - torchvision的ViT_B_16默认无法获取dropout层梯度,且该模型默认使用
scaled_dot_product_attention优化实现,找不到类似fused_attn的属性。
1. 为何禁用fused_attn才能获取正常的注意力图和梯度?
fused_attn是DeiT中基于CUDA实现的融合注意力算子,它将注意力计算的多步流程(QKV投影、缩放点积、softmax、dropout)打包成单一底层CUDA kernel执行,以此提升性能。- 这种融合实现为了减少内存开销和加快速度,不会保留注意力权重这类中间计算张量,同时反向传播的梯度计算路径也被简化——底层kernel没有暴露这些中间变量,导致你无法提取注意力图,dropout层的梯度也无法被追踪到。
- 禁用
fused_attn后,DeiT会切换到PyTorch原生算子实现的标准注意力流程,每个步骤的中间张量都会被保留,自动微分的梯度路径完整,因此能正常获取注意力图和dropout层的梯度。
2. 如何在torchvision的ViT_B_16中禁用scaled_dot_product_attention?
有两种可行方案:
方案一:重写注意力层实现
手动替换ViT的注意力模块,用原生PyTorch算子实现标准自注意力流程,替代SDPA:
import torch from torchvision.models.vision_transformer import VisionTransformerEncoderBlock def replace_attention(block): class CustomAttention(block.attn.__class__): def forward(self, x: torch.Tensor) -> torch.Tensor: B, N, C = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) q, k, v = qkv.unbind(0) # 手动计算注意力,不使用SDPA attn = (q @ k.transpose(-2, -1)) * self.scale attn = attn.softmax(dim=-1) attn = self.attn_drop(attn) x = (attn @ v).transpose(1, 2).reshape(B, N, C) x = self.proj(x) x = self.proj_drop(x) return x, attn # 返回注意力权重用于后续提取 block.attn = CustomAttention(**block.attn.__dict__) # 遍历所有encoder block替换注意力层 for block in pretrained_vit_1.encoder.layers: replace_attention(block)
方案二:全局禁用SDPA优化
通过PyTorch的后端开关,强制SDPA使用原生数学实现,而非融合算子:
torch.backends.cuda.enable_flash_sdp(False) torch.backends.cuda.enable_math_sdp(True)
注意:这是全局设置,会影响所有使用SDPA的模型,且不同PyTorch版本的参数可能略有差异。
3. 实现差异导致异常的原因是什么?
异常的核心是融合优化算子与标准PyTorch算子的设计目标冲突:
- 融合算子(DeiT的
fused_attn、torchvision ViT的SDPA)的核心目标是提升性能,通过合并计算步骤、减少内存读写来加速,但代价是牺牲了中间变量的可访问性和梯度追踪的灵活性——它们不会暴露注意力权重这类中间结果,反向传播时也可能跳过部分梯度的存储或计算,导致无法提取注意力图和dropout层梯度。 - 标准PyTorch算子实现的注意力流程,每一步计算都会保留独立张量,自动微分的梯度路径完整,所有中间变量都可被访问,这也是vit-explain这类工具依赖的基础实现方式。
- 此外,不同库的ViT设计细节不同:DeiT为
fused_attn预留了开关,可随时切换到标准实现;而torchvision的ViT默认直接使用SDPA,没有内置类似的开关,因此需要手动修改或全局设置来禁用优化。
内容的提问来源于stack exchange,提问作者Jose Ramon
相关产品推荐
相关产品推荐

