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

如何基于大段文本训练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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 00:43:25