Phi-3模型注意力层提取与可视化的正确方法及疑问
哈喽,我来帮你梳理清楚这个问题——你当前提取注意力的方式确实有问题,这也是后续维度混乱、可视化异常的核心原因,咱们一步步来解决:
一、先把注意力提取的方法改对
Phi-3是纯Decoder架构的模型,默认调用model.model(input_ids)的时候,不会主动返回注意力权重。你必须显式加一个参数output_attentions=True,才能让模型把每一层的注意力都输出给你。修改这部分代码:
# 关键:添加output_attentions=True,让模型返回注意力权重 model_output = model.model(input_ids, output_attentions=True) # 现在提取的attentions是一个tuple,每个元素对应模型某一层的注意力 attentions = model_output.attentions
现在给你拆解下这个attentions的结构,你就明白为什么之前的维度不对了:
attentions的长度等于模型的总层数(比如Phi-3-medium有40层,mini版是12层,你可以用print(model.config.num_hidden_layers)确认)- 每个tuple元素(也就是某一层的注意力),形状是
(batch_size, num_heads, seq_len, seq_len)batch_size:你的输入批次大小(这里是1,因为你只输入了一个prompt)num_heads:该层的注意力头数量(Phi-3-medium是16个,mini是12个,用print(model.config.num_attention_heads)可以查)seq_len:你输入prompt的token数量(比如你输入的prompt转成token后是15个,那这里就是15)
举个例子,如果你想拿第5层(从1开始数)的注意力,就用attn_layer = attentions[4](编程里索引从0开始),这时候attn_layer的形状是(1, 16, 15, 15)(假设16头、15个token)。而你要的n_tokens × n_tokens矩阵,就是某一个头的注意力:比如第3个头,就取attn_layer[0, 2, :, :],这就是标准的注意力矩阵了。
二、你之前提取的attention变量为啥维度奇怪?
你之前用model_output[-1]拿到的内容,根本不是注意力权重——大概率是模型最后一层的隐藏状态或者其他中间输出,和注意力完全不沾边,所以才会出现1x40x40x15x15这种不符合预期的维度,这一步必须改过来。
三、调整你的可视化代码,让它正常工作
基于正确提取的注意力,我给你调整了可视化代码,适配Phi-3的因果注意力(纯Decoder模型的注意力是带mask的,下三角全为0,因为模型不能“偷看”未来的token):
import matplotlib.pyplot as plt def save_attention_image(attentions, tokens, layer_idx, filename='attention.png'): """ 可视化指定层的所有注意力头 :param attentions: 从model_output提取的attentions tuple :param tokens: 输入prompt对应的token列表(用tokenizer.convert_ids_to_tokens转) :param layer_idx: 要可视化的层索引(从0开始) :param filename: 保存的图片文件名 """ # 取出指定层的注意力,转成numpy数组 attn_layer = attentions[layer_idx].detach().cpu().numpy() batch_size, num_heads, seq_len, _ = attn_layer.shape # 处理Phi-3 tokenizer的特殊空格标记(比如Ġ开头的token,转成正常空格) cleaned_tokens = [token.replace('Ġ', ' ') for token in tokens] # 计算子图网格:按4列排列,自动算行数 n_rows = (num_heads + 3) // 4 fig, axes = plt.subplots(n_rows, 4, figsize=(20, 5 * n_rows)) axes = axes.flatten() # 逐个绘制注意力头的热力图 for head_idx in range(num_heads): ax = axes[head_idx] # 取出当前头的n_tokens×n_tokens注意力矩阵 attn_matrix = attn_layer[0, head_idx, :, :] # 绘制热力图,vmin/vmax固定在0-1更直观 cax = ax.matshow(attn_matrix, cmap='viridis', vmin=0, vmax=1) ax.set_title(f'Head {head_idx + 1}', fontsize=10) # 设置token标签,调整字体大小避免拥挤 ax.set_xticks(range(seq_len)) ax.set_yticks(range(seq_len)) ax.set_xticklabels(cleaned_tokens, rotation=90, fontsize=8) ax.set_yticklabels(cleaned_tokens, fontsize=8) # 隐藏刻度线,让图更清爽 ax.tick_params(axis='both', which='both', length=0) # 隐藏多余的空图 for ax in axes[num_heads:]: ax.axis('off') # 添加颜色条,统一展示注意力权重的数值范围 fig.colorbar(cax, ax=axes.ravel().tolist(), shrink=0.6) plt.suptitle(f'Attention Weights - Layer {layer_idx + 1}', fontsize=16, y=1.02) plt.tight_layout() plt.savefig(filename, bbox_inches='tight') plt.close() # 调用示例:先把输入转成token列表 input_tokens = tokenizer.convert_ids_to_tokens(input_ids[0]) # 可视化第1层(索引0)的注意力 save_attention_image(attentions, input_tokens, layer_idx=0)
四、为什么很多注意力头看起来是均匀分布的?
这其实是纯Decoder模型里很常见的现象,不用太惊讶,主要有几个原因:
- 注意力头的分工不同:不是所有头都负责“死死盯着某几个token”——有些头可能负责全局上下文的整合、位置信息的对齐,或者维持因果注意力的基础分布,这类头的权重就会比较均匀。
- 输入序列过短:如果你的prompt只有十几个token(比如你这里的15个),很多头还没机会展现出明显的聚焦模式,看起来就会偏均匀。你可以试试更长的prompt,比如30-50个token,就能看到更多头出现聚焦的情况。
- 检查是否是因果注意力:如果你的可视化里没有看到下三角全为0(也就是模型能“看到”未来的token),那说明你提取的还是错误的注意力,得回头再检查提取代码。
最后给你个小技巧:可以先打印模型配置print(model.config),确认层数、头数这些参数,这样你对attentions的结构就更有底了~
备注:内容来源于stack exchange,提问作者Jose Ramon

