You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

大文本(>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")

关键说明

  1. 文本分段逻辑:

    • 先计算问题的token长度,确保每个文本片段和问题一起编码后不超过512token
    • 片段之间保留overlap(默认50个token),避免答案被拆分在两个片段的交界处
  2. 置信度筛选:

    • 对每个片段的答案计算start_score + end_score,作为答案的置信度
    • 只保留置信度最高的有效答案(排除[CLS])
  3. 边界处理:

    • 最后一个片段直接取到文本末尾,避免遗漏内容
    • 如果所有片段都返回[CLS],则输出"Geen antwoord gevonden"(荷兰语“未找到答案”)

这样修改后,模型就能遍历长文本的所有部分,找到正确的答案了。

内容的提问来源于stack exchange,提问作者Liza Darwesh

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.11 08:04:32