如何在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
相关产品推荐
相关产品推荐

