You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 14:56:13