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

如何在MarianMTModel生成输出时可视化Cross-Attention矩阵

解决MarianMTModel生成过程中Cross-Attention矩阵的提取与可视化问题

问题分析

你当前用前向钩子捕获Cross-Attention时,得到的注意力头数只有4个而非预期的8个,核心原因有两点:

  • 钩子捕获的是投影层的输入而非输出,没有拿到经过线性投影后的Q/K张量
  • 未将投影后的Q/K张量拆分为多头维度,直接计算的是整体注意力而非单头注意力

方案一:改进前向钩子方法

修正钩子逻辑,捕获投影层输出,并正确拆分多头:

from transformers import MarianMTModel, MarianTokenizer
import torch
import matplotlib.pyplot as plt
from torch.nn import functional as F

model_name = "Helsinki-NLP/opus-mt-en-de"
tokenizer = MarianTokenizer.from_pretrained(model_name)
model = MarianMTModel.from_pretrained(model_name)
model.eval()

# 存储各层的Q/K投影输出
keys = {}
queries = {}

def get_key_hook(layer_idx):
    def hook(module, input, output):
        # 捕获投影后的key(output是投影层的输出)
        keys[layer_idx] = output
    return hook

def get_query_hook(layer_idx):
    def hook(module, input, output):
        # 捕获投影后的query(output是投影层的输出)
        queries[layer_idx] = output
    return hook

# 注册钩子到decoder的encoder_attn投影层
hooks = []
for i, layer in enumerate(model.model.decoder.layers):
    hooks.append(layer.encoder_attn.k_proj.register_forward_hook(get_key_hook(i)))
    hooks.append(layer.encoder_attn.q_proj.register_forward_hook(get_query_hook(i)))

input_text = "Please translate this to German."
inputs = tokenizer(input_text, return_tensors="pt")

# 生成时禁用缓存,确保每一步都重新计算注意力
with torch.no_grad():
    translated_tokens = model.generate(**inputs, use_cache=False, max_new_tokens=50)

translated_text = tokenizer.decode(translated_tokens[0], skip_special_tokens=True)
input_tokens = tokenizer.convert_ids_to_tokens(inputs["input_ids"][0])
output_tokens = tokenizer.convert_ids_to_tokens(translated_tokens[0])

attentions = []
num_layers = len(keys)
# 获取模型的注意力头数和头维度
num_heads = model.config.num_heads
head_dim = model.config.d_model // num_heads

for layer_idx in range(num_layers):
    K = keys[layer_idx]  # shape: [bsz, src_seq_len, d_model]
    Q = queries[layer_idx]  # shape: [bsz, tgt_seq_len, d_model]
    
    # 拆分多头:[bsz, seq_len, d_model] -> [bsz, num_heads, seq_len, head_dim]
    K = K.view(K.size(0), K.size(1), num_heads, head_dim).transpose(1, 2)
    Q = Q.view(Q.size(0), Q.size(1), num_heads, head_dim).transpose(1, 2)
    
    # 计算注意力分数并softmax
    attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(head_dim, dtype=torch.float32))
    attn_weights = F.softmax(attn_scores, dim=-1)
    attentions.append(attn_weights)

# 堆叠成[num_layers, num_heads, tgt_seq_len, src_seq_len],去掉batch维度
attentions = torch.stack(attentions, dim=0).squeeze(0)

print("形状:(层数, 头数, 目标序列长度, 源序列长度)")
print(attentions.shape)

# 可视化第一层第一个头的注意力矩阵
plt.figure(figsize=(12, 8))
plt.imshow(attentions[0, 0].cpu().numpy(), cmap='viridis')
plt.colorbar()

plt.xticks(range(len(input_tokens)), input_tokens, rotation=90)
plt.yticks(range(len(output_tokens)), output_tokens)

plt.xlabel("源语言Token")
plt.ylabel("目标语言Token")
plt.title("Cross-Attention矩阵(第一层,第一个注意力头)")
plt.tight_layout()
plt.show()

# 移除钩子
for hook in hooks:
    hook.remove()

方案二:使用generate内置参数直接获取注意力

更简洁的方式是利用generate的return_dict_in_generate和output_attentions参数,直接获取生成过程中的注意力权重:

from transformers import MarianMTModel, MarianTokenizer
import torch
import matplotlib.pyplot as plt

model_name = "Helsinki-NLP/opus-mt-en-de"
tokenizer = MarianTokenizer.from_pretrained(model_name)
model = MarianMTModel.from_pretrained(model_name)
model.eval()

input_text = "Please translate this to German."
inputs = tokenizer(input_text, return_tensors="pt")

with torch.no_grad():
    outputs = model.generate(
        **inputs,
        use_cache=False,
        max_new_tokens=50,
        return_dict_in_generate=True,
        output_attentions=True
    )

# 提取cross-attention:Marian的cross_attentions按生成步骤存储,每一步对应一个新增token的注意力
cross_attentions = []
for step_attn in outputs.cross_attentions:
    # 取第一层第一个头,当前生成token对应的注意力(-1表示最新生成的token)
    cross_attentions.append(step_attn[0, 0, -1, :].cpu().numpy())

# 拼接成完整的注意力矩阵:[tgt_seq_len, src_seq_len]
full_attn = torch.tensor(cross_attentions)

input_tokens = tokenizer.convert_ids_to_tokens(inputs["input_ids"][0])
output_tokens = tokenizer.convert_ids_to_tokens(outputs.sequences[0])

# 可视化
plt.figure(figsize=(12, 8))
plt.imshow(full_attn, cmap='viridis')
plt.colorbar()

plt.xticks(range(len(input_tokens)), input_tokens, rotation=90)
plt.yticks(range(len(output_tokens)), output_tokens)

plt.xlabel("源语言Token")
plt.ylabel("目标语言Token")
plt.title("Cross-Attention矩阵(第一层,第一个注意力头)")
plt.tight_layout()
plt.show()

关键说明

  • 方案一适合需要精细控制注意力捕获过程的场景,能获取所有层所有头的完整注意力矩阵
  • 方案二更简洁,利用库内置功能,避免手动处理钩子,但需要注意MarianMTModel返回的注意力结构(cross_attentions按生成步骤存储)
  • 两种方案都需要设置use_cache=False,确保生成过程中每一步都重新计算完整的注意力矩阵,而非复用缓存的键值对

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 01:16:09