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

如何向HuggingFace BertModel传入二维注意力掩码?

关于BertModel传入自定义token级注意力掩码的实现方案

可行,你需要的这种控制每个token对其他token的注意力权限的需求,完全可以通过传入三维注意力掩码实现,之前实验无效果大概率是掩码格式或用法不正确导致的。

核心原理

HuggingFace Transformers中,BertModel的attention_mask参数支持三维张量(形状为[batch_size, target_seq_len, source_seq_len]),用于直接定义每个token能关注到哪些其他token。区别于二维padding掩码(仅标记是否为有效token),三维掩码可以实现更精细的注意力控制。

正确操作步骤

  1. 构造三维掩码
    每个样本对应一个二维矩阵,矩阵中1表示允许当前token关注对应位置的token,0表示禁止关注。批量场景下将多个二维矩阵堆叠为三维张量,形状为[batch_size, seq_len, seq_len]。

  2. 掩码格式转换
    将掩码转换为布尔型(torch.bool)或浮点型,确保模型能正确识别允许/禁止的位置。模型内部会自动将禁止位置的注意力分数设为-inf,从而屏蔽该位置的注意力计算。

  3. 传入模型并验证
    在调用BertModel.forward时,将三维掩码传入attention_mask参数,同时确保掩码与输入张量在同一设备(CPU/GPU)。可以通过开启output_attentions=True来输出注意力分数,验证掩码是否生效。

示例代码

import torch
from transformers import BertModel, BertTokenizer

# 初始化模型和分词器
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
model = BertModel.from_pretrained('bert-base-uncased')

# 构造输入文本
input_texts = ["hello world test", "foo bar baz"]
inputs = tokenizer(input_texts, return_tensors="pt", padding=True, truncation=True)

# 构造自定义三维注意力掩码
# 样本1:前两个token可关注所有,最后一个仅能关注自身
mask_sample1 = torch.tensor([
    [1, 1, 1],
    [1, 1, 1],
    [0, 0, 1]
], dtype=torch.bool)
# 样本2:前两个token不能关注第三个,第三个可关注所有
mask_sample2 = torch.tensor([
    [1, 1, 0],
    [1, 1, 0],
    [1, 1, 1]
], dtype=torch.bool)
attention_mask = torch.stack([mask_sample1, mask_sample2])

# 对齐设备
attention_mask = attention_mask.to(model.device)

# 前向传播,开启注意力分数输出以验证
outputs = model(**inputs, attention_mask=attention_mask, output_attentions=True)

# 查看最后一层的注意力分数(形状:[num_heads, batch_size, seq_len, seq_len])
last_layer_attentions = outputs.attentions[-1]
print("最后一层注意力分数示例:", last_layer_attentions[0, 0])  # 第一个head、第一个样本的注意力矩阵

常见问题排查

  • 无效果原因:若之前实验无变化,可能是使用了旧版本Transformers(早期版本对三维掩码支持不完善),建议升级到最新稳定版;或掩码未转换为正确格式(比如用了整数型而非布尔/浮点型)。
  • 设备不匹配:确保掩码与模型、输入张量在同一设备,否则会被自动忽略或报错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 17:05:31