AWS · ML 博客

探索基于Amazon Nova的自蒸馏推理在监督微调中的应用

Exploring self-distilled reasoning for supervised fine-tuning with Amazon Nova

二〇二六年七月二十三日 · 英文原文

Amazon Nova 2 系列模型在 SFT 中面临推理抑制问题:无推理轨迹的数据集训练会导致模型丧失推理能力,数学准确率从 70% 降至 0%。为此,提出自蒸馏推理(SDR),利用基础 Nova 2 Lite 模型为 SFT 数据集生成推理轨迹并前置到输出中。在 MedMCQA、CoCoHD、Invoice-OCR 三个 benchmark 上,SDR 将数学性能恢复至约 68%,同时目标性能比最佳模型融合提升超 6.5%,无需人工标注或事后插值。

当你使用监督微调(Supervised Fine-Tuning, SFT)微调模型时,为训练数据创建高质量的思维链(chain-of-thought, CoT)推理轨迹通常不切实际,且成本高昂。因此,你可能会选择在 SFT 过程中跳过推理,仅使用输入和输出进行训练。然而,推理是 Amazon Nova 2 系列模型的一项关键能力,并且已被证明能显著提升 Nova 及其他前沿模型的预测性能,尤其是在编码和数学等难题上。借助 Amazon Nova 2 定制化,你可以通过 SFT 和强化微调(Reinforcement Fine-Tuning, RFT)等技术,在你的特定领域运用思维模型的能力。对于 SFT,释放这些增益的一个关键要求是 SFT 数据集中包含高质量的黄金 CoT 轨迹,即来自强教师模型并经过进一步验证和清洗的思维轨迹。在本文中,我们探讨了一种为缺乏推理轨迹的 SFT 定制化数据集生成思维 token 的思路。我们首先研究推理抑制问题,然后介绍自蒸馏推理(Self-Distilled Reasoning, SDR),在三个 benchmark 上对其进行验证,并提供实用建议。SDR 将基础 Amazon Nova 2 Lite 模型的思维链复用于非推理数据集。通过实验,我们观察到引入 SDR 的一些影响,从提升目标性能到缓解灾难性遗忘。这与越来越多发现训练期间将信息从模型蒸馏到自身(即自蒸馏)重要性的研究工作一致。下图展示了 SDR 与普通 SFT 和模型融合在目标性能和数学保持方面的对比。模型融合是一种直接保留先前习得技能的方法,它获取 SFT 检查点并将其融合回基础模型(SFT 之前)。然而,融合通常以牺牲 SFT 检查点上获得的增益为代价,因为灾难性遗忘表现为目标性能与通用性能之间的权衡。另一方面,SDR 缓解了这个问题,其目标性能与纯 SFT(无融合)相似或更好,同时保留了大部分通用性能。由于灾难性遗忘,我们观察到在数据集上进行 SFT 可能导致先前知识完全丢失。例如,我们发现普通 SFT 导致数学性能平均从 70%(基础)下降到 6%(表 2)。使用 SDR,我们几乎可以完全恢复它。自蒸馏推理使我们能够在不牺牲目标性能的情况下保留基础模型的能力,并且在许多情况下还能提升目标性能。模型融合是缓解灾难性遗忘的常见解决方案,但 SDR 更有效地实现了这种保留。SDR 通过用模型自身的推理轨迹增强数据集,提供了训练中的正则化。它不需要单独的教师模型或事后插值。这种方法不需要人工标注,适用于任何领域的现有 SFT 数据集,并且在保留目标性能和通用性能方面比模型融合更有效。与模型融合相比,自蒸馏在数学性能上几乎相同(70% 对比 68%),同时目标性能平均提升超过 6.5%。这一发现与近期主张自蒸馏是缓解灾难性遗忘强大机制的研究一致。

背景:在无推理 token 的数据集上进行 SFT 可能导致推理能力丧失

在介绍该方法之前,理解为什么在无推理轨迹的数据集上进行 SFT 会导致模型性能下降至关重要。本节探讨推理抑制问题,量化推理对目标性能的影响,并讨论模型融合作为一种现有的恢复机制。

推理抑制问题

在推理模式激活的情况下,对非推理数据集进行 SFT 训练会导致一个关键问题:即使在推理时明确开启推理模式,模型也会失去推理能力。我们假设这种现象发生是因为训练目标同时计算推理 token 和输出 token 的损失。当训练数据仅包含输入-输出对而没有中间推理步骤时,模型在生成连贯推理轨迹方面得不到任何监督信号。损失函数实际上惩罚了那些不直接贡献于最终输出的 token,从而训练模型绕过其推理机制。这种行为体现了捷径学习,即模型依赖虚假相关性而非学习稳健的推理模式。

量化推理的影响

尽管存在抑制问题,但在训练和推理时都开启推理会在目标性能上产生显著收益。为了确认公平比较,我们保持训练和推理设置之间的一致性:比较训练和推理时都开启推理与都关闭推理的情况。我们的实验表明,在训练和推理时开启推理会带来显著的改进。下表展示了在 LLaVA CoT 数据集上使用不同 LoRA(Low-Rank Adaptation)融合权重的效果。

融合权重 目标性能(推理开启) 目标性能(推理关闭) 差值
0.0(基础) 12.30% 12.30% 0%
0.3 34.38% 35.28% -0.9%
0.5 60.51% 48.80% +11.7%
0.7 66.67% 54.05% +12.6%
1.0(完全融合) 65.17% 47.90% +17.3%

表 1:推理和融合权重对目标性能的影响。训练和推理推理模式保持一致。

模型融合作为恢复机制

我们目前推荐的解决无推理 token 的 SFT 用例中推理抑制问题的方法是使用模型融合作为恢复推理技能的机制。类似于模型融合可以在领域特定微调后恢复数学或编码能力,我们发现推理能力可以通过微调模型和基础模型之间的加权插值来恢复。这种恢复与通用性能指标(如数学和编码 benchmark)的恢复密切相关,表明推理是支撑多种下游技能的基本能力。

为了更好地理解这一点,我们在本实验中使用以下 benchmark:

我们研究推理是否随模型融合而出现。在一个指令遵循 benchmark 上,我们统计了一个在没有推理的情况下训练的模型产生的推理 token 数量。x 轴显示定制模型与基础模型之间的融合权重,其中 0.0 对应基础模型,1.0 是未融合的 SFT 模型。

对定制化的影响

如果你正在部署定制化模型,这些发现具有重要意义。首先,在训练和推理期间保持相同的推理模式对于获得最佳性能至关重要。其次,如果不加干预,领域特定的 SFT 会显著降低通用推理能力。第三,模型融合提供了一种事后解决方案,但需要仔细调整融合权重以平衡相互竞争的目标。这些权衡促使我们开发一种无需事后干预即可保留推理能力的训练方法。

推理在 SFT 中的作用

SFT 中的推理轨迹具有三个关键功能,它们连接了强化微调(RFT)、正则化和自蒸馏范式。首先,推理轨迹通过约束微调模型使其接近基础模型的策略,充当隐式策略正则化。这以类似于 RLHF 中 KL 正则化的方式防止灾难性遗忘和能力退化。其次,推理提供了过程监督,类似于 RFT 中的逐步奖励塑形,而不仅仅是结果监督。这导致了更稳健的泛化。第三,自蒸馏,即模型从其自身预测中学习,已被证明在从重生网络到宪法 AI 等多个领域都有效。推理轨迹通过捕捉模型的问题解决过程而不仅仅是最终输出,提供了特别丰富的自监督信号,正如近期持续学习研究所展示的那样。近期工作表明,有效的微调模型在参数空间中保持与基础模型的接近,我们的方法利用这一见解,将基础模型自身的推理轨迹作为训练目标,创建一个与模型内部表示一致的学习信号,同时在整个推理过程中提供密集的 token 级监督。

自蒸馏推理

我们提出了一种有效的方法,使用自蒸馏推理(SDR)对没有推理轨迹的数据集进行 SFT 定制化。该过程分三个阶段进行。首先,你使用基础推理模型(如 Amazon Nova 2 Lite)为 SFT 数据集中的每个示例生成推理轨迹。接下来,你通过将生成的推理轨迹前置到原始输出来增强训练数据。最后,你在开启推理模式的情况下微调模型,对推理 token 和输出 token 都提供监督。

SDR 为你的 SFT 工作流提供了几个优势。零标注成本,因为不需要人工创建推理轨迹。推理轨迹与基础模型自身的问题解决方法一致。该技术可扩展到任何领域的现有 SFT 数据集,并允许你灵活使用不同的教师模型。

为 SFT 构建推理轨迹

我们查询 Amazon Bedrock 以获取给定数据集的推理轨迹。有两种构建思维链推理的方法,各有不同的权衡。

实现流水线

以下代码展示了如何使用 Amazon Bedrock Converse API 实现引导推理。该流水线读取你的 SFT 训练数据,生成思维链(CoT)推理轨迹,并使用 reasoningContent 块组装最终数据集。

设置和 CoT 提示生成

import json, re, boto3, time

MODEL_ID = "amazon.nova-lite-v1:0"
bedrock = boto3.client("bedrock-runtime", region_name="us-east-1")

COT_SYSTEM = """You are an expert evaluator. Provide a detailed Chain-of-Thought explaining why the given ground truth answer is correct and makes sense."""

COT_USER_TEMPLATE = """Given the task context and ground truth annotation, provide a Chain-of-Thought.

**Original System Context:**
{system_prompt}

**Original User Request:**
{user_prompt}

**Ground Truth Annotation:**
```json
{ground_truth}

ONLY provide your justification within tags."""


**调用 Amazon Bedrock 并组装最终数据集**

```python
# Invoke Bedrock for each CoT prompt
for req in cot_requests:
    response = bedrock.converse(
        modelId=MODEL_ID,
        system=[{"text": req["system"]}],
        messages=[{"role":"user", "content":[{"text":req["query"]}]}],
        inferenceConfig={"maxTokens":4096,"temperature":0.3,"topP": 0.9,"topK": 50})
    output = response["output"]["message"]["content"][0]["text"]
    time.sleep(2) # rate limit

# Parse tags and build SFT dataset
for idx, item in enumerate(original_data):
    tags = extract_xml_tag(inference[idx], "justification")
    reasoning_block = {"reasoningContent": {"reasoningText": {"text": "Based on the provided context, lets think"
                                                                     " step by step " + tags[0].strip()}}}
    sft_dataset.append({
        "schemaVersion": "bedrock-conversation-2024",
        "system": item["system"],
        "messages": [item["messages"][0],
                     {"role": "assistant", "content": [reasoning_block, original_text]}]
    })

推理配置参数

样本输入(无推理)

{
  "schemaVersion": "bedrock-conversation-2024",
  "system": [{"text": "You are a legal expert..."}],
  "messages": [
    {"role": "user", "content": [{"text": "Review this clause: Company shall not disclose data to third parties."}]},
    {"role": "assistant", "content": [{"text": "[{\"comment_id\": \"001\", \"severity\": \"Medium\", \"comment\": \"Add standard confidentiality exceptions...\", ...}]"}]}
  ]
}

样本输出(前置推理)

{
  "role": "assistant",
  "content": [
    {"reasoningContent": {"reasoningText": {"text": "Based on the provided context, lets think step by step: 1. Correctness: The annotation correctly identifies an issue with the clause by suggesting standard confidentiality exceptions... 2. Severity: Medium is appropriate because... 3. Category: Privacy fits because data protection is a core privacy concern... 6. References: std-clauses.pdf validates the need for exceptions."}}},
    {"text": "[{\"comment_id\": \"001\", ...}]"}
  ]
}

我们的实验(在下一节讨论)表明,这种方法在目标领域性能和通用能力保留方面都带来了显著改进,验证了自蒸馏推理在微调期间提供有价值的归纳偏置的假设。

实验

我们在三个没有推理的 benchmark 上运行实验。MedMCQA 是一个问答风格的数据集,包含 10k 训练样本。CoCoHD 是一个国会听证会数据集,具有长上下文(20-60k token),任务是生成包含特定字段的 JSON。Invoice-OCR 是一个多模态数据集,包含发票图像,需要输出包含图像中字段的 JSON。

下表总结了我们在三个 benchmark 上的 SDR 结果。每列捕捉实验设置的一个特定维度:

Dataset Reasoning Traces Training Reasoning Inference Reasoning Merging Target Math ↑ (control)
Nova 2 Lite None No No 55.60% 12.90%
Nova 2 Lite None No Yes 55.40% 70%
MedMCQA None No Yes Yes 60.50% 8.30%
MedMCQA None No Yes No 63.80% 0%
MedMCQA Nova 2 Lite Basic Yes Yes No 66.60% 67.90%
MedMCQA Nova 2 Lite Guided Yes Yes No 65.70% 59.60%
Nova 2 Lite None No No 45% 12.90%
Nova 2 Lite None No Yes 45% 70%
CoCoHD None No Yes Yes 59.30% 70%
CoCoHD None No Yes No 61.30% 6.30%
CoCoHD Nova 2 Lite Basic Yes Yes No 55.40% 73.80%
CoCoHD Nova 2 Lite Supervised Yes Yes No 61.30% 65.80%
Nova 2 Lite None No No 82.10% 12.90%
Nova 2 Lite None No Yes 81.40% 70%
Invoice-OCR None No Yes Yes 82.30% 70.80%
Invoice-OCR None No Yes No 88.10% 9.60%
Invoice-OCR Nova 2 Lite Basic Yes Yes No 86.10% 67.10%
Invoice-OCR Nova 2 Lite Supervised Yes Yes No 87.90% 67.90%

表 2:三个 benchmark 上的 SDR 结果。行按数据集分组。所有 SDR 行使用融合权重 1.0(无需融合)。

具有部分或缺失推理轨迹的 SFT

推理数据可能昂贵或难以整理,因此我们经常遇到只有部分 SFT 数据集包含推理的情况。为了更好地理解 SFT 在这种情况下的表现,我们接下来进行一项消融研究。我们还探讨了 SDR 如何作为一种低开销且经济高效的替代方案来帮助解决这种情况。

仅使用部分推理轨迹进行 SFT 训练

在这项研究中,我们考察了当只有数据集的子集包含推理时 SFT 的效果,这模拟了许多用户的常见现实场景。我们对来自 LLaVA CoT benchmark 的 10k 训练样本进行了 LoRA SFT 的消融实验,使用了推理轨迹。此处的评估 benchmark 是 MathVista,基础模型在此 benchmark 上表现不佳(12.3%),主要原因是格式不一致,而 LoRA SFT 显示出巨大改进(34.4%)。

数学(控制)benchmark — 灾难性遗忘

数学(控制)benchmark 衡量模型在 SFT 后是否保留了通用数学推理能力。这是一个显著的发现:没有 SDR 预填充,随着训练中推理数据百分比的降低,通用数学能力从 68% 崩溃到 3%。使用 SDR,无论你包含多少推理数据,它都稳定在 65–72%。

关键洞察

推理时的中位数推理 token

此图表显示了推理期间用于推理轨迹的 token 预算。SDR 模型产生更多的推理 token(400–860)。使用少于或等于(≤)25% 推理数据且无预填充训练的模型产生零个推理 token。这意味着它们已经完全失去了逐步推理的能力。

关键洞察

MathVista(目标)— 推理开启与关闭

此图表比较了在推理时启用和禁用推理时的 MathVista 准确率。在推理时启用推理影响最大。无论训练数据组成如何,它大致使性能翻倍。

关键洞察

结论与建议

我们在三个 benchmark(MedMCQA、CoCoHD、Invoice-OCR)上的实验以及对部分推理数据集(LLaVA-CoT)的消融研究得出了一个一致的发现:自蒸馏推理(SDR)提供了一种零成本的方法,能够同时提升目标性能并保留通用能力,而无需模型融合。

这些发现的出发点是推理在普通 SFT 下是脆弱的。在没有推理轨迹的数据上进行训练会导致推理抑制:模型走捷径直接给出答案,并失去通用技能,数学准确率从 70% 崩溃到 0%。模型融合可以部分恢复这些技能,但它引入了目标性能与通用性能之间的权衡,难以调整,并且向训练流水线添加了一个检查点混合步骤。

SDR 消除了这种权衡。通过使用基础模型自身的推理轨迹增强数据集,SDR 将通用性能(数学)保持在约 68%,同时在目标性能上比最佳模型融合检查点提供超过 6.5% 的相对增益,且无额外标注成本。即使基础模型在目标领域较弱,该方法也有效:在 MedMCQA 上,基础模型得分为 55.6%,SDR 将目标准确率提升了约 3 个百分点(63.8% 到 66.6%),并将数学从 0% 恢复到 67.9%。换句话说,无论基础模型是否具有潜在的领域技能需要解锁,SDR 都有效。

这种益处扩展到混合推理训练,其中只有一部分 SFT 样本带有推理轨迹。这是一个常见的现实场景,因为推理数据整理成本高昂,并且我们的 LLaVA-CoT 消融研究表明它也很脆弱。通用能力在 75% 的推理覆盖率下保持良好(数学 66.6%),但在此之下急剧崩溃:在 50% 时数学下降到 21.7%,在 25% 时下降到 3.3%,在 0% 时下降到 2.5%。目标准确率在这些混合比例中非单调变化,因此没有安全的稀释点。用 Nova 2 Lite 预填充缺失的轨迹完全消除了这个悬崖:数学在所有混合比例下都保持在 65–71% 的范围内,并且 50% 的 SDR 达到了 45.2% 的目标准确率,高于使用 100% 人工推理观察到的 34.4%。这种增益也传递到推理禁用的推理中,因此团队可以一次性修补他们的数据集并保持现有的服务路径。

一个注意事项影响了关于使用哪个教师模型的建议。更强的教师并不总是更好:Amazon Nova 2 Pro(预览版)轨迹提高了目标准确率,但导致通用性能成比例下降,这与先前关于师生差距大不利于蒸馏的发现一致。Lite 生成的轨迹在目标性能和通用性能之间提供了最平衡的结果,是我们推荐的默认选择。

综合来看,这些结果指向一条实用的实践者路径:使用同族轨迹(Amazon Nova 2 Lite 客户使用 Amazon Nova 2 Lite)的 SDR 作为默认的 SFT 配方,每当推理覆盖率低于约 75% 时将其用作缺口填充器,并为基础模型在目标领域较弱且通用能力成本可接受的情况保留更强教师或引导推理变体。更广泛的意义在于,SDR 是一种有用的归纳偏置,适用于持续学习场景,其中新技能随时间出现,需要保留先前学到的技能。将其与数据混合和轻度融合相结合是下一步自然要探索的方向。

实践者决策指南

使用此指南根据你的数据组成和延迟约束,确定何时使用自蒸馏推理(SDR)、模型融合或标准 SFT。

场景 推荐
1 默认或新的 SFT 任务没有推理,但观察到通用性能回归 如果对延迟敏感:使用模型融合恢复通用性能回归。如果延迟约束允许推理:使用带有 Nova 2 Lite 生成推理轨迹的 SDR。这提供了目标性能和通用能力保留的最佳平衡(通常高于模型融合)。
2 基础模型在目标领域较弱 考虑使用引导推理在轨迹生成期间提供额外信号。
3 数据集已有部分推理(≥ 50%) 我们建议在开启推理的情况下进行训练。现有的轨迹足以保留通用技能。
4 数据集有 如果延迟约束允许推理,用 Nova 2 Lite 推理预填充缺失的轨迹(一次性,离线)。否则,通用性能将崩溃。
5 你不需要通用能力 推理数据对目标性能影响最小。标准 SFT 即可。
6 持续学习或多技能保留 考虑将 SDR 与数据混合和小(或零)融合权重结合使用,以获得最佳保留先前学到的技能。
7 你仍想使用模型融合 我们建议在应用 SDR 时使用较小的融合权重。SDR 已经提供了融合所要补偿的正则化。

有关 Amazon Nova 模型和定制化选项的更多信息,请参阅 Amazon Bedrock 文档。要开始使用 SFT 定制化,请参阅 Amazon Bedrock 模型定制化指南。如果您有问题或想分享您使用 SDR 的经验,请联系您的 AWS 客户团队。

关于作者

Rushil Anirudh Rushil 是 Forge 团队的应用科学家,致力于 SFT 背后的科学——使 LoRA 和全参数微调对前沿模型定制化更准确、更可靠。他的研究涵盖视觉和语言领域的生成模型、模型鲁棒性和不确定性量化。工作之余,他通常在湾区寻找美食,或阅读历史和第一接触科幻小说。

Shiva Mahalingam Shiva 是 AWS 工程、建筑、房地产和运输(ECRT)领域的高级解决方案架构师。他与企业客户合作设计和实施云原生架构,重点关注代理数据、模型定制化和代理 AI 解决方案。他帮助组织应对微调、提示工程和领域适应的复杂性,以从其数据中释放业务价值。他热衷于帮助客户在 AWS 上从实验转向生产就绪的 AI 工作负载。不工作时,他听音乐和演奏打击乐器。

Anupam Dewan Anupam 是 Amazon Nova 团队的高级解决方案架构师,对生成式 AI 及其实际应用充满热情。他专注于 Nova 定制化和 Amazon Nova Forge,帮助企业通过定制化的力量释放 LLM 的真正潜力。他还热衷于教授数据科学和分析,并帮助企业构建适合其业务的 LLM。工作之余,你会发现他在远足、做志愿者或享受大自然。

译自 AWS · ML 博客 · 录于 二〇二六年七月二十三日