如何解决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']
触发的错误
- 初始运行时错误:
AssertionError: only bool and floating types of attn_mask are supported
- 将
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
相关产品推荐
相关产品推荐

