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

修改后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.
问题分析

你当前提取的并不是注意力权重矩阵,而是注意力层输出的上下文向量,原因如下:

  1. 维度不符:标准注意力权重的维度应该是[nbr_layers, num_heads, seq_len, seq_len](包含注意力头、查询/键序列长度两个维度),而你得到的[nbr_layers, seq_length, hidden_layer_dim]和hidden_states的维度一致,说明你提取的是注意力层处理后的特征,而非权重。
  2. 模型特性:这个模型是去填充(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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 01:53:09