修改后BERT模型提取注意力维度异常:获[层数,seq_len,隐层维]而非[seq_len,seq_len]
问题描述
我在从一个无注意力输出的修改版BERT模型(基于DNABERT2的修改实现)中提取注意力权重。通过修改BertEncoder及其中间类(ModelLayer、SelfUnpaddedAttention等)后,得到的注意力维度为[nbr_layers, seq_length, hidden_layer_dim],但我需要的是可用于可视化的[seq_length, seq_length]注意力矩阵,不确定是提取过程有误,还是需要调整代码来得到目标维度。
修改后的BertEncoder核心提取代码如下:
class BertEncoder(nn.Module): """A stack of BERT layers providing the backbone of Mosaic BERT. This module is modeled after the Hugging Face BERT's :class:`~transformers.model.bert.modeling_bert.BertEncoder`, but with substantial modifications to implement unpadding and ALiBi. Compared to the analogous Hugging Face BERT module, this module handles unpadding to reduce unnecessary computation at padded tokens, and pre-computes attention biases to implement ALiBi. """ # ... 省略其他代码 ... # PART WHERE I EXTRACT ATTENTION all_encoder_layers = [] all_attention_weights = [] # List to store attention weights if subset_mask is None: for layer_module in self.layer: # Since we get now attention too, we need to unpack 2 elements instead of 1. hidden_states, attention_weights = layer_module(hidden_states, cu_seqlens, seqlen, None, indices, attn_mask=attention_mask, bias=alibi_attn_mask) all_attention_weights.append(attention_weights) # Store attention weights if output_all_encoded_layers: all_encoder_layers.append(hidden_states) # Pad inputs and mask. It will insert back zero-padded tokens. # Assume ntokens is total number of tokens (padded and non-padded) # and ntokens_unpad is total number of non-padded tokens. # Then padding performs the following de-compression: # hidden_states[ntokens_unpad,hidden] -> hidden_states[ntokens,hidden] hidden_states = pad_input(hidden_states, indices, batch, seqlen) else: for i in range(len(self.layer) - 1): layer_module = self.layer[i] # Since we get now attention too, we need to unpack 2 elements instead of 1. hidden_states, attention_weights = layer_module(hidden_states, cu_seqlens, seqlen, None, indices, attn_mask=attention_mask, bias=alibi_attn_mask) all_attention_weights.append(attention_weights) # Store attention weights if output_all_encoded_layers: all_encoder_layers.append(hidden_states) subset_idx = torch.nonzero(subset_mask[attention_mask_bool], as_tuple=False).flatten() # Since we get now attention too, we need to unpack 2 elements instead of 1. hidden_states, attention_weights = self.layer[-1](hidden_states, cu_seqlens, seqlen, subset_idx=subset_idx, indices=indices, attn_mask=attention_mask, bias=alibi_attn_mask) all_attention_weights.append(attention_weights) # appending the attention of different layers together. if not output_all_encoded_layers: all_encoder_layers.append(hidden_states) # Since we now return both, we need to handle them wherever BertEncoder forward is called. return all_encoder_layers, all_attention_weights # Return both hidden states and attention weights # return all_encoder_layers # original return.
问题分析
你当前提取的并不是注意力权重矩阵,而是注意力层输出的上下文向量,原因如下:
- 维度不符:标准注意力权重的维度应该是
[nbr_layers, num_heads, seq_len, seq_len](包含注意力头、查询/键序列长度两个维度),而你得到的[nbr_layers, seq_length, hidden_layer_dim]和hidden_states的维度一致,说明你提取的是注意力层处理后的特征,而非权重。 - 模型特性:这个模型是去填充(Unpadded)版本的BERT,核心优化是跳过padding token的计算,原始的
SelfUnpaddedAttention类可能原本就没有返回注意力权重,只返回了处理后的特征;或者你在ModelLayer中错误地把上下文向量当成了注意力权重返回。
解决步骤
1. 定位真正的注意力权重张量
打开SelfUnpaddedAttention类的实现,找到注意力计算的核心逻辑:
- 找到
attn_scores = torch.matmul(query, key.transpose(-1, -2))这一步(计算查询和键的相似度) - 找到
attn_weights = nn.functional.softmax(attn_scores, dim=-1)这一步(得到归一化后的注意力权重矩阵,维度应为[num_heads, unpadded_seq_len, unpadded_seq_len]) - 修改
SelfUnpaddedAttention的forward方法,让它返回这个attn_weights,而不是返回上下文向量。
2. 修改ModelLayer的返回逻辑
确保ModelLayer在调用SelfUnpaddedAttention后,正确接收并返回注意力权重,而非把上下文向量当作权重返回。比如:
# 在ModelLayer的forward中 attn_output, attn_weights = self.attention(...) # 后续处理... return attn_output, attn_weights # 确保第二个返回值是真正的注意力权重
3. 还原去填充的注意力矩阵
因为模型用了去填充优化,得到的attn_weights是针对无padding的精简序列的,需要用模型中的indices参数(记录非padding token在原始序列中的位置)把矩阵还原为原始带padding的维度:
- 初始化一个全0的矩阵
full_attn_weights,维度为[orig_seq_len, orig_seq_len] - 把精简序列的注意力权重填充到
full_attn_weights中对应indices的位置,padding位置保持0
4. 处理注意力头维度
如果得到的权重包含注意力头维度([num_heads, seq_len, seq_len]),可以选择对所有头取平均,或者单独可视化某个头的权重,得到最终的[seq_len, seq_len]矩阵用于可视化。
内容的提问来源于stack exchange,提问作者Wassim Jaoui
相关产品推荐
相关产品推荐

