BERT预训练MLM+NSP任务出现张量维度不匹配RuntimeError求助
错误原因
- 核心触发点是输入模型的序列长度超过了BERT默认的最大位置嵌入维度512,两个张量维度不匹配触发RuntimeError。
- 你的代码存在三个问题:
TextDatasetForNextSentencePrediction是transformers旧版本已废弃的类,本身无自动截断超长序列的逻辑:如果输入单句tokenize后长度超过你设置的BLOCK_SIZE,该类不会做截断,直接返回原长度的token序列。- 你提前过滤长句的逻辑应该是按字符数判断,没有用对应tokenizer做校验,存在单句token数超过阈值的情况,两个句子拼接加特殊token后总长度达到882,远超BERT的512长度限制。
- 你设置
tokenizer.model_max_length = MAX_LEN的方式不生效,该参数不会被旧版数据集类读取用于截断逻辑。
修复方案
- 调整
BLOCK_SIZE取值:NSP任务输入格式为<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> 句1 [SEP] 句2 [SEP],共占用3个特殊token位置,BLOCK_SIZE最大可设为510,不要用过小的128:
BLOCK_SIZE = 510
- 显式配置BERT的最大位置嵌入,避免默认参数歧义:
# 替换原来的BertConfig初始化逻辑 config = BertConfig( max_position_embeddings=512, vocab_size=30000 if TOKENIZER_PATH == "hukuk_tokenizer" else 30522 )
- 提前过滤所有单句的token长度:对输入文件的每一行用你的tokenizer做编码,过滤掉单句token数超过250的样本,保证两个句子拼接后不会超过512的限制:
# 提前过滤长句示例 filtered_lines = [] with open(DATA_PATH, 'r', encoding='utf-8') as f: for line in f: line = line.strip() if not line: continue tokens = tokenizer.encode(line, add_special_tokens=False) if len(tokens) <= 250: filtered_lines.append(line) with open("data/filtered_sentences.txt", 'w', encoding='utf-8') as f: f.write('\n'.join(filtered_lines)) # 之后将DATA_PATH替换为过滤后的文件路径
- 若不想自己处理数据,可升级transformers到最新版本,使用官方推荐的
datasets库加载数据后调用map方法做批量截断、NSP样本生成,逻辑更可控。
内容的提问来源于stack exchange,提问作者O.Sahin
相关产品推荐
相关产品推荐

