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

如何解决HuggingFace Tokenizer与自定义PyTorch Transformer解码器适配问题?

问题:HuggingFace Tokenizer与自定义PyTorch Transformer解码器配合使用的错误解决

问题场景

使用PyTorch自定义Transformer解码器时,配合HuggingFace Tokenizer生成的attention_mask出现类型和形状不匹配的错误,以下是相关代码及错误信息:

自定义解码器定义

decoder_layer = TransformerDecoderLayer(embedding_size, num_heads, hidden_size, dropout, batch_first=True)
self.decoder = TransformerDecoder(decoder_layer, 1)

解码器调用代码

output = self.decoder(output, embedded, tgt_mask=attention_mask)

HuggingFace Tokenizer生成掩码代码

batch = tokenizer(example['text'], return_tensors="pt", truncation=True, max_length=1024, padding='max_length')
inputs = batch['input_ids']
attention_mask = batch['attention_mask']

触发的错误

  1. 初始运行时错误:

AssertionError: only bool and floating types of attn_mask are supported

  1. 将attention_mask改为attention_mask = batch['attention_mask'].bool()后,出现新错误:

RuntimeError: The shape of the 2D attn_mask is torch.Size([4, 1024]), but should be (1024, 1024)

问题根源

PyTorch的TransformerDecoder中,tgt_mask和HuggingFace Tokenizer生成的attention_mask是两种完全不同的掩码:

  • tgt_mask:用于自注意力的序列遮挡(比如防止模型看到未来token),要求形状为[序列长度, 序列长度](或批次维度扩展后的[批次大小, 序列长度, 序列长度])。
  • Tokenizer生成的attention_mask:用于标记哪些是有效token、哪些是padding,形状为[批次大小, 序列长度],对应解码器的tgt_key_padding_mask参数。

解决方案

方案1:仅处理padding token(最常用场景)

将Tokenizer生成的掩码转换类型后,传入tgt_key_padding_mask参数,而非tgt_mask:

# 转换掩码类型为bool
tgt_key_padding_mask = batch['attention_mask'].bool()
# 正确调用解码器
output = self.decoder(output, embedded, tgt_key_padding_mask=tgt_key_padding_mask)

方案2:同时处理padding和序列遮挡(自回归任务场景)

如果需要防止模型看到未来token,需额外生成下三角遮挡掩码,同时传入两种掩码:

import torch

# 生成自注意力遮挡掩码(下三角矩阵,遮挡未来token)
seq_len = embedded.size(1)
tgt_mask = torch.tril(torch.ones(seq_len, seq_len, device=embedded.device)).bool()

# 处理padding的掩码
tgt_key_padding_mask = batch['attention_mask'].bool()

# 同时传入两种掩码
output = self.decoder(output, embedded, tgt_mask=tgt_mask, tgt_key_padding_mask=tgt_key_padding_mask)

关键参数说明

  • tgt_mask:控制自注意力的序列内可见性,形状需匹配序列长度维度,默认无掩码(所有token互相可见)。
  • tgt_key_padding_mask:过滤padding token的注意力计算,形状与Tokenizer生成的attention_mask一致。
  • 由于解码器设置了batch_first=True,tgt_mask也可以扩展为[批次大小, 序列长度, 序列长度],但通常单样本的[序列长度, 序列长度]掩码会自动广播到整个批次。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 23:35:17