如何基于大段文本训练BERT问答模型?解决张量尺寸不匹配报错
解决BERT长文本问答训练的长度限制问题
报错原因
BERT-base/bert-large默认最大输入序列长度为1024,当输入的「问题+上下文」总token数超过该值时,模型预设的位置编码、注意力层维度与输入张量维度不匹配,就会触发RuntimeError: The size of tensor a (xxx) must match the size of tensor b (1024)报错。
长文本处理方案
1. 滑动窗口文本分块
将超长上下文分割为多个不超过1024长度的子块,每个子块与问题拼接后单独输入模型,最后融合各子块的预测结果。该方案无需更换模型,适合对全局上下文依赖较低的场景。
2. 使用长序列兼容的BERT变体
直接采用原生支持超长序列的模型,比如:
- Longformer:支持最长4096/16384序列长度,采用稀疏注意力机制降低计算量
- RoBERTa-Large-LM-Long:支持最长2048序列长度
这类模型无需手动分块,可直接处理长文本。
批量处理支持
两种方案均支持批量处理,只需保证每个batch内的样本序列长度统一:
- 分块方案:每个子块长度控制在1024内,batch内子块padding到当前batch最大长度
- 长序列模型:padding到模型支持的最大长度(如4096)或当前batch最大长度
代码示例
方案1:滑动窗口分块处理(基于Hugging Face Transformers)
import torch from transformers import BertTokenizer, BertForQuestionAnswering # 加载基础BERT模型与分词器 tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') model = BertForQuestionAnswering.from_pretrained('bert-base-uncased') def split_long_context(question, context, max_seq_len=1024): # 计算问题的token长度(含[CLS]和[SEP]) question_tokens = tokenizer.encode(question, add_special_tokens=False) question_total_len = len(question_tokens) + 2 # [CLS] + question + [SEP] # 单个上下文块的最大可用长度 max_context_chunk_len = max_seq_len - question_total_len - 1 # 预留最后一个[SEP]位置 # 分割上下文token context_tokens = tokenizer.encode(context, add_special_tokens=False) chunks = [] for i in range(0, len(context_tokens), max_context_chunk_len): chunk_tokens = context_tokens[i:i+max_context_chunk_len] # 构建模型输入格式:[CLS] question [SEP] chunk [SEP] input_ids = tokenizer.build_inputs_with_special_tokens(question_tokens, chunk_tokens) attention_mask = [1] * len(input_ids) token_type_ids = [0]*question_total_len + [1]*(len(chunk_tokens)+1) chunks.append({ 'input_ids': torch.tensor([input_ids]), 'attention_mask': torch.tensor([attention_mask]), 'token_type_ids': torch.tensor([token_type_ids]) }) return chunks def predict_long_qa(question, context): chunks = split_long_context(question, context) all_start_logits = [] all_end_logits = [] with torch.no_grad(): for chunk in chunks: outputs = model(**chunk) all_start_logits.append(outputs.start_logits) all_end_logits.append(outputs.end_logits) # 融合所有块的预测结果(取概率最高的起止位置) start_logits = torch.cat(all_start_logits, dim=1) end_logits = torch.cat(all_end_logits, dim=1) start_idx = torch.argmax(start_logits) end_idx = torch.argmax(end_logits) # 映射回原始文本并解码答案 all_tokens = tokenizer.encode(question, context, add_special_tokens=False) answer_tokens = all_tokens[start_idx:end_idx+1] return tokenizer.decode(answer_tokens) # 测试调用 question = "What is the core mechanism of photosynthesis?" long_context = "..." # 替换为你的超长上下文文本 answer = predict_long_qa(question, long_context) print(answer)
方案2:使用Longformer处理长文本(含批量训练)
import torch from transformers import LongformerTokenizer, LongformerForQuestionAnswering, Trainer, TrainingArguments import datasets # 加载Longformer模型与分词器(支持最长4096序列) tokenizer = LongformerTokenizer.from_pretrained('allenai/longformer-base-4096') model = LongformerForQuestionAnswering.from_pretrained('allenai/longformer-base-4096') # 预处理数据集(需符合SQuAD格式:每个样本包含question、context、answers字段) def preprocess_function(examples): questions = [q.strip() for q in examples["question"]] inputs = tokenizer( questions, examples["context"], max_length=4096, truncation="only_second", # 仅截断上下文,保留完整问题 padding="max_length", return_offsets_mapping=True, ) offset_mapping = inputs.pop("offset_mapping") answers = examples["answers"] start_positions = [] end_positions = [] for i, offset in enumerate(offset_mapping): answer = answers[i] start_char = answer["answer_start"][0] end_char = start_char + len(answer["text"][0]) sequence_ids = inputs.sequence_ids(i) # 定位上下文在token序列中的起止索引 idx = 0 while sequence_ids[idx] != 1: idx += 1 context_start = idx while sequence_ids[idx] == 1: idx += 1 context_end = idx - 1 # 映射字符位置到token位置 if offset[context_start][0] > end_char or offset[context_end][1] < start_char: start_positions.append(0) end_positions.append(0) else: idx = context_start while idx <= context_end and offset[idx][0] <= start_char: idx += 1 start_positions.append(idx - 1) idx = context_end while idx >= context_start and offset[idx][1] >= end_char: idx -= 1 end_positions.append(idx + 1) inputs["start_positions"] = start_positions inputs["end_positions"] = end_positions return inputs # 加载并预处理自定义数据集 dataset = datasets.load_dataset('json', data_files='your_qa_dataset.json') tokenized_dataset = dataset.map(preprocess_function, batched=True) # 设置训练参数(根据显存调整batch size) training_args = TrainingArguments( output_dir="./longformer_qa_model", per_device_train_batch_size=2, per_device_eval_batch_size=2, num_train_epochs=3, logging_dir="./logs", logging_steps=10, save_steps=100, evaluation_strategy="epoch" ) # 初始化Trainer并启动训练 trainer = Trainer( model=model, args=training_args, train_dataset=tokenized_dataset["train"], eval_dataset=tokenized_dataset["validation"], ) trainer.train() # 单样本推理示例 def long_qa_inference(question, context): inputs = tokenizer(question, context, return_tensors='pt', max_length=4096, truncation=True) with torch.no_grad(): outputs = model(**inputs) start_idx = torch.argmax(outputs.start_logits) end_idx = torch.argmax(outputs.end_logits) answer = tokenizer.decode(inputs['input_ids'][0][start_idx:end_idx+1]) return answer # 测试推理 question = "What is the core mechanism of photosynthesis?" long_context = "..." # 替换为你的超长上下文文本 answer = long_qa_inference(question, long_context) print(answer)
内容的提问来源于stack exchange,提问作者Siddharth Kumar Shukla
相关产品推荐
相关产品推荐

