用自蒸馏推理补全 SFT 数据:缓解 Amazon Nova 的推理抑制

2026-07-22 24 预计阅读时间: 1 分钟
来源: aws.amazon.com AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:10 分钟

监督微调数据经常只包含“问题—答案”,没有展示如何从输入走到结论。对于需要多步推导的任务,这种数据结构可能让模型学会跳过推理过程,甚至压制原本具备的推理能力。Self-Distilled Reasoning(SDR,自蒸馏推理)提供了一条实用路径:让模型为现有样本生成推理轨迹,再把经过筛选的轨迹加入 SFT 数据。

原文围绕三个环节展开:识别推理抑制问题、构造 SDR 数据,以及在三个基准上验证这一方法。这里重点说明它为什么有效,以及如何把它落到可审计的数据流水线中。

为什么答案正确,训练信号仍可能不够

假设原始数据只有下面两个字段:

{"question":"一辆车以 60 km/h 行驶 2.5 小时,共行驶多少千米?","answer":"150 千米"}

这条样本能监督最终答案,却没有明确告诉模型:

  1. 应识别“路程 = 速度 × 时间”;
  2. 应执行 60 × 2.5
  3. 应保留并检查单位。

当大量 SFT 样本都只奖励短答案时,模型可能把“立即输出结论”当成目标格式。这就是推理抑制值得关注的地方:问题不一定是模型不知道怎么推,而是微调数据持续告诉它不要展示或使用中间推导。

这并不意味着所有任务都需要长篇思维过程。分类、抽取和固定格式转换通常更看重输出稳定性。SDR 更适合数学、逻辑、规划、代码分析,以及其他需要多步决策的任务。

SDR 如何把模型变成自己的数据教师

SDR 的核心不是引入一个完全不同的教师模型,而是利用目标模型或同系列模型,为缺少推理轨迹的数据生成候选推导。一个完整流程通常包含四步:

  1. 生成:输入原问题和参考答案,让模型补充能够导向该答案的推理。
  2. 验证:检查候选推理的最终结论是否与参考答案一致。
  3. 筛选:移除自相矛盾、泄漏无关信息、格式损坏或过度冗长的样本。
  4. 微调:把问题、推理和最终答案组织成统一训练格式。

这里最容易出现的错误是把模型生成的每一条轨迹都当成真值。自蒸馏会继承模型自身的偏差;如果候选推理碰巧得到正确答案,但中间逻辑错误,SFT 反而会强化这种错误模式。因此,SDR 的价值取决于验证器,而不仅取决于生成器。

在评估上,也不能只看最终准确率。原文在三个基准上进行验证,这类实验至少应同时观察:

  • 最终答案准确率是否提高或保持;
  • 面对需要推理的问题时,模型是否仍会稳定地产生推导;
  • 简单任务是否因输出变长而增加延迟和成本;
  • 推理文本是否忠实、可读,而不是事后拼接的解释。

可以这样实践:用 Amazon Bedrock 生成候选轨迹

下面是一个可改造的最小 Python 示例。它假设你已经在 Amazon Bedrock 中获得目标 Amazon Nova 模型的调用权限,并配置了 AWS 凭证。模型 ID、区域和提示词应按实际环境调整;这段代码用于演示 SDR 数据构造流程,不代表原文提供的固定训练接口或提示模板。

安装依赖:

python -m pip install boto3
export AWS_REGION=us-east-1
export NOVA_MODEL_ID=amazon.nova-pro-v1:0

创建 build_sdr.py

import json
import os
from pathlib import Path

import boto3

region = os.getenv("AWS_REGION", "us-east-1")
model_id = os.getenv("NOVA_MODEL_ID", "amazon.nova-pro-v1:0")
client = boto3.client("bedrock-runtime", region_name=region)

samples = [
    {
        "id": "distance-001",
        "question": "A car travels at 60 km/h for 2.5 hours. How far does it travel?",
        "answer": "150 km",
    },
    {
        "id": "sequence-001",
        "question": "What is the next number in 2, 6, 12, 20, 30?",
        "answer": "42",
    },
]

system_prompt = """You create concise training rationales for supervised fine-tuning.
Given a question and a verified reference answer, produce a short derivation that
supports the answer. Do not introduce facts absent from the question. Return JSON
with exactly two string fields: rationale and final_answer."""


def generate_trace(sample):
    prompt = json.dumps(
        {"question": sample["question"], "reference_answer": sample["answer"]},
        ensure_ascii=False,
    )
    response = client.converse(
        modelId=model_id,
        system=[{"text": system_prompt}],
        messages=[{"role": "user", "content": [{"text": prompt}]}],
        inferenceConfig={"maxTokens": 512, "temperature": 0.1, "topP": 0.9},
    )
    text = response["output"]["message"]["content"][0]["text"]
    return json.loads(text)


with Path("sdr_candidates.jsonl").open("w", encoding="utf-8") as output:
    for sample in samples:
        try:
            trace = generate_trace(sample)
            accepted = trace["final_answer"].strip() == sample["answer"].strip()
            record = {
                **sample,
                "rationale": trace["rationale"],
                "generated_answer": trace["final_answer"],
                "accepted": accepted,
            }
        except (KeyError, json.JSONDecodeError) as exc:
            record = {**sample, "accepted": False, "error": str(exc)}

        output.write(json.dumps(record, ensure_ascii=False) + "\n")

print("Wrote sdr_candidates.jsonl")

运行:

python build_sdr.py
python -m json.tool < <(head -n 1 sdr_candidates.jsonl)

示例中的字符串相等检查只适合演示。生产流水线应使用任务相关验证器:数学题可以重新计算表达式,代码题可以执行单元测试,结构化抽取可以用 JSON Schema 校验,开放式问答则可以组合规则、另一个评审模型和人工抽检。

把候选轨迹转换成训练数据

只有通过验证的样本才应进入训练集。可以继续用一个小脚本生成通用的消息格式:

import json

with open("sdr_candidates.jsonl", encoding="utf-8") as source, open(
    "sft_train.jsonl", "w", encoding="utf-8"
) as target:
    for line in source:
        item = json.loads(line)
        if not item.get("accepted"):
            continue

        training_item = {
            "messages": [
                {"role": "user", "content": item["question"]},
                {
                    "role": "assistant",
                    "content": (
                        "Reasoning:\n"
                        + item["rationale"].strip()
                        + "\n\nAnswer:\n"
                        + item["answer"].strip()
                    ),
                },
            ]
        }
        target.write(json.dumps(training_item, ensure_ascii=False) + "\n")

真正提交训练前,需要把消息字段映射到所用定制服务要求的 schema。还要固定推理与答案的分隔格式,否则模型可能在推理文本、最终答案和停止条件之间产生混淆。

上线前的决策清单

SDR 不应被当成“自动补齐思维链”的一次性脚本。更稳妥的采用方式是从小规模实验开始:

  • 保留一组完全不参与轨迹生成和微调的评估集,防止数据污染;
  • 将答案型 SFT 与 SDR-SFT 作为独立实验组,比较准确率、输出长度、延迟和成本;
  • 为每种任务建立可执行或可重复的验证规则;
  • 对候选轨迹去重,避免模型反复学习同一种模板;
  • 抽检“答案正确但理由错误”的样本,这类错误最难被简单匹配发现;
  • 不把内部推理文本直接暴露给不可信用户或下游系统,敏感场景应输出简洁、可审计的解释;
  • 保留生成模型版本、提示词、采样参数和拒绝原因,确保数据可追溯。

自蒸馏推理的关键收益,是把缺失的过程监督重新放回 SFT 数据;它的主要代价,则是生成成本、验证复杂度和错误自我强化的风险。只有当验证链路足够可靠,并且评估证明目标任务确实依赖多步推导时,SDR 才值得进入正式训练流水线。


相关推荐