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

基于TinyLlama的QA任务微调报错求助:修复RuntimeError及代码问题

问题修复方案:TinyLlama-1.1B-Chat-v1.0 微调MILQA QA任务

1. 解决LlamaForQuestionAnswering权重未初始化问题

TinyLlama-1.1B-Chat-v1.0是对话模型,自带的预训练权重不包含QA任务的输出头,直接加载LlamaForQuestionAnswering会生成未初始化的新增参数。可通过以下方式修复:

from transformers import AutoModelForQuestionAnswering, AutoTokenizer
import torch.nn as nn

model_name = "TinyLlama/TinyLlama-1.1B-Chat-v1.0"
tokenizer = AutoTokenizer.from_pretrained(model_name)

# 加载模型时忽略不匹配参数(QA头为新增层),并指定数据类型
model = AutoModelForQuestionAnswering.from_pretrained(
    model_name,
    ignore_mismatched_sizes=True,
    torch_dtype=torch.float16
)

# 手动初始化QA输出头参数,让初始化更合理
with torch.no_grad():
    model.qa_outputs.weight.data.normal_(mean=0.0, std=model.config.initializer_range)
    if model.qa_outputs.bias is not None:
        model.qa_outputs.bias.data.zero_()

2. 解决张量维度不匹配问题(349 vs 327)

该问题源于数据预处理时,输入序列与标签的start/end位置未对齐,或序列长度设置冲突。以下是修复后的预处理代码:

def preprocess_function(examples):
    questions = [q.strip() for q in examples["question"]]
    contexts = [c.strip() for c in examples["context"]]
    answers = examples["answers"]

    # 分词时保留偏移量,用于计算正确的start/end标签位置
    tokenized_examples = tokenizer(
        questions,
        contexts,
        truncation="only_second",  # 仅截断上下文,保留完整问题
        max_length=512,  # 需小于模型的max_position_embeddings
        stride=128,
        return_overflowing_tokens=True,
        return_offsets_mapping=True,
        padding="max_length",
    )

    offset_mapping = tokenized_examples.pop("offset_mapping")
    sample_map = tokenized_examples.pop("overflow_to_sample_mapping")
    start_positions = []
    end_positions = []

    for i, offset in enumerate(offset_mapping):
        sample_idx = sample_map[i]
        answer = answers[sample_idx]
        start_char = answer["answer_start"][0]
        end_char = start_char + len(answer["text"][0])
        sequence_ids = tokenized_examples.sequence_ids(i)

        # 定位上下文对应的token区间
        context_start = 0
        while sequence_ids[context_start] != 1:
            context_start += 1
        context_end = len(sequence_ids) - 1
        while sequence_ids[context_end] != 1:
            context_end -= 1

        # 过滤答案超出上下文范围的无效样本
        if offset[context_start][0] > end_char or offset[context_end][1] < start_char:
            start_positions.append(0)
            end_positions.append(0)
        else:
            # 计算start token位置
            idx = context_start
            while idx <= context_end and offset[idx][0] <= start_char:
                idx += 1
            start_positions.append(idx - 1)

            # 计算end token位置
            idx = context_end
            while idx >= context_start and offset[idx][1] >= end_char:
                idx -= 1
            end_positions.append(idx + 1)

    tokenized_examples["start_positions"] = start_positions
    tokenized_examples["end_positions"] = end_positions
    return tokenized_examples

加载并过滤数据集:

from datasets import load_dataset

dataset = load_dataset("SzegedAI/MILQA")
tokenized_dataset = dataset.map(
    preprocess_function,
    batched=True,
    remove_columns=dataset["train"].column_names,
)
# 过滤无效样本(start/end为0的样本)
tokenized_dataset = tokenized_dataset.filter(
    lambda x: x["start_positions"] != 0 and x["end_positions"] != 0
)

3. 解决历史遗留问题

  • 传递字典列表问题:上述预处理函数返回的是符合Hugging Face Dataset格式的字典,每个键对应一个列表,避免了嵌套字典的问题。
  • "Scalar tensor has no len()"问题:确保start_positions和end_positions为列表类型,而非单个张量。上述代码中通过循环逐个添加元素,生成的是列表格式的标签,不会触发该错误。

内容的提问来源于stack exchange,提问作者Levente Ledenyk lev4922

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 01:45:23