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

如何在ViT中通过Hook获取attention值而非层的返回值?

如何通过Forward Hook获取ViT注意力层的attn值

当然可以,有两种实用方法:

方法一:给注意力层加虚拟中转模块

最规范的做法是修改注意力层的结构,在attn计算完成后,用一个虚拟的nn.Identity模块中转一下张量,这样就能给这个虚拟模块挂Hook来捕获attn值。

修改后的注意力层代码:

class Attention(nn.Module):
    def __init__(self, dim, num_heads=8, qkv_bias=False, attn_drop=0., proj_drop=0.):
        super().__init__()
        self.num_heads = num_heads
        head_dim = dim // num_heads
        self.scale = head_dim ** -0.5

        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
        self.attn_drop = nn.Dropout(attn_drop)
        self.proj = nn.Linear(dim, dim)
        self.proj_drop = nn.Dropout(proj_drop)
        # 新增虚拟模块,仅用于中转attn张量,不做任何运算
        self.attn_catcher = nn.Identity()

    def forward(self, x):
        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[0], qkv[1], qkv[2]

        attn = (q @ k.transpose(-2, -1)) * self.scale
        attn = attn.softmax(dim=-1)
        attn = self.attn_drop(attn)
        
        # 让attn经过虚拟模块,方便挂Hook
        attn = self.attn_catcher(attn)
        
        x = (attn @ v).transpose(1, 2).reshape(B, N, C)
        x = self.proj(x)
        x = self.proj_drop(x)
        
        return x

之后给虚拟模块注册Hook:

def get_activations(name):
    activations = {}
    def hook(module, input, output):
        activations[name] = output.detach()
    return hook, activations

# 注册Hook到虚拟模块
hook, attn_store = get_activations('attention_map')
h_attn = model.blocks[-1].attn.attn_catcher.register_forward_hook(hook)

# 前向传播后即可获取attn值
model(input_tensor)
print(attn_store['attention_map'].shape)  # 输出形状应为 (B, num_heads, N, N)

方法二:用猴子补丁临时修改forward方法(无需改原类)

如果不想重新定义注意力类,也可以用猴子补丁临时修改目标注意力层的forward方法,把attn存为模块的属性,再通过Hook读取:

attn_store = {}

# 定义Hook函数,读取模块的attn属性
def attn_hook(module, input, output):
    attn_store['attention_map'] = module.temp_attn.detach()

# 给目标注意力层注册Hook
h_attn = model.blocks[-1].attn.register_forward_hook(attn_hook)

# 保存原forward方法,替换成自定义的
original_forward = model.blocks[-1].attn.forward
def modified_forward(x):
    B, N, C = x.shape
    qkv = model.blocks[-1].attn.qkv(x).reshape(B, N, 3, model.blocks[-1].attn.num_heads, C // model.blocks[-1].attn.num_heads).permute(2, 0, 3, 1, 4)
    q, k, v = qkv[0], qkv[1], qkv[2]

    attn = (q @ k.transpose(-2, -1)) * model.blocks[-1].attn.scale
    attn = attn.softmax(dim=-1)
    attn = model.blocks[-1].attn.attn_drop(attn)
    
    # 把attn存为模块的临时属性
    model.blocks[-1].attn.temp_attn = attn
    
    x = (attn @ v).transpose(1, 2).reshape(B, N, C)
    x = model.blocks[-1].attn.proj(x)
    x = model.blocks[-1].attn.proj_drop(x)
    
    return x

model.blocks[-1].attn.forward = modified_forward

# 前向传播获取结果
model(input_tensor)
print(attn_store['attention_map'].shape)

第一种方法更稳定,适合长期使用;第二种方法无需改动原模型类,适合快速调试。

内容的提问来源于stack exchange,提问作者Mitch

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 13:05:07