如何解决PyTorch中Bangla-BERT文本分类的张量维度不匹配错误?
解决Bangla-BERT文本分类时的RuntimeError维度不匹配问题
错误原因分析
这个报错的核心是输入序列长度(1296)超过了sagorsarker/bangla-bert-base预训练模型的最大输入长度限制(512)。BERT类模型在预训练阶段固定了位置嵌入的维度(对应最大序列长度512),当输入序列过长时,模型无法扩展位置嵌入的维度,直接触发维度不匹配的RuntimeError。
具体解决方案
1. 强制限制Tokenize阶段的序列长度
加载Tokenizer时必须明确设置max_length=512,同时开启截断和填充,确保所有输入序列被严格限制在模型支持的长度范围内:
from transformers import AutoTokenizer # 加载预训练Tokenizer tokenizer = AutoTokenizer.from_pretrained("sagorsarker/bangla-bert-base") # 对训练/验证文本进行tokenize,强制截断过长序列 train_encodings = tokenizer( list(train_text), truncation=True, padding=True, max_length=512, return_tensors="pt" ) val_encodings = tokenizer( list(val_text), truncation=True, padding=True, max_length=512, return_tensors="pt" )
2. 移除自定义词干提取操作
你当前使用的BanglaStemmer词干处理会破坏预训练模型的分词逻辑:预训练Tokenizer已经针对孟加拉语词汇做了优化,自定义词干提取会改变原有词汇形态,导致Tokenizer生成的序列长度异常膨胀,甚至出现大量未登录词。建议先注释掉该步骤,验证问题是否解决:
# 暂时注释词干提取代码 # stemmer = BanglaStemmer() # df["comment"] = df["comment"].apply(lambda x: " ".join([stemmer.stem(word) for word in word_tokenize(x)]))
3. 验证批次数据维度
构建DataLoader后,打印批次输入的维度,确认序列长度是否为512:
for batch in train_loader: print("输入ID维度:", batch['input_ids'].shape) # 预期输出: (32, 512),32为你的batch_size print("注意力掩码维度:", batch['attention_mask'].shape) # 预期输出: (32, 512) break
补充说明
如果业务场景必须使用词干提取,建议先对比词干处理前后的词汇在预训练Tokenizer词汇表中的覆盖率,避免因大量未登录词导致序列长度失控。
内容的提问来源于stack exchange,提问作者Md Mahadi Hasan Sany
相关产品推荐
相关产品推荐

