大文本(>512词)下BERT问答模型运行问题及解决方案
解决BERT问答模型处理长文本(超过512token)的问题
你的问题根源很清晰:当输入文本过长时,tokenizer.encode_plus会截断文本以满足512token的限制,但如果答案恰好被截断掉了,模型就找不到有效答案,只能返回默认的[CLS] token。
要解决这个问题,核心思路是把长文本拆分成多个重叠的片段,让模型逐个处理每个片段,然后从所有片段的结果中选出置信度最高的答案。这样就能保证答案所在的片段被模型处理到。
修改后的完整代码
from transformers import AutoTokenizer, AutoModelForQuestionAnswering import torch max_seq_length = 512 tokenizer = AutoTokenizer.from_pretrained("henryk/bert-base-multilingual-cased-finetuned-dutch-squad2") model = AutoModelForQuestionAnswering.from_pretrained("henryk/bert-base-multilingual-cased-finetuned-dutch-squad2") # 读取长文本 with open("test.txt", "r") as f: text = f.read() questions = [ "Wat is de hoofdstad van Nederland?", "Van welk automerk is een Cayenne?", "In welk jaar is pindakaas geproduceerd?", ] def split_long_text(text, question_token_len, max_seq_len, overlap=50): """ 将长文本拆分为多个重叠的片段,确保每个片段和问题编码后不超过max_seq_len """ # 计算每个文本片段的最大token数(预留问题和特殊token的位置) max_text_len = max_seq_len - question_token_len - 3 # 3是[CLS] + [SEP] + [SEP]的数量 text_tokens = tokenizer.tokenize(text) total_tokens = len(text_tokens) chunks = [] start = 0 while start < total_tokens: end = start + max_text_len # 最后一个片段直接取到末尾 if end >= total_tokens: end = total_tokens chunks.append(tokenizer.convert_tokens_to_string(text_tokens[start:end])) # 移动起始位置,保留重叠部分 start = end - overlap return chunks for question in questions: # 计算问题的token长度 question_tokens = tokenizer.tokenize(question) question_token_len = len(question_tokens) # 拆分长文本为片段 text_chunks = split_long_text(text, question_token_len, max_seq_length) best_answer = "" best_score = -float("inf") for chunk in text_chunks: # 编码问题和当前片段 inputs = tokenizer.encode_plus( question, chunk, add_special_tokens=True, max_length=max_seq_length, truncation=True, return_tensors="pt" ) input_ids = inputs["input_ids"].tolist()[0] answer_start_scores, answer_end_scores = model(**inputs, return_dict=False) # 找到当前片段的候选答案,并计算置信度(start和end分数之和) answer_start = torch.argmax(answer_start_scores) answer_end = torch.argmax(answer_end_scores) + 1 current_score = answer_start_scores[0][answer_start] + answer_end_scores[0][answer_end-1] # 转换为可读文本 current_answer = tokenizer.convert_tokens_to_string( tokenizer.convert_ids_to_tokens(input_ids[answer_start:answer_end]) ) # 跳过无意义的[CLS]答案 if current_answer != "[CLS]" and current_score > best_score: best_score = current_score best_answer = current_answer # 输出结果,如果没有找到有效答案,提示无结果 print(f"Question: {question}") print(f"Answer: {best_answer if best_answer else 'Geen antwoord gevonden'}\n")
关键说明
文本分段逻辑:
- 先计算问题的token长度,确保每个文本片段和问题一起编码后不超过512token
- 片段之间保留
overlap(默认50个token),避免答案被拆分在两个片段的交界处
置信度筛选:
- 对每个片段的答案计算
start_score + end_score,作为答案的置信度 - 只保留置信度最高的有效答案(排除
[CLS])
- 对每个片段的答案计算
边界处理:
- 最后一个片段直接取到文本末尾,避免遗漏内容
- 如果所有片段都返回
[CLS],则输出"Geen antwoord gevonden"(荷兰语“未找到答案”)
这样修改后,模型就能遍历长文本的所有部分,找到正确的答案了。
内容的提问来源于stack exchange,提问作者Liza Darwesh
相关产品推荐
相关产品推荐

