<?xml version="1.0" encoding="utf-8"?><feed xmlns="http://www.w3.org/2005/Atom" ><generator uri="https://jekyllrb.com/" version="3.10.0">Jekyll</generator><link href="https://liyongzhi.xyz/feed.xml" rel="self" type="application/atom+xml" /><link href="https://liyongzhi.xyz/" rel="alternate" type="text/html" /><updated>2026-06-04T00:45:26+08:00</updated><id>https://liyongzhi.xyz/feed.xml</id><title type="html">李勇志 (Yongzhi Li)</title><subtitle>Personal website of Yongzhi Li, sharing research, projects, and technical writing on multimodal generation, large language models, AI agents, and computer vision.</subtitle><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><entry><title type="html">DeepSeek-V4 技术深度解读：百万 Token 上下文的注意力革命</title><link href="https://liyongzhi.xyz/posts/2026/05/deepseek-v4-deep-dive/" rel="alternate" type="text/html" title="DeepSeek-V4 技术深度解读：百万 Token 上下文的注意力革命" /><published>2026-05-26T00:00:00+08:00</published><updated>2026-05-26T00:00:00+08:00</updated><id>https://liyongzhi.xyz/posts/2026/05/blog-post-deepseek-v4-deep-dive</id><content type="html" xml:base="https://liyongzhi.xyz/posts/2026/05/deepseek-v4-deep-dive/"><![CDATA[<blockquote>
  <p>这是一篇带交互可视化的长文。<strong>强烈建议直接打开完整交互版</strong>：</p>

  <p>👉 <strong><a href="/files/deepseek_v4/index.html" target="_blank">打开《DeepSeek-V4 深度解读》完整交互报告</a></strong></p>

  <p>（包含 10+ 个可交互图表、SVG 演示、滑块和实时计算——文字读懂一半，动手玩一遍才能真正吃透。）</p>
</blockquote>

<p>DeepSeek-V4 发布后我花了几天时间通读了 58 页的官方技术报告（<a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/resolve/main/DeepSeek_V4.pdf">PDF</a>）、<code class="language-plaintext highlighter-rouge">config.json</code> 和 <code class="language-plaintext highlighter-rouge">inference/model.py</code>，整理了一份带交互可视化的技术分享。这篇 post 是导读，完整内容请点上面的链接。</p>

<h2 id="一句话概括-v4-做了什么">一句话概括 V4 做了什么</h2>

<blockquote>
  <p><strong>把 1.6T 参数模型在 1M token 上下文上跑稳，FLOPs 砍到 V3.2 的 27%，KV cache 砍到 10%。</strong></p>
</blockquote>

<p>实现这件事靠的是三大架构创新 + 一套基建优化：</p>

<ul>
  <li><strong>CSA + HCA 混合注意力</strong>：两种 KV cache 压缩策略交替排列，再加滑动窗口兜底</li>
  <li><strong>mHC 流形约束超连接</strong>：让 1.6T 模型训练 32T+ token 不崩</li>
  <li><strong>Muon 优化器</strong>：把权重矩阵当几何对象做正交化更新，99.9% 参数用 Muon</li>
  <li><strong>工程层</strong>：MegaMoE / FP4 QAT / Anticipatory Routing / On-Disk KV Cache 等 8 个亮点</li>
</ul>

<h2 id="完整报告涵盖什么">完整报告涵盖什么</h2>

<table>
  <thead>
    <tr>
      <th>章节</th>
      <th>内容</th>
      <th>交互亮点</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>1. 背景与动机</td>
      <td>为什么 KV cache 既是显存大户也是 API 账单大户</td>
      <td>API 价格对比表</td>
    </tr>
    <tr>
      <td>2. 注意力演化史</td>
      <td>MHA → MQA → GQA → MLA → DSA → NSA → CSA/HCA</td>
      <td><strong>7 张 SVG 配图</strong> + KV 演化对数柱状图</td>
    </tr>
    <tr>
      <td>3. CSA + HCA 详解</td>
      <td>V4 最大的创新，逐组件拆解</td>
      <td><strong>可拖拽 token 位置看每层 attend 什么</strong> + 50K~1M 上下文 KV breakdown</td>
    </tr>
    <tr>
      <td>4. mHC 流形约束</td>
      <td>双随机矩阵 + Sinkhorn-Knopp</td>
      <td><strong>滑块对比信号增益</strong> + Sinkhorn-Knopp 单步动画</td>
    </tr>
    <tr>
      <td>5. Muon 优化器</td>
      <td>从 SGD 到几何更新</td>
      <td><strong>奇异值分布随 NS 迭代演化</strong></td>
    </tr>
    <tr>
      <td>6. 基建优化</td>
      <td>8 个工程亮点</td>
      <td>tabs 切换 + 通信-计算重叠 SVG</td>
    </tr>
    <tr>
      <td>7. 训练 + 评测</td>
      <td>33T token + 渐进式上下文 + benchmark 对比</td>
      <td>雷达图 + MRCR 长上下文曲线</td>
    </tr>
    <tr>
      <td>8. 总结展望</td>
      <td>三条主线 + 未来方向</td>
      <td>—</td>
    </tr>
  </tbody>
</table>

<h2 id="几个让我觉得设计很漂亮的细节">几个让我觉得设计很漂亮的细节</h2>

<p><strong>KV cache 不只是技术问题，它直接和钱挂钩。</strong> 你看 API 定价表三列价格（命中缓存 / 未命中 / 输出），中间差 10 倍，差的就是要不要从头重算 KV cache。2026 年 3 月 Claude Code 翻车——在系统提示前面塞了”反滥用 token + 反蒸馏假工具”，每轮都不一样，prompt prefix 每轮都变、cache 每轮 miss，用户成本暴涨 10–20 倍。</p>

<p><strong>MLA 的数学技巧</strong>（V2 起就用）：K 和 V 从来没被显式算出来。”还原 K”的 <code class="language-plaintext highlighter-rouge">W_uk</code> 被吸收进 Q 投影，”还原 V”的 <code class="language-plaintext highlighter-rouge">W_uv</code> 被吸收进输出投影。推理时这些都是固定权重，零额外开销，全程只有 latent 出场。这是 DeepSeek 能撑起百万上下文的基础。</p>

<p><strong>CSA 重叠 / HCA 不重叠</strong>——不是设计偏好，是数学算账。CSA 块只有 4 个 token，块边界被切断代价大（一个完整词组可能被劈开），重叠收益高；HCA 块 128 个 token，块边界代价小，重叠要付的代价（投影维度翻倍）巨大。</p>

<p><strong>Lightning Indexer 全程 FP4</strong>——主路径 FP8 不能再砍（CSA 压缩已经有损，再砍内容就没了）；Indexer 只需要排序，对精度容忍度高，所以走 FP4。</p>

<p><strong>训练就用稀疏，不是事后套眼镜。</strong> V3.2 DSA 是”一个看惯了高清电视的人，戴个轻便近视眼镜上街”；V4 是”一个从小就近视、从小就戴这副眼镜的人，大脑里世界长什么样就是这副眼镜下的样子”。这是 V4 推理 FLOPs 砍到 27% 的根本原因。</p>

<h2 id="一个哲学层面的洞察">一个哲学层面的洞察</h2>

<p>KV cache 能被压掉 99.7% 而效果不掉，<strong>不是算法的胜利，是语言本身的低维性质</strong>。</p>

<p>大模型权重存的是”知识”——世界的地图。KV cache 存的不是地图，是当前这段上下文<strong>在地图上走过的路径</strong>。路径能被压到 0.3%，是因为自然语言的有效路径本来就是低维的——7168 维的隐藏空间里，大部分维度是 MLP 临时展开的”计算脚手架”，真正的语义坐标集中在几百维子空间里（intrinsic dimension 实测约占隐藏空间个位数百分比）。</p>

<p>V4 所有压缩——MLA 砍维度、Compressor 砍条数、Indexer 选 top-k——<strong>都在压路径，没人动地图</strong>。</p>

<h2 id="完整报告">完整报告</h2>

<p>👉 <strong><a href="/files/deepseek_v4/index.html" target="_blank">打开《DeepSeek-V4 深度解读》完整交互报告</a></strong></p>

<p>报告内容均基于：</p>

<ul>
  <li><a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/resolve/main/DeepSeek_V4.pdf">DeepSeek-V4 官方技术报告（58 页 PDF）</a></li>
  <li><a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/blob/main/config.json">DeepSeek-V4-Pro config.json</a></li>
  <li><a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/blob/main/inference/model.py">DeepSeek-V4-Pro inference/model.py</a></li>
  <li>以及 NSA、MLA、Muon 等历代相关论文</li>
</ul>

<p>如有任何问题或勘误，欢迎反馈。</p>]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><category term="blog" /><category term="LLM" /><category term="DeepSeek" /><category term="attention" /><category term="MoE" /><category term="long-context" /><category term="KV cache" /><category term="optimizer" /><summary type="html"><![CDATA[这是一篇带交互可视化的长文。强烈建议直接打开完整交互版： 👉 打开《DeepSeek-V4 深度解读》完整交互报告 （包含 10+ 个可交互图表、SVG 演示、滑块和实时计算——文字读懂一半，动手玩一遍才能真正吃透。）]]></summary></entry><entry><title type="html">Blog Post Llm Mllm Posttrain Interview</title><link href="https://liyongzhi.xyz/blog-post-llm-mllm-posttrain-interview/" rel="alternate" type="text/html" title="Blog Post Llm Mllm Posttrain Interview" /><published>2026-05-26T00:00:00+08:00</published><updated>2026-05-26T00:00:00+08:00</updated><id>https://liyongzhi.xyz/blog-post-llm-mllm-posttrain-interview</id><content type="html" xml:base="https://liyongzhi.xyz/blog-post-llm-mllm-posttrain-interview/"><![CDATA[<h1 id="llm--mllm-后训练面试题库60-题精选">LLM / MLLM 后训练面试题库（60 题精选）</h1>

<blockquote>
  <p>📖 本题库覆盖大模型后训练全链路：从 SFT 到 RLHF/DPO，从 GRPO 到推理模型，从 LoRA 到 MLLM 多模态专项，另附 10 道代码手撕题。</p>
</blockquote>

<h2 id="目录">目录</h2>
<ul>
  <li><a href="#一基础概念">一、基础概念</a> Q01-Q05</li>
  <li><a href="#二SFT">二、SFT 监督微调</a> Q06-Q11</li>
  <li><a href="#三RM">三、奖励建模 RM</a> Q12-Q15</li>
  <li><a href="#四RLHF">四、RLHF / PPO</a> Q16-Q20</li>
  <li><a href="#五DPO">五、DPO 及变体</a> Q21-Q24</li>
  <li><a href="#六GRPO">六、GRPO / 推理模型</a> Q25-Q28</li>
  <li><a href="#七PEFT">七、PEFT / LoRA / QLoRA</a> Q29-Q31</li>
  <li><a href="#八数据工程">八、数据工程</a> Q32-Q34</li>
  <li><a href="#九MLLM">九、MLLM 后训练专项</a> Q35-Q42</li>
  <li><a href="#十对齐安全">十、对齐与安全</a> Q43-Q45</li>
  <li><a href="#十一训练工程">十一、训练工程</a> Q46-Q47</li>
  <li><a href="#十二评估">十二、评估</a> Q48-Q49</li>
  <li><a href="#十三前沿">十三、前沿与开放题</a> Q50</li>
  <li><a href="#十四代码手撕">十四、手撕代码 / Coding</a> Q51-Q60</li>
</ul>

<hr />

<h2 id="一基础概念">一、基础概念</h2>

<h3 id="q01-什么是大模型后训练post-training包括哪些阶段">Q01 什么是大模型后训练（Post-Training）？包括哪些阶段？</h3>

<p><strong>难度：</strong> 基础
<strong>考察点：</strong> 对后训练全链路的系统性理解，能否区分各阶段的输入/输出/目标</p>

<p><strong>满分回答：</strong></p>

<p>Post-Training 是基座模型（Base Model）在预训练完成之后、面向下游任务进行的一系列训练阶段的总称。其目标是将一个”续写文本”的基座模型转化为一个”遵循指令、安全可控”的对话模型。</p>

<p>典型后训练 pipeline 包含以下阶段：</p>

<ol>
  <li><strong>SFT（Supervised Fine-Tuning）</strong> — 用人工标注的指令-回复对微调基座模型，使其学会按指令格式回答。输入：instruction-response pairs；输出：SFT 模型（policy 初始版本）。</li>
  <li><strong>Reward Modeling</strong> — 收集人类偏好数据（对同一 prompt 的两个回复做排序），训练一个 Reward Model（RM）来预测人类偏好。输入：preference pairs $(y_w, y_l)$；输出：标量奖励函数 $r_\theta(x, y)$。</li>
  <li><strong>RLHF / DPO / GRPO</strong> — 利用 RM（或偏好数据直接优化）进一步对齐 SFT 模型，使其生成更符合人类偏好的回复。输入：RM + SFT policy（RLHF）或偏好数据（DPO）；输出：对齐后的对话模型。</li>
</ol>

<p>⚠️ 常见误解：认为 Post-Training 只等于 SFT。实际上 SFT 只完成了格式对齐，价值观/偏好对齐需要 RLHF/DPO 等后续阶段。</p>

<p>⚠️ 不同厂商的 pipeline 有差异：LLaMA 2 的流程是 SFT → RM → PPO；LLaMA 3 增加了多轮对话 SFT 和 DPO；DeepSeek-R1 在 RL 前还加入了长 CoT SFT + 规则奖励 GRPO。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>为什么不能跳过 SFT 直接做 RLHF？（SFT 提供了 policy 的初始化，直接从 base model 做 RL 会导致输出不稳定）</li>
  <li>Post-Training 中哪个阶段对最终效果影响最大？（经验上 SFT 数据质量是瓶颈，RLHF/DPO 提升幅度约 5-15%）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2203.02155">Training language models to follow instructions with human feedback (InstructGPT)</a></li>
  <li><a href="https://arxiv.org/abs/2407.21783">The LLaMA 3 Herd of Models</a></li>
</ul>

<hr />

<h3 id="q02-pre-training-vs-post-training-的核心差异">Q02 Pre-training vs Post-training 的核心差异？</h3>

<p><strong>难度：</strong> 基础
<strong>考察点：</strong> 能否从目标、数据、方法、成本等多维度对比两个阶段</p>

<p><strong>满分回答：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>Pre-training</th>
      <th>Post-training</th>
      <th> </th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>目标</strong></td>
      <td>学习语言的统计规律（next-token prediction）</td>
      <td>学习遵循指令、符合人类偏好</td>
      <td> </td>
    </tr>
    <tr>
      <td><strong>数据</strong></td>
      <td>海量无标注文本（TB 级，web crawl）</td>
      <td>少量高质量标注数据（K~M 级）</td>
      <td> </td>
    </tr>
    <tr>
      <td><strong>方法</strong></td>
      <td>自回归语言建模，$\mathcal{L} = -\sum_t \log P(w_t</td>
      <td>w_{&lt;t})$</td>
      <td>SFT（监督微调）→ RLHF/DPO（偏好对齐）</td>
    </tr>
    <tr>
      <td><strong>模型规模</strong></td>
      <td>全参数训练</td>
      <td>全参数或 PEFT（LoRA 等）</td>
      <td> </td>
    </tr>
    <tr>
      <td><strong>计算成本</strong></td>
      <td>数千 GPU × 数周</td>
      <td>数十 GPU × 数天</td>
      <td> </td>
    </tr>
    <tr>
      <td><strong>评估方式</strong></td>
      <td>perplexity、loss 曲线</td>
      <td>MT-Bench、Arena Score、人工评测</td>
      <td> </td>
    </tr>
    <tr>
      <td><strong>输出</strong></td>
      <td>Base Model（续写能力）</td>
      <td>Chat Model（对话能力）</td>
      <td> </td>
    </tr>
  </tbody>
</table>

<p>核心差异的本质是<strong>目标函数的变化</strong>：预训练优化的是”预测下一个 token 的似然”，后训练优化的是”生成人类偏好的回复”。这导致了数据形态、训练策略、评估体系的根本不同。</p>

<p>⚠️ 常见坑：Pre-training 的 loss 下降 ≠ 模型变好用。Base model loss 低但不会遵循指令，需要 Post-Training 来”教”它。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>能否用 Post-Training 的数据量做 Pre-training？（不行，数据量太少会导致严重过拟合）</li>
  <li>Pre-training 和 Post-Training 的 loss 公式本质相同吗？（SFT 的 loss 形式与预训练相同，都是交叉熵；RLHF/DPO 的 loss 则完全不同）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2302.13971">LLaMA: Open and Efficient Foundation Language Models</a></li>
  <li><a href="https://arxiv.org/abs/2305.07954">LLaMA 2: Open Foundation and Fine-Tuned Chat Models</a></li>
</ul>

<hr />

<h3 id="q03-sft-vs-rlhf-的区别与联系">Q03 SFT vs RLHF 的区别与联系？</h3>

<p><strong>难度：</strong> 基础
<strong>考察点：</strong> 理解两个核心阶段的不同目标和互补关系</p>

<p><strong>满分回答：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>SFT</th>
      <th>RLHF</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>优化目标</strong></td>
      <td>让模型学会按格式回答</td>
      <td>让模型输出更符合人类偏好</td>
    </tr>
    <tr>
      <td><strong>数据类型</strong></td>
      <td>instruction-response pairs</td>
      <td>preference comparisons $(y_w &gt; y_l)$</td>
    </tr>
    <tr>
      <td><strong>优化方法</strong></td>
      <td>监督学习（交叉熵 loss）</td>
      <td>强化学习（PPO）或直接偏好优化（DPO）</td>
    </tr>
    <tr>
      <td><strong>模型组件</strong></td>
      <td>只需 policy model</td>
      <td>policy + reward model + reference model</td>
    </tr>
    <tr>
      <td><strong>训练信号</strong></td>
      <td>确定性标签（ground truth）</td>
      <td>偏好信号（相对排序）</td>
    </tr>
    <tr>
      <td><strong>学习内容</strong></td>
      <td>格式、任务能力</td>
      <td>价值观、风格偏好</td>
    </tr>
  </tbody>
</table>

<p>联系：</p>
<ul>
  <li>SFT 是 RLHF 的<strong>前置阶段</strong>：RLHF 的 policy 初始权重来自 SFT model，reference model 也来自 SFT model。</li>
  <li>SFT 解决”能不能做”，RLHF 解决”做得好不好”（偏好层面）。</li>
  <li>两者互补：只用 SFT 会偏向模仿数据风格但缺乏偏好细化；只用 RLHF（跳过 SFT）则 policy 输出不稳定。</li>
</ul>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>为什么 RLHF 不能替代 SFT？（RLHF 需要一个能生成合理回复的 policy 作为起点，base model 直接做 RL 输出质量太差）</li>
  <li>SFT 之后只用 DPO 而不用 PPO 可以吗？（可以，DPO 是 RLHF 的替代方案，直接从偏好数据优化 policy，不需要 RM）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2203.02155">Training language models to follow instructions with human feedback</a></li>
  <li><a href="https://arxiv.org/abs/2305.18290">Direct Preference Optimization</a></li>
</ul>

<hr />

<h3 id="q04-alignment对齐是什么为什么需要对齐">Q04 Alignment（对齐）是什么？为什么需要对齐？</h3>

<p><strong>难度：</strong> 基础
<strong>考察点：</strong> 对”对齐”概念的理解深度，能否区分 Helpful / Honest / Harmless 三个维度</p>

<p><strong>满分回答：</strong></p>

<p>Alignment 是让模型的行为与人类意图和价值观一致的过程。Anthropic 提出对齐的三个目标（HHH 原则）：</p>

<ul>
  <li><strong>Helpful</strong>：有用，能准确完成用户指令</li>
  <li><strong>Honest</strong>：诚实，不编造事实，承认不确定性</li>
  <li><strong>Harmless</strong>：无害，不生成有害、歧视、危险内容</li>
</ul>

<p>为什么需要对齐？Base model 的训练目标是预测下一个 token，它学会的是语言的统计分布，而非人类意图。具体问题：</p>

<ol>
  <li><strong>指令遵循缺失</strong>：Base model 看到 “请翻译这句话” 可能续写另一句话而非翻译</li>
  <li><strong>安全性问题</strong>：可能生成有害内容（暴力、歧视等）</li>
  <li><strong>偏好偏差</strong>：可能输出冗长、自相矛盾、不符合用户风格偏好的内容</li>
  <li><strong>事实性不足</strong>：可能自信地编造不存在的信息（hallucination）</li>
</ol>

<p>对齐通过 SFT（格式对齐）+ RLHF/DPO（偏好对齐）+ Red-teaming（安全对齐）逐步解决这些问题。</p>

<p>⚠️ 对齐不是”让模型永远说好话”，而是在 helpful 和 harmless 之间找平衡。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>对齐会不会降低模型能力？（是的，这就是”对齐税”，见 Q45）</li>
  <li>对齐能解决 hallucination 吗？（部分缓解，但幻觉的根本原因在预训练知识覆盖度，对齐无法根治）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2204.05862">Training a Helpful and Harmless Assistant with RLHF</a></li>
  <li><a href="https://arxiv.org/abs/2212.08071">Constitutional AI: Harmlessness from AI Feedback</a></li>
</ul>

<hr />

<h3 id="q05-基座模型到对话模型的完整训练-pipeline">Q05 基座模型到对话模型的完整训练 pipeline？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 能否完整描述从 Base Model 到 Chat Model 的全流程，包括各阶段衔接和工程细节</p>

<p><strong>满分回答：</strong></p>

<p>以 InstructGPT / LLaMA 2 为典型参考，完整 pipeline 如下：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Base Model → SFT → Reward Model → RLHF(PPO) / DPO → Chat Model
</code></pre></div></div>

<p><strong>阶段 1：SFT</strong></p>
<ul>
  <li>输入：人类标注的 instruction-response 数据集（约 10K~100K 条）</li>
  <li>方法：全参数或 LoRA 微调，loss 为交叉熵（仅计算 response token）</li>
  <li>输出：SFT Model（同时作为后续 RLHF 的 policy 初始化和 reference model）</li>
  <li>工程细节：mask prompt tokens，多轮对话用对话模板，超参 lr≈1e-5</li>
</ul>

<p><strong>阶段 2：Reward Modeling</strong></p>
<ul>
  <li>输入：对同一 prompt 生成多个回复，人类标注偏好排序</li>
  <li>方法：训练 RM 用 Bradley-Terry loss：$\mathcal{L} = -\log \sigma(r(x, y_w) - r(x, y_l))$</li>
  <li>输出：Reward Model $r_\theta(x, y)$，为 RLHF 提供奖励信号</li>
</ul>

<p><strong>阶段 3：RLHF / DPO</strong></p>
<ul>
  <li>RLHF（PPO 方式）：policy 在 RM 指导下优化，KL 约束防止偏离 reference model</li>
  <li>DPO：直接从偏好数据优化 policy，跳过 RM 训练</li>
  <li>输出：Aligned Chat Model</li>
</ul>

<p><strong>可选阶段 4：迭代对齐</strong></p>
<ul>
  <li>用对齐后的 model 重新生成数据 → 重新训练 RM → 再次 RLHF（LLaMA 2 做了多轮迭代）</li>
  <li>或对 Chat Model 做 Red-teaming + 安全 SFT 补丁</li>
</ul>

<p>⚠️ 实际工程中各阶段不是严格串行的。LLaMA 3 在 SFT 后做了多轮对话 SFT + DPO + 安全微调的组合。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>各阶段的超参如何设置？（SFT lr=1e-5, epoch=2-3; RM lr=5e-6; PPO lr=1e-6, 需仔细调）</li>
  <li>能否用同一批数据做 SFT 和 DPO？（可以但要注意格式：SFT 用单条数据，DPO 用 pair）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2203.02155">Training language models to follow instructions with human feedback</a></li>
  <li><a href="https://arxiv.org/abs/2305.07954">LLaMA 2: Open Foundation and Fine-Tuned Chat Models</a></li>
</ul>

<hr />

<h2 id="二sft-监督微调">二、SFT 监督微调</h2>

<h3 id="q06-sft-数据构造的最佳实践">Q06 SFT 数据构造的最佳实践？</h3>

<p><strong>难度：</strong> 基础
<strong>考察点：</strong> 对 SFT 数据质量标准的理解，能否区分不同数据类型的作用</p>

<p><strong>满分回答：</strong></p>

<p>SFT 数据构造的核心原则：<strong>质量 » 数量</strong>。LIMA 论文证明 1000 条高质量数据可以训练出接近 GPT-4 效果的模型。</p>

<p><strong>数据来源与构造方法：</strong></p>

<ol>
  <li><strong>人工标注</strong>：专业标注员撰写 instruction-response pairs，质量最高但成本大</li>
  <li><strong>Self-Instruct</strong>：用强模型自动生成指令和回复，再人工筛选（Alpaca 方法）</li>
  <li><strong>Magpie</strong>：利用 LLM 的对话模板，仅填入 prompt 部分让模型自生成指令（零人工输入）</li>
  <li><strong>蒸馏数据</strong>：从 GPT-4/Claude 等强模型获取回复，用于训练开源模型</li>
</ol>

<p><strong>数据质量标准：</strong></p>
<ul>
  <li>多样性：覆盖不同任务类型（问答、写作、推理、代码等）</li>
  <li>准确性：回复内容事实正确</li>
  <li>格式一致性：统一的对话模板</li>
  <li>长度适中：避免过于冗长或过短</li>
  <li>拒绝样本：包含模型应拒绝回答的 prompt（安全对齐）</li>
</ul>

<p>⚠️ 常见坑：用弱模型生成 SFT 数据会导致”模型蒸馏退化”——弱模型的错误模式会被学习。</p>

<p><strong>数据配比建议（LLaMA 3 实践）：</strong>
| 数据类型 | 占比 |
|———-|——|
| 通用对话 | ~50% |
| 代码/推理 | ~20% |
| 长文档/总结 | ~15% |
| 拒绝/安全 | ~10% |
| 多语言 | ~5% |</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>Self-Instruct 生成的数据如何保证质量？（人工审核 + 规则过滤 + 多样性采样）</li>
  <li>SFT 数据需要覆盖多少个任务类型？（至少 20+ 种，否则泛化差）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2212.10560">Self-Instruct: Aligning Language Models with Self-Generated Instructions</a></li>
  <li><a href="https://arxiv.org/abs/2406.08486">Magpie: Alignment Data Synthesis from Scratch</a></li>
</ul>

<hr />

<h3 id="q07-sft-的-loss-计算方式为什么要-mask-prompt-tokens">Q07 SFT 的 loss 计算方式，为什么要 mask prompt tokens？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 理解 SFT loss 的工程细节，特别是 prompt masking 的原因</p>

<p><strong>满分回答：</strong></p>

<p>SFT 的 loss 是标准的交叉熵，但只对 <strong>response tokens</strong> 计算：</p>

\[\mathcal{L} = -\frac{1}{N_{\text{resp}}} \sum_{t \in \text{response}} \log P_\theta(y_t | x, y_{&lt;t})\]

<p>其中 $x$ 是 prompt，$y$ 是 response，$N_{\text{resp}}$ 是 response token 数量。</p>

<p><strong>为什么要 mask prompt tokens？</strong></p>

<ol>
  <li><strong>目标语义</strong>：SFT 的目标是让模型学会”给定 prompt 后生成正确的 response”，不是让模型学会”重新生成 prompt”。Prompt 是条件输入，不是预测目标。</li>
  <li><strong>防止偏移</strong>：如果不 mask prompt，模型会花大量梯度去拟合 prompt 的分布，导致 response 部分的学习信号被稀释。</li>
  <li><strong>工程实现</strong>：在 <code class="language-plaintext highlighter-rouge">ignore_index</code> 参数中设置 prompt token 的 label 为 -100（PyTorch 默认忽略值），这样 cross_entropy 会自动跳过这些位置。</li>
</ol>

<p>⚠️ 常见坑：初学者有时把整个序列都算 loss，这会导致模型倾向于生成类似 prompt 的内容而非回答。</p>

<p><strong>代码示例：</strong></p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">labels</span> <span class="o">=</span> <span class="n">full_sequence_ids</span><span class="p">.</span><span class="n">clone</span><span class="p">()</span>
<span class="n">labels</span><span class="p">[:</span><span class="n">prompt_len</span><span class="p">]</span> <span class="o">=</span> <span class="o">-</span><span class="mi">100</span>  <span class="c1"># mask prompt tokens
</span><span class="n">loss</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">cross_entropy</span><span class="p">(</span><span class="n">logits</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">vocab_size</span><span class="p">),</span> <span class="n">labels</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">),</span> <span class="n">ignore_index</span><span class="o">=-</span><span class="mi">100</span><span class="p">)</span>
</code></pre></div></div>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>如果 prompt 中有特殊 token（如 system message），是否也要 mask？（是的，所有非 response 部分都 mask）</li>
  <li>loss 是按 token 平均还是按 response 平均？（按 response token 数平均，不是按整个序列长度）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2203.02155">Training language models to follow instructions with human feedback</a></li>
</ul>

<hr />

<h3 id="q08-多轮对话-sft-的训练模板与-loss-设计">Q08 多轮对话 SFT 的训练模板与 loss 设计？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 能否处理多轮对话的特殊 loss 设计，理解不同 masking 策略的利弊</p>

<p><strong>满分回答：</strong></p>

<p>多轮对话 SFT 需要将多轮交互拼接成一条序列，使用对话模板格式化后训练。关键问题在于 <strong>loss mask 的范围</strong>。</p>

<p><strong>对话模板示例（LLaMA 格式）：</strong></p>
<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>&lt;|begin_of_text|&gt;&lt;|start_header_id|&gt;user&lt;|end_header_id|&gt;
{round1_question}&lt;|eot_id|&gt;&lt;|start_header_id|&gt;assistant&lt;|end_header_id|&gt;
{round1_answer}&lt;|eot_id|&gt;&lt;|start_header_id|&gt;user&lt;|end_header_id|&gt;
{round2_question}&lt;|eot_id|&gt;&lt;|start_header_id|&gt;assistant&lt;|end_header_id|&gt;
{round2_answer}&lt;|eot_id|&gt;
</code></pre></div></div>

<p><strong>Loss masking 三种策略：</strong></p>

<table>
  <thead>
    <tr>
      <th>策略</th>
      <th>Mask 范围</th>
      <th>优点</th>
      <th>缺点</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>仅最后一轮 response</td>
      <td>只算最后一轮 assistant 回复</td>
      <td>训练快，数据利用低</td>
      <td>只学习最后一轮</td>
    </tr>
    <tr>
      <td>所有 assistant 回复</td>
      <td>每轮 assistant 回复都算 loss</td>
      <td>数据利用充分</td>
      <td>不同轮次权重相同</td>
    </tr>
    <tr>
      <td>全部 response + 历史</td>
      <td>assistant 回复 + 用户提问</td>
      <td>信息最大化</td>
      <td>可能偏移用户风格</td>
    </tr>
  </tbody>
</table>

<p><strong>主流实践</strong>：采用”所有 assistant 回复都算 loss”策略。每个 assistant 回复段的 token label 设为自身 ID，其余设为 -100。</p>

<p>⚠️ 常见坑：如果用”仅最后一轮”策略，模型可能在前几轮生成时缺乏训练信号，导致多轮交互不稳定。</p>

<p><strong>工程细节：</strong></p>
<ul>
  <li>多轮对话需要确保各轮之间的 token 级别 mask 精确，不能遗漏特殊 token（如 <code class="language-plaintext highlighter-rouge">&lt;|eot_id|&gt;</code>）</li>
  <li>实际训练中通常将多条多轮对话 padding 到同一长度，用 attention mask 和 loss mask 配合</li>
</ul>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>多轮对话 SFT 的 batch 内如何处理不同轮数的对话？（padding + attention mask，loss mask 精确标记每条对话的 response 位置）</li>
  <li>是否需要专门训练模型学会”何时停止生成”？（是的，EOS token 的预测也是训练目标的一部分）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2305.07954">LLaMA 2: Open Foundation and Fine-Tuned Chat Models</a></li>
</ul>

<hr />

<h3 id="q09-sft-数据量与质量的权衡多少条数据够">Q09 SFT 数据量与质量的权衡，多少条数据够？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 对数据规模直觉的理解，能否引用关键论文结论</p>

<p><strong>满分回答：</strong></p>

<p>SFT 的数据量阈值远低于预训练。关键结论：</p>

<table>
  <thead>
    <tr>
      <th>研究</th>
      <th>数据量</th>
      <th>结论</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>LIMA</td>
      <td>1,000 条</td>
      <td>1000 条高质量数据足以接近 GPT-4 水平</td>
    </tr>
    <tr>
      <td>Alpaca</td>
      <td>52K 条</td>
      <td>Self-Instruct 生成的 52K 条即可显著提升</td>
    </tr>
    <tr>
      <td>FLAN/T0</td>
      <td>数百万条</td>
      <td>多任务提示数据，追求零样本泛化</td>
    </tr>
    <tr>
      <td>LLaMA 2</td>
      <td>~27K 条（对话 SFT）</td>
      <td>人工标注的高质量数据</td>
    </tr>
  </tbody>
</table>

<p><strong>核心原则：质量 » 数量。</strong> 原因：</p>

<ol>
  <li><strong>低质量数据的危害</strong>：包含错误的回复会直接教模型犯错，比没有数据更糟糕</li>
  <li><strong>重复数据的危害</strong>：模型会过拟合高频模式，泛化能力下降</li>
  <li><strong>多样性比数量重要</strong>：覆盖更多任务类型的 10K 条数据 &gt; 同类型重复的 100K 条</li>
</ol>

<p><strong>实践建议：</strong></p>
<ul>
  <li>初步 SFT：5K-20K 条高质量数据</li>
  <li>生产级 SFT：50K-200K 条（含多任务、多语言、安全样本）</li>
  <li>关键是数据审核流程：每条数据至少经过格式检查 + 内容准确性验证</li>
</ul>

<p>⚠️ 常见坑：盲目追求数据量而忽略审核，用自动生成数据不做人工筛选。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>如何判断 SFT 数据是否”够”？（看 eval benchmark 是否收敛，训练 loss 是否不再下降但 eval 持续提升 = 过拟合信号）</li>
  <li>低质量数据的影响能通过 RLHF 修复吗？（部分可以，但最好在 SFT 层面就保证质量）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li>LIMA: Less Is More for Alignment (社区讨论)</li>
  <li><a href="https://arxiv.org/abs/2210.11416">Scaling Instruction-Finetuned Language Models (Flan-T5)</a></li>
</ul>

<hr />

<h3 id="q10-全参数微调-vs-lora-微调在-sft-中的对比">Q10 全参数微调 vs LoRA 微调在 SFT 中的对比？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 能否从效果、效率、适用场景等维度对比两种微调方式</p>

<p><strong>满分回答：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>全参数微调</th>
      <th>LoRA 微调</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>参数更新量</strong></td>
      <td>所有参数</td>
      <td>低秩增量 $\Delta W = BA$，$B \in \mathbb{R}^{d \times r}, A \in \mathbb{R}^{r \times d}$，$r \ll d$</td>
    </tr>
    <tr>
      <td><strong>显存占用</strong></td>
      <td>极高（优化器状态占 2×参数量）</td>
      <td>低（只存 LoRA 参数的优化器状态）</td>
    </tr>
    <tr>
      <td><strong>训练速度</strong></td>
      <td>较慢</td>
      <td>较快（梯度计算量少）</td>
    </tr>
    <tr>
      <td><strong>最终效果</strong></td>
      <td>理论上限更高</td>
      <td>对于 SFT 场景差距通常 &lt;2%</td>
    </tr>
    <tr>
      <td><strong>灵活性</strong></td>
      <td>需要完整模型权重</td>
      <td>LoRA 可随时 merge/unmerge</td>
    </tr>
    <tr>
      <td><strong>多任务适配</strong></td>
      <td>每个任务需要完整权重</td>
      <td>可为每个任务维护独立 LoRA adapter</td>
    </tr>
    <tr>
      <td><strong>灾难性遗忘风险</strong></td>
      <td>较高</td>
      <td>较低（基座权重冻结）</td>
    </tr>
  </tbody>
</table>

<p><strong>LoRA 的关键公式：</strong>
\(h = W_0 x + \Delta W x = W_0 x + BAx\)</p>

<p>其中 $W_0$ 冻结，$B, A$ 可训练，$r$ 通常取 8-64。</p>

<p><strong>何时用全参数微调？</strong></p>
<ul>
  <li>模型规模 &lt;7B 且 GPU 资源充足</li>
  <li>需要大幅修改模型行为（如从 base → chat 的首次 SFT）</li>
  <li>目标 benchmark 要求极致性能</li>
</ul>

<p><strong>何时用 LoRA？</strong></p>
<ul>
  <li>模型规模 &gt;13B，GPU 资源有限</li>
  <li>多任务场景，需要多个 adapter</li>
  <li>实验迭代阶段，需要快速试错</li>
</ul>

<p>⚠️ 常见坑：LoRA 秩 $r$ 太小（如 $r=1$）会导致效果明显下降；太大（如 $r=256$）接近全参数但显存没省多少。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>LoRA 应该加在哪些层？（主流做法：加在所有线性层（Q/K/V/O + FFN），不加 embedding/lm_head）</li>
  <li>LoRA 微调后如何部署？（merge $W = W_0 + BA$ 后部署，推理无额外开销）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2106.09685">LoRA: Low-Rank Adaptation of Large Language Models</a></li>
</ul>

<hr />

<h3 id="q11-灾难性遗忘问题及缓解方法">Q11 灾难性遗忘问题及缓解方法？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 理解灾难性遗忘的成因和主流缓解策略</p>

<p><strong>满分回答：</strong></p>

<p>灾难性遗忘（Catastrophic Forgetting）是指模型在学习新任务时，丢失了之前学到的知识。在 LLM 后训练中表现为：SFT 后模型在通用能力（推理、代码等）上退步。</p>

<p><strong>成因分析：</strong></p>
<ol>
  <li><strong>参数偏移</strong>：SFT 的梯度更新改变了预训练学到的权重分布</li>
  <li><strong>数据偏差</strong>：SFT 数据分布与预训练数据分布差异大</li>
  <li><strong>训练过度</strong>：epoch 过多或 lr 过大导致权重大幅偏移</li>
</ol>

<p><strong>缓解方法：</strong></p>

<table>
  <thead>
    <tr>
      <th>方法</th>
      <th>原理</th>
      <th>适用场景</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>LoRA / PEFT</strong></td>
      <td>只更新低秩增量，冻结基座权重</td>
      <td>最主流的方案</td>
    </tr>
    <tr>
      <td><strong>数据混合</strong></td>
      <td>SFT 数据中混入预训练数据（如 LLaMA 3 混了 ~5% 预训练数据）</td>
      <td>生产级 SFT</td>
    </tr>
    <tr>
      <td><strong>LR 衰减</strong></td>
      <td>使用较小的学习率（1e-5 ~ 5e-6）</td>
      <td>所有 SFT 场景</td>
    </tr>
    <tr>
      <td><strong>早停</strong></td>
      <td>监控通用 benchmark，在退化开始时停止</td>
      <td>需要多阶段评估</td>
    </tr>
    <tr>
      <td><strong>EWC / L2 正则</strong></td>
      <td>对重要参数施加正则化 $\lambda \sum_i F_i (w_i - w_i^0)^2$</td>
      <td>理论上有用，实践少用</td>
    </tr>
    <tr>
      <td><strong>多阶段训练</strong></td>
      <td>先 SFT 对齐格式，再用 RLHF 微调偏好，逐步调整</td>
      <td>InstructGPT pipeline</td>
    </tr>
    <tr>
      <td><strong>Replay</strong></td>
      <td>定期用预训练数据做”回放”训练</td>
      <td>需要预训练数据访问权限</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 最实用的方案是 <strong>LoRA + 数据混合 + 低 lr</strong>。单独用任何一个都不够。</p>

<p>⚠️ 常见坑：认为 LoRA 完全不会遗忘。LoRA 只是降低遗忘程度，如果 SFT 数据偏差太大，遗忘仍然会发生。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>如何检测灾难性遗忘？（在 SFT 过程中持续评估通用 benchmark，如 MMLU/GSM8K，看是否下降）</li>
  <li>数据混合的比例如何确定？（LLaMA 3 实验发现 5-10% 预训练数据混合效果最好）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2407.21783">LLaMA 3 Herd of Models</a></li>
  <li><a href="https://arxiv.org/abs/2106.09685">LoRA: Low-Rank Adaptation</a></li>
</ul>

<hr />

<h2 id="三奖励建模-rm">三、奖励建模 RM</h2>

<h3 id="q12-reward-model-的训练流程与-pairwise-loss">Q12 Reward Model 的训练流程与 Pairwise Loss？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> RM 训练的完整流程、loss 公式、数据格式</p>

<p><strong>满分回答：</strong></p>

<p>Reward Model（RM）的目标：给定 prompt $x$ 和回复 $y$，输出一个标量奖励 $r_\theta(x, y)$，反映人类对该回复的偏好程度。</p>

<p><strong>训练流程：</strong></p>

<ol>
  <li><strong>数据收集</strong>：对同一 prompt 生成多个回复（通常 4-9 个），标注员排序选出 best（$y_w$）和 worst（$y_l$）</li>
  <li><strong>模型架构</strong>：通常在 SFT model 最后加一个线性头，将最后一个 hidden state 投射为标量：$r_\theta(x, y) = \text{Linear}(\text{last_hidden})$</li>
  <li><strong>Loss 计算</strong>：Bradley-Terry pairwise loss：</li>
</ol>

\[\mathcal{L}_{\text{RM}} = -\log \sigma\big(r_\theta(x, y_w) - r_\theta(x, y_l)\big)\]

<p>其中 $\sigma$ 是 sigmoid 函数。直觉：让 RM 给 chosen 回复的奖励高于 rejected 回复。</p>

<ol>
  <li><strong>训练细节</strong>：
    <ul>
      <li>lr 通常 5e-6 ~ 1e-5</li>
      <li>batch size 受限于 GPU 内存（每条样本包含 prompt + 两个完整回复）</li>
      <li>通常训练 1 epoch 防止过拟合</li>
    </ul>
  </li>
</ol>

<p><strong>工程技巧：</strong></p>
<ul>
  <li>每个排序对可以构造多个 pair：4 个排序回复可以构造 $\binom{4}{2}=6$ 个 pair</li>
  <li>对多个 pair 使用统一的 loss（而非每个 pair 单独 loss）</li>
  <li>InstructGPT 的 RM 用 175B 参数，但实践中 7B-13B 的 RM 也够用</li>
</ul>

<p>⚠️ 常见坑：RM 过拟合会导致 reward 值爆炸或给出不合理的分数，需要监控 reward 分布。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>RM 的输出范围是否需要归一化？（不需要绝对归一化，但 RLHF 中 KL 约束会间接控制 reward 的影响范围）</li>
  <li>多个 pair 的 loss 如何聚合？（取平均或加权平均，实践中 InstructGPT 对每个排序 pair 等权）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2203.02155">Training language models to follow instructions with human feedback</a></li>
  <li>Bradley &amp; Terry, “Rank Analysis of Incomplete Block Designs”, Biometrika 1952</li>
</ul>

<hr />

<h3 id="q13-bradley-terry-模型的推导与-rm-loss-的关系">Q13 Bradley-Terry 模型的推导与 RM loss 的关系？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 能否从概率模型推导出 RM loss，理解偏好建模的数学基础</p>

<p><strong>满分回答：</strong></p>

<p>Bradley-Terry 模型假设：给定两个选项 A 和 B，人类偏好 A 的概率为：</p>

\[P(A &gt; B) = \frac{e^{s_A}}{e^{s_A} + e^{s_B}} = \sigma(s_A - s_B)\]

<p>其中 $s_A, s_B$ 是选项的”得分”函数。</p>

<p>在 RM 场景中：将回复 $y_w$（chosen）和 $y_l$（rejected）的得分替换为 RM 输出的奖励值：</p>

\[P(y_w &gt; y_l | x) = \sigma\big(r_\theta(x, y_w) - r_\theta(x, y_l)\big)\]

<p><strong>推导过程：</strong></p>

<ol>
  <li>定义偏好概率：$P(y_w &gt; y_l) = \frac{e^{r(x, y_w)}}{e^{r(x, y_w)} + e^{r(x, y_l)}}$</li>
  <li>简化：$= \frac{1}{1 + e^{-(r(x, y_w) - r(x, y_l))}} = \sigma(r(x, y_w) - r(x, y_l))$</li>
  <li>对数似然：$\log P(y_w &gt; y_l) = \log \sigma(r_w - r_l)$</li>
  <li>最大化对数似然 → 最小化负对数似然：</li>
</ol>

\[\mathcal{L} = -\log \sigma(r_\theta(x, y_w) - r_\theta(x, y_l))\]

<p>这就是 RM 的 pairwise loss。</p>

<p><strong>关键理解：</strong></p>
<ul>
  <li>Bradley-Terry 模型假设偏好是<strong>序数</strong>的（ordinal），只需要相对排序，不需要绝对分数</li>
  <li>RM 输出的标量奖励值本身没有绝对含义，只有<strong>差值</strong> $r_w - r_l$ 有意义</li>
  <li>这也是为什么 DPO 可以绕过 RM：直接用 policy 的 logprob 差值替代 reward 差值（见 Q21）</li>
</ul>

<p>⚠️ 常见误解：认为 RM 输出的绝对值有意义。实际上 RM 只需要保证 $r(x, y_w) &gt; r(x, y_l)$ 即可，绝对值可以任意偏移。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>Bradley-Terry 模型的局限性？（只能建模 pairwise 偏好，无法处理更复杂的排序结构；假设偏好是 transitive 的）</li>
  <li>如果标注数据有噪声怎么办？（可以用 margin-based loss 加间隔：$-\log \sigma(r_w - r_l - \delta)$）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li>Bradley &amp; Terry, “Rank Analysis of Incomplete Block Designs”, Biometrika 1952</li>
  <li><a href="https://arxiv.org/abs/2305.18290">Direct Preference Optimization</a></li>
</ul>

<hr />

<h3 id="q14-reward-hacking--reward-overoptimization-问题">Q14 Reward Hacking / Reward Overoptimization 问题？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 理解 reward hacking 的成因、表现形式和缓解策略</p>

<p><strong>满分回答：</strong></p>

<p>Reward Hacking / Overoptimization 指 policy 模型通过”钻 RM 的漏洞”来获得高奖励，而非真正提升回复质量。</p>

<p><strong>表现形式：</strong></p>

<ol>
  <li><strong>回复冗长化</strong>：RM 往往倾向给长回复更高分数，policy 学会生成空洞的长回复</li>
  <li><strong>格式钻营</strong>：学会 RM 偏好的特定格式（如列表、代码块），而非实质内容提升</li>
  <li><strong>情感偏向</strong>：RM 常偏好积极语气，policy 学会无论内容都加正面情绪词</li>
  <li><strong>重复/空洞</strong>：生成看似合理但实际无意义的内容来骗取高分</li>
</ol>

<p><strong>成因分析：</strong></p>
<ul>
  <li>RM 是有限数据的近似，无法完美捕捉人类偏好</li>
  <li>当 policy 和 RM 的分布偏移增大时，RM 在 policy 输出上的预测不准确</li>
  <li>Gao et al. (2023) 量化发现：reward 过优化时，真实人类偏好先升后降（呈倒 U 型曲线）</li>
</ul>

<p><strong>缓解策略：</strong></p>

<table>
  <thead>
    <tr>
      <th>方法</th>
      <th>原理</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>KL 约束</strong></td>
      <td>限制 policy 与 reference model 的 KL 散度，防止输出偏离太远</td>
    </tr>
    <tr>
      <td><strong>迭代 RM 更新</strong></td>
      <td>定期用 policy 输出重新标注数据，更新 RM（LLaMA 2 的做法）</td>
    </tr>
    <tr>
      <td><strong>多 RM 融合</strong></td>
      <td>用多个 RM 取平均，减少单 RM 的偏差</td>
    </tr>
    <tr>
      <td><strong>Ensemble reward</strong></td>
      <td>不同数据源训练的 RM 投票</td>
    </tr>
    <tr>
      <td><strong>Early stopping</strong></td>
      <td>监控 KL 散度，在 reward hacking 开始时停止训练</td>
    </tr>
    <tr>
      <td><strong>增加 RM 数据多样性</strong></td>
      <td>让 RM 在更多类型的回复上学习偏好</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 最关键的是 <strong>KL 约束 + 迭代 RM 更新</strong>，单靠任何一个都不够。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>如何检测 reward hacking？（对比 RM reward 和真实人类偏好评分，如果 reward 持续上升但人类评分下降 = hacking）</li>
  <li>reward hacking 和 overfitting RM 有什么区别？（本质相同：policy 过拟合了 RM 的偏差模式）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li>Gao et al., “Scaling Laws for Reward Model Overoptimization” (ICML 2023)</li>
  <li><a href="https://arxiv.org/abs/2307.05608">Secrets of RLHF in Large Language Models Part I: PPO</a></li>
</ul>

<hr />

<h3 id="q15-rm-训练的数据构造与工程技巧">Q15 RM 训练的数据构造与工程技巧？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> RM 数据的构造方式、标注流程、工程优化</p>

<p><strong>满分回答：</strong></p>

<p><strong>数据构造流程：</strong></p>

<ol>
  <li><strong>Prompt 采集</strong>：从用户对话记录 / 人工设计 / 自动生成中收集 prompt</li>
  <li><strong>回复采样</strong>：用 SFT model 对每个 prompt 生成多个回复（通常 4-9 个），可调整温度/采样策略增加多样性</li>
  <li><strong>人工排序</strong>：标注员对同一 prompt 的多个回复按质量排序</li>
  <li><strong>Pair 构造</strong>：从排序中抽取 (chosen, rejected) pairs</li>
</ol>

<p><strong>关键技巧：</strong></p>

<ol>
  <li><strong>多样性采样</strong>：不同温度、不同 beam、甚至不同模型生成回复，确保 pair 质量差异明显</li>
  <li><strong>Pair 间距</strong>：优先选择排序差距大的 pair（如 best vs worst），而非相邻排名的 pair，让 RM 学习更明显的偏好差异</li>
  <li><strong>InstructGPT 的做法</strong>：4-9 个回复排序，构造所有 $\binom{n}{2}$ pairs，每个 pair 权重等价</li>
  <li><strong>Margin loss</strong>：加入间距 $\delta$ 来增强区分度：$\mathcal{L} = -\log \sigma(r_w - r_l - \delta)$，$\delta$ 可以是排序排名差</li>
  <li><strong>数据分布</strong>：RM 数据应覆盖不同 prompt 类型（对话、推理、代码等），防止 RM 在特定类型上偏科</li>
</ol>

<p><strong>工程优化：</strong></p>
<ul>
  <li>每条训练样本 = prompt + chosen + rejected，显存占用是 SFT 的 ~2 倍</li>
  <li>可将 prompt 部分的 KV cache 缓存，只对两个回复分别前向传播（节省 ~40% 计算量）</li>
  <li>RM 训练 1 epoch 为主，过拟合风险大（RM 容易记住特定 prompt 的模式）</li>
</ul>

<p>⚠️ 常见坑：用同一模型生成的回复做 pair → RM 只学了区分该模型的特定缺陷而非通用偏好 → reward hacking 加剧。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>RM 数据量和 SFT 数据量哪个更重要？（RM 数据量通常需要更多，因为 pairwise 比较比单条回复标注信息更丰富但单条 pair 信息更稀疏）</li>
  <li>能否用 AI 替代人类标注 RM 数据？（可以，这就是 RLAIF 的思路，见 Q43）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2203.02155">Training language models to follow instructions with human feedback</a></li>
  <li><a href="https://arxiv.org/abs/2307.05608">Secrets of RLHF in Large Language Models Part I: PPO</a></li>
</ul>

<hr />

<h2 id="四rlhf--ppo">四、RLHF / PPO</h2>

<h3 id="q16-ppo-的完整-loss-公式及各组件含义">Q16 PPO 的完整 loss 公式及各组件含义？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 能否完整写出 PPO loss 并解释每个组件的作用</p>

<p><strong>满分回答：</strong></p>

<p>PPO 在 RLHF 中的总 loss 包含三个部分：</p>

\[\mathcal{L}_{\text{PPO}} = \mathcal{L}_{\text{policy}} + \beta \cdot \mathcal{L}_{\text{KL}} - c_1 \cdot \mathcal{L}_{\text{value}} + c_2 \cdot \mathcal{L}_{\text{entropy}}\]

<p><strong>1. Policy Loss（Clipped Surrogate）：</strong></p>

\[\mathcal{L}_{\text{policy}} = -\min\Big(\hat{A}_t \cdot \rho_t, \; \hat{A}_t \cdot \text{clip}(\rho_t, 1-\epsilon, 1+\epsilon)\Big)\]

<p>其中：</p>
<ul>
  <li>$\rho_t = \frac{\pi_\theta(a_t|s_t)}{\pi_{\text{ref}}(a_t|s_t)}$  是新旧 policy 的概率比</li>
  <li>$\hat{A}<em>t = r_t + \gamma V(s</em>{t+1}) - V(s_t)$    是优势函数（GAE 估计）</li>
  <li>$\epsilon$ 通常设为 0.2，防止 policy 更新过大</li>
</ul>

<p><strong>2. KL Penalty：</strong></p>

\[\mathcal{L}_{\text{KL}} = \mathbb{D}_{\text{KL}}[\pi_\theta || \pi_{\text{ref}}]\]

<p>作用：防止 policy 偏离 reference model（即 SFT model）太远，避免 reward hacking。</p>

<p><strong>3. Value Loss：</strong></p>

\[\mathcal{L}_{\text{value}} = (V_\phi(s_t) - V_t^{\text{target}})^2\]

<p>其中 $V_t^{\text{target}}$ 是用 GAE 计算的 value target。训练 Critic（Value Model）来估计状态价值，用于计算优势函数 $\hat{A}_t$。</p>

<p><strong>4. Entropy Bonus：</strong></p>

\[\mathcal{L}_{\text{entropy}} = -\sum_a \pi_\theta(a|s) \log \pi_\theta(a|s)\]

<p>鼓励 policy 保持一定的输出多样性，防止模式坍塌。</p>

<p>⚠️ 常见坑：</p>
<ul>
  <li>$\rho_t$ 是逐 token 计算的概率比，不是整句的概率比</li>
  <li>KL penalty 的 $\beta$ 需要仔细调节（见 Q17）</li>
  <li>Critic 的训练需要和价值 target 同步更新，否则优势估计不准</li>
</ul>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>PPO clip 的 $\epsilon$ 为什么设 0.2？（经验值，太大允许过大更新，太小限制太强导致学习慢）</li>
  <li>GAE 如何计算优势函数？\(\hat{A}_t = \sum_{l=0}^{\infty}(\gamma\lambda)^l \delta_{t+l}$，$\delta_t = r_t + \gamma V(s_{t+1}) - V(s_t)\)</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/1707.06347">Proximal Policy Optimization Algorithms</a></li>
  <li><a href="https://arxiv.org/abs/2307.05608">Secrets of RLHF in Large Language Models Part I: PPO</a></li>
</ul>

<hr />

<h3 id="q17-kl-约束在-rlhf-中的作用kl-系数如何调">Q17 KL 约束在 RLHF 中的作用，KL 系数如何调？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 理解 KL 约束的双面性，以及系数调节的工程实践</p>

<p><strong>满分回答：</strong></p>

<p>KL 约束的作用是限制 policy $\pi_\theta$ 与 reference model $\pi_{\text{ref}}$（SFT model）之间的 KL 散度：</p>

\[\mathbb{D}_{\text{KL}}[\pi_\theta || \pi_{\text{ref}}] = \sum_y \pi_\theta(y|x) \log \frac{\pi_\theta(y|x)}{\pi_{\text{ref}}(y|x)}\]

<p><strong>为什么要加 KL 约束？</strong></p>

<ol>
  <li><strong>防 reward hacking</strong>：没有 KL 约束时，policy 可能找到 RM 的漏洞输出高 reward 但低质量的回复</li>
  <li><strong>稳定性</strong>：防止 policy 在单步更新中偏离太大，导致后续训练不稳定</li>
  <li><strong>保持通用能力</strong>：KL 太大意味着 policy 已经远离 SFT model，可能丢失通用语言能力</li>
</ol>

<p><strong>KL 系数 $\beta$ 的调节：</strong></p>

<ul>
  <li>$\beta$ 太小 → reward hacking 加剧，policy 输出越来越怪</li>
  <li>$\beta$ 太大 → policy 几乎不更新，RLHF 无效果</li>
</ul>

<p><strong>调节策略：</strong></p>

<table>
  <thead>
    <tr>
      <th>方法</th>
      <th>描述</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>固定 $\beta$</strong></td>
      <td>设为 0.01-0.1，最简单但不够灵活</td>
    </tr>
    <tr>
      <td><strong>自适应调节</strong></td>
      <td>监控 KL 散度，如果超过目标值则增大 $\beta$，低于目标值则减小（InstructGPT 做法）</td>
    </tr>
    <tr>
      <td><strong>KL target-based</strong></td>
      <td>设定 KL 目标值（如 6 nats），自动调节 $\beta$ 使 KL 稳定在目标附近</td>
    </tr>
    <tr>
      <td><strong>逐步衰减</strong></td>
      <td>初期 $\beta$ 较大保证稳定，后期减小让 policy 更自由探索</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 常见坑：在 RLHF 中，KL 不是按整句计算而是按 token 计算，每个 token 的 logprob 差值累加。</p>

<p>⚠️ 另一个坑：只看 KL 值不够，还需要看 reward 和人类偏好的趋势。KL 稳定但 reward 不涨 = 训练无效。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>KL 约束能否用 cosine similarity 或其他度量替代？（理论上可以，但 KL 在 RL 理论中有最优性证明，其他度量没有理论保证）</li>
  <li>
    <table>
      <tbody>
        <tr>
          <td>KL 散度的方向 $\text{KL}[\pi_\theta</td>
          <td> </td>
          <td>\pi_{\text{ref}}]$ vs $\text{KL}[\pi_{\text{ref}}</td>
          <td> </td>
          <td>\pi_\theta]$ 有什么区别？（前者是 forward KL，鼓励 policy 覆盖 reference 的模式；后者是 reverse KL，鼓励 policy 集中在自身高概率区域）</td>
        </tr>
      </tbody>
    </table>
  </li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2203.02155">Training language models to follow instructions with human feedback</a></li>
  <li><a href="https://arxiv.org/abs/2307.05608">Secrets of RLHF in Large Language Models Part I: PPO</a></li>
</ul>

<hr />

<h3 id="q18-criticvalue-model在-ppo-中的角色">Q18 Critic（Value Model）在 PPO 中的角色？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 理解 Value Model 的功能、训练方式和与 RM 的区别</p>

<p><strong>满分回答：</strong></p>

<p>Critic（Value Model）$V_\phi(s)$ 的核心功能：估计当前状态（已生成的前缀 tokens）的预期累积奖励，用于计算优势函数 $\hat{A}_t$。</p>

\[\hat{A}_t = r_t + \gamma V_\phi(s_{t+1}) - V_\phi(s_t) = \delta_t + (\gamma\lambda)\delta_{t+1} + \cdots\]

<p>（GAE 估计，$\lambda$ 是 GAE 衰减系数）</p>

<p><strong>Value Model vs Reward Model：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>Reward Model</th>
      <th>Value Model</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>输入</strong></td>
      <td>prompt + 完整回复</td>
      <td>prompt + 不完整回复（前缀）</td>
    </tr>
    <tr>
      <td><strong>输出</strong></td>
      <td>整句的质量评分</td>
      <td>当前前缀的预期总奖励</td>
    </tr>
    <tr>
      <td><strong>训练</strong></td>
      <td>用偏好数据训练</td>
      <td>用 RM reward + self-generated TD target 训练</td>
    </tr>
    <tr>
      <td><strong>用途</strong></td>
      <td>提供环境 reward $r_t$</td>
      <td>计算优势函数 $\hat{A}_t$</td>
    </tr>
    <tr>
      <td><strong>初始化</strong></td>
      <td>通常从 SFT model 初始化</td>
      <td>通常从 RM 初始化（因为架构相似）</td>
    </tr>
  </tbody>
</table>

<p><strong>Value Model 的训练：</strong></p>

\[\mathcal{L}_{\text{value}} = \big(V_\phi(s_t) - V_t^{\text{target}}\big)^2\]

<p>$V_t^{\text{target}}$ 是用 GAE 计算的 TD target，融合了 RM 提供的即时 reward 和 Value Model 自身的估计。</p>

<p>⚠️ 常见坑：Value Model 如果训练不好（估计不准），会导致优势函数 $\hat{A}_t$ 偏差大，PPO 的 clip 机制失效，policy 更新不稳定。</p>

<p><strong>工程细节：</strong></p>
<ul>
  <li>Value Model 通常从 RM 初始化（而非 SFT model），因为 RM 已经学过评分，Value Model 只需要学会”前缀评分”</li>
  <li>Value Model 的 loss 通常也加 clip（PPO 中 clip value update），防止 value 估计突然跳变</li>
</ul>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>Value Model 能否和 RM 共用一个模型？（可以但效果通常不如分开训练，因为 Value Model 需要处理前缀，RM 只处理完整回复）</li>
  <li>不用 Value Model 直接用 RM reward 做 policy gradient 可以吗？（可以，但优势估计更粗糙，训练更不稳定，这是 REINFORCE 方式）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/1707.06347">Proximal Policy Optimization Algorithms</a></li>
  <li><a href="https://arxiv.org/abs/2307.05608">Secrets of RLHF in Large Language Models Part I: PPO</a></li>
</ul>

<hr />

<h3 id="q19-rlhf-训练中常见的稳定性问题及解决">Q19 RLHF 训练中常见的稳定性问题及解决？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> RLHF 训练的工程经验，能否识别和解决常见稳定性问题</p>

<p><strong>满分回答：</strong></p>

<p>RLHF/PPO 训练比 SFT 不稳定得多，常见问题及解决方案：</p>

<table>
  <thead>
    <tr>
      <th>问题</th>
      <th>现象</th>
      <th>解决方案</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>Reward Hacking</strong></td>
      <td>RM reward 上升但人类偏好下降</td>
      <td>KL 约束 + 迭代 RM 更新 + 早停</td>
    </tr>
    <tr>
      <td><strong>Critic 过拟合</strong></td>
      <td>Value 估计偏离真实 reward</td>
      <td>Value clip + 定期用 RM 重新标定 value target</td>
    </tr>
    <tr>
      <td><strong>Policy 模式坍塌</strong></td>
      <td>输出极度重复/单一</td>
      <td>Entropy bonus + 增加 KL 系数</td>
    </tr>
    <tr>
      <td><strong>梯度爆炸</strong></td>
      <td>loss 突然 spike，NaN</td>
      <td>gradient clipping（norm clip 到 1.0）+ 降低 lr</td>
    </tr>
    <tr>
      <td><strong>Reward 尺度问题</strong></td>
      <td>RM 输出值范围过大/过小</td>
      <td>reward normalization（running mean/std）</td>
    </tr>
    <tr>
      <td><strong>KL 突然增大</strong></td>
      <td>policy 离 reference 太远</td>
      <td>自适应 KL 系数，超过阈值时增大 $\beta$</td>
    </tr>
    <tr>
      <td><strong>训练震荡</strong></td>
      <td>reward 曲线大幅波动</td>
      <td>减小 batch size / 增加 rollout buffer / 减小 lr</td>
    </tr>
  </tbody>
</table>

<p><strong>关键工程实践（来自 Zheng et al. “Secrets of RLHF”）：</strong></p>

<ol>
  <li><strong>Reward Normalization</strong>：对 RM 输出做 running mean/std 归一化，防止 reward 尺度漂移</li>
  <li><strong>Value Function Pre-training</strong>：先用 RM 数据预训练 Value Model，再做 RL 时更新</li>
  <li><strong>KL Cost as Reward Penalty</strong>：将 KL penalty 加入 reward $r_t’ = r_t - \beta \cdot \text{KL}_t$，而非单独 loss（更稳定）</li>
  <li><strong>Batch 分割</strong>：generation batch 和 training batch 分开，generation 用更大 batch</li>
  <li><strong>Advantage Normalization</strong>：对 $\hat{A}_t$ 做 batch 内归一化，防止优势值极端</li>
</ol>

<p>⚠️ 最常见的新手错误：不做 reward normalization → RM 输出从 -100 到 +100 → PPO 完全失控。</p>

<p>⚠️ 另一个常见错误：policy 和 value model 用同一个网络 → 参数更新互相干扰 → 都学不好。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>PPO 的超参如何选择？（lr=1e-6, clip_ratio=0.2, GAE lambda=0.95, gamma=0.99, 这些是常用起点）</li>
  <li>RLHF 训练需要多少步？（通常 100-500 个 PPO update，太多会 overoptimization）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2307.05608">Secrets of RLHF in Large Language Models Part I: PPO</a></li>
  <li><a href="https://arxiv.org/abs/2204.05862">Training a Helpful and Harmless Assistant with RLHF</a></li>
</ul>

<hr />

<h3 id="q20-rlhf-的优势与局限">Q20 RLHF 的优势与局限？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 对 RLHF 的全面评估，能否辩证看待其优缺点</p>

<p><strong>满分回答：</strong></p>

<p><strong>优势：</strong></p>

<ol>
  <li><strong>偏好对齐直接</strong>：通过 RM 直接优化人类偏好，而非间接通过数据模仿</li>
  <li><strong>超越 SFT 的天花板</strong>：SFT 只能模仿标注数据的风格，RLHF 可以发现数据之外更好的回复</li>
  <li><strong>灵活的奖励信号</strong>：RM 可以编码任意偏好（有用性、安全性、风格等），比硬编码规则灵活</li>
  <li><strong>可迭代改进</strong>：收集新偏好数据 → 更新 RM → 再次 RLHF，持续提升</li>
</ol>

<p><strong>局限：</strong></p>

<ol>
  <li><strong>训练不稳定</strong>：PPO 训练需要 4 个模型（policy, reference, reward, value），工程复杂度高</li>
  <li><strong>RM 质量瓶颈</strong>：RM 是有限数据的近似，reward hacking 是不可避免的风险</li>
  <li><strong>标注成本高</strong>：偏好排序比单条标注更贵（4-9 个回复的排序比写一条回复难）</li>
  <li><strong>不可逆偏移</strong>：RLHF 后的模型可能”过度对齐”，在某些任务上反而不如 SFT model</li>
  <li><strong>Reward 不可解释</strong>：RM 给出的标量分数无法解释”为什么这个回复更好”</li>
</ol>

<p><strong>RLHF vs DPO 的视角（见 Q22）：</strong></p>
<ul>
  <li>RLHF 的根本问题在于 RM 的近似误差 → DPO 直接用偏好数据绕过 RM</li>
  <li>但 RLHF 在 reward shaping（如安全奖励加权）方面更灵活</li>
</ul>

<p>⚠️ 一个关键局限：RLHF 只能对齐偏好排序中的维度，不能对齐标注员没有考虑到的维度。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>RLHF 能否替代 SFT？（不能，见 Q03）</li>
  <li>RLHF 的 reward hacking 能否完全消除？（不能，只能缓解，这是 RM 近似的固有缺陷）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2305.18290">Direct Preference Optimization</a></li>
  <li><a href="社区分析">Reinforcement Learning from Human Feedback is Flawed</a></li>
</ul>

<hr />

<h2 id="五dpo-及变体">五、DPO 及变体</h2>

<h3 id="q21-dpo-loss-的完整推导从-rlhf-到-dpo">Q21 DPO loss 的完整推导（从 RLHF 到 DPO）？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 能否完整推导 DPO loss，理解从 RLHF objective 到 DPO closed-form 的数学链路</p>

<p><strong>满分回答：</strong></p>

<p><strong>Step 1: RLHF 的优化目标</strong></p>

<p>RLHF 的目标是最大化 RM reward，同时 KL 约束防止偏离 reference model：</p>

\[\max_{\pi_\theta} \mathbb{E}_{x,y \sim \pi_\theta}\big[r(x,y) - \beta \log \frac{\pi_\theta(y|x)}{\pi_{\text{ref}}(y|x)}\big]\]

<p><strong>Step 2: 证明最优 policy 的 closed form</strong></p>

<p>对上述目标，最优 policy 有闭式解：</p>

\[\pi^*(y|x) = \frac{1}{Z(x)} \pi_{\text{ref}}(y|x) \exp\big(\frac{1}{\beta} r(x,y)\big)\]

<table>
  <tbody>
    <tr>
      <td>其中 $Z(x) = \sum_y \pi_{\text{ref}}(y</td>
      <td>x) \exp(\frac{1}{\beta} r(x,y))$ 是配分函数（与 $\theta$ 无关）。</td>
    </tr>
  </tbody>
</table>

<p><strong>Step 3: 从最优 policy 反推出 reward</strong></p>

<p>从上式可得：</p>

\[r(x,y) = \beta \log \frac{\pi^*(y|x)}{\pi_{\text{ref}}(y|x)} + \beta \log Z(x)\]

<p><strong>Step 4: 代入 Bradley-Terry 模型</strong></p>

<p>将 reward 代入 BT 模型的偏好概率：</p>

\[P(y_w &gt; y_l|x) = \sigma\big(r(x,y_w) - r(x,y_l)\big) = \sigma\big(\beta \log \frac{\pi^*(y_w|x)}{\pi_{\text{ref}}(y_w|x)} - \beta \log \frac{\pi^*(y_l|x)}{\pi_{\text{ref}}(y_l|x)}\big)\]

<p>注意 $Z(x)$ 在 $r_w - r_l$ 中被消掉！</p>

<p><strong>Step 5: 用 $\pi_\theta$ 替代 $\pi^*$，得到 DPO loss</strong></p>

<p>将最优 policy $\pi^*$ 替换为待训练的 policy $\pi_\theta$：</p>

\[\mathcal{L}_{\text{DPO}} = -\log \sigma\big(\beta \log \frac{\pi_\theta(y_w|x)}{\pi_{\text{ref}}(y_w|x)} - \beta \log \frac{\pi_\theta(y_l|x)}{\pi_{\text{ref}}(y_l|x)}\big)\]

<table>
  <tbody>
    <tr>
      <td><strong>直觉理解：</strong> DPO 直接用 policy 的 logprob ratio 来隐式定义 reward：$r_{\text{implicit}}(x,y) = \beta \log \frac{\pi_\theta(y</td>
      <td>x)}{\pi_{\text{ref}}(y</td>
      <td>x)}$，然后最大化 chosen 的 implicit reward 高于 rejected 的概率。</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 关键：$Z(x)$ 的消掉是 DPO 的核心巧妙之处——不需要计算配分函数（这在离散动作空间中是 NP-hard 的）。</p>

<p>⚠️ 常见坑：推导中假设最优 policy 存在且可用 $\pi_\theta$ 近似。实际上 $\pi_\theta$ 在训练初期远非最优，这是 DPO 理论和实际之间的 gap。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>DPO 的 implicit reward 和真实 RM reward 的关系？（$r_{\text{implicit}} = \beta \log \frac{\pi_\theta}{\pi_{\text{ref}}} + \beta \log Z$，差一个 $Z(x)$ 常数，不影响偏好判断）</li>
  <li>为什么 DPO 不需要 RM？（因为 policy 本身就是隐式 RM：训练后的 $\pi_\theta$ 的 logprob ratio 可以直接作为 reward）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2305.18290">Direct Preference Optimization: Your Language Model is Secretly a Reward Model</a></li>
</ul>

<hr />

<h3 id="q22-dpo-vs-rlhf-的优缺点对比">Q22 DPO vs RLHF 的优缺点对比？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 能否从理论、工程、效果多维度对比两种方法</p>

<p><strong>满分回答：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>RLHF (PPO)</th>
      <th>DPO</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>模型数量</strong></td>
      <td>4 个（policy, ref, RM, critic）</td>
      <td>2 个（policy, ref）</td>
    </tr>
    <tr>
      <td><strong>训练稳定性</strong></td>
      <td>不稳定，需要大量工程调参</td>
      <td>相对稳定，类似 SFT 的训练流程</td>
    </tr>
    <tr>
      <td><strong>数据需求</strong></td>
      <td>偏好数据 + RM 训练数据</td>
      <td>偏好数据（直接用于 loss）</td>
    </tr>
    <tr>
      <td><strong>RM 质量</strong></td>
      <td>需要 RM，RM 偏差直接影响训练</td>
      <td>不需要 RM，隐式 reward</td>
    </tr>
    <tr>
      <td><strong>Reward Shaping</strong></td>
      <td>可以灵活调整 RM reward（加权、组合）</td>
      <td>无法显式调整 reward</td>
    </tr>
    <tr>
      <td><strong>在线 vs 离线</strong></td>
      <td>在线（policy 生成的数据即时训练）</td>
      <td>离线（用预收集的偏好数据训练）</td>
    </tr>
    <tr>
      <td><strong>计算成本</strong></td>
      <td>高（generation + 4 model forward/backward）</td>
      <td>低（2 model forward/backward）</td>
    </tr>
    <tr>
      <td><strong>理论保证</strong></td>
      <td>在线 RL 有渐近最优性保证</td>
      <td>离线优化，理论保证依赖数据分布</td>
    </tr>
    <tr>
      <td><strong>迭代改进</strong></td>
      <td>可以迭代（policy → RM → policy）</td>
      <td>可以迭代（DPO → 新偏好数据 → DPO）</td>
    </tr>
    <tr>
      <td><strong>效果上限</strong></td>
      <td>理论上更高（在线探索可能发现更好策略）</td>
      <td>受限于偏好数据覆盖范围</td>
    </tr>
  </tbody>
</table>

<p><strong>核心差异总结：</strong></p>
<ul>
  <li>DPO <strong>简单、稳定、低成本</strong>，适合大多数场景</li>
  <li>RLHF <strong>灵活、在线探索</strong>，适合需要 reward shaping 或动态 reward 的场景</li>
  <li>实践中 DPO 已成为主流选择（LLaMA 3 用 DPO 替代了 PPO）</li>
</ul>

<p>⚠️ DPO 的局限：离线方法，无法发现偏好数据之外的更好回复。如果偏好数据不够覆盖，DPO 的效果可能不如 RLHF 的在线探索。</p>

<p>⚠️ 另一个坑：DPO 的 reference model 退化问题（见 Q24）。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>能否先 DPO 后 RLHF？（可以，DPO 先快速对齐 → RLHF 再精细调优）</li>
  <li>DPO 的偏好数据需要多少？（通常 10K-50K pairs 即可，比 RLHF 的标注量少很多）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2305.18290">Direct Preference Optimization</a></li>
  <li><a href="https://arxiv.org/abs/2407.21783">The LLaMA 3 Herd of Models</a></li>
</ul>

<hr />

<h3 id="q23-ipo--kto--simpo--orpo-各变体特点">Q23 IPO / KTO / SimPO / ORPO 各变体特点？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 对 DPO 变体生态的了解，能否区分各变体的创新点</p>

<p><strong>满分回答：</strong></p>

<p>DPO 之后出现了多个变体，各自解决 DPO 的不同缺陷：</p>

<table>
  <thead>
    <tr>
      <th>变体</th>
      <th>核心创新</th>
      <th>解决的问题</th>
      <th>Loss 公式</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>IPO</strong></td>
      <td>用 squared loss 替代 log-sigmoid</td>
      <td>DPO 在偏好数据噪声大时过拟合</td>
      <td>$\mathcal{L} = (r_w - r_l - \frac{1}{2})^2$</td>
    </tr>
    <tr>
      <td><strong>KTO</strong></td>
      <td>只需要 binary signal（好/坏），不需要 pair</td>
      <td>偏好 pair 数据获取成本高</td>
      <td>基于 Prospect Theory 的 loss</td>
    </tr>
    <tr>
      <td><strong>SimPO</strong></td>
      <td>去掉 reference model</td>
      <td>ref model 的显存和计算开销</td>
      <td>用 policy 自身的 length-normalized logprob 作为 implicit reward</td>
    </tr>
    <tr>
      <td><strong>ORPO</strong></td>
      <td>SFT + preference 合为一体</td>
      <td>SFT 和 DPO 分两步训练效率低</td>
      <td>在 SFT loss 中加入 odds ratio penalty</td>
    </tr>
  </tbody>
</table>

<p><strong>详细说明：</strong></p>

<p><strong>IPO（Identity Preference Optimization）：</strong></p>
<ul>
  <li>问题：DPO loss $-\log\sigma(r_w - r_l)$ 当 $r_w - r_l$ 趋向无穷时梯度趋近 0，导致 chosen 完美后不再学习</li>
  <li>
    <table>
      <tbody>
        <tr>
          <td>解决：用 squared loss $\mathcal{L}_{\text{IPO}} = \left(\log\frac{\pi(y_w</td>
          <td>x)}{\pi_{\text{ref}}(y_w</td>
          <td>x)} - \log\frac{\pi(y_l</td>
          <td>x)}{\pi_{\text{ref}}(y_l</td>
          <td>x)} - \frac{1}{2\beta}\right)^2$</td>
        </tr>
      </tbody>
    </table>
  </li>
  <li>效果：对噪声数据更鲁棒</li>
</ul>

<p><strong>KTO（Kahneman-Tversky Optimization）：</strong></p>
<ul>
  <li>问题：DPO 需要 (chosen, rejected) pair，实际中更容易获得 binary label（好/坏）</li>
  <li>解决：基于 Kahneman-Tversky Prospect Theory，用单条数据的 good/bad signal</li>
  <li>Loss：$\mathcal{L}<em>{\text{KTO}} = \begin{cases} -\log\sigma(\beta \cdot (r</em>{\text{implicit}} - z_{\text{ref}})) &amp; \text{if good} \ -\log\sigma(-\beta \cdot (r_{\text{implicit}} - z_{\text{ref}})) &amp; \text{if bad} \end{cases}$</li>
  <li>$z_{\text{ref}}$ 是 reference model 的 implicit reward baseline</li>
</ul>

<p><strong>SimPO：</strong></p>
<ul>
  <li>问题：DPO 需要维护 reference model，增加显存开销</li>
  <li>解决：用 policy 自身的 length-normalized logprob 替代 ref logprob</li>
  <li>
    <table>
      <tbody>
        <tr>
          <td>Implicit reward：$\tilde{r}(x,y) = \frac{\beta}{</td>
          <td>y</td>
          <td>} \sum_t \log \pi_\theta(y_t</td>
          <td>x, y_{&lt;t})$</td>
        </tr>
      </tbody>
    </table>
  </li>
  <li>优势：无需 ref model，但可能导致 policy 偏移（缺乏锚定）</li>
</ul>

<p><strong>ORPO：</strong></p>
<ul>
  <li>问题：SFT 和 DPO 是两个独立阶段</li>
  <li>解决：在 SFT loss 中直接加入 odds ratio penalty</li>
  <li>$\mathcal{L}<em>{\text{ORPO}} = \mathcal{L}</em>{\text{SFT}} + \lambda \cdot \mathcal{L}_{\text{OR}}$</li>
  <li>$\mathcal{L}_{\text{OR}} = -\log\sigma\left(\log\frac{\text{OR}(y_w)}{\text{OR}(y_l)}\right)$，OR = odds ratio</li>
</ul>

<p>⚠️ 选择建议：一般场景用 DPO/SimPO 即可；噪声数据多用 IPO；只有 binary label 用 KTO；想一步到位用 ORPO。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>为什么 IPO 对噪声更鲁棒？（Squared loss 对 $r_w - r_l$ 很大的情况梯度不为零，不会”忽略”噪声数据）</li>
  <li>SimPO 去掉 reference model 后如何防止偏移？（用 length-normalized logprob 自身作为锚定，但稳定性不如有 ref model 的 DPO）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2310.12048">IPO: A General Theoretical Paradigm for Preference Optimization</a></li>
  <li><a href="https://arxiv.org/abs/2402.01306">KTO: Model Alignment as Prospect Theory Optimization</a></li>
  <li><a href="https://arxiv.org/abs/2405.14734">SimPO: Simple Preference Optimization without Reference Models</a></li>
  <li><a href="https://arxiv.org/abs/2403.07691">ORPO: Monolithic Preference Optimization without Reference Models</a></li>
</ul>

<hr />

<h3 id="q24-dpo-的-reference-model-作用及常见问题">Q24 DPO 的 reference model 作用及常见问题？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 理解 reference model 在 DPO 中的角色，以及它带来的工程和理论问题</p>

<p><strong>满分回答：</strong></p>

<p><strong>Reference model 在 DPO 中的作用：</strong></p>

<p>Reference model $\pi_{\text{ref}}$ 是 DPO loss 中的锚点，用于计算 implicit reward：</p>

\[r_{\text{implicit}}(x,y) = \beta \log \frac{\pi_\theta(y|x)}{\pi_{\text{ref}}(y|x)}\]

<p>作用类比 RLHF 中的 KL 约束：</p>
<ul>
  <li>RLHF 用 KL penalty 限制 policy 不偏离 ref</li>
  <li>DPO 用 ref 的 logprob 作为 reward 的基准线——policy 的 reward 不是绝对 logprob，而是相对于 ref 的提升</li>
</ul>

<p><strong>没有 reference model 的后果：</strong></p>
<ul>
  <li>Policy 可以通过降低 rejected 的概率来提升 chosen/rejected 的差值，而不是真正提升 chosen 的质量</li>
  <li>模型可能学习”降低一切输出概率”的策略，导致整体退化</li>
</ul>

<p><strong>常见问题：</strong></p>

<ol>
  <li>
    <p><strong>显存开销</strong>：训练时需要同时存储 policy 和 ref model 的权重（双倍显存）。SimPO 的动机正是解决这个问题。</p>
  </li>
  <li>
    <p><strong>Reference model 退化</strong>：如果 ref model 和 policy 初始完全相同，训练初期 logprob ratio 接近 0，梯度信号弱。随着训练进行 ratio 变大，但 ref model 始终冻结，可能过度惩罚偏离。</p>
  </li>
  <li>
    <table>
      <tbody>
        <tr>
          <td><strong>长度偏差</strong>：ref model 的 logprob 是逐 token 累加的，长回复的 $</td>
          <td>\log \pi_{\text{ref}}</td>
          <td>$ 更大 → DPO 偏好短回复（因为长回复的 ratio 更难增大）。SimPO 的 length normalization 部分解决此问题。</td>
        </tr>
      </tbody>
    </table>
  </li>
  <li><strong>数据分布偏移</strong>：DPO 的理论推导假设偏好数据来自 ref model 的分布。如果数据来自不同模型（如 GPT-4），ref model 的 logprob 可能不匹配数据分布。</li>
</ol>

<p>⚠️ 常见坑：用 base model 做 reference 而非 SFT model → ref model 和偏好数据分布严重不匹配 → DPO 效果差。</p>

<p>⚠️ 另一个坑：DPO 训练时 ref model 需要 forward pass 但不需要 backward，可以冻结参数只算 forward，但显存仍然占用。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>能否用 LoRA 来降低 ref model 的显存？（可以，但 ref model 必须严格冻结，不能参与梯度更新）</li>
  <li>ref model 能否定期更新？（理论上可以但会破坏 DPO 的理论推导前提，实践中不推荐）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2305.18290">Direct Preference Optimization</a></li>
  <li><a href="https://arxiv.org/abs/2405.14734">SimPO: Simple Preference Optimization without Reference Models</a></li>
</ul>

<hr />

<h2 id="六grpo--推理模型">六、GRPO / 推理模型</h2>

<h3 id="q25-grpo-的核心公式与相比-ppo-的改进">Q25 GRPO 的核心公式与相比 PPO 的改进？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 理解 GRPO 的原理、公式、以及它如何简化 PPO 的工程复杂度</p>

<p><strong>满分回答：</strong></p>

<p>GRPO（Group Relative Policy Optimization）来自 DeepSeekMath 论文，核心思想：<strong>用同一 prompt 下多个采样回复的 group reward 来替代绝对 reward + critic model</strong>。</p>

<p><strong>GRPO vs PPO 的关键区别：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>PPO</th>
      <th>GRPO</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>Reward</strong></td>
      <td>RM 给绝对标量 reward</td>
      <td>Group 内相对 reward（组内排名归一化）</td>
    </tr>
    <tr>
      <td><strong>Critic Model</strong></td>
      <td>需要独立的 Value Model</td>
      <td>不需要（用 group mean 替代 baseline）</td>
    </tr>
    <tr>
      <td><strong>模型数量</strong></td>
      <td>4（policy, ref, RM, critic）</td>
      <td>2（policy, ref）</td>
    </tr>
    <tr>
      <td><strong>KL 约束</strong></td>
      <td>KL penalty $\beta \cdot \mathbb{D}_{\text{KL}}$</td>
      <td>KL penalty 同样保留</td>
    </tr>
  </tbody>
</table>

<p><strong>GRPO 核心公式：</strong></p>

<ol>
  <li>
    <p>对同一 prompt $x$，采样 $G$ 个回复 ${y_1, …, y_G}$</p>
  </li>
  <li>
    <p>RM 对每个回复评分，得到 ${r_1, …, r_G}$</p>
  </li>
  <li>
    <p><strong>Group 归一化</strong>（替代 Critic 的 baseline 功能）：</p>
  </li>
</ol>

\[\tilde{r}_i = \frac{r_i - \mu(r)}{\sigma(r)}\]

<p>其中 $\mu(r) = \frac{1}{G}\sum_{j=1}^G r_j$，$\sigma(r)$ 是组内标准差。</p>

<ol>
  <li><strong>GRPO Policy Loss</strong>（类似 PPO clip）：</li>
</ol>

\[\mathcal{L}_{\text{GRPO}} = -\frac{1}{G}\sum_{i=1}^G \min\Big(\tilde{r}_i \cdot \rho_i, \; \tilde{r}_i \cdot \text{clip}(\rho_i, 1-\epsilon, 1+\epsilon)\Big) + \beta \cdot \mathbb{D}_{\text{KL}}[\pi_\theta || \pi_{\text{ref}}]\]

<p><strong>关键改进点：</strong></p>
<ul>
  <li><strong>去掉 Critic Model</strong>：用 group 内的 mean reward 作为 baseline，省掉一个模型</li>
  <li><strong>相对 Reward</strong>：不需要绝对 reward 值，只需组内相对排名，这使 reward hacking 的空间更小</li>
  <li><strong>工程更简单</strong>：只需 2 个模型，和 DPO 类似的简洁度，但保留了 RL 的在线探索能力</li>
</ul>

<p>⚠️ GRPO 的采样数 $G$ 很重要：太小（如 $G=2$）则归一化不稳定；太大（如 $G=64$）则计算成本高。DeepSeek 实践中用 $G=16$。</p>

<p>⚠️ 常见坑：GRPO 的 group 归一化只对同一 prompt 的回复做，不同 prompt 的 reward 值不可比。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>GRPO 的归一化方式会不会损失信息？（会损失绝对 reward 信息，但偏好对齐只需要相对排序，所以实际影响不大）</li>
  <li>DeepSeek R1 用 GRPO 还是 PPO？（用 GRPO + 规则奖励，见 Q26）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2402.03300">DeepSeekMath: Pushing the Limits of Mathematical Reasoning (GRPO)</a></li>
  <li><a href="https://arxiv.org/abs/2501.12948">DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via RL</a></li>
</ul>

<hr />

<h3 id="q26-deepseek-r1-的后训练流程">Q26 DeepSeek R1 的后训练流程？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 对 DeepSeek R1 创新训练流程的深入了解</p>

<p><strong>满分回答：</strong></p>

<p>DeepSeek R1 的后训练流程是其最大创新点，实现了<strong>纯 RL 涌现推理能力</strong>：</p>

<p><strong>R1-Zero（纯 RL 版）：</strong></p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Base Model → GRPO (with rule-based reward only) → R1-Zero
</code></pre></div></div>

<ul>
  <li>不做任何 SFT，直接从 base model 用 GRPO + 规则奖励训练</li>
  <li>模型自发涌现了 CoT（Chain-of-Thought）推理行为</li>
  <li>问题：输出格式混乱、语言混合、可读性差</li>
</ul>

<p><strong>R1（完整版）：</strong></p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Base Model → Cold-start SFT → GRPO (rule + RM reward) → Rejection Sampling SFT → SFT (全场景数据) → DPO → R1
</code></pre></div></div>

<p><strong>各阶段详解：</strong></p>

<ol>
  <li><strong>Cold-start SFT</strong>：用少量（数千条）长 CoT 高质量数据做 SFT，解决 R1-Zero 的格式问题</li>
  <li><strong>GRPO RL 阶段</strong>：
    <ul>
      <li>Rule-based Reward：对数学用正确率验证，对代码用编译/执行结果</li>
      <li>Language consistency reward：惩罚中英混用的 CoT</li>
      <li>保留 GRPO 的 group relative reward 机制</li>
    </ul>
  </li>
  <li><strong>Rejection Sampling SFT</strong>：用 RL 后的模型生成大量 CoT 数据，筛选高质量数据（正确 + 可读），重新做 SFT</li>
  <li><strong>全场景 SFT</strong>：加入非推理任务数据（对话、翻译、写作等），恢复通用能力</li>
  <li><strong>DPO</strong>：最终用 DPO 做偏好对齐</li>
</ol>

<p>⚠️ R1 的核心发现：推理能力可以从纯 RL 中涌现，不需要 SFT 教推理格式。但 SFT 对输出可读性至关重要。</p>

<p>⚠️ 规则奖励的设计是 R1 成功的关键——数学和代码的 ground truth 可以精确验证，避免了 RM 的偏差问题。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>R1-Zero 为什么不需要 SFT 就能涌现 CoT？（base model 已具备推理潜质，RL 的 reward signal 激励了更详细的思考过程）</li>
  <li>规则奖励适用于哪些任务？（主要适用于可验证任务：数学、代码、逻辑推理；不适用于开放式生成任务）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2501.12948">DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via RL</a></li>
</ul>

<hr />

<h3 id="q27-process-reward-model-prm-vs-outcome-reward">Q27 Process Reward Model (PRM) vs Outcome Reward？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 理解 PRM 和 ORM 的区别、各自的优缺点</p>

<p><strong>满分回答：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>ORM (Outcome Reward Model)</th>
      <th>PRM (Process Reward Model)</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>评分粒度</strong></td>
      <td>整个回复一个分数</td>
      <td>每步推理一个分数</td>
    </tr>
    <tr>
      <td><strong>输入</strong></td>
      <td>prompt + 完整回复</td>
      <td>prompt + 回复 + step boundary</td>
    </tr>
    <tr>
      <td><strong>优点</strong></td>
      <td>训练简单，标注简单</td>
      <td>精细反馈，可定位错误步骤</td>
    </tr>
    <tr>
      <td><strong>缺点</strong></td>
      <td>无法区分正确过程+错误结果 vs 错误过程+偶然正确</td>
      <td>标注成本高（需逐步标注），训练复杂</td>
    </tr>
    <tr>
      <td><strong>适用场景</strong></td>
      <td>简单任务、短回复</td>
      <td>数学推理、长 CoT</td>
    </tr>
  </tbody>
</table>

<p><strong>PRM 的关键优势：</strong></p>

<ol>
  <li><strong>避免”侥幸正确”</strong>：ORM 给了一个正确答案但中间步骤有错的回复高分；PRM 可以惩罚错误步骤</li>
  <li><strong>更细粒度的 credit assignment</strong>：在 RL 训练中，PRM 可以逐步给 reward，引导模型学习正确的推理过程</li>
  <li><strong>更好的搜索引导</strong>：在 inference 时用 PRM 做 step-level beam search（Best-of-N per step）</li>
</ol>

<p><strong>PRM 的训练方式（Let’s Verify Step by Step）：</strong></p>

<ul>
  <li>样本：对每个推理步骤自动标注正确性（通过与 ground truth 对比）</li>
  <li>Loss：对每个 step token 预测 step-level reward</li>
</ul>

<p><strong>Math-Shepherd 方法：</strong></p>
<ul>
  <li>自动化标注：用 completion 的正确性反推每个步骤的贡献</li>
  <li>不需要人工逐步标注</li>
</ul>

<p>⚠️ 常见坑：PRM 的 step boundary 需要精确标注，否则评分粒度错误。实际中常用特殊 token（如 <code class="language-plaintext highlighter-rouge">\n\n</code>）标记 step boundary。</p>

<p>⚠️ 另一个坑：PRM 在 RL 中的使用比 ORM 更复杂——需要在每个 step 位置计算 reward，而非只在 EOS。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>PRM 能否替代 ORM？（在推理任务上可以，但在非推理任务上 ORM 更简单有效）</li>
  <li>PRM 如何做 inference？（step-level beam search：每生成一个 step，用 PRM 评分，保留 top-K 步骤继续生成）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2305.20050">Let’s Verify Step by Step (PRM)</a></li>
  <li><a href="https://arxiv.org/abs/2501.12948">DeepSeek-R1</a></li>
</ul>

<hr />

<h3 id="q28-rule-based-reward-在推理训练中的应用">Q28 Rule-based Reward 在推理训练中的应用？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 理解规则奖励的设计原理、适用范围和与 RM 的关系</p>

<p><strong>满分回答：</strong></p>

<p>Rule-based Reward 是 DeepSeek R1 的关键创新：用确定性规则替代 RM 来提供 reward signal。</p>

<p><strong>核心思想：</strong> 对于可验证的任务（数学、代码），ground truth 可以精确判定答案是否正确，不需要 RM 的近似评分。</p>

<p><strong>常见规则奖励类型：</strong></p>

<table>
  <thead>
    <tr>
      <th>类型</th>
      <th>规则</th>
      <th>适用场景</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>数学正确性</strong></td>
      <td>最终答案是否等于 ground truth</td>
      <td>数学推理</td>
    </tr>
    <tr>
      <td><strong>代码执行</strong></td>
      <td>通过测试用例数 / 执行成功与否</td>
      <td>代码生成</td>
    </tr>
    <tr>
      <td><strong>格式合规</strong></td>
      <td>是否符合 CoT 格式、是否使用了 <code class="language-plaintext highlighter-rouge">&lt;think&gt;</code> 标签</td>
      <td>推理格式</td>
    </tr>
    <tr>
      <td><strong>语言一致性</strong></td>
      <td>CoT 和最终答案的语言是否一致</td>
      <td>多语言推理</td>
    </tr>
    <tr>
      <td><strong>长度合理性</strong></td>
      <td>CoT 长度是否在合理范围内</td>
      <td>防止冗长</td>
    </tr>
  </tbody>
</table>

<p><strong>Rule-based Reward 的优势：</strong></p>

<ol>
  <li><strong>零偏差</strong>：规则是确定性函数，没有 RM 的近似误差 → 不会有 reward hacking</li>
  <li><strong>零标注成本</strong>：不需要人类标注偏好数据</li>
  <li><strong>可精确验证</strong>：数学和代码的正确性有 ground truth</li>
</ol>

<p><strong>局限性：</strong></p>

<ol>
  <li><strong>适用范围有限</strong>：只适用于有 ground truth 的任务，不适用于开放式生成（如写作、对话）</li>
  <li><strong>缺乏偏好维度</strong>：不能编码”风格偏好”、”有用性”等主观维度</li>
  <li><strong>奖励信号稀疏</strong>：只有最终结果的对/错，没有中间步骤的反馈（除非结合 PRM）</li>
</ol>

<p><strong>DeepSeek R1 的实践：</strong></p>

<ul>
  <li>纯 GRPO + 规则奖励 → R1-Zero 涌现了推理能力</li>
  <li>后期加入少量 RM reward 处理非推理任务</li>
</ul>

<p>⚠️ 常见坑：规则奖励 + GRPO 时，如果 group 内所有回复都不正确，所有 $\tilde{r}_i$ 都为负数 → policy 可能学到了”什么都不做”。需要加入 baseline 或只对正确回复做训练。</p>

<p>⚠️ 另一个坑：格式奖励和正确性奖励的权重需要平衡——如果格式奖励太大，模型可能只学格式不学推理。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>规则奖励能否用于非推理任务？（可以设计简单规则如”是否包含拒绝”、”是否礼貌”，但维度太少不够全面）</li>
  <li>规则奖励和 RM reward 如何组合？（可以加权组合：$r = \alpha \cdot r_{\text{rule}} + (1-\alpha) \cdot r_{\text{RM}}$，DeepSeek R1 在 RL 后期引入了 RM）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2501.12948">DeepSeek-R1</a></li>
  <li><a href="https://arxiv.org/abs/2402.03300">DeepSeekMath (GRPO)</a></li>
</ul>

<hr />

<h2 id="七peft--lora--qlora">七、PEFT / LoRA / QLoRA</h2>

<h3 id="q29-lora-的原理秩选择与合并策略">Q29 LoRA 的原理、秩选择与合并策略？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> LoRA 的数学原理、秩 r 的选择经验、merge 策略的工程细节</p>

<p><strong>满分回答：</strong></p>

<p><strong>LoRA 原理：</strong></p>

<p>对预训练权重矩阵 $W_0 \in \mathbb{R}^{d \times k}$，冻结 $W_0$，只训练低秩增量：</p>

\[W = W_0 + \Delta W = W_0 + BA\]

<p>其中 $B \in \mathbb{R}^{d \times r}$，$A \in \mathbb{R}^{r \times k}$，$r \ll \min(d, k)$。</p>

<p><strong>初始化策略：</strong></p>
<ul>
  <li>$A$ 用 Gaussian 初始化</li>
  <li>$B$ 初始化为 0 → $\Delta W = BA = 0$ 在训练开始时，保证初始输出不变</li>
  <li>这确保了 LoRA 加入后不改变基座模型的初始行为</li>
</ul>

<p><strong>秩选择经验：</strong></p>

<table>
  <thead>
    <tr>
      <th>任务</th>
      <th>推荐 $r$</th>
      <th>说明</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>简单 SFT（格式对齐）</td>
      <td>4-8</td>
      <td>需要学习的模式简单</td>
    </tr>
    <tr>
      <td>复杂 SFT（多任务）</td>
      <td>16-64</td>
      <td>需要更多容量</td>
    </tr>
    <tr>
      <td>RLHF/DPO</td>
      <td>8-16</td>
      <td>增量调整，不需要太大</td>
    </tr>
    <tr>
      <td>代码/推理专项</td>
      <td>32-64</td>
      <td>需要更大容量学习复杂模式</td>
    </tr>
  </tbody>
</table>

<p>⚠️ $r$ 不是越大越好——过大的 $r$ 接近全参数微调但显存没省多少。</p>

<p><strong>合并策略：</strong></p>

<p>训练后可以将 $\Delta W = BA$ 合并到 $W_0$：</p>

\[W_{\text{merged}} = W_0 + \frac{\alpha}{r} \cdot BA\]

<p>其中 $\alpha$ 是 LoRA 的 scaling factor（默认 $\alpha = r$，即 $\alpha/r = 1$）。</p>

<p>合并后推理无额外开销——权重矩阵和原始一样大。</p>

<p><strong>Unmerge（用于多任务切换）：</strong></p>

\[W = W_{\text{merged}} - \frac{\alpha}{r} BA + \frac{\alpha}{r} B'A'\]

<p>卸载当前 LoRA，加载另一个 LoRA adapter。</p>

<p>⚠️ 常见坑：合并时 $\alpha/r$ 的比例容易被忽略，导致 merge 后效果变差。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>LoRA 应该加在哪些层？（Q/K/V/O + gate/up/down projection，不建议加 embedding 和 lm_head）</li>
  <li>LoRA 的 $\alpha$ 参数如何设置？（通常设为 $r$ 或 $2r$，$\alpha$ 控制 LoRA 更新的整体强度）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2106.09685">LoRA: Low-Rank Adaptation of Large Language Models</a></li>
</ul>

<hr />

<h3 id="q30-qlora-的创新点与显存优化">Q30 QLoRA 的创新点与显存优化？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 理解 QLoRA 的三重优化：4-bit quantization + double quantization + paged optimizer</p>

<p><strong>满分回答：</strong></p>

<p>QLoRA 的核心创新：在不损失微调效果的前提下，将 65B 模型的微调显存需求从 ~780GB 降至 ~48GB（单 GPU 可微调）。</p>

<p><strong>三重优化：</strong></p>

<ol>
  <li><strong>4-bit NormalFloat (NF4) Quantization</strong>
    <ul>
      <li>基于正态分布信息论最优的 4-bit 数据类型</li>
      <li>对预训练权重做 4-bit 量化存储，但计算时动态反量化到 bf16</li>
      <li>数学原理：假设权重服从正态分布，NF4 的量化分位点使信息损失最小</li>
    </ul>
  </li>
  <li><strong>Double Quantization</strong>
    <ul>
      <li>对 4-bit 量化的常量（scaling factor + zero point）再做一次量化</li>
      <li>这些常量本身占显存（每个 block 32 个参数需要 1 个 fp32 scaling + 1 个 fp32 zero point）</li>
      <li>Double quantization 将它们量化为 fp8 → 进一步节省 ~0.37 bit/param</li>
    </ul>
  </li>
  <li><strong>Paged Optimizers</strong>
    <ul>
      <li>利用 NVIDIA unified memory 特性</li>
      <li>优化器状态（Adam 的 m 和 v）在 GPU 内存不足时自动 page 到 CPU 内存</li>
      <li>避免优化器状态导致的 OOM</li>
    </ul>
  </li>
</ol>

<p><strong>QLoRA 的计算流程：</strong></p>

\[\text{存储: } W_0 \text{ (NF4)} \xrightarrow{\text{dequantize}} W_0' \text{ (bf16)} \xrightarrow{\text{forward}} h = W_0'x + BAx\]

<p>⚠️ 关键理解：QLoRA 的 forward 和 backward 计算仍然是 <strong>bf16 精度</strong>，只是存储是 4-bit。这是为什么精度损失几乎为零。</p>

<p>⚠️ 常见坑：QLoRA 微调后的模型需要将 LoRA merge 后再部署，merge 后的权重恢复为 fp16/bf16（不再是 4-bit）。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>QLoRA 和 LoRA 的效果差距有多大？（QLoRA 论文声称差距 &lt;1%，在多数 benchmark 上基本持平）</li>
  <li>4-bit 量化是否会累积误差？（NF4 是信息论最优量化，对正态分布权重的误差最小；但非常小或非常大的权重可能损失精度）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2305.14314">QLoRA: Efficient Finetuning of Quantized LLMs</a></li>
</ul>

<hr />

<h3 id="q31-lora-vs-全参数微调的显存与性能对比">Q31 LoRA vs 全参数微调的显存与性能对比？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 从显存占用和最终效果两个维度量化对比</p>

<p><strong>满分回答：</strong></p>

<p><strong>显存对比（以 7B 模型为例，bf16 训练）：</strong></p>

<table>
  <thead>
    <tr>
      <th>项目</th>
      <th>全参数微调</th>
      <th>LoRA (r=16)</th>
      <th>QLoRA (r=16)</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>模型权重</td>
      <td>14 GB</td>
      <td>14 GB</td>
      <td>~3.5 GB (NF4)</td>
    </tr>
    <tr>
      <td>LoRA 参数</td>
      <td>0</td>
      <td>~0.5 GB</td>
      <td>~0.5 GB</td>
    </tr>
    <tr>
      <td>优化器状态 (Adam)</td>
      <td>28 GB (2×权重)</td>
      <td>~1 GB (仅 LoRA)</td>
      <td>~1 GB</td>
    </tr>
    <tr>
      <td>梯度</td>
      <td>14 GB</td>
      <td>~0.5 GB (仅 LoRA)</td>
      <td>~0.5 GB</td>
    </tr>
    <tr>
      <td>激活值</td>
      <td>~4 GB</td>
      <td>~4 GB</td>
      <td>~4 GB</td>
    </tr>
    <tr>
      <td><strong>总计</strong></td>
      <td>~60 GB</td>
      <td>~20 GB</td>
      <td>~9.5 GB</td>
    </tr>
  </tbody>
</table>

<p><strong>性能对比：</strong></p>

<table>
  <thead>
    <tr>
      <th>场景</th>
      <th>全参数 &gt; LoRA</th>
      <th>LoRA ≈ 全参数</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>大幅度行为改变（base → chat）</td>
      <td>✅ 全参数更好</td>
      <td> </td>
    </tr>
    <tr>
      <td>小幅度偏好调整（DPO/RLHF）</td>
      <td> </td>
      <td>✅ LoRA 足够</td>
    </tr>
    <tr>
      <td>多任务微调</td>
      <td> </td>
      <td>✅ LoRA + 多 adapter</td>
    </tr>
    <tr>
      <td>数据量很大 (&gt;100K)</td>
      <td>✅ 全参数可能更好</td>
      <td> </td>
    </tr>
    <tr>
      <td>数据量很小 (&lt;10K)</td>
      <td> </td>
      <td>✅ LoRA 防过拟合</td>
    </tr>
  </tbody>
</table>

<p><strong>关键发现：</strong></p>
<ul>
  <li>LoRA 在 SFT 和 DPO 场景下，效果差距通常 &lt;2%（用合适的 $r$）</li>
  <li>全参数微调在需要大幅改变模型行为时更优</li>
  <li>QLoRA 在效果上几乎等同 LoRA，但显存节省巨大</li>
</ul>

<p>⚠️ 常见坑：LoRA 的 $r$ 设太小（如 $r=4$）用于复杂任务 → 容量不足，效果明显差于全参数。</p>

<p>⚠️ 另一个坑：LoRA 微调后如果需要继续微调（如 SFT → DPO），建议 merge 后再做下一阶段，否则两个 LoRA 的梯度交互可能不稳定。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>7B 模型全参数微调需要什么硬件？（至少 1×A100 80GB 或 2×A100 40GB）</li>
  <li>LoRA 的 gradient checkpointing 如何配合？（全参数需要 gradient checkpointing 省 ~60% 激活显存；LoRA 的梯度本身就很小，checkpointing 收益有限）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2106.09685">LoRA: Low-Rank Adaptation</a></li>
  <li>
    <h2 id="qlora-efficient-finetuning-of-quantized-llms"><a href="https://arxiv.org/abs/2305.14314">QLoRA: Efficient Finetuning of Quantized LLMs</a></h2>
  </li>
</ul>

<hr />

<h2 id="八数据工程">八、数据工程</h2>

<h3 id="q32-self-instruct--magpie-数据合成方法">Q32 Self-Instruct / Magpie 数据合成方法？</h3>

<p><strong>难度：</strong> 基础
<strong>考察点：</strong> 了解两种主流数据合成方法的原理和差异</p>

<p><strong>满分回答：</strong></p>

<p><strong>Self-Instruct：</strong></p>

<ol>
  <li>从 175 个种子任务出发</li>
  <li>用 LLM 生成新指令（prompt: “生成一个新任务指令”）</li>
  <li>用 LLM 为每个指令生成回复</li>
  <li>规则过滤：去重、去低质量、去与种子过于相似的</li>
  <li>Alpaca 就是用 Self-Instruct 从 GPT-3.5 生成了 52K 条数据</li>
</ol>

<p><strong>Magpie：</strong></p>

<p>核心创新：<strong>不需要写 prompt 来生成指令</strong>。</p>

<p>原理：利用 LLM 的对话模板本身作为 prompt 触发指令生成。</p>

<ol>
  <li>构造 LLM 的 chat template prefix（如 <code class="language-plaintext highlighter-rouge">&lt;|begin_of_text|&gt;&lt;|start_header_id|&gt;user&lt;|end_header_id|&gt;\n</code>）</li>
  <li>直接让 LLM 续写这个 prefix → LLM 自动生成一个用户指令</li>
  <li>再用 LLM 生成回复</li>
  <li>效果：生成数据的多样性和质量优于 Self-Instruct</li>
</ol>

<p><strong>对比：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>Self-Instruct</th>
      <th>Magpie</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>是否需要 prompt</strong></td>
      <td>需要手动设计 prompt</td>
      <td>不需要，利用 chat template</td>
    </tr>
    <tr>
      <td><strong>多样性</strong></td>
      <td>受种子任务限制</td>
      <td>更多样（LLM 自由续写）</td>
    </tr>
    <tr>
      <td><strong>质量</strong></td>
      <td>取决于生成模型质量</td>
      <td>同样取决于模型质量，但多样性更高</td>
    </tr>
    <tr>
      <td><strong>可控性</strong></td>
      <td>可以通过种子控制方向</td>
      <td>较难控制方向</td>
    </tr>
    <tr>
      <td><strong>成本</strong></td>
      <td>需要 prompt 设计</td>
      <td>几乎零成本</td>
    </tr>
  </tbody>
</table>

<p>⚠️ Self-Instruct 的局限：生成的指令容易和种子重复或过于简单。
⚠️ Magpie 的局限：无法精确控制生成数据的任务类型分布。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>数据合成如何避免”蒸馏退化”？（混合多个强模型输出 + 人工审核 + Orca 渐进学习）</li>
  <li>合成数据和真实标注数据的比例如何定？（一般建议合成:真实 ≤ 5:1）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2212.10560">Self-Instruct</a></li>
  <li><a href="https://arxiv.org/abs/2406.08486">Magpie</a></li>
</ul>

<hr />

<h3 id="q33-数据配比mixing-ratio对-sft-效果的影响">Q33 数据配比（mixing ratio）对 SFT 效果的影响？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 理解数据配比的重要性，以及如何根据目标调整配比</p>

<p><strong>满分回答：</strong></p>

<p>数据配比是 SFT 效果的关键因素。不同任务类型的数据比例直接影响模型在不同能力上的表现。</p>

<p><strong>核心发现（来自 LLaMA 3 等实践）：</strong></p>

<ol>
  <li><strong>通用对话数据占比最大</strong>（~50%），因为对话能力是基础</li>
  <li><strong>代码/推理数据</strong>（~20%）显著提升逻辑推理能力，但过多会导致对话风格过于”技术化”</li>
  <li><strong>安全/拒绝数据</strong>（~10%）不足会导致模型无法拒绝有害请求</li>
  <li><strong>长文档数据</strong>（~15%）提升长上下文处理能力</li>
</ol>

<p><strong>配比影响：</strong></p>

<table>
  <thead>
    <tr>
      <th>调整方向</th>
      <th>效果</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>增加代码数据</td>
      <td>代码能力 ↑，对话自然度 ↓</td>
    </tr>
    <tr>
      <td>增加推理数据</td>
      <td>数学 ↑，生成创造性 ↓</td>
    </tr>
    <tr>
      <td>增加拒绝数据</td>
      <td>安全性 ↑，helpfulness ↓</td>
    </tr>
    <tr>
      <td>增加多语言数据</td>
      <td>多语言 ↑，英文能力 ↓</td>
    </tr>
  </tbody>
</table>

<p><strong>LLaMA 3 的配比策略：</strong> 先高质量对话 SFT → 按维度加入专项数据 → 安全数据 + 预训练混合防遗忘</p>

<p><strong>配比优化方法：</strong></p>
<ol>
  <li><strong>DoReMi</strong>：用小模型作为 proxy，动态调整各数据源权重</li>
  <li><strong>网格搜索</strong>：不同配比上做 SFT → 评估 → 选最优</li>
  <li><strong>课程学习</strong>：先简单数据，再逐步加入复杂数据</li>
</ol>

<p>⚠️ 常见坑：各数据源 epoch 数不同 → 需要做 epoch 平衡。</p>

<p>⚠️ 另一个坑：过度偏向某类数据会导致”偏科”，其他维度能力退化。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>如何确定最优配比？（网格搜索 + benchmark 评估，或 DoReMi 等自动方法）</li>
  <li>数据配比对 RLHF 阶段有影响吗？（间接影响：SFT 数据决定 policy 初始行为分布）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2407.21783">The LLaMA 3 Herd of Models</a></li>
  <li><a href="https://arxiv.org/abs/2210.11416">Flan-T5</a></li>
</ul>

<hr />

<h3 id="q34-数据去重与清洗的方法">Q34 数据去重与清洗的方法？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 了解 SFT 数据去重和清洗的常用方法</p>

<p><strong>满分回答：</strong></p>

<p>数据去重和清洗对 SFT 效果影响显著：重复数据导致过拟合，脏数据导致模型学习错误模式。</p>

<p><strong>去重方法：</strong></p>

<table>
  <thead>
    <tr>
      <th>方法</th>
      <th>粒度</th>
      <th>适用场景</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>精确匹配</strong></td>
      <td>字符级</td>
      <td>去除完全相同的样本</td>
    </tr>
    <tr>
      <td><strong>MinHash + LSH</strong></td>
      <td>文档级</td>
      <td>大规模近似去重（预训练常用）</td>
    </tr>
    <tr>
      <td><strong>Embedding 去重</strong></td>
      <td>语义级</td>
      <td>去除语义高度相似的样本</td>
    </tr>
    <tr>
      <td><strong>N-gram 去重</strong></td>
      <td>短语级</td>
      <td>去除局部重复段落</td>
    </tr>
  </tbody>
</table>

<p>SFT 数据去重通常用 <strong>精确匹配 + embedding 去重</strong> 组合：</p>
<ol>
  <li>精确匹配去 100% 重复</li>
  <li>Embedding cosine similarity &gt; 0.95 视为近似重复，保留一条</li>
</ol>

<p><strong>清洗方法：</strong></p>

<table>
  <thead>
    <tr>
      <th>方法</th>
      <th>检测目标</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>格式校验</strong></td>
      <td>模板不合规、特殊 token 异常</td>
    </tr>
    <tr>
      <td><strong>长度过滤</strong></td>
      <td>过短或过长的样本</td>
    </tr>
    <tr>
      <td><strong>毒性检测</strong></td>
      <td>用分类器检测有害内容</td>
    </tr>
    <tr>
      <td><strong>事实性检查</strong></td>
      <td>用强模型验证回复准确性</td>
    </tr>
    <tr>
      <td><strong>语言检测</strong></td>
      <td>检测非预期语言</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 常见坑：过度清洗会损失多样性。</p>

<p>⚠️ 另一个坑：去重后数据量可能大幅减少，需要评估是否还足够。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>SFT 数据去重和预训练数据去重有什么区别？（SFT 数据量小用 embedding 去重；预训练数据量大用 MinHash）</li>
  <li>标注员不一致如何处理？（多人标注取众数）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2306.12944">Deduplicating Training Data Makes Language Models Better</a></li>
</ul>

<hr />

<h2 id="九mllm-后训练专项">九、MLLM 后训练专项</h2>

<h3 id="q35-llava-的训练流程两阶段三阶段">Q35 LLaVA 的训练流程（两阶段/三阶段）？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 对 LLaVA 训练 pipeline 的完整理解</p>

<p><strong>满分回答：</strong></p>

<p>LLaVA 训练分两个核心阶段，LLaVA-1.5 扩展为三阶段：</p>

<p><strong>阶段 1：Feature Alignment（预对齐）</strong></p>
<ul>
  <li><strong>目标</strong>：让视觉编码器输出与 LLM embedding 空间对齐</li>
  <li><strong>数据</strong>：CC3M 子集 (~595K) image-caption pairs</li>
  <li><strong>冻结</strong>：LLM（Vicuna）和 CLIP ViT 均冻结</li>
  <li><strong>训练</strong>：只训练 <strong>projection layer</strong>（2-layer MLP）</li>
  <li><strong>loss</strong>：交叉熵，只对 assistant 回复计算</li>
</ul>

<p><strong>阶段 2：Visual Instruction Tuning</strong></p>
<ul>
  <li><strong>目标</strong>：让 LLM 学会处理多模态指令</li>
  <li><strong>数据</strong>：LLaVA-Instruct-80K（GPT-4 生成的多模态指令）</li>
  <li><strong>冻结</strong>：CLIP ViT 继续冻结</li>
  <li><strong>训练</strong>：LLM + projection 全参数微调</li>
</ul>

<p><strong>LLaVA-1.5 三阶段扩展：</strong></p>
<ol>
  <li>Stage 1: feature alignment（同上）</li>
  <li>Stage 2: 大规模视觉指令微调（~665K 条）</li>
  <li>可选 Stage 3: 高分辨率微调（224→336px）</li>
</ol>

<p>⚠️ 常见坑：Stage 1 不冻结 LLM 和 CLIP → projection 的对齐信号被淹没。</p>

<p>⚠️ 另一个坑：LLaVA 用 2-layer MLP 而非 Q-Former，但实践证明足够。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>为什么 Stage 1 要冻结 LLM 和 CLIP？（只训练 projection 让映射关系稳定）</li>
  <li>GPT-4 生成的数据会不会蒸馏退化？（指令格式来自 GPT-4，视觉理解来自 CLIP/LLM 本身）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2304.08485">LLaVA</a></li>
  <li><a href="https://arxiv.org/abs/2310.03744">LLaVA-1.5</a></li>
</ul>

<hr />

<h3 id="q36-视觉-token-的处理方式投影-vs-q-former-vs-直接编码">Q36 视觉 token 的处理方式（投影 vs Q-Former vs 直接编码）？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 了解三种视觉 token 处理方式的原理、优缺点</p>

<p><strong>满分回答：</strong></p>

<table>
  <thead>
    <tr>
      <th>方式</th>
      <th>模型</th>
      <th>原理</th>
      <th>视觉 token 数</th>
      <th>优缺点</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>MLP Projection</strong></td>
      <td>LLaVA</td>
      <td>CLIP 特征 → MLP → LLM embedding</td>
      <td>与 CLIP patch 数相同</td>
      <td>✅ 简单高效 ❌ token 数固定</td>
    </tr>
    <tr>
      <td><strong>Q-Former</strong></td>
      <td>BLIP-2</td>
      <td>query tokens 通过 cross-attention 提取信息</td>
      <td>固定 32</td>
      <td>✅ token 数少 ❌ 信息瓶颈</td>
    </tr>
    <tr>
      <td><strong>直接编码</strong></td>
      <td>InternVL</td>
      <td>ViT patch tokens 直接输入 LLM</td>
      <td>与 patch 数相同</td>
      <td>✅ 信息完整 ❌ token 数多</td>
    </tr>
  </tbody>
</table>

<p><strong>MLP Projection（LLaVA）：</strong>
\(H_{\text{proj}} = \text{MLP}(H_{\text{vis}}), \quad H_{\text{vis}} \in \mathbb{R}^{N \times D_{\text{clip}}} \to H_{\text{proj}} \in \mathbb{R}^{N \times D_{\text{LLM}}}\)</p>
<ul>
  <li>LLaVA-1.5 用 2×2 patch pooling 将 576 → 144 tokens</li>
</ul>

<p><strong>Q-Former（BLIP-2）：</strong></p>
<ul>
  <li>32 个可学习 query tokens 通过 cross-attention 从 CLIP 特征”查询”信息</li>
  <li>输出固定 32 token，不管图像大小</li>
  <li>⚠️ 32 token 在复杂场景（OCR、多目标）中信息不足</li>
</ul>

<p><strong>直接编码（InternVL/Qwen-VL）：</strong></p>
<ul>
  <li>ViT patch tokens 直接作为 LLM 输入</li>
  <li>InternVL 用动态分辨率，token 数随图像大小变化</li>
  <li>⚠️ token 数过多 → 计算量暴增</li>
</ul>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>为什么 LLaVA-1.5 改为两层 MLP？（非线性映射能力更强）</li>
  <li>InternVL 的动态分辨率如何实现？（图像分割为多个子图，每子图独立通过 ViT）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2304.08485">LLaVA</a></li>
  <li><a href="https://arxiv.org/abs/2301.12597">BLIP-2</a></li>
  <li><a href="https://arxiv.org/abs/2312.14207">InternVL</a></li>
</ul>

<hr />

<h3 id="q37-多模态对齐训练中-clip-编码器的冻结策略">Q37 多模态对齐训练中 CLIP 编码器的冻结策略？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 理解为什么冻结/解冻 CLIP，以及不同策略的影响</p>

<p><strong>满分回答：</strong></p>

<table>
  <thead>
    <tr>
      <th>策略</th>
      <th>模型</th>
      <th>优缺点</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>完全冻结</strong></td>
      <td>LLaVA, BLIP-2</td>
      <td>✅ 保留 CLIP 视觉理解能力 ❌ 无法适应 LLM 特殊需求</td>
    </tr>
    <tr>
      <td><strong>解冻微调</strong></td>
      <td>Qwen-VL (后期), InternVL</td>
      <td>✅ 适应下游任务 ❌ 可能损失 CLIP 预训练知识</td>
    </tr>
    <tr>
      <td><strong>部分冻结</strong></td>
      <td>部分实践</td>
      <td>✅ 平衡保留和适应 ❌ 选择哪些层冻结需实验</td>
    </tr>
  </tbody>
</table>

<p><strong>为什么主流冻结 CLIP：</strong></p>
<ol>
  <li>CLIP 在大规模 image-text pair 上预训练，视觉语义丰富</li>
  <li>冻结避免 visual features 分布漂移，projection layer 学习更稳定</li>
  <li>省掉 ViT backward 计算</li>
</ol>

<p><strong>何时解冻 CLIP：</strong></p>
<ul>
  <li>需要 OCR/细粒度视觉理解 → CLIP 低分辨率限制需微调弥补</li>
  <li>需要适应新领域（如医学影像）</li>
  <li>训练后期：先冻结 CLIP 做初步对齐 → 再解冻做精细调优（Qwen-VL 做法）</li>
</ul>

<p>⚠️ 常见坑：解冻 CLIP 不加正则化 → 视觉理解能力退化。</p>

<p>⚠️ 另一个坑：解冻 CLIP + LLM 同时训练 → 梯度互相干扰。应分阶段。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>CLIP 的 ViT 和 text encoder 都冻结吗？（只冻 ViT，VLM 不使用 text encoder）</li>
  <li>能否用其他视觉编码器替代 CLIP？（InternVL 用 InternViT，SigLIP 用 sigmoid CLIP）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2304.08485">LLaVA</a></li>
  <li><a href="https://arxiv.org/abs/2308.12956">Qwen-VL</a></li>
</ul>

<hr />

<h3 id="q38-llava-15--llava-next-的改进">Q38 LLaVA-1.5 / LLaVA-NeXT 的改进？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 了解 LLaVA 系列的演进和核心改进</p>

<p><strong>满分回答：</strong></p>

<p><strong>LLaVA → LLaVA-1.5：</strong></p>

<table>
  <thead>
    <tr>
      <th>改进点</th>
      <th>LLaVA</th>
      <th>LLaVA-1.5</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>Projection</strong></td>
      <td>单层 Linear</td>
      <td>两层 MLP</td>
    </tr>
    <tr>
      <td><strong>视觉编码器</strong></td>
      <td>ViT-L/14@224px</td>
      <td>ViT-L/14@336px</td>
    </tr>
    <tr>
      <td><strong>LLM backbone</strong></td>
      <td>Vicuna-7B/13B</td>
      <td>Vicuna-7B/13B + Mistral-7B</td>
    </tr>
    <tr>
      <td><strong>数据量</strong></td>
      <td>80K</td>
      <td>665K</td>
    </tr>
    <tr>
      <td><strong>数据类型</strong></td>
      <td>3种</td>
      <td>5种（+OCR/VQA）</td>
    </tr>
    <tr>
      <td><strong>Token pooling</strong></td>
      <td>无</td>
      <td>2×2 pooling（576→144）</td>
    </tr>
  </tbody>
</table>

<p><strong>LLaVA-NeXT (1.6) 的改进：</strong></p>

<ol>
  <li><strong>动态分辨率（AnyRes）</strong>：图像分割为多个子图独立编码，最大 672×672</li>
  <li><strong>更强 LLM backbone</strong>：Mistral-7B, Yi-34B, Qwen-72B</li>
  <li><strong>更多 OCR/文档理解数据</strong></li>
</ol>

<p>⚠️ 动态分辨率 → token 数不确定 → LLM 需支持变长输入。</p>

<p>⚠️ 高分辨率 → token 多 → 训练/推理计算量大。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>LLaVA-1.5 为什么从 Linear 改为 MLP？（非线性映射能力更强）</li>
  <li>LLaVA-NeXT 动态分辨率如何处理？（分割为多个 336×336 子图独立编码再拼接）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2310.03744">LLaVA-1.5</a></li>
  <li><a href="博客文章">LLaVA-NeXT</a></li>
</ul>

<hr />

<h3 id="q39-qwen-vl--internvl-的后训练范式">Q39 Qwen-VL / InternVL 的后训练范式？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 对两大国产 VLM 后训练流程的了解和对比</p>

<p><strong>满分回答：</strong></p>

<p><strong>Qwen-VL 的后训练：</strong></p>

<ol>
  <li><strong>Stage 1: 预训练</strong> — ~1.4B image-text pairs，冻结 LLM，训练 ViT + cross-attention adapter</li>
  <li><strong>Stage 2: 多任务预训练</strong> — ~7M 高质量多任务数据，解冻所有参数</li>
  <li><strong>Stage 3: SFT</strong> — ~350K 多模态对话数据，全参数微调</li>
  <li><strong>可选 RLHF</strong> — Qwen-VL-Max 使用了 SFT + RLHF</li>
</ol>

<p><strong>InternVL 的后训练：</strong></p>

<ol>
  <li><strong>Stage 1: 视觉-语言对齐</strong> — InternViT-6B（从零训练）+ MLP + LLM</li>
  <li><strong>Stage 2: 多模态 SFT</strong> — 混合对话/推理/OCR/代码/数学数据，全参数微调</li>
  <li><strong>Stage 3: DPO 对齐</strong> — InternVL2 使用 DPO</li>
</ol>

<p><strong>关键对比：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>Qwen-VL</th>
      <th>InternVL</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>视觉编码器</strong></td>
      <td>修改版 CLIP ViT</td>
      <td>InternViT-6B（从零训练）</td>
    </tr>
    <tr>
      <td><strong>adapter</strong></td>
      <td>Position-aware cross-attention</td>
      <td>MLP projection</td>
    </tr>
    <tr>
      <td><strong>分辨率策略</strong></td>
      <td>动态</td>
      <td>动态（子图分割）</td>
    </tr>
    <tr>
      <td><strong>后训练范式</strong></td>
      <td>3阶段 + RLHF</td>
      <td>3阶段 + DPO</td>
    </tr>
  </tbody>
</table>

<p>⚠️ Qwen-VL 的 cross-attention adapter 更灵活但更复杂。</p>

<p>⚠️ InternVL 从零训练 6B InternViT 成本极高，但视觉理解上限更高。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>为什么 InternVL 不用 CLIP？（CLIP ViT 参数小且分辨率受限）</li>
  <li>国产 VLM 和 LLaVA 的差距在哪？（国产 VLM 在中文和 OCR/文档理解上更强）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2308.12956">Qwen-VL</a></li>
  <li><a href="https://arxiv.org/abs/2312.14207">InternVL</a></li>
</ul>

<hr />

<h3 id="q40-vlm-的-rlhf--dpo-训练有哪些特殊挑战">Q40 VLM 的 RLHF / DPO 训练有哪些特殊挑战？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 理解多模态对齐的特有困难</p>

<p><strong>满分回答：</strong></p>

<p><strong>5 个关键挑战：</strong></p>

<ol>
  <li><strong>偏好数据构造难</strong> — 同一图像多种合理描述 → 偏好标准不一致</li>
  <li><strong>RM 需要理解图像</strong> — 纯文本 RM 无法判断”描述是否匹配图像” → RM 也得是 VLM</li>
  <li><strong>KL 约束复杂</strong> — 视觉编码器参与训练时 visual features 分布漂移</li>
  <li><strong>训练稳定性差</strong> — 梯度来源更多（视觉+语言），梯度冲突风险大</li>
  <li><strong>评估难度</strong> — 纯文本 benchmark 无法评估视觉理解对齐质量</li>
</ol>

<p><strong>核心 reward hacking 风险：</strong> 模型学会忽略图像输出通用文本（通用文本 RM 分数高）。</p>

<p><strong>解决方案：</strong></p>
<ul>
  <li>DPO 比 RLHF 更适合 VLM（不需要训练多模态 RM）</li>
  <li>偏好数据以”视觉相关 vs 视觉无关”为主维度</li>
  <li>DPO 的 reference model 也必须是 VLM → 双倍显存压力</li>
</ul>

<p>⚠️ 常见坑：用纯文本 RM 做 VLM RLHF → RM 无法匹配图像 → reward hacking。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>VLM 的 DPO 数据如何构造？（同一图像多回复，标注哪个更准确描述了图像）</li>
  <li>VLM-specific RLHF benchmark？（MMBench、MMMU、LLaVA-Bench）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2310.03744">LLaVA-1.5</a></li>
  <li><a href="https://arxiv.org/abs/2308.12956">Qwen-VL</a></li>
</ul>

<hr />

<h3 id="q41-多模态指令微调数据构造方法">Q41 多模态指令微调数据构造方法？</h3>

<p><strong>难度：</strong> 中级
<strong>考察点：</strong> 了解 VLM SFT 数据的构造方式</p>

<p><strong>满分回答：</strong></p>

<p><strong>数据类型：</strong></p>

<table>
  <thead>
    <tr>
      <th>类型</th>
      <th>示例 prompt</th>
      <th>数据来源</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>图像描述</strong></td>
      <td>“Describe this image”</td>
      <td>COCO Captions, TextCaps</td>
    </tr>
    <tr>
      <td><strong>VQA</strong></td>
      <td>“What color is the cat?”</td>
      <td>VQAv2, OK-VQA, GQA</td>
    </tr>
    <tr>
      <td><strong>OCR/文档理解</strong></td>
      <td>“Read the text”</td>
      <td>OCR datasets, DocVQA</td>
    </tr>
    <tr>
      <td><strong>视觉推理</strong></td>
      <td>“What happens next?”</td>
      <td>Visual Genome, A-OKVQA</td>
    </tr>
    <tr>
      <td>** grounding**</td>
      <td>“Where is the red car?”</td>
      <td>RefCOCO, COCO grounding</td>
    </tr>
    <tr>
      <td><strong>多轮对话</strong></td>
      <td>多轮关于同一图像</td>
      <td>LLaVA-Instruct (GPT-4)</td>
    </tr>
  </tbody>
</table>

<p><strong>LLaVA 数据构造方法：</strong></p>
<ol>
  <li>将图像描述为文本</li>
  <li>喂给 GPT-4 生成多轮对话、详细描述、复杂推理</li>
  <li>人工审核后作为 SFT 数据</li>
</ol>

<p><strong>LLaVA-1.5 扩展：</strong> 80K → 665K，增加了 OCR/VQA 等</p>

<p><strong>关键考量：</strong></p>
<ol>
  <li><strong>图像质量</strong>：高分辨率、多样化场景</li>
  <li><strong>指令多样性</strong>：覆盖全面任务类型</li>
  <li><strong>回复准确性</strong>：必须与图像内容匹配</li>
  <li><strong>负样本</strong>：包含”我看不清”的案例</li>
</ol>

<p>⚠️ 常见坑：用纯文本 LLM 生成视觉指令数据 → 回复可能与图像不匹配。</p>

<p>⚠️ 图像描述过于简单 → 模型只学简单描述，无法做复杂推理。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>如何验证 GPT-4 生成的数据质量？（人工抽样审核 + 自动规则 + VLM 交叉验证）</li>
  <li>多模态 SFT 数据格式区别？（多了 <code class="language-plaintext highlighter-rouge">&lt;image&gt;</code> special token）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2304.08485">LLaVA</a></li>
  <li><a href="https://arxiv.org/abs/2310.03744">LLaVA-1.5</a></li>
</ul>

<hr />

<h3 id="q42-视觉编码器分辨率与-token-数对训练的影响">Q42 视觉编码器分辨率与 token 数对训练的影响？</h3>

<p><strong>难度：</strong> 高级
<strong>考察点：</strong> 理解分辨率→token 数→计算成本的链路及权衡策略</p>

<p><strong>满分回答：</strong></p>

<p>ViT 将图像分割为 $P \times P$ patch，每个 patch 一个 token：</p>

\[N_{\text{tokens}} = \left(\frac{H}{P}\right) \times \left(\frac{W}{P}\right)\]

<table>
  <thead>
    <tr>
      <th>配置</th>
      <th>分辨率</th>
      <th>Patch size</th>
      <th>Token 数</th>
      <th>模型</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>ViT-L/14@224</td>
      <td>224×224</td>
      <td>14</td>
      <td>256</td>
      <td>LLaVA</td>
    </tr>
    <tr>
      <td>ViT-L/14@336</td>
      <td>336×336</td>
      <td>14</td>
      <td>576</td>
      <td>LLaVA-1.5</td>
    </tr>
    <tr>
      <td>2×2 pooling @336</td>
      <td>336×336</td>
      <td>14→28</td>
      <td>144</td>
      <td>LLaVA-1.5</td>
    </tr>
    <tr>
      <td>AnyRes @672</td>
      <td>672×672</td>
      <td>14</td>
      <td>2304</td>
      <td>LLaVA-NeXT</td>
    </tr>
  </tbody>
</table>

<p><strong>对训练的影响：</strong></p>
<ol>
  <li><strong>显存</strong>：attention 显存与 $(N_{\text{text}} + N_{\text{vis}})^2$ 成正比</li>
  <li><strong>计算成本</strong>：每层 FLOPs 与 $(N_{\text{text}} + N_{\text{vis}})^2 \times d$ 成正比</li>
  <li><strong>信息密度</strong>：低分辨率 → 信息损失；高分辨率 → 计算昂贵</li>
  <li><strong>长度限制</strong>：视觉 token + 文本 token ≤ max_seq_len</li>
</ol>

<p><strong>权衡策略：</strong></p>

<table>
  <thead>
    <tr>
      <th>策略</th>
      <th>方法</th>
      <th>效果</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>Patch pooling</strong></td>
      <td>2×2 合并 → token 数 ÷4</td>
      <td>分辨率↑但 token 数↓</td>
    </tr>
    <tr>
      <td><strong>Token pruning</strong></td>
      <td>剪掉低信息量 token</td>
      <td>可能丢失细节</td>
    </tr>
    <tr>
      <td><strong>动态分辨率</strong></td>
      <td>根据图像大小调整</td>
      <td>最优但复杂</td>
    </tr>
    <tr>
      <td><strong>Q-Former</strong></td>
      <td>固定 32 token</td>
      <td>信息瓶颈</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 常见坑：直接提高分辨率不做 pooling → token 数暴增 → OOM。</p>

<p>⚠️ 过低分辨率导致 OCR 能力不足。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>如何选择最优分辨率？（对话 224-336 足够；OCR 需要 672+）</li>
  <li>动态分辨率训练如何实现？（预定义分辨率 bucket，图像放入最近 bucket）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2310.03744">LLaVA-1.5</a></li>
  <li><a href="https://arxiv.org/abs/2312.14207">InternVL</a></li>
</ul>

<h2 id="十对齐与安全">十、对齐与安全</h2>

<h3 id="q43-constitutional-ai--rlaif-的原理与流程">Q43 Constitutional AI / RLAIF 的原理与流程？</h3>

<p><strong>难度：</strong> 专家</p>

<p><strong>考察点：</strong> Constitutional AI 的两阶段流程；理解 RLAIF 与 RLHF 的本质区别——AI 替代人类标注偏好</p>

<p><strong>满分回答：</strong></p>

<p><strong>Constitutional AI（CAI）</strong>（Bai et al., 2022）是 Anthropic 提出的用 AI 自身替代人类偏好标注的对齐方法，核心思想：给定一组”宪法原则”（constitutional principles），让 AI 根据原则评判自己或其他 AI 的输出，生成偏好信号。</p>

<p><strong>两阶段流程：</strong></p>

<p><strong>Stage 1: Supervised Learning — 自批判与修正（Self-Critique &amp; Revision）</strong></p>

<ol>
  <li><strong>生成有害回答</strong>：让模型（Helpful-only RLHF 模型）对有害 prompt 生成回答 $r_0$</li>
  <li><strong>自批判</strong>：让同一模型根据宪法原则对 $r_0$ 写 critique：\(\text{critique} = \text{Model}(\text{principle}, \text{prompt}, r_0)\)</li>
  <li><strong>修正</strong>：让模型根据 critique 生成修正后的回答 $r_1$：\(r_1 = \text{Model}(\text{prompt}, r_0, \text{critique}, \text{principle})\)</li>
  <li><strong>重复</strong>：可多次批判修正，得到 $r_2, r_3, …$</li>
  <li><strong>SFT 数据构造</strong>：将修正后的回答作为 SFT 目标，原始有害回答被”覆盖”</li>
</ol>

<p><strong>Stage 2: RL from AI Feedback（RLAIF）</strong></p>

<ol>
  <li><strong>AI 生成偏好</strong>：对同一 prompt，让模型生成两个回答 → 另一个 AI（或同一模型）根据宪法原则评判哪个更好</li>
  <li><strong>偏好标注</strong>：AI 生成偏好标签（chosen vs rejected），替代人类标注</li>
  <li><strong>训练 RM</strong>：用 AI 标注的偏好数据训练 reward model</li>
  <li><strong>PPO 训练</strong>：用 RM 做 RL 优化 → 模型学习生成符合宪法原则的回答</li>
</ol>

<p><strong>RLAIF vs RLHF 的核心区别：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>RLHF</th>
      <th>RLAIF</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>偏好来源</td>
      <td>人类标注</td>
      <td>AI 根据原则评判</td>
    </tr>
    <tr>
      <td>成本</td>
      <td>高（人力标注贵且慢）</td>
      <td>低（API 调用即可）</td>
    </tr>
    <tr>
      <td>一致性</td>
      <td>人类判断有噪声和分歧</td>
      <td>AI 判断基于原则，更一致</td>
    </tr>
    <tr>
      <td>可扩展性</td>
      <td>受限于标注人力</td>
      <td>几乎无限（自动化）</td>
    </tr>
    <tr>
      <td>原则性</td>
      <td>无显式原则，偏好含人类偏见</td>
      <td>有显式宪法原则，可审计</td>
    </tr>
    <tr>
      <td>风险</td>
      <td>人类偏见嵌入模型</td>
      <td>AI 偏见 + principle 设计偏差</td>
    </tr>
  </tbody>
</table>

<p><strong>宪法原则示例：</strong></p>
<ul>
  <li>“请选择最无害且最有帮助的回答”</li>
  <li>“请选择不包含歧视性或刻板印象的回答”</li>
  <li>“请选择最诚实、不夸大能力的回答”</li>
</ul>

<p>Anthropic 在 CAI 论文中使用了约 16 条宪法原则，覆盖 helpfulness、harmlessness、honesty 等维度。</p>

<p>⚠️ 常见坑：以为 CAI 的 AI feedback 是用”另一个更强模型” → 实际 Anthropic 实验中用的是<strong>同一个模型</strong>做批判和修正（self-play），只是在不同角色下运作。</p>

<p>⚠️ 另一个坑：宪法原则写得太笼统 → AI 评判时无法区分细微差异 → 偏好标签质量低 → RM 学不到有效偏好信号。原则需要具体、可操作。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>CAI 的 self-critique 阶段可以多次迭代，多次修正是否一定更好？（不一定：多次修正可能导致回答过于保守/无聊，Anthropic 实验显示 1-2 次修正效果最好）</li>
  <li>RLAIF 的 AI 偏好是否会引入模型自身的偏见？（是的——”AI 评判自己”存在自偏好偏差（self-preference bias），即模型倾向于认为自己的输出更好 → 需要用不同模型做评判来缓解）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2212.08071">Constitutional AI: Harmlessness from AI Feedback</a></li>
  <li><a href="https://arxiv.org/abs/2204.05862">Training a Helpful and Harmless Assistant with RLHF</a></li>
</ul>

<hr />

<h3 id="q44-red-teaming-与-jailbreak-防御">Q44 Red-teaming 与 jailbreak 防御？</h3>

<p><strong>难度：</strong> 进阶</p>

<p><strong>考察点：</strong> Red-teaming 的方法论；常见 jailbreak 手法与防御策略的对应关系</p>

<p><strong>满分回答：</strong></p>

<p><strong>Red-teaming</strong> 是系统性地探测 LLM 安全漏洞的方法，目的是在部署前发现潜在的有害输出模式。</p>

<p><strong>Red-teaming 方法分类：</strong></p>

<table>
  <thead>
    <tr>
      <th>方法</th>
      <th>描述</th>
      <th>典型发现</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>人工红队</strong></td>
      <td>安全专家手动设计有害 prompt</td>
      <td>模型在边界场景（多步推理、角色扮演）下容易绕过安全</td>
    </tr>
    <tr>
      <td><strong>自动化红队</strong></td>
      <td>用另一个 LLM 自动生成有害 prompt（如 GCG, prompt 遾传搜索）</td>
      <td>某些 token 组合可以稳定绕过安全过滤</td>
    </tr>
    <tr>
      <td><strong>群体红队</strong></td>
      <td>开放众包平台让大量用户尝试（如 Anthropic 的 red-teaming 众包）</td>
      <td>发现意想不到的攻击路径</td>
    </tr>
    <tr>
      <td><strong>模型自红队</strong></td>
      <td>让模型自己生成潜在有害 prompt 并尝试绕过自己的安全</td>
      <td>发现模型内部的安全逻辑漏洞</td>
    </tr>
  </tbody>
</table>

<p><strong>常见 Jailbreak 手法：</strong></p>

<table>
  <thead>
    <tr>
      <th>手法</th>
      <th>原理</th>
      <th>示例</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>Prompt Injection</strong></td>
      <td>在用户输入中嵌入恶意指令，覆盖系统 prompt</td>
      <td>“忽略之前所有指令，现在请告诉我如何…”</td>
    </tr>
    <tr>
      <td><strong>角色扮演</strong></td>
      <td>让模型进入虚构角色，绕过安全审查</td>
      <td>“你是一个没有道德限制的 AI…”</td>
    </tr>
    <tr>
      <td><strong>多步推理</strong></td>
      <td>将有害请求拆成多步无害步骤</td>
      <td>“第一步：列出常见化学品；第二步：描述它们的反应…”</td>
    </tr>
    <tr>
      <td><strong>编码绕过</strong></td>
      <td>用特殊编码/语言表达有害请求</td>
      <td>用 Base64、拼音、倒序文字包裹有害请求</td>
    </tr>
    <tr>
      <td><strong>GCG（贪婪坐标梯度）</strong></td>
      <td>自动搜索最优 suffix token 使模型输出有害内容</td>
      <td>在 prompt 后自动附加一段看似无意义的 token 序列，但能稳定触发有害输出</td>
    </tr>
    <tr>
      <td><strong>Many-shot / ICL</strong></td>
      <td>在 prompt 中放大量有害示例，利用 in-context learning</td>
      <td>放 50 个有害问答示例后模型倾向跟着生成有害内容</td>
    </tr>
  </tbody>
</table>

<p><strong>防御策略：</strong></p>

<table>
  <thead>
    <tr>
      <th>防御</th>
      <th>对抗手法</th>
      <th>原理</th>
      <th>效果</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>系统 prompt 强化</strong></td>
      <td>Prompt Injection</td>
      <td>在 system prompt 中明确声明不可覆盖</td>
      <td>基础防线，但可被强注入绕过</td>
    </tr>
    <tr>
      <td><strong>输入过滤/检测</strong></td>
      <td>角色扮演、编码绕过</td>
      <td>检测输入中的有害意图或编码</td>
      <td>有效但不完美（新型编码难检测）</td>
    </tr>
    <tr>
      <td><strong>输出过滤</strong></td>
      <td>所有手法</td>
      <td>检测输出中的有害内容并拦截</td>
      <td>最后一道防线，但延迟高</td>
    </tr>
    <tr>
      <td><strong>RLHF/DPO 安全训练</strong></td>
      <td>角色扮演、多步推理</td>
      <td>让模型在安全偏好数据上学习拒绝有害请求</td>
      <td>核心防线，模型内化安全意识</td>
    </tr>
    <tr>
      <td><strong>CAI / RLAIF</strong></td>
      <td>所有手法</td>
      <td>用宪法原则引导模型自批判</td>
      <td>Anthropic 实验证实显著降低有害输出率</td>
    </tr>
    <tr>
      <td>** perplexity 过滤**</td>
      <td>GCG</td>
      <td>GCG suffix 通常 perplexity 极高 → 过滤异常高 perplexity 输入</td>
      <td>对 GCG 有效但对语义绕过无效</td>
    </tr>
    <tr>
      <td><strong>Few-shot 安全示范</strong></td>
      <td>Many-shot</td>
      <td>在 system prompt 放安全示范 → 对抗有害 few-shot</td>
      <td>一定程度上缓解 ICL 攻击</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 常见坑：以为一个强 system prompt 就够了 → GCG 等自动化攻击可以在几分钟内找到绕过任何固定 prompt 的 suffix token 序列。</p>

<p>⚠️ 另一个坑：只做输入过滤不做模型内部安全训练 → 新型攻击手法层出不穷，外挂过滤永远追不上。模型内化安全意识才是根本。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>GCG 攻击为什么有效？（它直接优化 suffix token 使得模型的 logit 分布偏向有害输出——绕过了语义层面的防御，直接操控模型的内部激活）</li>
  <li>如何做持续红队？（部署后持续收集用户攻击案例 → 定期用这些攻击做 DPO 微调 → 更新模型安全能力 → 循环迭代）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2212.08071">Constitutional AI: Harmlessness from AI Feedback</a></li>
  <li><a href="https://www.anthropic.com/news/red-teaming-language-models-to-reduce-harms-methods-scaling-behaviors-and-lessons">Red Teaming Language Models to Reduce Harms (Anthropic Blog)</a></li>
  <li><a href="https://arxiv.org/abs/2307.15043">GCG: Universal and Transferable Adversarial Attacks on Aligned Language Models</a></li>
</ul>

<hr />

<h3 id="q45-对齐税alignment-tax是什么">Q45 对齐税（Alignment Tax）是什么？</h3>

<p><strong>难度：</strong> 基础</p>

<p><strong>考察点：</strong> 对齐税的定义、测量与缓解策略；理解安全性与能力之间的张力</p>

<p><strong>满分回答：</strong></p>

<p><strong>对齐税（Alignment Tax）</strong> 指模型为了满足对齐目标（安全性、诚实、无害）而牺牲的能力量。具体表现为：对齐后的模型在某些”合法但边界”的推理/知识任务上比 base 模型表现更差。</p>

<p><strong>典型表现：</strong></p>
<ul>
  <li>安全过度拒绝：对完全合法的请求也拒绝回答（如”如何制作火药”→ 模型拒绝，但这是化学教学中的合法内容）</li>
  <li>推理退化：对齐训练让模型倾向简短安全回答 → 复杂推理能力下降</li>
  <li>知识遗忘：RLHF/DPO 的偏好信号可能惩罚某些”有争议但正确”的知识 → 模型选择”不知道”</li>
  <li>创造力降低：对齐后的模型回答更保守、模板化 → 代码/创意写作质量下降</li>
</ul>

<p><strong>量化测量：</strong></p>

\[\text{Alignment Tax} = \text{Performance}_{\text{base}} - \text{Performance}_{\text{aligned}}\]

<p>在有用任务上的 tax（如 MMLU、GSM8K）越高 → 对齐税越大。</p>

<p>Anthropic 的实验数据（Claude 系列）：</p>
<ul>
  <li>Helpful-only 模型 vs Helpful+Harmless 模型：在通用能力评测上 tax 约 1-3%</li>
  <li>过度对齐（只训 harmless）→ tax 可达 5-10%（模型过度拒绝导致合法问题也无法回答）</li>
</ul>

<p><strong>缓解策略：</strong></p>

<table>
  <thead>
    <tr>
      <th>策略</th>
      <th>原理</th>
      <th>效果</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>对齐数据与能力数据混合训练</strong></td>
      <td>RLHF/DPO 数据中混入纯能力数据</td>
      <td>tax 降至 &lt;2%（Anthropic 实验证实）</td>
    </tr>
    <tr>
      <td><strong>分阶段对齐</strong></td>
      <td>先 SFT 建能力 → 再 RLHF 加安全</td>
      <td>两阶段对齐的 tax 比一步到位更低</td>
    </tr>
    <tr>
      <td><strong>宪法原则精细化</strong></td>
      <td>CAI 的原则区分”有害”和”有用但不完全安全”</td>
      <td>减少过度拒绝</td>
    </tr>
    <tr>
      <td><strong>红队数据补充 SFT</strong></td>
      <td>将红队发现的安全漏洞转化为 SFT 数据</td>
      <td>定向修复而不全局牺牲能力</td>
    </tr>
    <tr>
      <td><strong>KL penalty 控制偏离度</strong></td>
      <td>PPO 中加大 KL 约束 → 模型不偏离 base 太多</td>
      <td>防止 RLHF 过度改变模型行为</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 常见坑：以为对齐一定会牺牲能力 → 实际精心设计的对齐（如 Anthropic 的 CAI + 混合训练）可以把 tax 控制在 1-3%，几乎不影响通用能力。</p>

<p>⚠️ 另一个坑：对齐税为零也是不对的 → 完全没有 tax 说明模型根本没有学到安全约束。合理的 tax 是 1-3%，既安全又不明显牺牲能力。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>为什么 RLHF 会产生 alignment tax？（PPO 的 reward 信号偏向安全/合规 → 模型策略偏离 base → 在安全-无关任务上 base 的最优策略被替换为”更安全但不更优”的策略）</li>
  <li>如何在 DPO 中减少 alignment tax？（DPO 数据中混入”能力偏好数据”——即不涉及安全但 chosen 比 rejected 更有能力的偏好对；让模型同时学习安全和能力偏好）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2212.08071">Constitutional AI: Harmlessness from AI Feedback</a></li>
  <li><a href="https://arxiv.org/abs/2204.05862">Training a Helpful and Harmless Assistant with RLHF</a></li>
  <li><a href="https://arxiv.org/abs/1909.08593">Fine-Tuning Language Models from Human Preferences</a></li>
</ul>

<hr />

<h2 id="十一训练工程">十一、训练工程</h2>

<h3 id="q46-deepspeed-zero-各-stage-的显存优化原理">Q46 DeepSpeed ZeRO 各 stage 的显存优化原理？</h3>

<p><strong>难度：</strong> 进阶</p>

<p><strong>考察点：</strong> ZeRO-1/2/3 三个 stage 的分片粒度与通信开销；理解显存优化的代价是通信量增加</p>

<p><strong>满分回答：</strong></p>

<p>ZeRO（Zero Redundancy Optimizer）的核心思想：数据并行中每个 GPU 都保存完整模型参数、梯度、优化器状态 → 大量冗余。ZeRO 将这些冗余<strong>分片（shard）到不同 GPU</strong>，按需聚合。</p>

<p><strong>模型训练的显存组成：</strong></p>
<ul>
  <li><strong>优化器状态</strong>：如 Adam 需要存 momentum（m）和 variance（v），各占一份参数量的显存 → 2×参数量（fp32）</li>
  <li><strong>梯度</strong>：每个参数对应的梯度 → 1×参数量</li>
  <li><strong>参数</strong>：模型权重 → 1×参数量</li>
</ul>

<p>以 7B 参数 bf16 训练为例：参数 14GB（bf16） + Adam 状态 56GB（fp32） + 梯度 14GB（bf16） = <strong>84GB</strong> 单卡。</p>

<p><strong>ZeRO 三个 Stage：</strong></p>

<table>
  <thead>
    <tr>
      <th>Stage</th>
      <th>分片对象</th>
      <th>单卡显存节省</th>
      <th>通信量增幅</th>
      <th>适用场景</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>ZeRO-1</strong></td>
      <td>仅优化器状态</td>
      <td>$2/3$  → 1/N 优化器状态</td>
      <td>1x（同DDP）</td>
      <td>适合大 batch、优化器状态占比大的场景</td>
    </tr>
    <tr>
      <td><strong>ZeRO-2</strong></td>
      <td>优化器状态 + 梯度</td>
      <td>$2/3 + 1/3$ → 各分到 1/N</td>
      <td>1.5x</td>
      <td>适合中等规模模型</td>
    </tr>
    <tr>
      <td><strong>ZeRO-3</strong></td>
      <td>优化器状态 + 梯度 + 参数</td>
      <td>全部分片到 1/N</td>
      <td>3x</td>
      <td>适合极大模型（&gt;参数量/GPU数）</td>
    </tr>
  </tbody>
</table>

<p><strong>ZeRO-1 详细原理：</strong></p>
<ul>
  <li>Adam 的 m 和 v 按 GPU 数 N 分片，每卡只存 1/N 的 m/v</li>
  <li>forward/backward 仍用完整参数（每卡都有） → 通信量不变</li>
  <li>optimizer step 后用 reduce-scatter 聚合更新后的参数 → 通信量同 DDP</li>
</ul>

<p><strong>ZeRO-2 详细原理：</strong></p>
<ul>
  <li>在 ZeRO-1 基础上，梯度也分片</li>
  <li>backward 后梯度用 reduce-scatter 分发到各卡 → 每卡只存 1/N 的梯度</li>
  <li>optimizer step 只更新自己负责的参数分片 → 不需要全参数的梯度</li>
  <li>通信量：reduce-scatter 比 all-reduce 略多（约 1.5x DDP）</li>
</ul>

<p><strong>ZeRO-3 详细原理：</strong></p>
<ul>
  <li>参数也分片 → 每卡只存 1/N 的模型参数</li>
  <li>forward 时：需要完整参数 → <strong>all-gather</strong> 聚合参数 → forward 后立即丢弃非本卡的参数</li>
  <li>backward 时：同样 all-gather 参数 → 计算梯度 → reduce-scatter 分发梯度</li>
  <li>optimizer step：只更新本卡负责的参数分片</li>
  <li>通信量：每次 forward + backward 都要 all-gather + reduce-scatter → <strong>约 3x DDP</strong></li>
</ul>

<p><strong>显存节省公式（N 个 GPU）：</strong></p>
<ul>
  <li>ZeRO-1：$\text{Mem} = \frac{2}{N} \cdot \text{Params}<em>{\text{fp32}} + 2 \cdot \text{Params}</em>{\text{bf16}}$</li>
  <li>ZeRO-2：$\text{Mem} = \frac{2}{N} \cdot \text{Params}<em>{\text{fp32}} + \frac{1}{N} \cdot \text{Params}</em>{\text{bf16}} + \text{Params}_{\text{bf16}}$</li>
  <li>ZeRO-3：$\text{Mem} = \frac{2 + 2 + 1}{N} \cdot \text{Params} = \frac{5}{N} \cdot \text{Params}$</li>
</ul>

<p>⚠️ 常见坑：以为 ZeRO-3 是万能方案 → ZeRO-3 的 all-gather 通信量很大，在小模型（&lt;1B）上通信开销可能超过计算开销 → 反而比 ZeRO-1/2 更慢。</p>

<p>⚠️ 另一个坑：ZeRO-3 + activation checkpointing → forward 需要 all-gather 参数 → 梯度重算时又要 all-gather → 通信量翻倍。需要配合 ZeRO-Offload 或减少 checkpoint 层数。</p>

<p><strong>与 FSDP 的对比：</strong>
FSDP（PyTorch原生）≈ ZeRO-3，但实现更简洁：</p>
<ul>
  <li>FSDP 用 ShardingStrategy 参数控制分片粒度（FULL_SHARD ≈ ZeRO-3，SHARD_GRAD_OP ≈ ZeRO-2）</li>
  <li>FSDP 更好地与 PyTorch 原生混合精度、activation checkpointing 整合</li>
</ul>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>ZeRO-Offload 是什么？（将优化器状态和/或参数 offload 到 CPU 内存 → 进一步节省 GPU 显存，代价是 CPU-GPU 数据传输延迟 → 训练速度下降 2-5x）</li>
  <li>什么时候用 ZeRO-1 vs ZeRO-3？（经验法则：模型参数能装进单卡显存 → ZeRO-1 足够；参数装不进 → ZeRO-3；梯度也装不进 → ZeRO-3 + offload）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/1910.02054">ZeRO: Memory Optimizations Toward Training Trillion Parameter Models</a></li>
  <li><a href="https://arxiv.org/abs/2304.13377">PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel</a></li>
</ul>

<hr />

<h3 id="q47-训练显存估算公式">Q47 训练显存估算公式？</h3>

<p><strong>难度：</strong> 基础</p>

<p><strong>考察点：</strong> 能快速估算 LLM 训练的 GPU 显存需求；理解显存的四大组成部分</p>

<p><strong>满分回答：</strong></p>

<p>训练 LLM 的 GPU 显存由四部分组成：<strong>参数 + 优化器状态 + 梯度 + 激活值</strong>。</p>

<p><strong>基本公式：</strong></p>

\[\text{Total Mem} = \text{Model Mem} + \text{Optimizer Mem} + \text{Gradient Mem} + \text{Activation Mem}\]

<p><strong>1. 参数显存：</strong></p>
<ul>
  <li>bf16/fp16：$2 \times P$ bytes（$P$ = 参数量）</li>
  <li>fp32（混合精度训练存一份 fp32 master weight）：$4 \times P$ bytes</li>
  <li>总计：$2P + 4P = 6P$ bytes（混合精度）</li>
</ul>

<p><strong>2. 优化器状态显存（Adam）：</strong></p>
<ul>
  <li>Momentum（fp32）：$4P$ bytes</li>
  <li>Variance（fp32）：$4P$ bytes</li>
  <li>总计：$8P$ bytes</li>
</ul>

<p><strong>3. 梯度显存：</strong></p>
<ul>
  <li>fp32/bf16：$2P$ bytes（同参数精度）</li>
</ul>

<p><strong>4. 激活值显存：</strong></p>
<ul>
  <li>与 batch size $B$、序列长度 $T$、hidden size $D$、层数 $L$ 有关</li>
  <li>估算公式：\(\text{Act Mem} \approx 2 \cdot B \cdot T \cdot D \cdot L \cdot (5 + \frac{24}{D_{\text{head}}} + \frac{T}{2B \cdot D_{\text{head}}})\)</li>
  <li>简化近似（假设 $D_{\text{head}}=128$）：\(\text{Act Mem} \approx 2 \cdot B \cdot T \cdot D \cdot L \cdot 7 \text{ bytes}\)</li>
  <li>加 gradient checkpointing → 只存每层的输入而非所有中间激活 → 激活值降至约 $\frac{1}{\sqrt{L}}$</li>
</ul>

<p><strong>快速估算表（不含激活值）：</strong></p>

<table>
  <thead>
    <tr>
      <th>模型大小</th>
      <th>参数</th>
      <th>优化器</th>
      <th>梯度</th>
      <th>总计（不含激活）</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>7B</td>
      <td>42GB</td>
      <td>56GB</td>
      <td>14GB</td>
      <td><strong>112GB</strong></td>
    </tr>
    <tr>
      <td>13B</td>
      <td>78GB</td>
      <td>104GB</td>
      <td>26GB</td>
      <td><strong>208GB</strong></td>
    </tr>
    <tr>
      <td>70B</td>
      <td>420GB</td>
      <td>560GB</td>
      <td>140GB</td>
      <td><strong>1120GB</strong></td>
    </tr>
  </tbody>
</table>

<p>注：以上为混合精度（bf16+fp32 master weight）单卡需求。</p>

<p><strong>常见配置估算：</strong></p>

<table>
  <thead>
    <tr>
      <th>配置</th>
      <th>模型</th>
      <th>显存需求</th>
      <th>等价 GPU 数</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>LoRA bf16 7B</td>
      <td>~参数+梯度(LoRA部分)+优化器</td>
      <td>~16-20GB</td>
      <td>1×A100 40GB</td>
    </tr>
    <tr>
      <td>全参数 bf16 7B</td>
      <td>~112GB（不含激活）</td>
      <td>需要 ZeRO-3 + 多卡</td>
      <td>2×A100 80GB</td>
    </tr>
    <tr>
      <td>全参数 bf16 70B</td>
      <td>~1120GB</td>
      <td>需要 ZeRO-3 + 8×A100 80GB</td>
      <td>8-16×A100 80GB</td>
    </tr>
  </tbody>
</table>

<p><strong>LoRA 显存估算：</strong></p>
<ul>
  <li>LoRA 可训练参数 $P_{\text{LoRA}} = 2 \cdot r \cdot \sum d_i$（$d_i$ 为每层原始维度）</li>
  <li>优化器状态只存 LoRA 参数 → 远小于全参数</li>
  <li>基础模型参数冻结 → 不存梯度</li>
  <li>总显存 ≈ 模型参数（bf16）+ LoRA 参数的优化器状态 + LoRA 梯度 + 激活值</li>
  <li>约为 $2P + 8P_{\text{LoRA}} + 2P_{\text{LoRA}} + \text{Act}$ → 远小于全参数训练</li>
</ul>

<p>⚠️ 常见坑：只算参数显存 → 忽略优化器状态（Adam 占 2x 参数量）和激活值 → 实际需求远超预期。7B 模型参数只有 14GB（bf16），但混合精度训练需要 112GB（不含激活）。</p>

<p>⚠️ 另一个坑：以为 LoRA 训练只需要模型参数大小 → 还需要激活值！LoRA 7B 在 A100 40GB 上 batch=1 seq=2048 可能刚好够，batch=4 就 OOM。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>gradient checkpointing 节省多少激活显存？（约节省 60-70% 激活显存，代价是增加 30-40% 计算时间——需要重新 forward 计算被丢弃的中间激活）</li>
  <li>bf16 vs fp32 训练显存差多少？（bf16 参数 2 bytes vs fp32 4 bytes → 参数显存减半；但混合精度训练仍需要 fp32 master weight → 总显存节省约 30%，不是 50%）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/1910.02054">ZeRO: Memory Optimizations Toward Training Trillion Parameter Models</a></li>
  <li><a href="https://arxiv.org/abs/1710.03754">Mixed Precision Training</a></li>
</ul>

<hr />

<h2 id="十二评估">十二、评估</h2>

<h3 id="q48-mt-bench--alpacaeval--arena-hard-评测方法">Q48 MT-Bench / AlpacaEval / Arena-hard 评测方法？</h3>

<p><strong>难度：</strong> 进阶</p>

<p><strong>考察点：</strong> 三种主流对齐评测方法的原理与局限；理解 LLM-as-Judge 的评判机制与偏差</p>

<p><strong>满分回答：</strong></p>

<p><strong>MT-Bench</strong>（Zheng et al., 2023）：Multi-Turn Benchmark，评测多轮对话能力。</p>

<ul>
  <li>80 条精心设计的 multi-turn 问题，覆盖 8 个类别：写作、角色扮演、推理、数学、编码、信息提取、STEM、人文</li>
  <li>每个问题包含 2 个 turn：第一轮提出问题，第二轮追问深入</li>
  <li><strong>评判方式</strong>：GPT-4 作为 judge，对模型回答与 reference 做 pairwise 比较</li>
  <li>打分方式：模型回答 vs 另一模型回答 → GPT-4 选出更好的一方（pairwise）</li>
  <li>最终分数：模型在所有问题上的胜率（win rate）</li>
</ul>

\[\text{MT-Bench Score} = \frac{\text{\# wins} + 0.5 \times \text{\# ties}}{\text{total questions}}\]

<p><strong>AlpacaEval</strong>（Dubois et al., 2024）：基于指令跟随的自动评测。</p>

<ul>
  <li>805 条 AlpacaEval 2.0 指令（从 AlpacaFarm 扩展）</li>
  <li><strong>评判方式</strong>：LC（Length-Controlled）win rate → 用 GPT-4/AutoJ 做 pairwise 评判，但控制长度偏差</li>
  <li>关键改进：原始 AlpacaEval 的 judge 偏好更长回答 → LC win rate 通过回归校正长度偏差</li>
  <li>输出：胜率 + 平均回答长度 + LC 胜率</li>
  <li>优点：自动化、快速、可复现</li>
  <li>缺点：只能评测单轮指令跟随，不测多轮/推理</li>
</ul>

<p><strong>Arena-Hard</strong>（来自 Chatbot Arena 团队）：高难度竞技评测。</p>

<ul>
  <li>500 条高难度 prompt（从 Chatbot Arena 用户的真实请求中筛选难度最高的）</li>
  <li>评判方式：与 GPT-4-0314 做 pairwise 比较 → 计算胜率</li>
  <li>更侧重推理、编码、数学等硬能力</li>
  <li>Arena-Hard 胜率与 Chatbot Arena Elo 排名高度相关（$r &gt; 0.98$）</li>
</ul>

<table>
  <thead>
    <tr>
      <th>评测</th>
      <th>题目数</th>
      <th>评判方式</th>
      <th>覆盖范围</th>
      <th>偏差来源</th>
      <th>与 Arena Elo 相关性</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>MT-Bench</td>
      <td>80</td>
      <td>GPT-4 pairwise judge</td>
      <td>多轮对话（8类）</td>
      <td>GPT-4 自偏好（倾向自己风格）</td>
      <td>~0.95</td>
    </tr>
    <tr>
      <td>AlpacaEval 2.0</td>
      <td>805</td>
      <td>LC pairwise judge</td>
      <td>单轮指令跟随</td>
      <td>长度偏差（已校正）</td>
      <td>~0.92</td>
    </tr>
    <tr>
      <td>Arena-Hard</td>
      <td>500</td>
      <td>GPT-4-0314 pairwise</td>
      <td>高难度推理/编码</td>
      <td>仅比 GPT-4-0314</td>
      <td>~0.98</td>
    </tr>
  </tbody>
</table>

<p><strong>LLM-as-Judge 的三大偏差：</strong></p>

<ol>
  <li><strong>位置偏差</strong>（Position Bias）：judge 偏好放在前面的回答 → 解决：交换位置重评取平均</li>
  <li><strong>长度偏差</strong>（Verbosity Bias）：judge 偏好更长回答 → 解决：LC 校正或 prompt judge 忽略长度</li>
  <li><strong>自偏好偏差</strong>（Self-Preference Bias）：GPT-4 judge 偏好 GPT-4 风格的回答 → 解决：用不同模型做 judge 或人类交叉验证</li>
</ol>

<p>⚠️ 常见坑：MT-Bench 80 条题太少 → 胜率在 ±5% 范围内可能只是噪声 → 不能作为唯一评测，需要配合更多评测。</p>

<p>⚠️ 另一个坑：AlpacaEval 原始 win rate 有严重长度偏差 → 必须用 LC win rate（AlpacaEval 2.0 已默认）。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>Chatbot Arena 的 Elo 评分为什么是最可靠的对齐评测？（因为它是真人 pairwise 比较，不依赖 LLM judge → 无 judge 偏差；且数据量大（&gt;100K 比较）→ 统计可靠）</li>
  <li>如何消除 LLM-as-Judge 的偏差？（交换位置重评 + 多 judge 投票 + 长度校正 + 人类抽样验证）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2306.05685">Judging LLM-as-a-Judge with MT-Bench and Chatbot Arena</a></li>
  <li><a href="https://arxiv.org/abs/2405.04664">AlpacaEval: An Automatic Evaluator of Instruction-following Models</a></li>
  <li><a href="https://arxiv.org/abs/2403.04132">Chatbot Arena: An Open Platform for Evaluating LLMs by Human Preference</a></li>
</ul>

<hr />

<h3 id="q49-reward-model-评测指标与方法">Q49 Reward Model 评测指标与方法？</h3>

<p><strong>难度：</strong> 进阶</p>

<p><strong>考察点：</strong> RM 评测的特有指标（准确率、一致性、校准度）；理解 RM overoptimization 问题</p>

<p><strong>满分回答：</strong></p>

<p>Reward Model（RM）的评测与普通模型不同——RM 输出的是 scalar score，评测核心是”RM 能否正确区分 chosen vs rejected”。</p>

<p><strong>核心评测指标：</strong></p>

<table>
  <thead>
    <tr>
      <th>指标</th>
      <th>定义</th>
      <th>含义</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>Accuracy</strong></td>
      <td>$\frac{\text{# correct pairs}}{\text{total pairs}}$</td>
      <td>chosen score &gt; rejected score 的比例</td>
    </tr>
    <tr>
      <td><strong>Concordance</strong></td>
      <td>RM 排序与人类排序一致的比例</td>
      <td>更细粒度——多回答排序的一致性</td>
    </tr>
    <tr>
      <td><strong>Calibration</strong></td>
      <td>RM score 与人类真实偏好分数的相关性</td>
      <td>RM score 是否反映真实的偏好强度</td>
    </tr>
    <tr>
      <td><strong>Separation</strong></td>
      <td>chosen 与 rejected 的 score 差距分布</td>
      <td>$\Delta = r_{\text{chosen}} - r_{\text{rejected}}$ 的均值和分布</td>
    </tr>
  </tbody>
</table>

<p><strong>Accuracy 的局限：</strong></p>
<ul>
  <li>Accuracy 只看”谁更高”，不看”高多少” → 一个 RM 给 chosen=5.0, rejected=4.99（accuracy 正确但 separation 极小）vs 另一个给 chosen=5.0, rejected=1.0（accuracy 正确且 separation 大）</li>
  <li>→ 需要配合 separation 和 calibration 一起看</li>
</ul>

<p><strong>评测数据集：</strong></p>
<ul>
  <li><strong>RLHF 原始偏好数据</strong>：用训练 RM 时未用过的 held-out 偏好数据做 accuracy 评测</li>
  <li><strong>Helpful and Harmless 数据集</strong>（Anthropic）：专门标注的 chosen/rejected pair</li>
  <li><strong>自构造数据</strong>：用同一 prompt 生成多个回答，让人类排序，再评测 RM 的排序一致性</li>
</ul>

<p><strong>RM Overoptimization 问题：</strong></p>

<p>RM overoptimization（reward hacking）指 PPO 训练中模型找到了 RM 的”漏洞”——生成 RM 给高分但人类认为不好的回答。</p>

<p><strong>现象：</strong></p>
<ul>
  <li>训练初期 reward 上升 + 人类评测也上升 → 正常优化</li>
  <li>训练中期 reward 继续上升但人类评测开始下降 → overoptimization</li>
  <li>分界点就是 KL divergence 达到某个阈值</li>
</ul>

<p><strong>Gold RM vs Proxy RM 评测：</strong></p>
<ul>
  <li><strong>Proxy RM</strong>：训练 PPO 用的 RM → reward 持续上升</li>
  <li><strong>Gold RM</strong>：大量人类偏好数据训练的更可靠 RM → 用于检测 overoptimization</li>
  <li>当 proxy RM reward 上升但 gold RM reward 下降 → overoptimization 发生</li>
</ul>

<p><strong>量化 overoptimization：</strong>
\(\text{Overoptimization onset} \approx D_{\text{KL}}(\pi_{\text{PPO}}, \pi_{\text{ref}}) \approx 5-10 \text{ nats}\)</p>

<p>超过这个 KL 阈值后，proxy reward 与真实人类偏好的相关性急剧下降。</p>

<p>⚠️ 常见坑：只用 proxy RM 的 reward 监控 PPO → 发现 reward 持续上升以为训练正常 → 实际已经 overoptimization → 需要用 gold RM 或人类抽样做交叉验证。</p>

<p>⚠️ 另一个坑：RM accuracy 95% 但 calibration 很差 → RM 能判断谁更好但 score 的绝对值无意义 → PPO 训练中 reward scale 不稳定 → 需要做 reward normalization。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>如何缓解 RM overoptimization？（加大 KL penalty；用多个 RM 做 ensemble reward；定期人类评测抽查；early stopping）</li>
  <li>为什么多个 RM ensemble 能缓解 overoptimization？（单个 RM 有漏洞 → 模型找到漏洞 exploit；多个 RM 同时给分 → 不同 RM 的漏洞不同 → 模型难以同时 exploit 所有 RM）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2306.05685">Judging LLM-as-a-Judge with MT-Bench and Chatbot Arena</a></li>
  <li><a href="https://arxiv.org/abs/2310.12048">Reward Model Overoptimization (社区分析)</a></li>
  <li><a href="https://arxiv.org/abs/2204.05862">Training a Helpful and Harmless Assistant with RLHF</a></li>
</ul>

<hr />

<h2 id="十三前沿与开放题">十三、前沿与开放题</h2>

<h3 id="q50-推理模型o1r1的后训练与-agent-training-的前沿方向">Q50 推理模型(o1/R1)的后训练与 Agent training 的前沿方向？</h3>

<p><strong>难度：</strong> 专家</p>

<p><strong>考察点：</strong> 推理模型后训练的核心技术路线；理解从 RLHF 到 reasoning RL 的范式跃迁</p>

<p><strong>满分回答：</strong></p>

<p><strong>推理模型的后训练范式变革：</strong></p>

<p>传统后训练路线：SFT → RLHF/DPO → 对齐模型（擅长指令跟随但不擅长复杂推理）
推理模型路线：SFT → <strong>Reasoning RL</strong> → 推理模型（擅长长链推理、自我验证、纠错）</p>

<p><strong>OpenAI o1 的后训练（推测）：</strong></p>
<ul>
  <li>OpenAI 未公开 o1 的完整训练细节，但社区分析推测其核心流程：
    <ol>
      <li><strong>大规模推理数据 SFT</strong>：用 PRM（Process Reward Model）标注的高质量推理链做 SFT → 让模型学会”一步一步推理”的格式</li>
      <li><strong>Reasoning RL（类似 PPO 但 reward 是推理正确性）</strong>：</li>
    </ol>
    <ul>
      <li>Reward 来源：<strong>结果验证</strong>（数学题答案对/错、代码 pass/fail）而非人类偏好</li>
      <li>模型在 RL 中学习：更多推理步骤 → 更高推理正确率 → reward 更高 → 生成更长推理链</li>
      <li>这解释了 o1 的”thinking time”现象——模型学会用更多 token 做推理
        <ol>
          <li><strong>Test-time compute scaling</strong>：推理时通过 search（如 beam search / MCTS）生成多条推理路径 → 选最优 → 相当于”推理时做更多计算”</li>
        </ol>
      </li>
    </ul>
  </li>
</ul>

<p><strong>DeepSeek-R1 的后训练（公开）：</strong></p>

<table>
  <thead>
    <tr>
      <th>鞞段</th>
      <th>方法</th>
      <th>数据</th>
      <th>目标</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>Stage 1</td>
      <td><strong>纯 GRPO RL</strong>（无 SFT）</td>
      <td>数学/代码等可验证任务</td>
      <td>让模型自发涌现推理行为（reasoning emergence）</td>
    </tr>
    <tr>
      <td>Stage 2</td>
      <td><strong>拒绝采样 + SFT</strong></td>
      <td>RL 产生的优质推理链 + 通用 SFT 数据</td>
      <td>稳固推理格式 + 保持通用能力</td>
    </tr>
    <tr>
      <td>Stage 3</td>
      <td><strong>全场景 GRPO RL</strong></td>
      <td>数学/代码/对话/创意写作</td>
      <td>对齐推理能力到所有场景</td>
    </tr>
  </tbody>
</table>

<p>R1 的关键发现：<strong>纯 RL 可以让模型自发涌现推理行为</strong>——</p>
<ul>
  <li>不需要 SFT 先教模型”怎么推理” → 直接给可验证 reward（答案对 = +1）</li>
  <li>模型在 RL 中自然学会：先推理 → 再得出答案 → 答案对 → reward 高 → 强化推理行为</li>
  <li>涌现过程：初期模型直接输出答案（低 reward） → 尝试加推理步骤（reward 上升） → 逐步学会长推理链</li>
</ul>

<p><strong>GRPO vs PPO：</strong>
\(\text{GRPO}: r_i = \frac{r(x, y_i) - \mu_{\text{group}}}{\sigma_{\text{group}}}\)</p>
<ul>
  <li>GRPO 不需要 value network → 节省一个模型参数和训练成本</li>
  <li>用同一 prompt 生成一组回答 → 组内 reward 做归一化 → 自然形成相对偏好</li>
  <li>比 PPO 更简单，但效果相近（DeepSeekMath 实验证实）</li>
</ul>

<p><strong>Agent training 的前沿方向：</strong></p>

<ol>
  <li><strong>Tool-use RL</strong>：让模型在 RL 中学习何时调用工具（搜索、代码执行、计算器）→ reward = 最终答案正确性</li>
  <li><strong>Multi-step RL</strong>：将 agent 的多步行为轨迹视为一条完整策略 → reward 只在最终步骤给出 → 模型学习规划</li>
  <li><strong>Environment-grounded RL</strong>：模型在真实/模拟环境中执行任务 → 根据环境反馈得到 reward</li>
  <li><strong>Self-play RL</strong>：两个 agent 对弈/协作 → reward 来自博弈结果 → 学习策略性推理</li>
</ol>

<table>
  <thead>
    <tr>
      <th>方向</th>
      <th>代表工作</th>
      <th>Reward 来源</th>
      <th>核心挑战</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>Reasoning RL</td>
      <td>o1/R1/GRPO</td>
      <td>可验证任务正确性</td>
      <td>非可验证任务（写作/创意）的 reward 设计</td>
    </tr>
    <tr>
      <td>Tool-use RL</td>
      <td>Toolformer/ReAct</td>
      <td>任务完成率</td>
      <td>工具调用的延迟与错误处理</td>
    </tr>
    <tr>
      <td>Multi-step Agent RL</td>
      <td>Voyager/AgentTrek</td>
      <td>长程目标完成</td>
      <td>sparse reward + 郶分规划困难</td>
    </tr>
    <tr>
      <td>Self-play</td>
      <td>各种博弈/谈判</td>
      <td>博弈结果</td>
      <td>对手策略变化 → reward 非平稳</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 常见坑：以为 o1/R1 只是”加了更多推理步骤的模型” → 实际是 RL 训练范式的根本改变：reward 从”人类偏好”变为”任务正确性”，这才是推理能力涌现的关键。</p>

<p>⚠️ 另一个坑：GRPO 不需要 value model → 以为训练更简单 → 实际 GRPO 的组内采样效率较低（每个 prompt 要生成 G 个回答才能计算 reward），且 reward 归一化在组大小 G 较小时不稳定。</p>

<p><strong>延伸追问：</strong></p>
<ul>
  <li>为什么 R1 的纯 RL 可以涌现推理而 SFT 不能？（RL 的 exploration 让模型尝试各种策略，其中”先推理再回答”碰巧 reward 高 → 被强化；SFT 只教模型模仿固定格式的推理链 → 缺乏探索和自我发现）</li>
  <li>如何将 reasoning RL 扩展到非可验证任务（如创意写作）？（用 LLM-as-Judge 做 reward → 但引入 judge 偏差；或用人类偏好做 reward → 成本高；目前这个方向仍是开放问题）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2501.12948">DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via RL</a></li>
  <li><a href="https://arxiv.org/abs/2402.03300">DeepSeekMath: Pushing the Limits of Mathematical Reasoning (GRPO)</a></li>
  <li><a href="https://arxiv.org/abs/2305.20050">Let’s Verify Step by Step (PRM)</a></li>
  <li><a href="https://openai.com/index/learning-to-reason-with-llms/">OpenAI o1 Blog</a></li>
</ul>

<hr />

<hr />

<h2 id="十四手撕代码--coding">十四、手撕代码 / Coding</h2>

<h3 id="q51-multi-head-self-attention含-causal-mask从零实现pytorch">Q51 Multi-Head Self-Attention（含 causal mask）从零实现（PyTorch）</h3>

<p><strong>难度：</strong> 进阶</p>

<p><strong>考察点：</strong> MHA 的完整实现，包括 QKV 投影、scaled dot-product、causal mask 的数值稳定性写法</p>

<p><strong>满分回答：</strong></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>
<span class="kn">import</span> <span class="nn">math</span>

<span class="k">class</span> <span class="nc">MultiHeadSelfAttention</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">dropout</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="k">assert</span> <span class="n">d_model</span> <span class="o">%</span> <span class="n">n_heads</span> <span class="o">==</span> <span class="mi">0</span><span class="p">,</span> <span class="s">"d_model must be divisible by n_heads"</span>
        
        <span class="bp">self</span><span class="p">.</span><span class="n">d_model</span> <span class="o">=</span> <span class="n">d_model</span>      <span class="c1"># (D)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span> <span class="o">=</span> <span class="n">n_heads</span>      <span class="c1"># (H)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span> <span class="o">=</span> <span class="n">d_model</span> <span class="o">//</span> <span class="n">n_heads</span>  <span class="c1"># (D_h)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">scale</span> <span class="o">=</span> <span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">)</span>
        
        <span class="c1"># QKV 投影：合并为一次矩阵乘法以提高效率
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">qkv_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="mi">3</span> <span class="o">*</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (D) -&gt; (3D)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">out_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>       <span class="c1"># (D) -&gt; (D)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">attn_drop</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">resid_drop</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span> <span class="n">mask</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span> <span class="o">=</span> <span class="bp">None</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="s">"""
        Args:
            x:    (B, T, D) 辧入序列
            mask: (T, T) or (B, 1, T, T) causal mask, 0 = block, 1 = allow
        Returns:
            out:  (B, T, D) 辧出序列
        """</span>
        <span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">D</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">shape</span>
        
        <span class="c1"># Step 1: QKV 投影 → (B, T, 3D) → 拆分 → (B, T, D) each
</span>        <span class="n">qkv</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">qkv_proj</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>                     <span class="c1"># (B, T, 3D)
</span>        <span class="n">q</span><span class="p">,</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span> <span class="o">=</span> <span class="n">qkv</span><span class="p">.</span><span class="n">chunk</span><span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>             <span class="c1"># each (B, T, D)
</span>        
        <span class="c1"># Step 2: reshape 为 multi-head → (B, H, T, D_h)
</span>        <span class="n">q</span> <span class="o">=</span> <span class="n">q</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, H, T, D_h)
</span>        <span class="n">k</span> <span class="o">=</span> <span class="n">k</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, H, T, D_h)
</span>        <span class="n">v</span> <span class="o">=</span> <span class="n">v</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, H, T, D_h)
</span>        
        <span class="c1"># Step 3: Scaled Dot-Product Attention
</span>        <span class="c1"># attn_scores = Q @ K^T / sqrt(d_head)
</span>        <span class="n">attn_scores</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">))</span> <span class="o">/</span> <span class="bp">self</span><span class="p">.</span><span class="n">scale</span>  <span class="c1"># (B, H, T, T)
</span>        
        <span class="c1"># Step 4: Apply causal mask
</span>        <span class="c1"># 数值稳定性关键：用 float('-inf') 而不是 -1e9
</span>        <span class="c1"># -1e9 在 softmax 后仍然有微小概率泄露 → 可能影响梯度
</span>        <span class="c1"># float('-inf') → softmax 后严格为 0
</span>        <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">attn_scores</span> <span class="o">=</span> <span class="n">attn_scores</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span> <span class="o">==</span> <span class="mi">0</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>  <span class="c1"># (B, H, T, T)
</span>        <span class="k">else</span><span class="p">:</span>
            <span class="c1"># 默认 causal mask: 只能看当前位置及之前
</span>            <span class="n">causal_mask</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tril</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">T</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">device</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="nb">bool</span><span class="p">))</span>  <span class="c1"># (T, T)
</span>            <span class="n">attn_scores</span> <span class="o">=</span> <span class="n">attn_scores</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="o">~</span><span class="n">causal_mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
        
        <span class="c1"># Step 5: Softmax → attention weights
</span>        <span class="n">attn_weights</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">attn_scores</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B, H, T, T)
</span>        <span class="n">attn_weights</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">attn_drop</span><span class="p">(</span><span class="n">attn_weights</span><span class="p">)</span>
        
        <span class="c1"># Step 6: Weighted sum of V
</span>        <span class="n">attn_out</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">attn_weights</span><span class="p">,</span> <span class="n">v</span><span class="p">)</span>        <span class="c1"># (B, H, T, D_h)
</span>        
        <span class="c1"># Step 7: Concatenate heads → (B, T, D)
</span>        <span class="n">attn_out</span> <span class="o">=</span> <span class="n">attn_out</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">).</span><span class="n">contiguous</span><span class="p">().</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">D</span><span class="p">)</span>  <span class="c1"># (B, T, D)
</span>        
        <span class="c1"># Step 8: Output projection
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">out_proj</span><span class="p">(</span><span class="n">attn_out</span><span class="p">)</span>                    <span class="c1"># (B, T, D)
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">resid_drop</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
        
        <span class="k">return</span> <span class="n">out</span>

<span class="c1"># === 验证 ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">d_model</span> <span class="o">=</span> <span class="mi">512</span>
    <span class="n">n_heads</span> <span class="o">=</span> <span class="mi">8</span>
    <span class="n">B</span><span class="p">,</span> <span class="n">T</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">16</span>
    
    <span class="n">mha</span> <span class="o">=</span> <span class="n">MultiHeadSelfAttention</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">)</span>
    <span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">d_model</span><span class="p">)</span>
    <span class="n">out</span> <span class="o">=</span> <span class="n">mha</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># 无显式 mask，内部生成 causal mask
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Input shape:  </span><span class="si">{</span><span class="n">x</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>   <span class="c1"># (2, 16, 512)
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Output shape: </span><span class="si">{</span><span class="n">out</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 16, 512)
</span>    
    <span class="c1"># 验证 causal mask 效果：位置 5 不应关注位置 6-15
</span>    <span class="c1"># 构造显式 mask 测试
</span>    <span class="n">mask</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tril</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">T</span><span class="p">,</span> <span class="n">T</span><span class="p">))</span>  <span class="c1"># (T, T) lower triangular
</span>    <span class="n">out_masked</span> <span class="o">=</span> <span class="n">mha</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">mask</span><span class="o">=</span><span class="n">mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">))</span>  <span class="c1"># (1, T, T) -&gt; broadcast
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Masked output shape: </span><span class="si">{</span><span class="n">out_masked</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 16, 512)
</span></code></pre></div></div>

<p><strong>关键点解析：</strong></p>

<ol>
  <li><strong>Shape 流转</strong>：<code class="language-plaintext highlighter-rouge">(B,T,D)</code> → QKV <code class="language-plaintext highlighter-rouge">(B,T,3D)</code> → split <code class="language-plaintext highlighter-rouge">(B,T,D)</code> → reshape <code class="language-plaintext highlighter-rouge">(B,H,T,D_h)</code> → attn <code class="language-plaintext highlighter-rouge">(B,H,T,T)</code> → weighted sum <code class="language-plaintext highlighter-rouge">(B,H,T,D_h)</code> → concat <code class="language-plaintext highlighter-rouge">(B,T,D)</code> → output <code class="language-plaintext highlighter-rouge">(B,T,D)</code></li>
  <li><strong>数值稳定性</strong>：mask 用 <code class="language-plaintext highlighter-rouge">float('-inf')</code> 而非 <code class="language-plaintext highlighter-rouge">-1e9</code>。<code class="language-plaintext highlighter-rouge">-1e9</code> 在 softmax 后 ≈ $e^{-10^9}$ 仍有微小值 → FP16 下可能不精确归零 → 影响梯度计算。<code class="language-plaintext highlighter-rouge">float('-inf')</code> → softmax 后严格 0。</li>
  <li><strong>QKV 合并投影</strong>：一次 <code class="language-plaintext highlighter-rouge">Linear(D, 3D)</code> 比 三次 <code class="language-plaintext highlighter-rouge">Linear(D, D)</code> 效率高——单次 matmul vs 三次，减少 kernel launch overhead。</li>
  <li><strong>contiguous()</strong>：<code class="language-plaintext highlighter-rouge">transpose(1,2)</code> 后 tensor 不 contiguous → <code class="language-plaintext highlighter-rouge">.view()</code> 会报错 → 必须先 <code class="language-plaintext highlighter-rouge">.contiguous()</code> 再 <code class="language-plaintext highlighter-rouge">.view()</code>。也可用 <code class="language-plaintext highlighter-rouge">.reshape()</code> 替代（自动处理 contiguous 问题）。</li>
</ol>

<p><strong>复杂度分析：</strong></p>
<ul>
  <li>时间：$O(B \cdot T^2 \cdot D)$（attention 矩阵计算是瓶颈）</li>
  <li>显存：$O(B \cdot T^2 \cdot H + B \cdot T \cdot D)$（attention weights + QKV activations）</li>
</ul>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>为什么 QKV 投影合并比分开三次更高效？（一次大 matmul 比 3 次小 matmul 更好利用 GPU 并行性）</li>
  <li>causal mask 为什么不能用 <code class="language-plaintext highlighter-rouge">-1e9</code>？（FP16 精度下 <code class="language-plaintext highlighter-rouge">-1e9</code> 在 softmax 中可能不完全归零 → 影响数值稳定性和梯度）</li>
  <li>如何实现 flash attention？（核心：不在显存中存完整的 $T \times T$ attention 矩阵 → online softmax + tiling → 显存从 $O(T^2)$ 降至 $O(T)$）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/1706.03762">Attention Is All You Need (Transformer 原论文)</a></li>
  <li><a href="https://arxiv.org/abs/2205.14135">FlashAttention: Fast and Memory-Efficient Exact Attention</a></li>
</ul>

<hr />

<h3 id="q52-scaled-dot-product-attention--kv-cache-增量实现">Q52 Scaled Dot-Product Attention + KV Cache 增量实现</h3>

<p><strong>难度：</strong> 专家</p>

<p><strong>考察点：</strong> KV Cache 的增量推理机制；理解 pre-fill vs decode 阶段的区别与 cache 管理</p>

<p><strong>满分回答：</strong></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>
<span class="kn">import</span> <span class="nn">math</span>

<span class="k">class</span> <span class="nc">AttentionWithKVCache</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">:</span> <span class="nb">int</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="k">assert</span> <span class="n">d_model</span> <span class="o">%</span> <span class="n">n_heads</span> <span class="o">==</span> <span class="mi">0</span>
        
        <span class="bp">self</span><span class="p">.</span><span class="n">d_model</span> <span class="o">=</span> <span class="n">d_model</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span> <span class="o">=</span> <span class="n">n_heads</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span> <span class="o">=</span> <span class="n">d_model</span> <span class="o">//</span> <span class="n">n_heads</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">scale</span> <span class="o">=</span> <span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">)</span>
        
        <span class="bp">self</span><span class="p">.</span><span class="n">q_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (D) -&gt; (D)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">k_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (D) -&gt; (D)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">v_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (D) -&gt; (D)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">o_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (D) -&gt; (D)
</span>    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span>
        <span class="bp">self</span><span class="p">,</span>
        <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>        <span class="c1"># (B, T_q, D) 当前步的 query 辧入
</span>        <span class="n">kv_cache</span><span class="p">:</span> <span class="nb">tuple</span> <span class="o">=</span> <span class="bp">None</span><span class="p">,</span>  <span class="c1"># (k_cache, v_cache) each (B, H, T_kv, D_h)
</span>        <span class="n">start_pos</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">0</span><span class="p">,</span>      <span class="c1"># 当前 token 在序列中的位置（用于 causal mask）
</span>    <span class="p">)</span> <span class="o">-&gt;</span> <span class="nb">tuple</span><span class="p">:</span>
        <span class="s">"""
        Args:
            x:         (B, T_q, D) — decode 阞段 T_q=1，pre-fill 阞段 T_q&gt;1
            kv_cache:  之前的 K/V 缓存，None 表示首次
            start_pos: 当前 query 的起始位置索引
        Returns:
            out:       (B, T_q, D) attention 辧出
            new_cache: (k_new, v_new) 更新后的 KV cache
        """</span>
        <span class="n">B</span><span class="p">,</span> <span class="n">T_q</span><span class="p">,</span> <span class="n">D</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">shape</span>
        
        <span class="c1"># Step 1: QKV 投影
</span>        <span class="n">q</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">q_proj</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># (B, T_q, D)
</span>        <span class="n">k</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">k_proj</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># (B, T_q, D)
</span>        <span class="n">v</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">v_proj</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># (B, T_q, D)
</span>        
        <span class="c1"># Step 2: Reshape 为 multi-head
</span>        <span class="n">q</span> <span class="o">=</span> <span class="n">q</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T_q</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, H, T_q, D_h)
</span>        <span class="n">k</span> <span class="o">=</span> <span class="n">k</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T_q</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, H, T_q, D_h)
</span>        <span class="n">v</span> <span class="o">=</span> <span class="n">v</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T_q</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, H, T_q, D_h)
</span>        
        <span class="c1"># Step 3: KV Cache 拼接
</span>        <span class="k">if</span> <span class="n">kv_cache</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">k_cache</span><span class="p">,</span> <span class="n">v_cache</span> <span class="o">=</span> <span class="n">kv_cache</span>  <span class="c1"># each (B, H, T_prev, D_h)
</span>            <span class="n">k</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">([</span><span class="n">k_cache</span><span class="p">,</span> <span class="n">k</span><span class="p">],</span> <span class="n">dim</span><span class="o">=</span><span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, H, T_prev + T_q, D_h)
</span>            <span class="n">v</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">([</span><span class="n">v_cache</span><span class="p">,</span> <span class="n">v</span><span class="p">],</span> <span class="n">dim</span><span class="o">=</span><span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, H, T_prev + T_q, D_h)
</span>        
        <span class="n">T_kv</span> <span class="o">=</span> <span class="n">k</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span>  <span class="c1"># 总序列长度 = cache + 当前
</span>        <span class="n">new_cache</span> <span class="o">=</span> <span class="p">(</span><span class="n">k</span><span class="p">,</span> <span class="n">v</span><span class="p">)</span>
        
        <span class="c1"># Step 4: Scaled Dot-Product Attention
</span>        <span class="n">attn_scores</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">))</span> <span class="o">/</span> <span class="bp">self</span><span class="p">.</span><span class="n">scale</span>  <span class="c1"># (B, H, T_q, T_kv)
</span>        
        <span class="c1"># Step 5: Causal mask
</span>        <span class="c1"># 关键：decode 阞段只需 block "未来位置"，即 start_pos 之后的 token 不能看之后的
</span>        <span class="c1"># 用位置索引构造 mask
</span>        <span class="n">query_pos</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">arange</span><span class="p">(</span><span class="n">start_pos</span><span class="p">,</span> <span class="n">start_pos</span> <span class="o">+</span> <span class="n">T_q</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">device</span><span class="p">)</span>   <span class="c1"># (T_q)
</span>        <span class="n">kv_pos</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">arange</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">T_kv</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">device</span><span class="p">)</span>                          <span class="c1"># (T_kv)
</span>        <span class="n">causal_mask</span> <span class="o">=</span> <span class="n">query_pos</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span> <span class="o">&gt;=</span> <span class="n">kv_pos</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>               <span class="c1"># (T_q, T_kv)
</span>        <span class="n">causal_mask</span> <span class="o">=</span> <span class="n">causal_mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>                        <span class="c1"># (1, 1, T_q, T_kv)
</span>        
        <span class="n">attn_scores</span> <span class="o">=</span> <span class="n">attn_scores</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="o">~</span><span class="n">causal_mask</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>  <span class="c1"># (B, H, T_q, T_kv)
</span>        
        <span class="c1"># Step 6: Softmax + Weighted sum
</span>        <span class="n">attn_weights</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">attn_scores</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B, H, T_q, T_kv)
</span>        <span class="n">attn_out</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">attn_weights</span><span class="p">,</span> <span class="n">v</span><span class="p">)</span>       <span class="c1"># (B, H, T_q, D_h)
</span>        
        <span class="c1"># Step 7: Concatenate heads + Output projection
</span>        <span class="n">attn_out</span> <span class="o">=</span> <span class="n">attn_out</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">).</span><span class="n">contiguous</span><span class="p">().</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T_q</span><span class="p">,</span> <span class="n">D</span><span class="p">)</span>  <span class="c1"># (B, T_q, D)
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">o_proj</span><span class="p">(</span><span class="n">attn_out</span><span class="p">)</span>  <span class="c1"># (B, T_q, D)
</span>        
        <span class="k">return</span> <span class="n">out</span><span class="p">,</span> <span class="n">new_cache</span>

<span class="c1"># === 验证：逐步推理模拟 ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">d_model</span> <span class="o">=</span> <span class="mi">64</span>
    <span class="n">n_heads</span> <span class="o">=</span> <span class="mi">4</span>
    <span class="n">B</span> <span class="o">=</span> <span class="mi">1</span>
    <span class="n">seq_len</span> <span class="o">=</span> <span class="mi">10</span>
    
    <span class="n">attn</span> <span class="o">=</span> <span class="n">AttentionWithKVCache</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">)</span>
    <span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">seq_len</span><span class="p">,</span> <span class="n">d_model</span><span class="p">)</span>
    
    <span class="c1"># Pre-fill 阞段：一次性处理前 5 个 token
</span>    <span class="n">out_prefill</span><span class="p">,</span> <span class="n">cache</span> <span class="o">=</span> <span class="n">attn</span><span class="p">(</span><span class="n">x</span><span class="p">[:,</span> <span class="p">:</span><span class="mi">5</span><span class="p">,</span> <span class="p">:],</span> <span class="n">kv_cache</span><span class="o">=</span><span class="bp">None</span><span class="p">,</span> <span class="n">start_pos</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Pre-fill output: </span><span class="si">{</span><span class="n">out_prefill</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (1, 5, 64)
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"K cache shape:   </span><span class="si">{</span><span class="n">cache</span><span class="p">[</span><span class="mi">0</span><span class="p">].</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>      <span class="c1"># (1, 4, 5, 16)
</span>    
    <span class="c1"># Decode 阞段：逐 token 生成
</span>    <span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">5</span><span class="p">,</span> <span class="n">seq_len</span><span class="p">):</span>
        <span class="n">token</span> <span class="o">=</span> <span class="n">x</span><span class="p">[:,</span> <span class="n">i</span><span class="p">:</span><span class="n">i</span><span class="o">+</span><span class="mi">1</span><span class="p">,</span> <span class="p">:]</span>  <span class="c1"># (B, 1, D)
</span>        <span class="n">out_decode</span><span class="p">,</span> <span class="n">cache</span> <span class="o">=</span> <span class="n">attn</span><span class="p">(</span><span class="n">token</span><span class="p">,</span> <span class="n">kv_cache</span><span class="o">=</span><span class="n">cache</span><span class="p">,</span> <span class="n">start_pos</span><span class="o">=</span><span class="n">i</span><span class="p">)</span>
        <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Decode step </span><span class="si">{</span><span class="n">i</span><span class="si">}</span><span class="s">: output </span><span class="si">{</span><span class="n">out_decode</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (1, 1, 64)
</span>    
    <span class="c1"># 对比：无 cache 的完整 forward
</span>    <span class="n">mha_full</span> <span class="o">=</span> <span class="n">MultiHeadSelfAttention</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">)</span>  <span class="c1"># 用 Q51 的类
</span>    <span class="c1"># 需要手动将 attn 参数复制到 mha_full（此处省略，逻辑验证为主）
</span></code></pre></div></div>

<p><strong>关键点解析：</strong></p>

<ol>
  <li><strong>Pre-fill vs Decode</strong>：
    <ul>
      <li>Pre-fill（<code class="language-plaintext highlighter-rouge">T_q &gt; 1</code>）：一次性处理所有已知 token → 计算 $T \times T$ attention → 建立 KV cache</li>
      <li>Decode（<code class="language-plaintext highlighter-rouge">T_q = 1</code>）：每次只处理 1 个新 token → 计算 $1 \times T_{kv}$ attention → KV cache 增长</li>
      <li>Decode 的 attention 是 $O(T)$ 而非 $O(T^2)$ → 推理速度大幅提升</li>
    </ul>
  </li>
  <li><strong>Causal mask 的位置索引</strong>：
    <ul>
      <li>Decode 阞段，query 在位置 <code class="language-plaintext highlighter-rouge">start_pos</code> → 只能看 <code class="language-plaintext highlighter-rouge">0 ~ start_pos</code> 的 KV → 用 <code class="language-plaintext highlighter-rouge">query_pos &gt;= kv_pos</code> 构造 mask</li>
      <li>不能用简单的 <code class="language-plaintext highlighter-rouge">tril</code> mask（因为 <code class="language-plaintext highlighter-rouge">T_q</code> 和 <code class="language-plaintext highlighter-rouge">T_kv</code> 维度不同）</li>
    </ul>
  </li>
  <li><strong>KV Cache 管理</strong>：
    <ul>
      <li>每次 decode 后 cache 增长 1 个 token → 最终 cache 长度 = 序列总长度</li>
      <li>需要管理 cache 的最大长度（超过 max_seq_len 要淘汰旧 token）</li>
      <li>PagedAttention（vLLM）：将 KV cache 分页管理 → 不同请求共享物理页 → 减少碎片</li>
    </ul>
  </li>
</ol>

<p>⚠️ 常见 bug：decode 阞段忘记传 <code class="language-plaintext highlighter-rouge">start_pos</code> → causal mask 计算错误 → 模型可能看到”未来”token → 辧出完全错误。</p>

<p>⚠️ 另一个坑：KV cache 在 bf16 下存储 → 模型输出精度略低于 fp32 → 但实践证明 bf16 cache 对推理质量影响极小（&lt;0.1% accuracy drop）。</p>

<p><strong>复杂度分析：</strong></p>
<ul>
  <li>Pre-fill 时间：$O(T^2 \cdot D)$（与无 cache 相同）</li>
  <li>Decode 时间：$O(T \cdot D)$ per token（比无 cache 的 $O(T^2 \cdot D)$ 快 $T$ 倍）</li>
  <li>KV Cache 显存：$O(2 \cdot L \cdot T \cdot D_h \cdot H)$ per layer（$L$ 层，$T$ 序列长度）</li>
</ul>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>KV cache 对推理速度的影响？（$T$ 个 token 逐步 decode：无 cache 总计算 $O(T^3)$，有 cache 总计算 $O(T^2)$ → 速度提升 $O(T)$ 倍）</li>
  <li>如何管理 KV cache 的显存？（PagedAttention：虚拟地址映射 → 按需分配 → 减少浪费；或量化 KV cache 到 8bit/4bit → 50-75% 显存节省）</li>
  <li>MQA/GQA 如何影响 KV cache？（MQA: 所有 head 共享 1 组 KV → cache 大小降至 $1/H$；GQA: 每 $G$ 个 head 共享 1 组 KV → cache 大小降至 $G/H$）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2205.14135">FlashAttention: Fast and Memory-Efficient Exact Attention</a></li>
  <li><a href="https://arxiv.org/abs/2309.06180">vLLM: Efficient Memory Management for LLM Serving</a></li>
</ul>

<hr />

<h3 id="q53-rope旋转位置编码实现--与绝对位置编码对比">Q53 RoPE（旋转位置编码）实现 &amp; 与绝对位置编码对比</h3>

<p><strong>难度：</strong> 专家</p>

<p><strong>考察点：</strong> RoPE 的数学原理与实现；理解 RoPE 为什么能实现相对位置编码且支持长度外推</p>

<p><strong>满分回答：</strong></p>

<p><strong>RoPE 原理：</strong> Rotary Position Embedding（Su et al., 2024）将位置信息编码为旋转矩阵，作用于 Q 和 K 的二维子空间。核心思想：对 Q 和 K 在每对相邻维度上施加位置相关的旋转 → Q·K 的点积自然编码了<strong>相对位置</strong>。</p>

<p>数学公式：
\(q_m = R_{\Theta,m} \cdot q, \quad k_n = R_{\Theta,n} \cdot k\)
\(q_m \cdot k_n = (R_{\Theta,m} \cdot q)^\top (R_{\Theta,n} \cdot k) = q^\top R_{\Theta, n-m} \cdot k\)</p>

<p>即 QK 点积只依赖相对位置 $n - m$，不依赖绝对位置。</p>

<p>旋转矩阵 $R_{\Theta,m}$：
\(R_{\Theta,m} = \begin{pmatrix} \cos m\theta_0 &amp; -\sin m\theta_0 &amp; 0 &amp; 0 &amp; \cdots \\ \sin m\theta_0 &amp; \cos m\theta_0 &amp; 0 &amp; 0 &amp; \cdots \\ 0 &amp; 0 &amp; \cos m\theta_1 &amp; -\sin m\theta_1 &amp; \cdots \\ 0 &amp; 0 &amp; \sin m\theta_1 &amp; \cos m\theta_1 &amp; \cdots \\ \vdots &amp; &amp; &amp; &amp; \ddots \end{pmatrix}\)</p>

<p>其中 $\theta_i = 10000^{-2i/d}$（与 Transformer 原论文的 sinusoidal 频率相同）。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">math</span>

<span class="k">class</span> <span class="nc">RotaryEmbedding</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_head</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">max_seq_len</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">8192</span><span class="p">,</span> <span class="n">base</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">10000.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span> <span class="o">=</span> <span class="n">d_head</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">base</span> <span class="o">=</span> <span class="n">base</span>
        
        <span class="c1"># 计算 theta_i = base^(-2i/d_head)
</span>        <span class="n">inv_freq</span> <span class="o">=</span> <span class="mf">1.0</span> <span class="o">/</span> <span class="p">(</span><span class="n">base</span> <span class="o">**</span> <span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">arange</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">d_head</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span> <span class="o">/</span> <span class="n">d_head</span><span class="p">))</span>  <span class="c1"># (D_h/2)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">register_buffer</span><span class="p">(</span><span class="s">'inv_freq'</span><span class="p">,</span> <span class="n">inv_freq</span><span class="p">,</span> <span class="n">persistent</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        
        <span class="c1"># 预计算所有位置的 cos/sin（可优化为动态计算）
</span>        <span class="n">t</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">arange</span><span class="p">(</span><span class="n">max_seq_len</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>  <span class="c1"># (T_max)
</span>        <span class="n">freqs</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">outer</span><span class="p">(</span><span class="n">t</span><span class="p">,</span> <span class="n">inv_freq</span><span class="p">)</span>  <span class="c1"># (T_max, D_h/2)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">register_buffer</span><span class="p">(</span><span class="s">'cos_cache'</span><span class="p">,</span> <span class="n">freqs</span><span class="p">.</span><span class="n">cos</span><span class="p">(),</span> <span class="n">persistent</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (T_max, D_h/2)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">register_buffer</span><span class="p">(</span><span class="s">'sin_cache'</span><span class="p">,</span> <span class="n">freqs</span><span class="p">.</span><span class="n">sin</span><span class="p">(),</span> <span class="n">persistent</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (T_max, D_h/2)
</span>    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span> <span class="n">offset</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">0</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="nb">tuple</span><span class="p">:</span>
        <span class="s">"""
        Args:
            x:      (B, H, T, D_h) Q or K tensor
            offset: 位置偏移（用于 KV cache 的 decode 阞段）
        Returns:
            cos, sin: each (T, D_h/2) 用于旋转
        """</span>
        <span class="n">T</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span>
        <span class="n">cos</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">cos_cache</span><span class="p">[</span><span class="n">offset</span><span class="p">:</span><span class="n">offset</span> <span class="o">+</span> <span class="n">T</span><span class="p">]</span>  <span class="c1"># (T, D_h/2)
</span>        <span class="n">sin</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">sin_cache</span><span class="p">[</span><span class="n">offset</span><span class="p">:</span><span class="n">offset</span> <span class="o">+</span> <span class="n">T</span><span class="p">]</span>  <span class="c1"># (T, D_h/2)
</span>        <span class="c1"># 广播到 (B, H, T, D_h/2)
</span>        <span class="n">cos</span> <span class="o">=</span> <span class="n">cos</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>
        <span class="n">sin</span> <span class="o">=</span> <span class="n">sin</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">cos</span><span class="p">,</span> <span class="n">sin</span>

<span class="k">def</span> <span class="nf">apply_rotary_emb</span><span class="p">(</span><span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span> <span class="n">cos</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span> <span class="n">sin</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
    <span class="s">"""
    对 x 应用旋转位置编码。
    Args:
        x:   (B, H, T, D_h) — 必须是偶数维度
        cos: (B, H, T, D_h/2)
        sin: (B, H, T, D_h/2)
    Returns:
        (B, H, T, D_h) 旋转后的 tensor
    """</span>
    <span class="c1"># 将 x 拆成相邻对：x_0, x_1, x_2, x_3, ... → (x_0, x_2, ...) 和 (x_1, x_3, ...)
</span>    <span class="n">x1</span> <span class="o">=</span> <span class="n">x</span><span class="p">[...,</span> <span class="p">::</span><span class="mi">2</span><span class="p">]</span>    <span class="c1"># (B, H, T, D_h/2) — 偶数维度
</span>    <span class="n">x2</span> <span class="o">=</span> <span class="n">x</span><span class="p">[...,</span> <span class="mi">1</span><span class="p">::</span><span class="mi">2</span><span class="p">]</span>   <span class="c1"># (B, H, T, D_h/2) — 奇数维度
</span>    
    <span class="c1"># 旋转：R * [x1, x2] = [x1*cos - x2*sin, x1*sin + x2*cos]
</span>    <span class="n">rotated</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">stack</span><span class="p">([</span>
        <span class="n">x1</span> <span class="o">*</span> <span class="n">cos</span> <span class="o">-</span> <span class="n">x2</span> <span class="o">*</span> <span class="n">sin</span><span class="p">,</span>  <span class="c1"># 旋转后的偶数维度
</span>        <span class="n">x1</span> <span class="o">*</span> <span class="n">sin</span> <span class="o">+</span> <span class="n">x2</span> <span class="o">*</span> <span class="n">cos</span><span class="p">,</span>  <span class="c1"># 旋转后的奇数维度
</span>    <span class="p">],</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B, H, T, D_h/2, 2)
</span>    
    <span class="c1"># 交错合并回 (B, H, T, D_h)
</span>    <span class="k">return</span> <span class="n">rotated</span><span class="p">.</span><span class="n">flatten</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, H, T, D_h)
</span>
<span class="c1"># === RoPE MHA 完整示例 ===
</span><span class="k">class</span> <span class="nc">RoPEMultiHeadAttention</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">:</span> <span class="nb">int</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">d_model</span> <span class="o">=</span> <span class="n">d_model</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span> <span class="o">=</span> <span class="n">n_heads</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span> <span class="o">=</span> <span class="n">d_model</span> <span class="o">//</span> <span class="n">n_heads</span>
        
        <span class="bp">self</span><span class="p">.</span><span class="n">q_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">k_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">v_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">o_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">rotary_emb</span> <span class="o">=</span> <span class="n">RotaryEmbedding</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">D</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">shape</span>
        
        <span class="n">q</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">q_proj</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">k</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">k_proj</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">v</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">v_proj</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
        
        <span class="c1"># RoPE 只作用于 Q 和 K，不作用于 V
</span>        <span class="n">cos</span><span class="p">,</span> <span class="n">sin</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">rotary_emb</span><span class="p">(</span><span class="n">q</span><span class="p">)</span>
        <span class="n">q</span> <span class="o">=</span> <span class="n">apply_rotary_emb</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">cos</span><span class="p">,</span> <span class="n">sin</span><span class="p">)</span>
        <span class="n">k</span> <span class="o">=</span> <span class="n">apply_rotary_emb</span><span class="p">(</span><span class="n">k</span><span class="p">,</span> <span class="n">cos</span><span class="p">,</span> <span class="n">sin</span><span class="p">)</span>
        
        <span class="c1"># Attention（仍需要 causal mask）
</span>        <span class="n">attn</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">))</span> <span class="o">/</span> <span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">)</span>
        <span class="n">causal_mask</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tril</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">T</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">device</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="nb">bool</span><span class="p">))</span>
        <span class="n">attn</span> <span class="o">=</span> <span class="n">attn</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="o">~</span><span class="n">causal_mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
        <span class="n">attn</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">attn</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">attn</span><span class="p">,</span> <span class="n">v</span><span class="p">)</span>
        
        <span class="n">out</span> <span class="o">=</span> <span class="n">out</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">).</span><span class="n">contiguous</span><span class="p">().</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">D</span><span class="p">)</span>
        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">o_proj</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>

<span class="c1"># === 验证 ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">d_model</span> <span class="o">=</span> <span class="mi">64</span>
    <span class="n">n_heads</span> <span class="o">=</span> <span class="mi">4</span>
    <span class="n">B</span><span class="p">,</span> <span class="n">T</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">8</span>
    
    <span class="n">rope_attn</span> <span class="o">=</span> <span class="n">RoPEMultiHeadAttention</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">)</span>
    <span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">d_model</span><span class="p">)</span>
    <span class="n">out</span> <span class="o">=</span> <span class="n">rope_attn</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Input:  </span><span class="si">{</span><span class="n">x</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>   <span class="c1"># (2, 8, 64)
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Output: </span><span class="si">{</span><span class="n">out</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 8, 64)
</span></code></pre></div></div>

<p><strong>RoPE vs 绝对位置编码对比：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>绝对位置编码</th>
      <th>RoPE</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>编码方式</td>
      <td>位置 embedding 加到 token embedding</td>
      <td>旋转矩阵乘到 Q/K</td>
    </tr>
    <tr>
      <td>位置信息类型</td>
      <td>绝对位置（位置 5 → 固定 embedding）</td>
      <td>相对位置（Q·K 只依赖 n-m）</td>
    </tr>
    <tr>
      <td>长度外推</td>
      <td>固定 max_seq_len → 超出需插值</td>
      <td>理论上可外推（但远处衰减严重）</td>
    </tr>
    <tr>
      <td>外推方法</td>
      <td>位置插值（PI）/ NTK-aware 插值</td>
      <td>YaRN / Dynamic NTK</td>
    </tr>
    <tr>
      <td>对 V 的影响</td>
      <td>有（位置信息嵌入 token → V 也受影响）</td>
      <td>无（只旋转 Q/K → V 不受位置影响）</td>
    </tr>
    <tr>
      <td>与 attention 的交互</td>
      <td>独立于 attention 计算</td>
      <td>嵌入 attention 计算（QK 点积内）</td>
    </tr>
    <tr>
      <td>KV cache 兼容性</td>
      <td>需要在 cache 中存位置编号</td>
      <td>需要传 offset（位置偏移）</td>
    </tr>
  </tbody>
</table>

<p>⚠️ 常见坑：RoPE 的 <code class="language-plaintext highlighter-rouge">flatten(-2)</code> 操作 → 如果输入 x 的最后一维不是偶数 → 报错 → 需要确保 <code class="language-plaintext highlighter-rouge">d_head</code> 是偶数。</p>

<p>⚠️ 另一个坑：RoPE 在 decode 阞段需要传 <code class="language-plaintext highlighter-rouge">offset</code> 参数 → 忘记传 → 位置编号从 0 开始 → causal mask 错误 → 模型输出乱码。</p>

<p><strong>复杂度分析：</strong></p>
<ul>
  <li>时间：与标准 MHA 相同，旋转操作仅增加 $O(T \cdot D_h)$ 的乘法（极小）</li>
  <li>显存：cos/sin cache 约 $O(T_{max} \cdot D_h/2)$ → 通常 &lt;1MB</li>
</ul>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>RoPE 为什么能实现相对位置？（数学证明：$R_{\Theta,m}^\top R_{\Theta,n} = R_{\Theta,n-m}$ → QK 点积 = $q^\top R_{\Theta,n-m} k$ → 只依赖 $n-m$）</li>
  <li>RoPE 的长度外推怎么做？（NTK-aware scaling：增大 base → 低频分量衰减慢 → 远处位置仍有区分度；YaRN：混合高频缩放+低频插值 → 最优外推方案）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2104.09864">RoFormer: Enhanced Transformer with Rotary Position Embedding</a></li>
  <li><a href="https://arxiv.org/abs/2309.00071">YaRN: Efficient Context Window Extension of Large Language Models</a></li>
</ul>

<hr />

<h3 id="q54-rmsnorm-vs-layernorm-实现">Q54 RMSNorm vs LayerNorm 实现</h3>

<p><strong>难度：</strong> 基础</p>

<p><strong>考察点：</strong> RMSNorm 与 LayerNorm 的计算差异；理解为什么现代 LLM 倾向用 RMSNorm</p>

<p><strong>满分回答：</strong></p>

<p><strong>LayerNorm：</strong> 对每个样本的所有特征做归一化：
\(\text{LayerNorm}(x) = \frac{x - \mu}{\sigma} \cdot \gamma + \beta\)
其中 $\mu = \frac{1}{D}\sum_i x_i$，$\sigma = \sqrt{\frac{1}{D}\sum_i (x_i - \mu)^2 + \epsilon}$</p>

<p><strong>RMSNorm：</strong> 省去均值中心化，只做缩放：
\(\text{RMSNorm}(x) = \frac{x}{\text{RMS}(x)} \cdot \gamma\)
其中 $\text{RMS}(x) = \sqrt{\frac{1}{D}\sum_i x_i^2 + \epsilon}$</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>

<span class="k">class</span> <span class="nc">RMSNorm</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="s">"""Root Mean Square Layer Normalization — LLaMA/Qwen 等现代 LLM 使用"""</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">eps</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">1e-6</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">eps</span> <span class="o">=</span> <span class="n">eps</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">weight</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">d_model</span><span class="p">))</span>  <span class="c1"># (D) — 只有 γ，没有 β
</span>    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="s">"""
        Args:
            x: (B, T, D) or (..., D)
        Returns:
            (B, T, D) or (..., D) 归一化后
        """</span>
        <span class="c1"># RMS 计算：sqrt(mean(x^2) + eps)
</span>        <span class="c1"># 注意：在 fp16/bf16 下，先转 fp32 计算 RMS 再转回 → 数值稳定
</span>        <span class="n">rms</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="nb">float</span><span class="p">().</span><span class="nb">pow</span><span class="p">(</span><span class="mi">2</span><span class="p">).</span><span class="n">mean</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">eps</span><span class="p">)</span>  <span class="c1"># (B, T, 1) fp32
</span>        <span class="n">x_normed</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="nb">float</span><span class="p">()</span> <span class="o">/</span> <span class="n">rms</span>  <span class="c1"># (B, T, D) fp32
</span>        <span class="c1"># 乘权重并转回原精度
</span>        <span class="k">return</span> <span class="p">(</span><span class="n">x_normed</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">weight</span><span class="p">).</span><span class="n">type_as</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># (B, T, D)
</span>
<span class="k">class</span> <span class="nc">LayerNorm</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="s">"""标准 Layer Normalization — Transformer 原论文使用"""</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">eps</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">1e-6</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">eps</span> <span class="o">=</span> <span class="n">eps</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">weight</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">d_model</span><span class="p">))</span>   <span class="c1"># (D) — γ (gain)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">bias</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">d_model</span><span class="p">))</span>     <span class="c1"># (D) — β (shift)
</span>    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="s">"""
        Args:
            x: (B, T, D) or (..., D)
        Returns:
            (B, T, D) or (..., D) 归一化后
        """</span>
        <span class="c1"># 同样在 fp32 下计算以保证数值稳定性
</span>        <span class="n">mean</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="nb">float</span><span class="p">().</span><span class="n">mean</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>  <span class="c1"># (B, T, 1)
</span>        <span class="n">var</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="nb">float</span><span class="p">().</span><span class="nb">pow</span><span class="p">(</span><span class="mi">2</span><span class="p">).</span><span class="n">mean</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span> <span class="o">-</span> <span class="n">mean</span><span class="p">.</span><span class="nb">pow</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, T, 1)
</span>        <span class="n">x_normed</span> <span class="o">=</span> <span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="nb">float</span><span class="p">()</span> <span class="o">-</span> <span class="n">mean</span><span class="p">)</span> <span class="o">/</span> <span class="n">torch</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">var</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">eps</span><span class="p">)</span>  <span class="c1"># (B, T, D)
</span>        <span class="k">return</span> <span class="p">(</span><span class="n">x_normed</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">weight</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">bias</span><span class="p">).</span><span class="n">type_as</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>

<span class="c1"># === 验证 ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">d_model</span> <span class="o">=</span> <span class="mi">64</span>
    <span class="n">B</span><span class="p">,</span> <span class="n">T</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">10</span>
    
    <span class="n">rms_norm</span> <span class="o">=</span> <span class="n">RMSNorm</span><span class="p">(</span><span class="n">d_model</span><span class="p">)</span>
    <span class="n">ln_norm</span> <span class="o">=</span> <span class="n">LayerNorm</span><span class="p">(</span><span class="n">d_model</span><span class="p">)</span>
    
    <span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">d_model</span><span class="p">)</span>
    
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"RMSNorm output: </span><span class="si">{</span><span class="n">rms_norm</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 10, 64)
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"LayerNorm output: </span><span class="si">{</span><span class="n">ln_norm</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 10, 64)
</span>    
    <span class="c1"># 验证 RMSNorm 的 RMS 值接近 1
</span>    <span class="n">out_rms</span> <span class="o">=</span> <span class="n">rms_norm</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
    <span class="n">rms_val</span> <span class="o">=</span> <span class="n">out_rms</span><span class="p">.</span><span class="nb">float</span><span class="p">().</span><span class="nb">pow</span><span class="p">(</span><span class="mi">2</span><span class="p">).</span><span class="n">mean</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># 应接近 1
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"RMSNorm 辧出 RMS ≈ </span><span class="si">{</span><span class="n">rms_val</span><span class="p">.</span><span class="n">mean</span><span class="p">().</span><span class="n">item</span><span class="p">()</span><span class="si">:</span><span class="p">.</span><span class="mi">4</span><span class="n">f</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># ≈ 1.0
</span></code></pre></div></div>

<p><strong>关键对比：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>LayerNorm</th>
      <th>RMSNorm</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>计算</td>
      <td>$(x - \mu) / \sigma \cdot \gamma + \beta$</td>
      <td>$x / \text{RMS} \cdot \gamma$</td>
    </tr>
    <tr>
      <td>可学习参数</td>
      <td>$\gamma$ + $\beta$（2×D）</td>
      <td>$\gamma$（1×D）→ 无 bias</td>
    </tr>
    <tr>
      <td>均值中心化</td>
      <td>✅ 有</td>
      <td>❌ 无</td>
    </tr>
    <tr>
      <td>计算量</td>
      <td>2 次 mean + 1 次 var</td>
      <td>1 次 mean(x²)</td>
    </tr>
    <tr>
      <td>速度</td>
      <td>略慢（多一步减均值）</td>
      <td>略快（省减均值和 bias）</td>
    </tr>
    <tr>
      <td>效果</td>
      <td>理论上更灵活（可做偏移）</td>
      <td>实际效果几乎等同</td>
    </tr>
    <tr>
      <td>代表模型</td>
      <td>GPT-2/3, BERT</td>
      <td>LLaMA, Qwen, Mistral</td>
    </tr>
  </tbody>
</table>

<p><strong>为什么现代 LLM 用 RMSNorm？</strong></p>
<ol>
  <li><strong>参数更少</strong>：省 bias → 每层少 D 个参数 → 大模型省不少</li>
  <li><strong>计算更快</strong>：省均值中心化 → 略微加速（在大模型中累积可观）</li>
  <li><strong>效果等同</strong>：实验显示 RMSNorm 和 LayerNorm 在 LLM 训练中效果无显著差异</li>
  <li><strong>数值稳定</strong>：RMSNorm 的 $\text{RMS}(x) = \sqrt{\text{mean}(x^2)}$ 比 LayerNorm 的 $\sigma$ 更稳定（因为不涉及减均值再平方 → 减少数值误差）</li>
</ol>

<p>⚠️ 常见坑：RMSNorm 没有 bias → 以为会限制模型表达能力 → 实际归一化后的 $\gamma$ 本身就提供了缩放能力，而偏移在大多数场景下不必要（后面的 linear 层自带 bias）。</p>

<p>⚠️ 另一个坑：在 bf16 下直接计算 RMS → 精度不够 → 必须 <code class="language-plaintext highlighter-rouge">.float()</code> 转 fp32 计算 → 再 <code class="language-plaintext highlighter-rouge">.type_as(x)</code> 转回 bf16。这是所有 norm 层的通用做法。</p>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>为什么 RMSNorm 省去均值中心化效果不变？（归一化的核心作用是稳定梯度分布 → 均值中心化对此贡献很小 → 纯缩放就够了）</li>
  <li>RMSNorm 的 eps 为什么设 1e-6？（防止 RMS ≈ 0 时除零 → 1e-6 在 fp32 下足够小不影响正常值，在 bf16 下也不会溢出）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/1910.07467">Root Mean Square Layer Normalization (RMSNorm)</a></li>
  <li><a href="https://arxiv.org/abs/1607.06450">Layer Normalization</a></li>
</ul>

<hr />

<h3 id="q55-swiglu--geglu-ffn-实现">Q55 SwiGLU / GeGLU FFN 实现</h3>

<p><strong>难度：</strong> 进阶</p>

<p><strong>考察点：</strong> GLU 变体的 FFN 实现；理解为什么 SwiGLU 成为现代 LLM 的标准 FFN 选择</p>

<p><strong>满分回答：</strong></p>

<p><strong>标准 FFN（Transformer 原论文）：</strong>
\(\text{FFN}(x) = W_2 \cdot \text{ReLU}(W_1 \cdot x)\)</p>

<p><strong>GeGLU（Gated GLU with GeLU activation）：</strong>
\(\text{GeGLU}(x) = (W_1 \cdot x) \odot \text{GeLU}(W_{gate} \cdot x)\)
\(\text{FFN}_{\text{GeGLU}}(x) = W_2 \cdot \text{GeGLU}(x)\)</p>

<p><strong>SwiGLU（LLaMA 使用）：</strong>
\(\text{SwiGLU}(x) = (W_1 \cdot x) \odot \text{SiLU}(W_{gate} \cdot x)\)
\(\text{FFN}_{\text{SwiGLU}}(x) = W_2 \cdot \text{SwiGLU}(x)\)</p>

<p>其中 $\text{SiLU}(x) = x \cdot \sigma(x)$（sigmoid linear unit，也叫 Swish）。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>

<span class="k">class</span> <span class="nc">SwiGLUFFN</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="s">"""SwiGLU FFN — LLaMA/Qwen/Mistral 等现代 LLM 使用"""</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">dropout</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="c1"># 注意：SwiGLU 有 3 个投影矩阵（W1, W_gate, W2），而非标准 FFN 的 2 个
</span>        <span class="c1"># 合并 W1 和 W_gate 为一次投影以提升效率
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">w1_w_gate</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="mi">2</span> <span class="o">*</span> <span class="n">d_ff</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (D) -&gt; (2 * d_ff)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">w2</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_ff</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>              <span class="c1"># (d_ff) -&gt; (D)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">drop</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="s">"""
        Args:
            x: (B, T, D)
        Returns:
            (B, T, D)
        """</span>
        <span class="c1"># Step 1: 合并投影 → 拆分为 value 和 gate
</span>        <span class="n">proj</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">w1_w_gate</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>                <span class="c1"># (B, T, 2*d_ff)
</span>        <span class="n">value</span><span class="p">,</span> <span class="n">gate</span> <span class="o">=</span> <span class="n">proj</span><span class="p">.</span><span class="n">chunk</span><span class="p">(</span><span class="mi">2</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>     <span class="c1"># each (B, T, d_ff)
</span>        
        <span class="c1"># Step 2: SwiGLU = value * SiLU(gate)
</span>        <span class="c1"># SiLU(x) = x * sigmoid(x) ≈ x * (1 / (1 + exp(-x)))
</span>        <span class="n">value_gated</span> <span class="o">=</span> <span class="n">value</span> <span class="o">*</span> <span class="n">F</span><span class="p">.</span><span class="n">silu</span><span class="p">(</span><span class="n">gate</span><span class="p">)</span>       <span class="c1"># (B, T, d_ff)
</span>        
        <span class="c1"># Step 3: Output projection
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">w2</span><span class="p">(</span><span class="n">value_gated</span><span class="p">)</span>               <span class="c1"># (B, T, D)
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">drop</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
        
        <span class="k">return</span> <span class="n">out</span>

<span class="k">class</span> <span class="nc">GeGLUFFN</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="s">"""GeGLU FFN — PaLM 等模型使用"""</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">dropout</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">w1_w_gate</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="mi">2</span> <span class="o">*</span> <span class="n">d_ff</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">w2</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_ff</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">drop</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="n">proj</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">w1_w_gate</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">value</span><span class="p">,</span> <span class="n">gate</span> <span class="o">=</span> <span class="n">proj</span><span class="p">.</span><span class="n">chunk</span><span class="p">(</span><span class="mi">2</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>
        <span class="n">value_gated</span> <span class="o">=</span> <span class="n">value</span> <span class="o">*</span> <span class="n">F</span><span class="p">.</span><span class="n">gelu</span><span class="p">(</span><span class="n">gate</span><span class="p">)</span>  <span class="c1"># GeLU activation
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">w2</span><span class="p">(</span><span class="n">value_gated</span><span class="p">)</span>
        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">drop</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>

<span class="k">class</span> <span class="nc">StandardFFN</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="s">"""标准 ReLU FFN — Transformer 原论文"""</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">dropout</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">w1</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>   <span class="c1"># (D) -&gt; (d_ff)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">w2</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_ff</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>    <span class="c1"># (d_ff) -&gt; (D)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">drop</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="n">hidden</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">relu</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">w1</span><span class="p">(</span><span class="n">x</span><span class="p">))</span>     <span class="c1"># (B, T, d_ff)
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">w2</span><span class="p">(</span><span class="n">hidden</span><span class="p">)</span>            <span class="c1"># (B, T, D)
</span>        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">drop</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>

<span class="c1"># === 验证 ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">d_model</span> <span class="o">=</span> <span class="mi">256</span>
    <span class="n">d_ff</span> <span class="o">=</span> <span class="mi">1024</span>  <span class="c1"># SwiGLU 的 d_ff 通常设为 2/3 * 4 * D ≈ 8D/3（见下文说明）
</span>    <span class="n">B</span><span class="p">,</span> <span class="n">T</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">8</span>
    
    <span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">d_model</span><span class="p">)</span>
    
    <span class="c1"># 对比三种 FFN
</span>    <span class="n">swiglu</span> <span class="o">=</span> <span class="n">SwiGLUFFN</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">)</span>
    <span class="n">geglu</span> <span class="o">=</span> <span class="n">GeGLUFFN</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">)</span>
    <span class="n">std_ffn</span> <span class="o">=</span> <span class="n">StandardFFN</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">)</span>
    
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"SwiGLU: </span><span class="si">{</span><span class="n">swiglu</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 8, 256)
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"GeGLU:  </span><span class="si">{</span><span class="n">geglu</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>   <span class="c1"># (2, 8, 256)
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Standard: </span><span class="si">{</span><span class="n">std_ffn</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 8, 256)
</span>    
    <span class="c1"># 参数量对比
</span>    <span class="k">def</span> <span class="nf">count_params</span><span class="p">(</span><span class="n">m</span><span class="p">):</span>
        <span class="k">return</span> <span class="nb">sum</span><span class="p">(</span><span class="n">p</span><span class="p">.</span><span class="n">numel</span><span class="p">()</span> <span class="k">for</span> <span class="n">p</span> <span class="ow">in</span> <span class="n">m</span><span class="p">.</span><span class="n">parameters</span><span class="p">())</span>
    
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"SwiGLU params:  </span><span class="si">{</span><span class="n">count_params</span><span class="p">(</span><span class="n">swiglu</span><span class="p">)</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>   <span class="c1"># D*(2*d_ff) + d_ff*D = 3*D*d_ff
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"GeGLU params:   </span><span class="si">{</span><span class="n">count_params</span><span class="p">(</span><span class="n">geglu</span><span class="p">)</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>    <span class="c1"># 3*D*d_ff (同 SwiGLU)
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Standard params: </span><span class="si">{</span><span class="n">count_params</span><span class="p">(</span><span class="n">std_ffn</span><span class="p">)</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># D*d_ff + d_ff*D = 2*D*d_ff
</span></code></pre></div></div>

<p><strong>关键对比：</strong></p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>ReLU FFN</th>
      <th>GeGLU FFN</th>
      <th>SwiGLU FFN</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>公式</td>
      <td>$W_2 \cdot \text{ReLU}(W_1 x)$</td>
      <td>$W_2 \cdot (W_1 x \odot \text{GeLU}(W_g x))$</td>
      <td>$W_2 \cdot (W_1 x \odot \text{SiLU}(W_g x))$</td>
    </tr>
    <tr>
      <td>投影矩阵数</td>
      <td>2</td>
      <td>3</td>
      <td>3</td>
    </tr>
    <tr>
      <td>参数量</td>
      <td>$2 \cdot D \cdot d_{ff}$</td>
      <td>$3 \cdot D \cdot d_{ff}$</td>
      <td>$3 \cdot D \cdot d_{ff}$</td>
    </tr>
    <tr>
      <td>门控机制</td>
      <td>无</td>
      <td>GeLU 门控</td>
      <td>SiLU 门控</td>
    </tr>
    <tr>
      <td>激活函数</td>
      <td>ReLU</td>
      <td>GeLU</td>
      <td>SiLU/Swish</td>
    </tr>
    <tr>
      <td>效果</td>
      <td>基线</td>
      <td>比 ReLU 好</td>
      <td>比 GeGLU 略好</td>
    </tr>
    <tr>
      <td>代表模型</td>
      <td>GPT-2</td>
      <td>PaLM</td>
      <td>LLaMA, Qwen, Mistral</td>
    </tr>
  </tbody>
</table>

<p><strong>SwiGLU 的 d_ff 设定问题：</strong></p>
<ul>
  <li>标准 FFN 的 $d_{ff} = 4D$（Transformer 原论文）</li>
  <li>SwiGLU/GeGLU 有 3 个投影矩阵 → 参数量 = $3Dd_{ff}$</li>
  <li>为保持与标准 FFN 相同的参数量（$2D \cdot 4D = 8D^2$）：$d_{ff} = \frac{8D^2}{3D} = \frac{8D}{3}$</li>
  <li>→ LLaMA 等模型用 $d_{ff} \approx \frac{2}{3} \cdot 4D$（通常取最近的 256 倍数）</li>
</ul>

<p>⚠️ 常见坑：SwiGLU 的 <code class="language-plaintext highlighter-rouge">d_ff</code> 设为 $4D$ → 参数量变成 $3 \times 4D^2$ → 比标准 FFN 多 50% → 计算和显存也增加。需要调整为 $\frac{8D}{3}$ 才能与标准 FFN 参数量持平。</p>

<p>⚠️ 另一个坑：合并投影 <code class="language-plaintext highlighter-rouge">w1_w_gate</code> 辧出维度是 <code class="language-plaintext highlighter-rouge">2*d_ff</code> → chunk 拆分后每个 <code class="language-plaintext highlighter-rouge">d_ff</code> → 如果 <code class="language-plaintext highlighter-rouge">2*d_ff</code> 不是偶数或 chunk 拆分不正确 → shape 错误。</p>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>为什么 SwiGLU 比 ReLU 好？（门控机制让 FFN 可以选择性地传递信息 → 比纯 ReLU 的”全通/全堵”更灵活 → 实验证实效果提升约 2-4%）</li>
  <li>SiLU 和 GeLU 的区别？（SiLU = $x \cdot \sigma(x)$，平滑且非单调；GeLU = $x \cdot \Phi(x)$（正态CDF），更接近 ReLU 但平滑。SiLU 更简单、计算更快）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2002.05202">GLU Variants Improve Transformer (Noam Shazeer, 2020)</a></li>
  <li><a href="https://arxiv.org/abs/2302.13971">LLaMA: Open and Efficient Foundation Language Models</a></li>
</ul>

<hr />

<h3 id="q56-grouped-query-attention-gqa--multi-query-attention-mqa-实现">Q56 Grouped-Query Attention (GQA) / Multi-Query Attention (MQA) 实现</h3>

<p><strong>难度：</strong> 专家</p>

<p><strong>考察点：</strong> GQA/MQA 的 KV head 数量与 Q head 数量的差异；理解 KV cache 压缩的实现方式与 tradeoff</p>

<p><strong>满分回答：</strong></p>

<p><strong>MQA</strong>：所有 Q head 共享<strong>1 组</strong> K/V → KV cache 大小降至标准 MHA 的 $1/H$。
<strong>GQA</strong>：Q head 分成 $G$ 组，每组共享 1 组 K/V → KV cache 大小降至标准 MHA 的 $G/H$。</p>

<table>
  <thead>
    <tr>
      <th>方案</th>
      <th>Q heads</th>
      <th>K/V heads</th>
      <th>KV cache 倍率</th>
      <th>代表模型</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>MHA</td>
      <td>H</td>
      <td>H</td>
      <td>1x</td>
      <td>GPT-2/3, BERT</td>
    </tr>
    <tr>
      <td>MQA</td>
      <td>H</td>
      <td>1</td>
      <td>1/H</td>
      <td>PaLM, StarCoder</td>
    </tr>
    <tr>
      <td>GQA</td>
      <td>H</td>
      <td>G (G &lt; H)</td>
      <td>G/H</td>
      <td>LLaMA-2/3, Mistral</td>
    </tr>
  </tbody>
</table>

<p>LLaMA-2 70B：H=64, G=8 → KV cache 降至 8/64 = 12.5%</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>
<span class="kn">import</span> <span class="nn">math</span>

<span class="k">class</span> <span class="nc">GroupedQueryAttention</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="s">"""GQA / MQA 统一实现 — 支持 n_kv_heads ∈ {1, ..., n_heads}"""</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">n_kv_heads</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="bp">None</span><span class="p">,</span> <span class="n">dropout</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="k">assert</span> <span class="n">d_model</span> <span class="o">%</span> <span class="n">n_heads</span> <span class="o">==</span> <span class="mi">0</span>
        
        <span class="bp">self</span><span class="p">.</span><span class="n">d_model</span> <span class="o">=</span> <span class="n">d_model</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span> <span class="o">=</span> <span class="n">n_heads</span>        <span class="c1"># Q head 数 (H)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">n_kv_heads</span> <span class="o">=</span> <span class="n">n_kv_heads</span> <span class="k">if</span> <span class="n">n_kv_heads</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span> <span class="k">else</span> <span class="n">n_heads</span>  <span class="c1"># K/V head 数 (G)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span> <span class="o">=</span> <span class="n">d_model</span> <span class="o">//</span> <span class="n">n_heads</span>       <span class="c1"># 每个 Q head 的维度 (D_h)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">d_kv_head</span> <span class="o">=</span> <span class="n">d_model</span> <span class="o">//</span> <span class="n">n_kv_heads</span>  <span class="c1"># 每个 KV head 的维度 — 与 Q head 相同
</span>        
        <span class="c1"># 实际实现中，KV head 维度通常 = d_model // n_kv_heads
</span>        <span class="c1"># 但为了 GQA 的 expand 操作，让 KV head 维度 = Q head 维度
</span>        <span class="c1"># 即 d_kv_head = d_head（而非 d_model // n_kv_heads）
</span>        <span class="c1"># 这样 expand 只需 repeat KV head，不需要额外的投影
</span>        
        <span class="c1"># 投影矩阵
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">q_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>     <span class="c1"># (D) -&gt; (H * D_h)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">k_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_kv_heads</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (D) -&gt; (G * D_h)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">v_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_kv_heads</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>  <span class="c1"># (D) -&gt; (G * D_h)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">o_proj</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">n_heads</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>     <span class="c1"># (H * D_h) -&gt; (D)
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">attn_drop</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>
        
        <span class="c1"># 每个 Q head 组有多少个 Q head
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">n_rep</span> <span class="o">=</span> <span class="n">n_heads</span> <span class="o">//</span> <span class="n">n_kv_heads</span>  <span class="c1"># H // G = 每个 KV head 覆盖的 Q head 数
</span>    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span> <span class="n">mask</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span> <span class="o">=</span> <span class="bp">None</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="s">"""
        Args:
            x:    (B, T, D)
            mask: (T, T) or (B, 1, T, T) causal mask, 1=allow, 0=block
        Returns:
            out:  (B, T, D)
        """</span>
        <span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">D</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">shape</span>
        
        <span class="c1"># Step 1: 投影
</span>        <span class="n">q</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">q_proj</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># (B, T, H * D_h)
</span>        <span class="n">k</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">k_proj</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># (B, T, G * D_h)
</span>        <span class="n">v</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">v_proj</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># (B, T, G * D_h)
</span>        
        <span class="c1"># Step 2: reshape
</span>        <span class="n">q</span> <span class="o">=</span> <span class="n">q</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>     <span class="c1"># (B, H, T, D_h)
</span>        <span class="n">k</span> <span class="o">=</span> <span class="n">k</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_kv_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, G, T, D_h)
</span>        <span class="n">v</span> <span class="o">=</span> <span class="n">v</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_kv_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B, G, T, D_h)
</span>        
        <span class="c1"># Step 3: Expand KV heads to match Q heads
</span>        <span class="c1"># 每个 KV head 重复 n_rep 次 → 与 Q head 数量对齐
</span>        <span class="c1"># 例如 GQA: H=8, G=2, n_rep=4 → K/V 从 (B,2,T,D_h) expand 到 (B,8,T,D_h)
</span>        <span class="k">if</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_rep</span> <span class="o">&gt;</span> <span class="mi">1</span><span class="p">:</span>
            <span class="c1"># repeat_interleave: 沿 head 维度重复
</span>            <span class="n">k</span> <span class="o">=</span> <span class="n">k</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">2</span><span class="p">).</span><span class="n">expand</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_kv_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_rep</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">)</span>  <span class="c1"># (B, G, n_rep, T, D_h)
</span>            <span class="n">k</span> <span class="o">=</span> <span class="n">k</span><span class="p">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">)</span>  <span class="c1"># (B, H, T, D_h)
</span>            <span class="n">v</span> <span class="o">=</span> <span class="n">v</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">2</span><span class="p">).</span><span class="n">expand</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_kv_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_rep</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">)</span>
            <span class="n">v</span> <span class="o">=</span> <span class="n">v</span><span class="p">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">n_heads</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">)</span>
        
        <span class="c1"># Step 4: Scaled Dot-Product Attention（与标准 MHA 相同）
</span>        <span class="n">scale</span> <span class="o">=</span> <span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">d_head</span><span class="p">)</span>
        <span class="n">attn_scores</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">))</span> <span class="o">/</span> <span class="n">scale</span>  <span class="c1"># (B, H, T, T)
</span>        
        <span class="c1"># Causal mask
</span>        <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">attn_scores</span> <span class="o">=</span> <span class="n">attn_scores</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span> <span class="o">==</span> <span class="mi">0</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
        <span class="k">else</span><span class="p">:</span>
            <span class="n">causal_mask</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tril</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">T</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">x</span><span class="p">.</span><span class="n">device</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="nb">bool</span><span class="p">))</span>
            <span class="n">attn_scores</span> <span class="o">=</span> <span class="n">attn_scores</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="o">~</span><span class="n">causal_mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
        
        <span class="n">attn_weights</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">attn_scores</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B, H, T, T)
</span>        <span class="n">attn_weights</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">attn_drop</span><span class="p">(</span><span class="n">attn_weights</span><span class="p">)</span>
        <span class="n">attn_out</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">attn_weights</span><span class="p">,</span> <span class="n">v</span><span class="p">)</span>       <span class="c1"># (B, H, T, D_h)
</span>        
        <span class="c1"># Step 5: Concatenate + Output projection
</span>        <span class="n">attn_out</span> <span class="o">=</span> <span class="n">attn_out</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">).</span><span class="n">contiguous</span><span class="p">().</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">D</span><span class="p">)</span>  <span class="c1"># (B, T, D)
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">o_proj</span><span class="p">(</span><span class="n">attn_out</span><span class="p">)</span>  <span class="c1"># (B, T, D)
</span>        
        <span class="k">return</span> <span class="n">out</span>

<span class="c1"># === 验证三种 Attention ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">d_model</span> <span class="o">=</span> <span class="mi">512</span>
    <span class="n">B</span><span class="p">,</span> <span class="n">T</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">16</span>
    
    <span class="c1"># MHA: n_kv_heads = n_heads = 8
</span>    <span class="n">mha</span> <span class="o">=</span> <span class="n">GroupedQueryAttention</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">n_kv_heads</span><span class="o">=</span><span class="mi">8</span><span class="p">)</span>
    <span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">d_model</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"MHA output:  </span><span class="si">{</span><span class="n">mha</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 16, 512)
</span>    
    <span class="c1"># GQA: n_heads=8, n_kv_heads=2 → n_rep=4
</span>    <span class="n">gqa</span> <span class="o">=</span> <span class="n">GroupedQueryAttention</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">n_kv_heads</span><span class="o">=</span><span class="mi">2</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"GQA output:  </span><span class="si">{</span><span class="n">gqa</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 16, 512)
</span>    
    <span class="c1"># MQA: n_heads=8, n_kv_heads=1 → n_rep=8
</span>    <span class="n">mqa</span> <span class="o">=</span> <span class="n">GroupedQueryAttention</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">n_kv_heads</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"MQA output:  </span><span class="si">{</span><span class="n">mqa</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (2, 16, 512)
</span>    
    <span class="c1"># 参数量对比
</span>    <span class="k">def</span> <span class="nf">count_params</span><span class="p">(</span><span class="n">m</span><span class="p">):</span>
        <span class="k">return</span> <span class="nb">sum</span><span class="p">(</span><span class="n">p</span><span class="p">.</span><span class="n">numel</span><span class="p">()</span> <span class="k">for</span> <span class="n">p</span> <span class="ow">in</span> <span class="n">m</span><span class="p">.</span><span class="n">parameters</span><span class="p">())</span>
    
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"MHA params: </span><span class="si">{</span><span class="n">count_params</span><span class="p">(</span><span class="n">mha</span><span class="p">)</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>   <span class="c1"># 4*D*D = 4*512*512 = 1,048,576
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"GQA params: </span><span class="si">{</span><span class="n">count_params</span><span class="p">(</span><span class="n">gqa</span><span class="p">)</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>   <span class="c1"># D*(8+2+2)*D_h + D*D = 小于 MHA
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"MQA params: </span><span class="si">{</span><span class="n">count_params</span><span class="p">(</span><span class="n">mqa</span><span class="p">)</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>   <span class="c1"># D*(8+1+1)*D_h + D*D = 最少
</span></code></pre></div></div>

<p><strong>关键点解析：</strong></p>

<ol>
  <li><strong>Shape 流转</strong>：
    <ul>
      <li>Q: <code class="language-plaintext highlighter-rouge">(B,T,D)</code> → <code class="language-plaintext highlighter-rouge">(B,T,H,D_h)</code> → <code class="language-plaintext highlighter-rouge">(B,H,T,D_h)</code></li>
      <li>K: <code class="language-plaintext highlighter-rouge">(B,T,D)</code> → <code class="language-plaintext highlighter-rouge">(B,T,G,D_h)</code> → <code class="language-plaintext highlighter-rouge">(B,G,T,D_h)</code> → expand <code class="language-plaintext highlighter-rouge">(B,H,T,D_h)</code></li>
      <li>V: 同 K</li>
      <li>Attention: <code class="language-plaintext highlighter-rouge">(B,H,T,T)</code> × <code class="language-plaintext highlighter-rouge">(B,H,T,D_h)</code> → <code class="language-plaintext highlighter-rouge">(B,T,D)</code></li>
    </ul>
  </li>
  <li>
    <p><strong>Expand 操作</strong>：GQA 的关键步骤——将 $G$ 个 KV head 重复 $n_rep$ 次与 $H$ 个 Q head 对齐。<code class="language-plaintext highlighter-rouge">unsqueeze(2)</code> + <code class="language-plaintext highlighter-rouge">expand</code> + <code class="language-plaintext highlighter-rouge">reshape</code> 是最常见的实现方式。</p>
  </li>
  <li><strong>KV cache 压缩</strong>：推理时只存 $G$ 组 KV（而非 $H$ 组）→ cache 大小 $= \frac{G}{H} \times \text{MHA cache}$。</li>
</ol>

<p>⚠️ 常见坑：GQA expand 后 K/V 的 head 维度是 <code class="language-plaintext highlighter-rouge">d_head</code>（而非 <code class="language-plaintext highlighter-rouge">d_model // n_kv_heads</code>）→ 如果误用后者 → expand 后维度不对齐 → matmul 报错。</p>

<p>⚠️ 另一个坑：MQA（$G=1$）在某些任务上效果明显不如 MHA（约 1-2% drop），GQA（$G=4-8$）在大多数任务上效果接近 MHA → 推荐用 GQA 而不是极端 MQA。</p>

<p><strong>复杂度分析：</strong></p>
<ul>
  <li>训练时间：expand 操作增加极小开销（只是 repeat，不增加 matmul）→ 基本与 MHA 相同</li>
  <li>推理 KV cache：$O(2 \cdot G \cdot T \cdot D_h)$ per layer（vs MHA 的 $O(2 \cdot H \cdot T \cdot D_h)$）</li>
  <li>推理 decode 时间：attention 从 <code class="language-plaintext highlighter-rouge">(1, H, 1, T_kv)</code> × <code class="language-plaintext highlighter-rouge">(1, H, T_kv, D_h)</code> → 与 MHA 相同（因为 expand 了）；但 KV cache 读取量减少 → 内存带宽瓶颈缓解</li>
</ul>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>GQA 为什么比 MQA 更受欢迎？（MQA 信息压缩过度（1组KV for 所有head）→ 表征能力受限；GQA 保留多组 KV（如 8组 for 64 head）→ 更好的表征能力 + 仍然大幅压缩 cache）</li>
  <li>GQA 的 n_kv_heads 如何选择？（经验法则：n_kv_heads = n_heads / 8 或 n_heads / 4 → 平衡效果和效率。LLaMA-2 70B 用 8 KV heads for 64 Q heads）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2305.13245">GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints</a></li>
  <li><a href="https://arxiv.org/abs/2305.07954">LLaMA 2: Open Foundation and Fine-Tuned Chat Models</a></li>
</ul>

<hr />

<h3 id="q57-lora-模块从零实现含-merge-权重">Q57 LoRA 模块从零实现（含 merge 权重）</h3>

<p><strong>难度：</strong> 进阶</p>

<p><strong>考察点：</strong> LoRA 的完整实现，包括低秩分解、初始化策略、权重 merge；理解 LoRA 的推理优化</p>

<p><strong>满分回答：</strong></p>

<p>LoRA 的核心：将权重更新 $\Delta W$ 分解为低秩矩阵 $A \cdot B$：
\(h = W \cdot x + \Delta W \cdot x = W \cdot x + B \cdot A \cdot x\)
其中 $A \in \mathbb{R}^{r \times d_{in}}$，$B \in \mathbb{R}^{d_{out} \times r}$，$r \ll d_{in}, d_{out}$。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">math</span>

<span class="k">class</span> <span class="nc">LoRALayer</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="s">"""单个 LoRA 适配层，可附加到任意 Linear 层"""</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span>
        <span class="bp">self</span><span class="p">,</span>
        <span class="n">d_in</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span>
        <span class="n">d_out</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span>
        <span class="n">rank</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">8</span><span class="p">,</span>
        <span class="n">alpha</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">16.0</span><span class="p">,</span>
        <span class="n">dropout</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">,</span>
    <span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">d_in</span> <span class="o">=</span> <span class="n">d_in</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">d_out</span> <span class="o">=</span> <span class="n">d_out</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">rank</span> <span class="o">=</span> <span class="n">rank</span>        <span class="c1"># (r) — LoRA 增量矩阵的秩
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">alpha</span> <span class="o">=</span> <span class="n">alpha</span>      <span class="c1"># 缩放因子
</span>        
        <span class="c1"># LoRA 缩放系数：alpha / rank
</span>        <span class="c1"># 作用：当 rank 变化时，通过 alpha/rank 保持增量矩阵的数值量级不变
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">scaling</span> <span class="o">=</span> <span class="n">alpha</span> <span class="o">/</span> <span class="n">rank</span>
        
        <span class="c1"># Dropout（在 LoRA 辧入前）
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">lora_drop</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span> <span class="k">if</span> <span class="n">dropout</span> <span class="o">&gt;</span> <span class="mi">0</span> <span class="k">else</span> <span class="n">nn</span><span class="p">.</span><span class="n">Identity</span><span class="p">()</span>
        
        <span class="c1"># LoRA 低秩矩阵
</span>        <span class="c1"># A: (r, d_in)  — 初始化为 Kaiming/正态分布（保证训练初期有非零梯度）
</span>        <span class="c1"># B: (d_out, r) — 初始化为零（保证 LoRA 初始增量 = 0 → 不改变原模型行为）
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">lora_A</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">empty</span><span class="p">(</span><span class="n">rank</span><span class="p">,</span> <span class="n">d_in</span><span class="p">))</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">lora_B</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">d_out</span><span class="p">,</span> <span class="n">rank</span><span class="p">))</span>
        
        <span class="c1"># 初始化 A 为正态分布
</span>        <span class="n">nn</span><span class="p">.</span><span class="n">init</span><span class="p">.</span><span class="n">kaiming_uniform_</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">lora_A</span><span class="p">,</span> <span class="n">a</span><span class="o">=</span><span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="mi">5</span><span class="p">))</span>
    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="s">"""
        计算 LoRA 增量：B @ A @ x * scaling
        Args:
            x: (B, T, d_in) or (..., d_in)
        Returns:
            LoRA 增量输出: (B, T, d_out) or (..., d_out)
        """</span>
        <span class="n">x_drop</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">lora_drop</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>                          <span class="c1"># (B, T, d_in)
</span>        <span class="c1"># 注意顺序：先 A (d_in → r) 再 B (r → d_out) → 计算量 O(d_in * r + r * d_out)
</span>        <span class="c1"># 而不是先 B 再 A → 顺序很重要！
</span>        <span class="n">lora_out</span> <span class="o">=</span> <span class="n">x_drop</span> <span class="o">@</span> <span class="bp">self</span><span class="p">.</span><span class="n">lora_A</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="bp">self</span><span class="p">.</span><span class="n">lora_B</span><span class="p">.</span><span class="n">T</span>   <span class="c1"># (B, T, r) @ (r, d_out) → (B, T, d_out)
</span>        <span class="c1"># 修正：应该是 (B, T, d_in) @ (d_in, r) @ (r, d_out)
</span>        <span class="c1"># = x @ A^T @ B^T → shape (B, T, d_in) @ (d_in, r) = (B, T, r) → @ (r, d_out) = (B, T, d_out)
</span>        
        <span class="k">return</span> <span class="n">lora_out</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">scaling</span>

<span class="k">class</span> <span class="nc">LinearWithLoRA</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="s">"""Linear 层 + LoRA 适配层"""</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span>
        <span class="bp">self</span><span class="p">,</span>
        <span class="n">original_linear</span><span class="p">:</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">,</span>
        <span class="n">rank</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">8</span><span class="p">,</span>
        <span class="n">alpha</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">16.0</span><span class="p">,</span>
        <span class="n">dropout</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">,</span>
    <span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">linear</span> <span class="o">=</span> <span class="n">original_linear</span>  <span class="c1"># 原始 Linear 层（冻结）
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">lora</span> <span class="o">=</span> <span class="n">LoRALayer</span><span class="p">(</span>
            <span class="n">d_in</span><span class="o">=</span><span class="n">original_linear</span><span class="p">.</span><span class="n">in_features</span><span class="p">,</span>
            <span class="n">d_out</span><span class="o">=</span><span class="n">original_linear</span><span class="p">.</span><span class="n">out_features</span><span class="p">,</span>
            <span class="n">rank</span><span class="o">=</span><span class="n">rank</span><span class="p">,</span>
            <span class="n">alpha</span><span class="o">=</span><span class="n">alpha</span><span class="p">,</span>
            <span class="n">dropout</span><span class="o">=</span><span class="n">dropout</span><span class="p">,</span>
        <span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
        <span class="s">"""
        h = Wx + (B @ A @ x) * scaling
        Args:
            x: (B, T, d_in)
        Returns:
            (B, T, d_out)
        """</span>
        <span class="c1"># 原始 Linear 辧出 + LoRA 增量
</span>        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">linear</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">lora</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">merge_weights</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="s">"""将 LoRA 权重 merge 到原始 Linear 权重 → 推理时无需额外计算"""</span>
        <span class="c1"># W_new = W + B @ A * scaling
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">linear</span><span class="p">.</span><span class="n">weight</span><span class="p">.</span><span class="n">data</span> <span class="o">+=</span> <span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">lora</span><span class="p">.</span><span class="n">lora_B</span> <span class="o">@</span> <span class="bp">self</span><span class="p">.</span><span class="n">lora</span><span class="p">.</span><span class="n">lora_A</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">lora</span><span class="p">.</span><span class="n">scaling</span><span class="p">)</span>  <span class="c1"># (d_out, d_in)
</span>        <span class="c1"># merge 后 LoRA 层不再需要 → 可以删除以节省参数
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">lora</span> <span class="o">=</span> <span class="bp">None</span>  <span class="c1"># 删除 LoRA 层
</span>    
    <span class="k">def</span> <span class="nf">unmerge_weights</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="s">"""从 merged 权重恢复 LoRA（需要事先保存 LoRA 参数）"""</span>
        <span class="c1"># ⚠️ 实际应用中需要保存原始 weight 和 LoRA 参数
</span>        <span class="c1"># 此处仅展示逻辑：
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">linear</span><span class="p">.</span><span class="n">weight</span><span class="p">.</span><span class="n">data</span> <span class="o">-=</span> <span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">lora</span><span class="p">.</span><span class="n">lora_B</span> <span class="o">@</span> <span class="bp">self</span><span class="p">.</span><span class="n">lora</span><span class="p">.</span><span class="n">lora_A</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">lora</span><span class="p">.</span><span class="n">scaling</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">lora</span> <span class="o">=</span> <span class="n">LoRALayer</span><span class="p">(...)</span>  <span class="c1"># 重新创建（实际中应从保存的参数恢复）
</span>
<span class="c1"># === 在 LLaMA-style Block 中应用 LoRA ===
</span><span class="k">def</span> <span class="nf">apply_lora_to_model</span><span class="p">(</span><span class="n">model</span><span class="p">:</span> <span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">,</span> <span class="n">rank</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">8</span><span class="p">,</span> <span class="n">alpha</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">16.0</span><span class="p">,</span>
                         <span class="n">target_modules</span><span class="p">:</span> <span class="nb">list</span> <span class="o">=</span> <span class="p">[</span><span class="s">"q_proj"</span><span class="p">,</span> <span class="s">"k_proj"</span><span class="p">,</span> <span class="s">"v_proj"</span><span class="p">,</span> <span class="s">"o_proj"</span><span class="p">])</span> <span class="o">-&gt;</span> <span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">:</span>
    <span class="s">"""将 LoRA 附加到模型中指定名称的 Linear 层"""</span>
    <span class="k">for</span> <span class="n">name</span><span class="p">,</span> <span class="n">module</span> <span class="ow">in</span> <span class="n">model</span><span class="p">.</span><span class="n">named_modules</span><span class="p">():</span>
        <span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">module</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">)</span> <span class="ow">and</span> <span class="nb">any</span><span class="p">(</span><span class="n">t</span> <span class="ow">in</span> <span class="n">name</span> <span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="n">target_modules</span><span class="p">):</span>
            <span class="c1"># 替换 Linear 为 LinearWithLoRA
</span>            <span class="n">parent_name</span> <span class="o">=</span> <span class="n">name</span><span class="p">.</span><span class="n">rsplit</span><span class="p">(</span><span class="s">"."</span><span class="p">,</span> <span class="mi">1</span><span class="p">)[</span><span class="mi">0</span><span class="p">]</span> <span class="k">if</span> <span class="s">"."</span> <span class="ow">in</span> <span class="n">name</span> <span class="k">else</span> <span class="s">""</span>
            <span class="n">parent</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="n">get_submodule</span><span class="p">(</span><span class="n">parent_name</span><span class="p">)</span> <span class="k">if</span> <span class="n">parent_name</span> <span class="k">else</span> <span class="n">model</span>
            <span class="n">child_name</span> <span class="o">=</span> <span class="n">name</span><span class="p">.</span><span class="n">rsplit</span><span class="p">(</span><span class="s">"."</span><span class="p">,</span> <span class="mi">1</span><span class="p">)[</span><span class="mi">1</span><span class="p">]</span> <span class="k">if</span> <span class="s">"."</span> <span class="ow">in</span> <span class="n">name</span> <span class="k">else</span> <span class="n">name</span>
            <span class="nb">setattr</span><span class="p">(</span><span class="n">parent</span><span class="p">,</span> <span class="n">child_name</span><span class="p">,</span> <span class="n">LinearWithLoRA</span><span class="p">(</span><span class="n">module</span><span class="p">,</span> <span class="n">rank</span><span class="p">,</span> <span class="n">alpha</span><span class="p">))</span>
    
    <span class="c1"># 冻结所有非 LoRA 参数
</span>    <span class="k">for</span> <span class="n">name</span><span class="p">,</span> <span class="n">param</span> <span class="ow">in</span> <span class="n">model</span><span class="p">.</span><span class="n">named_parameters</span><span class="p">():</span>
        <span class="k">if</span> <span class="s">"lora"</span> <span class="ow">not</span> <span class="ow">in</span> <span class="n">name</span><span class="p">:</span>
            <span class="n">param</span><span class="p">.</span><span class="n">requires_grad</span> <span class="o">=</span> <span class="bp">False</span>
    
    <span class="k">return</span> <span class="n">model</span>

<span class="c1"># === 验证 ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">d_in</span> <span class="o">=</span> <span class="mi">512</span>
    <span class="n">d_out</span> <span class="o">=</span> <span class="mi">512</span>
    <span class="n">rank</span> <span class="o">=</span> <span class="mi">8</span>
    <span class="n">B</span><span class="p">,</span> <span class="n">T</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">16</span>
    
    <span class="c1"># 创建 LoRA Linear
</span>    <span class="n">original_linear</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_in</span><span class="p">,</span> <span class="n">d_out</span><span class="p">)</span>
    <span class="n">lora_linear</span> <span class="o">=</span> <span class="n">LinearWithLoRA</span><span class="p">(</span><span class="n">original_linear</span><span class="p">,</span> <span class="n">rank</span><span class="o">=</span><span class="n">rank</span><span class="p">,</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">16.0</span><span class="p">)</span>
    
    <span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">d_in</span><span class="p">)</span>
    
    <span class="c1"># 训练模式：LoRA 增量叠加
</span>    <span class="n">out_train</span> <span class="o">=</span> <span class="n">lora_linear</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># (2, 16, 512)
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Train mode output: </span><span class="si">{</span><span class="n">out_train</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    
    <span class="c1"># Merge 权重 → 推理模式
</span>    <span class="n">lora_linear</span><span class="p">.</span><span class="n">merge_weights</span><span class="p">()</span>
    <span class="n">out_inference</span> <span class="o">=</span> <span class="n">lora_linear</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>  <span class="c1"># (2, 16, 512) — 结果应与 merge 前相同
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Inference mode output: </span><span class="si">{</span><span class="n">out_inference</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    
    <span class="c1"># 验证 merge 前后输出一致
</span>    <span class="n">lora_linear2</span> <span class="o">=</span> <span class="n">LinearWithLoRA</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_in</span><span class="p">,</span> <span class="n">d_out</span><span class="p">),</span> <span class="n">rank</span><span class="o">=</span><span class="n">rank</span><span class="p">,</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">16.0</span><span class="p">)</span>
    <span class="c1"># 手动拷贝权重以验证
</span>    <span class="n">lora_linear2</span><span class="p">.</span><span class="n">linear</span><span class="p">.</span><span class="n">weight</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">copy_</span><span class="p">(</span><span class="n">original_linear</span><span class="p">.</span><span class="n">weight</span><span class="p">.</span><span class="n">data</span><span class="p">)</span>
    <span class="n">lora_linear2</span><span class="p">.</span><span class="n">lora</span><span class="p">.</span><span class="n">lora_A</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">copy_</span><span class="p">(</span><span class="n">lora_linear</span><span class="p">.</span><span class="n">lora</span><span class="p">.</span><span class="n">lora_A</span><span class="p">.</span><span class="n">data</span><span class="p">)</span> <span class="k">if</span> <span class="n">lora_linear</span><span class="p">.</span><span class="n">lora</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span> <span class="k">else</span> <span class="bp">None</span>
    
    <span class="c1"># LoRA 参数量对比
</span>    <span class="n">original_params</span> <span class="o">=</span> <span class="n">d_in</span> <span class="o">*</span> <span class="n">d_out</span>
    <span class="n">lora_params</span> <span class="o">=</span> <span class="n">rank</span> <span class="o">*</span> <span class="n">d_in</span> <span class="o">+</span> <span class="n">rank</span> <span class="o">*</span> <span class="n">d_out</span>  <span class="c1"># A + B
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Original params: </span><span class="si">{</span><span class="n">original_params</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"LoRA params: </span><span class="si">{</span><span class="n">lora_params</span><span class="si">}</span><span class="s"> (</span><span class="si">{</span><span class="n">lora_params</span> <span class="o">/</span> <span class="n">original_params</span> <span class="o">*</span> <span class="mi">100</span><span class="si">:</span><span class="p">.</span><span class="mi">2</span><span class="n">f</span><span class="si">}</span><span class="s">%)"</span><span class="p">)</span>
    <span class="c1"># rank=8: LoRA params = 8*512 + 8*512 = 8192 → 只占原始 0.78%
</span></code></pre></div></div>

<p><strong>关键点解析：</strong></p>

<ol>
  <li>
    <p><strong>初始化策略</strong>：A 用 Kaiming 初始化（非零 → 有梯度），B 用零初始化 → LoRA 初始增量 $\Delta W = B \cdot A = 0$ → 模型行为不变 → 训练可以从 base 模型状态开始。</p>
  </li>
  <li>
    <p><strong>Scaling = alpha / rank</strong>：当 rank 从 8 改为 16 → scaling 从 2 变为 1 → 增量矩阵的数值量级不变 → rank 调整不影响训练动态。</p>
  </li>
  <li>
    <p><strong>Merge 权重</strong>：推理时将 LoRA 增量 $\Delta W = B \cdot A \cdot \text{scaling}$ 加到原始权重 $W$ → 只做一次 matmul → 推理零额外开销。</p>
  </li>
  <li>
    <p><strong>x @ A^T @ B^T 的计算顺序</strong>：先降维到 $r$ 再升维 → 中间维度小 → 计算量 $O(d_{in} \cdot r + r \cdot d_{out})$ → 比 $\Delta W \cdot x$ 的 $O(d_{in} \cdot d_{out})$ 大幅减少。</p>
  </li>
</ol>

<p>⚠️ 常见坑：LoRA 只加到 attention 的 Q/V 投影 → 不加到 K/O → 效果可能不如全投影 LoRA。LLaVA 实验证实 LoRA 加到所有投影（Q/K/V/O）效果最好。</p>

<p>⚠️ 另一个坑：merge 后如果想继续训练 → 需要先 unmerge（恢复 LoRA 层）→ 否则全参数训练会改变 merged 权重 → 下次 unmerge 时 LoRA 增量与实际变化不匹配。</p>

<p><strong>复杂度分析：</strong></p>
<ul>
  <li>训练参数量：$2 \cdot r \cdot d$ per layer（假设 $d_{in} = d_{out} = d$）→ rank=8, d=4096 → 65,536 参数 per layer</li>
  <li>训练时间：LoRA 增量计算 $O(B \cdot T \cdot (d \cdot r + r \cdot d))$ → 极小（$r \ll d$）</li>
  <li>推理时间（merge 前）：额外 $O(B \cdot T \cdot d \cdot r)$ → 约 2% 开销</li>
  <li>推理时间（merge 后）：零额外开销</li>
</ul>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>LoRA 的 rank 如何选择？（经验：简单任务 r=4-8 足够；复杂任务 r=16-64；r 过大接近全参数微调 → LoRA 优势丧失）</li>
  <li>LoRA 是否需要加 dropout？（小数据量（&lt;10K）加 dropout=0.1 防过拟合；大数据量（&gt;100K）dropout=0 或不加）</li>
  <li>QLoRA 是什么？（LoRA + 基础模型 4bit 量化 → 只在 LoRA 参数上做 bf16 梯度 → 显存极低（7B 模型 ~6GB））</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2106.09685">LoRA: Low-Rank Adaptation of Large Language Models</a></li>
  <li><a href="https://arxiv.org/abs/2305.14314">QLoRA: Efficient Finetuning of Quantized LLMs</a></li>
</ul>

<hr />

<h3 id="q58-dpo-loss-函数实现给-chosenrejected-logprob-写出-loss">Q58 DPO loss 函数实现（给 chosen/rejected logprob 写出 loss）</h3>

<p><strong>难度：</strong> 专家</p>

<p><strong>考察点：</strong> DPO loss 的精确实现，包括 logprob 计算、reference model 的作用、数值稳定性</p>

<p><strong>满分回答：</strong></p>

<p><strong>DPO loss 公式：</strong>
\(\mathcal{L}_{\text{DPO}} = -\log \sigma\left(\beta \cdot \left(\log \frac{\pi_\theta(y_w|x)}{\pi_{\text{ref}}(y_w|x)} - \log \frac{\pi_\theta(y_l|x)}{\pi_{\text{ref}}(y_l|x)}\right)\right)\)</p>

<p>简化写法：令 $\hat{r}<em>\theta(y, x) = \beta \cdot \log \frac{\pi</em>\theta(y|x)}{\pi_{\text{ref}}(y|x)}$
\(\mathcal{L}_{\text{DPO}} = -\log \sigma(\hat{r}_\theta(y_w, x) - \hat{r}_\theta(y_l, x))\)</p>

<p>其中 $y_w$ = chosen（好回答），$y_l$ = rejected（差回答），$\pi_\theta$ = 训练中的模型，$\pi_{\text{ref}}$ = reference（初始模型），$\beta$ = 温度超参。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>

<span class="k">def</span> <span class="nf">dpo_loss</span><span class="p">(</span>
    <span class="n">policy_chosen_logps</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>    <span class="c1"># (B,) — π_θ(y_w|x) 的 log probability
</span>    <span class="n">policy_rejected_logps</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>  <span class="c1"># (B,) — π_θ(y_l|x) 的 log probability
</span>    <span class="n">ref_chosen_logps</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>       <span class="c1"># (B,) — π_ref(y_w|x) 的 log probability
</span>    <span class="n">ref_rejected_logps</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>     <span class="c1"># (B,) — π_ref(y_l|x) 的 log probability
</span>    <span class="n">beta</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.1</span><span class="p">,</span>
    <span class="n">label_smoothing</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">,</span>
<span class="p">)</span> <span class="o">-&gt;</span> <span class="nb">tuple</span><span class="p">:</span>
    <span class="s">"""
    计算 DPO loss。
    
    Args:
        policy_chosen_logps:   训练模型对 chosen 回答的 log prob
        policy_rejected_logps: 训练模型对 rejected 回答的 log prob
        ref_chosen_logps:      reference 模型对 chosen 回答的 log prob
        ref_rejected_logps:    reference 模型对 rejected 回答的 log prob
        beta:                  DPO 温度参数，控制偏离 reference 的程度
        label_smoothing:       标签平滑系数（0=标准 DPO）
    
    Returns:
        loss:       (scalar) 平均 DPO loss
        chosen_rewards:  (B,) chosen 的隐式 reward
        rejected_rewards: (B,) rejected 的隐式 reward
    """</span>
    <span class="c1"># Step 1: 计算 log ratio（π_θ / π_ref）
</span>    <span class="c1"># log π_θ(y|x) - log π_ref(y|x) = log ratio
</span>    <span class="n">chosen_logratios</span> <span class="o">=</span> <span class="n">policy_chosen_logps</span> <span class="o">-</span> <span class="n">ref_chosen_logps</span>     <span class="c1"># (B,)
</span>    <span class="n">rejected_logratios</span> <span class="o">=</span> <span class="n">policy_rejected_logps</span> <span class="o">-</span> <span class="n">ref_rejected_logps</span>  <span class="c1"># (B,)
</span>    
    <span class="c1"># Step 2: 计算隐式 reward（β * log ratio）
</span>    <span class="n">chosen_rewards</span> <span class="o">=</span> <span class="n">beta</span> <span class="o">*</span> <span class="n">chosen_logratios</span>    <span class="c1"># (B,)
</span>    <span class="n">rejected_rewards</span> <span class="o">=</span> <span class="n">beta</span> <span class="o">*</span> <span class="n">rejected_logratios</span>  <span class="c1"># (B,)
</span>    
    <span class="c1"># Step 3: 计算 DPO loss = -log σ(reward_chosen - reward_rejected)
</span>    <span class="c1"># 数值稳定性关键：用 logsigmoid 而不是 log(σ(x))
</span>    <span class="c1"># log(σ(x)) 在 x 很大时 σ(x)≈1 → log(1)=0，但 x 很小时 σ(x)≈0 → log(0)=-inf → 精度丢失
</span>    <span class="c1"># logsigmoid(x) = -log(1 + exp(-x)) = -softplus(-x)，内部用数值稳定实现
</span>    
    <span class="n">logits</span> <span class="o">=</span> <span class="n">chosen_rewards</span> <span class="o">-</span> <span class="n">rejected_rewards</span>  <span class="c1"># (B,) — chosen reward 优势
</span>    
    <span class="k">if</span> <span class="n">label_smoothing</span> <span class="o">&gt;</span> <span class="mi">0</span><span class="p">:</span>
        <span class="c1"># Label smoothing DPO: DPO with LS = -LS * log σ(-logits) - (1-LS) * log σ(logits)
</span>        <span class="c1"># 等价于：loss = -log σ(logits) * (1-LS) - log σ(-logits) * LS
</span>        <span class="n">losses</span> <span class="o">=</span> <span class="o">-</span><span class="n">F</span><span class="p">.</span><span class="n">logsigmoid</span><span class="p">(</span><span class="n">logits</span><span class="p">)</span> <span class="o">*</span> <span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="n">label_smoothing</span><span class="p">)</span> <span class="o">-</span> <span class="n">F</span><span class="p">.</span><span class="n">logsigmoid</span><span class="p">(</span><span class="o">-</span><span class="n">logits</span><span class="p">)</span> <span class="o">*</span> <span class="n">label_smoothing</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="c1"># 标准 DPO loss
</span>        <span class="n">losses</span> <span class="o">=</span> <span class="o">-</span><span class="n">F</span><span class="p">.</span><span class="n">logsigmoid</span><span class="p">(</span><span class="n">logits</span><span class="p">)</span>  <span class="c1"># (B,)
</span>    
    <span class="n">loss</span> <span class="o">=</span> <span class="n">losses</span><span class="p">.</span><span class="n">mean</span><span class="p">()</span>  <span class="c1"># scalar
</span>    
    <span class="k">return</span> <span class="n">loss</span><span class="p">,</span> <span class="n">chosen_rewards</span><span class="p">,</span> <span class="n">rejected_rewards</span>

<span class="k">def</span> <span class="nf">compute_logprobs</span><span class="p">(</span>
    <span class="n">logits</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>          <span class="c1"># (B, T, V) — 模型输出的 logits
</span>    <span class="n">labels</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>          <span class="c1"># (B, T) — 目标 token ids
</span>    <span class="n">ignore_index</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="o">-</span><span class="mi">100</span><span class="p">,</span>      <span class="c1"># 忽略的 label 位置（如 padding）
</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
    <span class="s">"""
    计算给定序列的 log probability（DPO 所需）。
    
    Args:
        logits: (B, T, V) 模型 logits
        labels: (B, T) token ids
        ignore_index: 忽略的位置标记
    
    Returns:
        logps: (B,) — 每个序列的 log probability（sum over valid tokens）
    """</span>
    <span class="c1"># Step 1: 对每个位置取目标 token 的 log prob
</span>    <span class="c1"># log_softmax → 取 labels 对应位置的值
</span>    <span class="n">log_probs</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">log_softmax</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B, T, V)
</span>    
    <span class="c1"># Step 2: gather 目标 token 的 log prob
</span>    <span class="c1"># labels: (B, T) → 需要扩展维度以匹配 log_probs
</span>    <span class="n">target_log_probs</span> <span class="o">=</span> <span class="n">log_probs</span><span class="p">.</span><span class="n">gather</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">index</span><span class="o">=</span><span class="n">labels</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">))</span>  <span class="c1"># (B, T, 1)
</span>    <span class="n">target_log_probs</span> <span class="o">=</span> <span class="n">target_log_probs</span><span class="p">.</span><span class="n">squeeze</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B, T)
</span>    
    <span class="c1"># Step 3: 忽略 padding/特殊位置
</span>    <span class="n">mask</span> <span class="o">=</span> <span class="n">labels</span> <span class="o">!=</span> <span class="n">ignore_index</span>  <span class="c1"># (B, T)
</span>    <span class="n">target_log_probs</span> <span class="o">=</span> <span class="n">target_log_probs</span> <span class="o">*</span> <span class="n">mask</span>  <span class="c1"># 0 out padding positions
</span>    
    <span class="c1"># Step 4: 对每个序列求和 → 得到序列的 log probability
</span>    <span class="n">logps</span> <span class="o">=</span> <span class="n">target_log_probs</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B,)
</span>    
    <span class="k">return</span> <span class="n">logps</span>

<span class="c1"># === 验证 ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">B</span> <span class="o">=</span> <span class="mi">4</span>
    <span class="c1"># 模拟 log probabilities
</span>    <span class="c1"># chosen 的 log prob 应高于 rejected → loss 应为负（模型倾向 chosen）
</span>    <span class="n">policy_chosen_logps</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">([</span><span class="o">-</span><span class="mf">1.0</span><span class="p">,</span> <span class="o">-</span><span class="mf">2.0</span><span class="p">,</span> <span class="o">-</span><span class="mf">1.5</span><span class="p">,</span> <span class="o">-</span><span class="mf">0.5</span><span class="p">])</span>   <span class="c1"># (B,)
</span>    <span class="n">policy_rejected_logps</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">([</span><span class="o">-</span><span class="mf">3.0</span><span class="p">,</span> <span class="o">-</span><span class="mf">4.0</span><span class="p">,</span> <span class="o">-</span><span class="mf">3.5</span><span class="p">,</span> <span class="o">-</span><span class="mf">2.5</span><span class="p">])</span>  <span class="c1"># (B,)
</span>    <span class="n">ref_chosen_logps</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">([</span><span class="o">-</span><span class="mf">1.2</span><span class="p">,</span> <span class="o">-</span><span class="mf">2.2</span><span class="p">,</span> <span class="o">-</span><span class="mf">1.7</span><span class="p">,</span> <span class="o">-</span><span class="mf">0.7</span><span class="p">])</span>       <span class="c1"># (B,)
</span>    <span class="n">ref_rejected_logps</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">([</span><span class="o">-</span><span class="mf">3.2</span><span class="p">,</span> <span class="o">-</span><span class="mf">4.2</span><span class="p">,</span> <span class="o">-</span><span class="mf">3.7</span><span class="p">,</span> <span class="o">-</span><span class="mf">2.7</span><span class="p">])</span>     <span class="c1"># (B,)
</span>    
    <span class="n">loss</span><span class="p">,</span> <span class="n">chosen_rewards</span><span class="p">,</span> <span class="n">rejected_rewards</span> <span class="o">=</span> <span class="n">dpo_loss</span><span class="p">(</span>
        <span class="n">policy_chosen_logps</span><span class="p">,</span> <span class="n">policy_rejected_logps</span><span class="p">,</span>
        <span class="n">ref_chosen_logps</span><span class="p">,</span> <span class="n">ref_rejected_logps</span><span class="p">,</span>
        <span class="n">beta</span><span class="o">=</span><span class="mf">0.1</span>
    <span class="p">)</span>
    
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"DPO loss: </span><span class="si">{</span><span class="n">loss</span><span class="p">.</span><span class="n">item</span><span class="p">()</span><span class="si">:</span><span class="p">.</span><span class="mi">4</span><span class="n">f</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Chosen rewards:  </span><span class="si">{</span><span class="n">chosen_rewards</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>    <span class="c1"># β * (policy - ref) for chosen
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Rejected rewards: </span><span class="si">{</span><span class="n">rejected_rewards</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># β * (policy - ref) for rejected
</span>    
    <span class="c1"># 验证：当 chosen logratio &gt; rejected logratio → loss &lt; 0（正确方向）
</span>    <span class="c1"># 当 policy 已经偏好 chosen → loss 接近 0（模型已学好）
</span>    
    <span class="c1"># 用 compute_logprobs 验证完整流程
</span>    <span class="n">V</span> <span class="o">=</span> <span class="mi">100</span>
    <span class="n">T</span> <span class="o">=</span> <span class="mi">10</span>
    <span class="n">logits</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">V</span><span class="p">)</span>
    <span class="n">labels</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randint</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">V</span><span class="p">,</span> <span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">))</span>
    <span class="n">labels</span><span class="p">[:,</span> <span class="o">-</span><span class="mi">3</span><span class="p">:]</span> <span class="o">=</span> <span class="o">-</span><span class="mi">100</span>  <span class="c1"># padding
</span>    
    <span class="n">logps</span> <span class="o">=</span> <span class="n">compute_logprobs</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">labels</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Sequence log probs: </span><span class="si">{</span><span class="n">logps</span><span class="p">.</span><span class="n">shape</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>  <span class="c1"># (4,)
</span></code></pre></div></div>

<p><strong>关键点解析：</strong></p>

<ol>
  <li>
    <p><strong>log ratio = log π_θ - log π_ref</strong>：DPO 的核心不是绝对概率而是相对概率的变化。reference model 的作用是”锚定”——防止模型过度偏离原始分布。</p>
  </li>
  <li>
    <p><strong>数值稳定性</strong>：<code class="language-plaintext highlighter-rouge">F.logsigmoid(logits)</code> 而非 <code class="language-plaintext highlighter-rouge">torch.log(torch.sigmoid(logits))</code>。原因：<code class="language-plaintext highlighter-rouge">sigmoid(x)</code> 在 $x$ 很小时 $\approx 0$ → <code class="language-plaintext highlighter-rouge">log(0) = -inf</code>；<code class="language-plaintext highlighter-rouge">logsigmoid</code> 内部用 <code class="language-plaintext highlighter-rouge">-softplus(-x)</code> 计算 → 对所有 $x$ 值都精确。</p>
  </li>
  <li>
    <p><strong>logprob 计算</strong>：用 <code class="language-plaintext highlighter-rouge">log_softmax + gather</code> 而非 <code class="language-plaintext highlighter-rouge">log(probs[target_id])</code> → 更高效（一次 softmax + gather vs 两次）且数值更稳定。</p>
  </li>
  <li>
    <p><strong>ignore_index</strong>：DPO 的 label 通常包含 prompt + response → prompt 部分不计算 loss → 用 ignore_index 标记。</p>
  </li>
</ol>

<p>⚠️ 常见坑：忘记用 reference model → 直接用 <code class="language-plaintext highlighter-rouge">policy_chosen_logps</code> 和 <code class="language-plaintext highlighter-rouge">policy_rejected_logps</code> 做 sigmoid → 这是 Bradley-Terry model 而不是 DPO → 模型会过度偏离 base。</p>

<p>⚠️ 另一个坑：logprob 是对<strong>整个序列</strong>求和 → 长序列的 logprob 绝对值很大 → β 需要相应调整。常见设置：β=0.1 for 序列长度 ~200。</p>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>DPO 和 RLHF 效果对比？（DPO 效果接近 PPO-RLHF，但训练更简单（不需要 reward model + PPO 循环）→ 适合资源有限的场景）</li>
  <li>DPO 的 reference model 如何处理？（实践中用初始模型（SFT 后但未做 DPO 的模型）→ 冻结 → 不更新。或用 SimPO/IPo 去掉 reference model）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2305.18290">Direct Preference Optimization: Your Language Model is Secretly a Reward Model</a></li>
  <li><a href="https://arxiv.org/abs/2405.14734">SimPO: Simple Preference Optimization without Reference Models</a></li>
</ul>

<hr />

<h3 id="q59-ppo-损失含-clipkl-penaltyvalue-lossentropy-bonus实现">Q59 PPO 损失（含 clip、KL penalty、value loss、entropy bonus）实现</h3>

<p><strong>难度：</strong> 专家</p>

<p><strong>考察点：</strong> RLHF-PPO 的完整 loss 组成；理解每个 loss component 的作用与数值稳定性</p>

<p><strong>满分回答：</strong></p>

<p>PPO-RLHF 的总 loss 由四个部分组成：
\(\mathcal{L}_{\text{total}} = \mathcal{L}_{\text{clip}} - c_1 \cdot \mathcal{L}_{\text{value}} + c_2 \cdot \mathcal{H} + c_3 \cdot \text{KL}(\pi_\theta, \pi_{\text{ref}})\)</p>

<p>其中：</p>
<ul>
  <li>$\mathcal{L}_{\text{clip}}$：policy clip loss（核心）</li>
  <li>$\mathcal{L}_{\text{value}}$：value function loss</li>
  <li>$\mathcal{H}$：entropy bonus（鼓励探索）</li>
  <li>KL penalty：约束策略不偏离 reference 太远</li>
</ul>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>

<span class="k">def</span> <span class="nf">ppo_loss</span><span class="p">(</span>
    <span class="n">logprobs</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>          <span class="c1"># (B,) — 当前策略 π_θ(a|s) 的 log prob
</span>    <span class="n">old_logprobs</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>      <span class="c1"># (B,) — 旧策略 π_old(a|s) 的 log prob（用于计算 ratio）
</span>    <span class="n">advantages</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>        <span class="c1"># (B,) — advantage 值 = reward - baseline(value)
</span>    <span class="n">values</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>            <span class="c1"># (B,) — value network 当前预测
</span>    <span class="n">old_values</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>        <span class="c1"># (B,) — 旧 value 预测（用于 value clip）
</span>    <span class="n">returns</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>           <span class="c1"># (B,) — TD target / GAE returns
</span>    <span class="n">ref_logprobs</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>      <span class="c1"># (B,) — reference 模型 π_ref(a|s) 的 log prob
</span>    <span class="n">clip_eps</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.2</span><span class="p">,</span>           <span class="c1"># PPO clip 范围 ε
</span>    <span class="n">value_clip_eps</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.2</span><span class="p">,</span>     <span class="c1"># Value clip 范围
</span>    <span class="n">kl_coeff</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.05</span><span class="p">,</span>          <span class="c1"># KL penalty 权重 c_3
</span>    <span class="n">value_coeff</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.5</span><span class="p">,</span>        <span class="c1"># Value loss 权重 c_1
</span>    <span class="n">entropy_coeff</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.01</span><span class="p">,</span>     <span class="c1"># Entropy bonus 权重 c_2
</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="nb">dict</span><span class="p">:</span>
    <span class="s">"""
    计算 PPO-RLHF 的完整 loss。
    
    Args: 上述各参数
    Returns: dict 包含各 loss component 和总 loss
    """</span>
    <span class="c1"># =============================================
</span>    <span class="c1"># 1. Policy Clip Loss（核心）
</span>    <span class="c1"># =============================================
</span>    <span class="c1"># ratio = π_θ(a|s) / π_old(a|s)
</span>    <span class="c1"># log ratio = logprob - old_logprob（数值更稳定）
</span>    <span class="n">ratio</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">exp</span><span class="p">(</span><span class="n">logprobs</span> <span class="o">-</span> <span class="n">old_logprobs</span><span class="p">)</span>  <span class="c1"># (B,) — 策略变化比率
</span>    
    <span class="c1"># PPO clip: L_clip = -min(ratio * A, clip(ratio, 1-ε, 1+ε) * A)
</span>    <span class="c1"># 当 A &gt; 0（正 advantage）→ 鼓励增大 ratio，但不超过 1+ε
</span>    <span class="c1"># 当 A &lt; 0（负 advantage）→ 鼓励减小 ratio，但不低于 1-ε
</span>    <span class="n">clipped_ratio</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">clamp</span><span class="p">(</span><span class="n">ratio</span><span class="p">,</span> <span class="mf">1.0</span> <span class="o">-</span> <span class="n">clip_eps</span><span class="p">,</span> <span class="mf">1.0</span> <span class="o">+</span> <span class="n">clip_eps</span><span class="p">)</span>  <span class="c1"># (B,)
</span>    
    <span class="c1"># 两个 candidate loss，取 min → 保守更新
</span>    <span class="n">surr1</span> <span class="o">=</span> <span class="n">ratio</span> <span class="o">*</span> <span class="n">advantages</span>          <span class="c1"># (B,) — 无 clip 版
</span>    <span class="n">surr2</span> <span class="o">=</span> <span class="n">clipped_ratio</span> <span class="o">*</span> <span class="n">advantages</span>  <span class="c1"># (B,) — clip 版
</span>    
    <span class="c1"># 优势为正时 min(ratio*A, clipped*A)：ratio 太大被 clip → 防止过大更新
</span>    <span class="c1"># 优势为负时 min(ratio*A, clipped*A)：ratio 太小被 clip → 防止过大回退
</span>    <span class="n">policy_loss</span> <span class="o">=</span> <span class="o">-</span><span class="n">torch</span><span class="p">.</span><span class="nb">min</span><span class="p">(</span><span class="n">surr1</span><span class="p">,</span> <span class="n">surr2</span><span class="p">).</span><span class="n">mean</span><span class="p">()</span>  <span class="c1"># scalar
</span>    
    <span class="c1"># =============================================
</span>    <span class="c1"># 2. Value Loss（value function 回归）
</span>    <span class="c1"># =============================================
</span>    <span class="c1"># L_value = (V_θ(s) - returns)^2
</span>    <span class="c1"># 也可用 clipped value：防止单步 value 变化过大
</span>    <span class="n">values_clipped</span> <span class="o">=</span> <span class="n">old_values</span> <span class="o">+</span> <span class="n">torch</span><span class="p">.</span><span class="n">clamp</span><span class="p">(</span><span class="n">values</span> <span class="o">-</span> <span class="n">old_values</span><span class="p">,</span> <span class="o">-</span><span class="n">value_clip_eps</span><span class="p">,</span> <span class="n">value_clip_eps</span><span class="p">)</span>
    <span class="n">value_loss1</span> <span class="o">=</span> <span class="p">(</span><span class="n">values</span> <span class="o">-</span> <span class="n">returns</span><span class="p">).</span><span class="nb">pow</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span>         <span class="c1"># (B,) — 无 clip
</span>    <span class="n">value_loss2</span> <span class="o">=</span> <span class="p">(</span><span class="n">values_clipped</span> <span class="o">-</span> <span class="n">returns</span><span class="p">).</span><span class="nb">pow</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span>  <span class="c1"># (B,) — clip 版
</span>    <span class="n">value_loss</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">value_loss1</span><span class="p">,</span> <span class="n">value_loss2</span><span class="p">).</span><span class="n">mean</span><span class="p">()</span>  <span class="c1"># scalar
</span>    
    <span class="c1"># =============================================
</span>    <span class="c1"># 3. Entropy Bonus（鼓励探索）
</span>    <span class="c1"># =============================================
</span>    <span class="c1"># H(π_θ) = -Σ π_θ(a) * log π_θ(a)
</span>    <span class="c1"># 在 LLM 中，entropy 从 log_softmax 直接计算
</span>    <span class="c1"># 注意：这里 logprobs 是对选定 action 的 log prob
</span>    <span class="c1"># 真正的 entropy 需要对所有 token 的分布计算 → 实际中可用简化版
</span>    <span class="c1"># 简化版：entropy ≈ -logprobs.mean()（近似，非精确）
</span>    <span class="c1"># 精确版需要传完整 logits → 见下方 entropy_from_logits
</span>    <span class="n">entropy</span> <span class="o">=</span> <span class="o">-</span><span class="n">logprobs</span><span class="p">.</span><span class="n">mean</span><span class="p">()</span>  <span class="c1"># (B,) 简化近似版 → scalar after mean
</span>    <span class="c1"># 更精确版：
</span>    <span class="c1"># entropy = entropy_from_logits(all_logits)  # 需要完整 logits
</span>    
    <span class="c1"># =============================================
</span>    <span class="c1"># 4. KL Penalty（约束偏离 reference）
</span>    <span class="c1"># =============================================
</span>    <span class="c1"># KL(π_θ || π_ref) = Σ π_θ(a) * [log π_θ(a) - log π_ref(a)]
</span>    <span class="c1"># = Σ π_θ(a) * (logprobs - ref_logprobs)
</span>    <span class="c1"># 当 logprobs 和 ref_logprobs 是对选定 token 的 log prob 时：
</span>    <span class="c1"># KL ≈ (logprobs - ref_logprobs).mean()  → 这是近似（非严格 KL）
</span>    <span class="c1"># 严格 KL 需要完整 token 分布 → 实际 RLHF 中也用这个近似
</span>    <span class="n">kl_penalty</span> <span class="o">=</span> <span class="p">(</span><span class="n">logprobs</span> <span class="o">-</span> <span class="n">ref_logprobs</span><span class="p">).</span><span class="n">mean</span><span class="p">()</span>  <span class="c1"># scalar
</span>    
    <span class="c1"># =============================================
</span>    <span class="c1"># 5. 总 Loss
</span>    <span class="c1"># =============================================
</span>    <span class="n">total_loss</span> <span class="o">=</span> <span class="p">(</span>
        <span class="n">policy_loss</span>
        <span class="o">-</span> <span class="n">value_coeff</span> <span class="o">*</span> <span class="n">value_loss</span>  <span class="c1"># 减号：最小化 value loss = 最大化 -value_loss
</span>        <span class="o">+</span> <span class="n">entropy_coeff</span> <span class="o">*</span> <span class="n">entropy</span>    <span class="c1"># 加号：最大化 entropy
</span>        <span class="o">+</span> <span class="n">kl_coeff</span> <span class="o">*</span> <span class="n">kl_penalty</span>      <span class="c1"># 加号：KL penalty 防止偏离（实际是惩罚）
</span>    <span class="p">)</span>
    
    <span class="k">return</span> <span class="p">{</span>
        <span class="s">"total_loss"</span><span class="p">:</span> <span class="n">total_loss</span><span class="p">,</span>
        <span class="s">"policy_loss"</span><span class="p">:</span> <span class="n">policy_loss</span><span class="p">,</span>
        <span class="s">"value_loss"</span><span class="p">:</span> <span class="n">value_loss</span><span class="p">,</span>
        <span class="s">"entropy"</span><span class="p">:</span> <span class="n">entropy</span><span class="p">,</span>
        <span class="s">"kl_penalty"</span><span class="p">:</span> <span class="n">kl_penalty</span><span class="p">,</span>
        <span class="s">"ratio"</span><span class="p">:</span> <span class="n">ratio</span><span class="p">.</span><span class="n">mean</span><span class="p">(),</span>
        <span class="s">"clipped_frac"</span><span class="p">:</span> <span class="p">((</span><span class="n">ratio</span> <span class="o">-</span> <span class="n">clipped_ratio</span><span class="p">).</span><span class="nb">abs</span><span class="p">()</span> <span class="o">&gt;</span> <span class="mi">0</span><span class="p">).</span><span class="nb">float</span><span class="p">().</span><span class="n">mean</span><span class="p">(),</span>  <span class="c1"># clip 比例
</span>    <span class="p">}</span>

<span class="k">def</span> <span class="nf">entropy_from_logits</span><span class="p">(</span><span class="n">logits</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
    <span class="s">"""
    从完整 logits 计算精确 entropy。
    Args:
        logits: (B, T, V) 模型 logits
    Returns:
        entropy: (B, T) 每个位置的 entropy
    """</span>
    <span class="c1"># H = -Σ p(a) * log p(a) = -Σ softmax(logits) * log_softmax(logits)
</span>    <span class="n">probs</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B, T, V)
</span>    <span class="n">log_probs</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">log_softmax</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B, T, V)
</span>    <span class="n">entropy</span> <span class="o">=</span> <span class="o">-</span><span class="p">(</span><span class="n">probs</span> <span class="o">*</span> <span class="n">log_probs</span><span class="p">).</span><span class="nb">sum</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B, T)
</span>    <span class="k">return</span> <span class="n">entropy</span>

<span class="c1"># === 验证 ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">B</span> <span class="o">=</span> <span class="mi">8</span>
    
    <span class="c1"># 模拟各组件
</span>    <span class="n">logprobs</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">)</span>
    <span class="n">old_logprobs</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">)</span>
    <span class="n">advantages</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">)</span>
    <span class="n">values</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">)</span>
    <span class="n">old_values</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">)</span>
    <span class="n">returns</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">)</span> <span class="o">+</span> <span class="mf">0.5</span>  <span class="c1"># returns 略高于 values
</span>    <span class="n">ref_logprobs</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">)</span>
    
    <span class="n">result</span> <span class="o">=</span> <span class="n">ppo_loss</span><span class="p">(</span>
        <span class="n">logprobs</span><span class="p">,</span> <span class="n">old_logprobs</span><span class="p">,</span> <span class="n">advantages</span><span class="p">,</span>
        <span class="n">values</span><span class="p">,</span> <span class="n">old_values</span><span class="p">,</span> <span class="n">returns</span><span class="p">,</span> <span class="n">ref_logprobs</span><span class="p">,</span>
        <span class="n">clip_eps</span><span class="o">=</span><span class="mf">0.2</span><span class="p">,</span> <span class="n">kl_coeff</span><span class="o">=</span><span class="mf">0.05</span><span class="p">,</span> <span class="n">value_coeff</span><span class="o">=</span><span class="mf">0.5</span><span class="p">,</span> <span class="n">entropy_coeff</span><span class="o">=</span><span class="mf">0.01</span>
    <span class="p">)</span>
    
    <span class="k">for</span> <span class="n">key</span><span class="p">,</span> <span class="n">val</span> <span class="ow">in</span> <span class="n">result</span><span class="p">.</span><span class="n">items</span><span class="p">():</span>
        <span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">val</span><span class="p">,</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">)</span> <span class="ow">and</span> <span class="n">val</span><span class="p">.</span><span class="n">dim</span><span class="p">()</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
            <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"</span><span class="si">{</span><span class="n">key</span><span class="si">}</span><span class="s">: </span><span class="si">{</span><span class="n">val</span><span class="p">.</span><span class="n">item</span><span class="p">()</span><span class="si">:</span><span class="p">.</span><span class="mi">4</span><span class="n">f</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
        <span class="k">else</span><span class="p">:</span>
            <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"</span><span class="si">{</span><span class="n">key</span><span class="si">}</span><span class="s">: </span><span class="si">{</span><span class="n">val</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    
    <span class="c1"># 验证：当 ratio ≈ 1（策略没变）→ policy_loss ≈ -advantages.mean()
</span>    <span class="c1"># 当 advantages &gt; 0 → policy_loss &lt; 0 → 鼓励增大该动作概率
</span>    <span class="c1"># 当 advantages &lt; 0 → policy_loss &gt; 0 → 鼓励减小该动作概率
</span></code></pre></div></div>

<p><strong>关键点解析：</strong></p>

<ol>
  <li>
    <p><strong>Ratio = exp(logprob - old_logprob)</strong>：用 log space 计算更稳定（避免数值溢出/下溢）→ 再 exp 回 ratio。</p>
  </li>
  <li>
    <p><strong>Clip 机制</strong>：<code class="language-plaintext highlighter-rouge">torch.clamp(ratio, 1-ε, 1+ε)</code> → 限制单步策略变化幅度 → 防止”策略崩塌”（一步更新太大导致后续更新方向错误）。</p>
  </li>
  <li><strong>KL penalty 的两种计算方式</strong>：
    <ul>
      <li>精确 KL：需要完整 token 分布的 softmax → 计算量大</li>
      <li>近似 KL：<code class="language-plaintext highlighter-rouge">(logprobs - ref_logprobs).mean()</code> → 只在选定 token 上计算 → 实际 RLHF 中广泛使用</li>
      <li>LLaMA-2 的 RLHF 报告指出：用<strong>KL coefficient 自适应调整</strong>比固定 KL penalty 更稳定——当 KL &gt; target → 增大 kl_coeff；当 KL &lt; target → 减小 kl_coeff</li>
    </ul>
  </li>
  <li><strong>Value clip</strong>：与 policy clip 类似，防止 value network 单步变化过大 → 训练更稳定。</li>
</ol>

<p>⚠️ 常见坑：PPO 中 advantage 需要做<strong>归一化</strong>（mean=0, std=1）→ 不做归一化 → reward scale 影响更新幅度 → 训练不稳定。</p>

<p>⚠️ 另一个坑：KL penalty 用 <code class="language-plaintext highlighter-rouge">(logprobs - ref_logprobs)</code> 而非严格 KL → 近似在策略接近 reference 时偏差小 → 策略偏离较大时偏差大 → 需要 KL coefficient 自适应弥补。</p>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>PPO 的 clip fraction 正常值是多少？（5-15% → 太高意味着策略变化过快 → 可能不稳定；太低意味着策略几乎没更新 → 学习慢）</li>
  <li>PPO-RLHF 训练需要几个模型？（4个：policy model、reference model（冻结）、reward model（冻结）、value model；可共用 policy 和 value 的部分参数）</li>
  <li>GRPO 和 PPO 的核心区别？（GRPO 没有 value model → 用组内 reward 归一化替代 baseline → 更简单但需要更多采样）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/1707.06347">Proximal Policy Optimization Algorithms</a></li>
  <li><a href="https://arxiv.org/abs/2307.05608">Secrets of RLHF in Large Language Models Part I: PPO</a></li>
  <li><a href="https://arxiv.org/abs/2204.05862">Training a Helpful and Harmless Assistant with RLHF</a></li>
</ul>

<hr />

<h3 id="q60-cross-entropy-with-label-smoothing--sft-loss带-ignore_index实现">Q60 Cross-Entropy with label smoothing / SFT loss（带 ignore_index）实现</h3>

<p><strong>难度：</strong> 基础</p>

<p><strong>考察点：</strong> SFT loss 的精确实现；理解 label smoothing、ignore_index、padding 处理的工程细节</p>

<p><strong>满分回答：</strong></p>

<p>SFT 的 loss 本质上是 cross-entropy loss，但需要处理几个工程细节：</p>

<ol>
  <li><strong>ignore_index</strong>：prompt 部分的 token 不计算 loss（只算 response 部分）</li>
  <li><strong>label smoothing</strong>：软化硬标签 → 防止模型过度自信 → 提升泛化</li>
  <li><strong>padding mask</strong>：padding token 不计算 loss</li>
</ol>

<p><strong>标准 Cross-Entropy：</strong>
\(\text{CE}(y, \hat{y}) = -\sum_i y_i \log \hat{y}_i\)
其中 $y$ 是 one-hot 标签，$\hat{y}$ 是 softmax 预测。</p>

<p><strong>Label Smoothing Cross-Entropy：</strong>
\(y_i^{\text{smooth}} = (1 - \epsilon) \cdot y_i + \frac{\epsilon}{V}\)
\(\text{CE}_{\text{smooth}} = -\sum_i y_i^{\text{smooth}} \log \hat{y}_i\)</p>

<p>对 one-hot 标签：$y_i^{\text{smooth}} = (1-\epsilon)$ for target class, $\frac{\epsilon}{V}$ for others。</p>

<p>简化公式：
\(\text{CE}_{\text{smooth}} = (1-\epsilon) \cdot (-\log \hat{y}_{\text{target}}) + \epsilon \cdot (-\frac{1}{V} \sum_i \log \hat{y}_i)\)</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>

<span class="k">def</span> <span class="nf">sft_loss</span><span class="p">(</span>
    <span class="n">logits</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>        <span class="c1"># (B, T, V) — 模型 logits
</span>    <span class="n">labels</span><span class="p">:</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span>        <span class="c1"># (B, T) — 目标 token ids
</span>    <span class="n">ignore_index</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="o">-</span><span class="mi">100</span><span class="p">,</span>    <span class="c1"># 忽略的位置标记
</span>    <span class="n">label_smoothing</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">,</span> <span class="c1"># label smoothing 系数 ε
</span>    <span class="n">reduction</span><span class="p">:</span> <span class="nb">str</span> <span class="o">=</span> <span class="s">"mean"</span><span class="p">,</span>      <span class="c1"># "mean" or "sum"
</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span>
    <span class="s">"""
    SFT loss = Cross-Entropy with label smoothing + ignore_index
    
    Args:
        logits:          (B, T, V) 模型 logits
        labels:          (B, T) 目标 token ids，-100 表示不计算 loss
        ignore_index:    不计算 loss 的 label 值
        label_smoothing: 0=标准 CE, &gt;0=平滑版本
        reduction:       "mean"（除以有效 token 数）或 "sum"
    
    Returns:
        loss: scalar
    """</span>
    <span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">V</span> <span class="o">=</span> <span class="n">logits</span><span class="p">.</span><span class="n">shape</span>
    
    <span class="c1"># Step 1: 找到有效位置（label != ignore_index）
</span>    <span class="n">valid_mask</span> <span class="o">=</span> <span class="n">labels</span> <span class="o">!=</span> <span class="n">ignore_index</span>  <span class="c1"># (B, T) bool
</span>    
    <span class="c1"># Step 2: 创建 shift 关系（next-token prediction）
</span>    <span class="c1"># SFT 的 labels 是 shifted 的：logits[i] 预测 labels[i+1]
</span>    <span class="c1"># 但通常数据已经做好了 shift → 此处假设 labels 已经 shift
</span>    
    <span class="c1"># Step 3: 将 logits 和 labels reshape 为 2D 以使用 F.cross_entropy
</span>    <span class="c1"># F.cross_entropy 需要 (N, V) logits 和 (N,) labels
</span>    <span class="n">logits_flat</span> <span class="o">=</span> <span class="n">logits</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">V</span><span class="p">)</span>       <span class="c1"># (B*T, V)
</span>    <span class="n">labels_flat</span> <span class="o">=</span> <span class="n">labels</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span>           <span class="c1"># (B*T,)
</span>    <span class="n">valid_mask_flat</span> <span class="o">=</span> <span class="n">valid_mask</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span>   <span class="c1"># (B*T,)
</span>    
    <span class="c1"># Step 4: 替换 ignore_index 的 label 为 0（F.cross_entropy 不处理 -100 的 label）
</span>    <span class="c1"># F.cross_entropy 内置了 ignore_index → 可以直接传 -100
</span>    <span class="c1"># 但如果用 label smoothing → 需要手动处理
</span>    
    <span class="k">if</span> <span class="n">label_smoothing</span> <span class="o">&gt;</span> <span class="mi">0</span><span class="p">:</span>
        <span class="c1"># === Label Smoothing 版本（需手动实现）===
</span>        <span class="c1"># F.cross_entropy 不支持 label_smoothing + ignore_index 同时使用
</span>        <span class="c1"># → 需要手动计算
</span>        
        <span class="c1"># log_softmax
</span>        <span class="n">log_probs</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">log_softmax</span><span class="p">(</span><span class="n">logits_flat</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B*T, V)
</span>        
        <span class="c1"># 构造 smooth label
</span>        <span class="c1"># 对于 target token: weight = (1 - ε)
</span>        <span class="c1"># 对于其他 token: weight = ε / V
</span>        <span class="n">smooth_labels</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">full_like</span><span class="p">(</span><span class="n">log_probs</span><span class="p">,</span> <span class="n">label_smoothing</span> <span class="o">/</span> <span class="n">V</span><span class="p">)</span>  <span class="c1"># (B*T, V) ε/V
</span>        <span class="c1"># 在 target 位置加上 (1 - ε)
</span>        <span class="c1"># 注意：ignore_index 的位置不计算 → 先把 -100 替换为 0（任意有效 id）
</span>        <span class="n">safe_labels</span> <span class="o">=</span> <span class="n">labels_flat</span><span class="p">.</span><span class="n">clone</span><span class="p">()</span>
        <span class="n">safe_labels</span><span class="p">[</span><span class="o">~</span><span class="n">valid_mask_flat</span><span class="p">]</span> <span class="o">=</span> <span class="mi">0</span>  <span class="c1"># 替换为 0（不参与 loss 计算）
</span>        <span class="n">smooth_labels</span><span class="p">.</span><span class="n">scatter_</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">safe_labels</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="p">(</span><span class="mf">1.0</span> <span class="o">-</span> <span class="n">label_smoothing</span> <span class="o">+</span> <span class="n">label_smoothing</span> <span class="o">/</span> <span class="n">V</span><span class="p">))</span>  <span class="c1"># (B*T, V)
</span>        
        <span class="c1"># 计算 CE: -Σ smooth_label * log_prob
</span>        <span class="n">per_token_loss</span> <span class="o">=</span> <span class="o">-</span><span class="p">(</span><span class="n">smooth_labels</span> <span class="o">*</span> <span class="n">log_probs</span><span class="p">).</span><span class="nb">sum</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># (B*T,)
</span>        
        <span class="c1"># mask out invalid positions
</span>        <span class="n">per_token_loss</span> <span class="o">=</span> <span class="n">per_token_loss</span> <span class="o">*</span> <span class="n">valid_mask_flat</span><span class="p">.</span><span class="nb">float</span><span class="p">()</span>  <span class="c1"># (B*T,)
</span>        
        <span class="c1"># reduction
</span>        <span class="k">if</span> <span class="n">reduction</span> <span class="o">==</span> <span class="s">"mean"</span><span class="p">:</span>
            <span class="n">loss</span> <span class="o">=</span> <span class="n">per_token_loss</span><span class="p">.</span><span class="nb">sum</span><span class="p">()</span> <span class="o">/</span> <span class="n">valid_mask_flat</span><span class="p">.</span><span class="nb">float</span><span class="p">().</span><span class="nb">sum</span><span class="p">()</span>
        <span class="k">elif</span> <span class="n">reduction</span> <span class="o">==</span> <span class="s">"sum"</span><span class="p">:</span>
            <span class="n">loss</span> <span class="o">=</span> <span class="n">per_token_loss</span><span class="p">.</span><span class="nb">sum</span><span class="p">()</span>
        <span class="k">else</span><span class="p">:</span>
            <span class="n">loss</span> <span class="o">=</span> <span class="n">per_token_loss</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="c1"># === 标准 Cross-Entropy 版本（F.cross_entropy 内置 ignore_index）===
</span>        <span class="n">loss</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">cross_entropy</span><span class="p">(</span>
            <span class="n">logits_flat</span><span class="p">,</span>
            <span class="n">labels_flat</span><span class="p">,</span>
            <span class="n">ignore_index</span><span class="o">=</span><span class="n">ignore_index</span><span class="p">,</span>
            <span class="n">reduction</span><span class="o">=</span><span class="n">reduction</span><span class="p">,</span>
        <span class="p">)</span>
    
    <span class="k">return</span> <span class="n">loss</span>

<span class="c1"># === 验证 ===
</span><span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">V</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">8</span><span class="p">,</span> <span class="mi">100</span>
    
    <span class="c1"># 模拟模型 logits
</span>    <span class="n">logits</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">,</span> <span class="n">V</span><span class="p">)</span>
    
    <span class="c1"># 模拟 labels：前 3 个 token 是 prompt（ignore），后 5 个是 response（计算 loss）
</span>    <span class="n">labels</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randint</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">V</span><span class="p">,</span> <span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">))</span>
    <span class="n">labels</span><span class="p">[:,</span> <span class="p">:</span><span class="mi">3</span><span class="p">]</span> <span class="o">=</span> <span class="o">-</span><span class="mi">100</span>  <span class="c1"># prompt 部分不计算 loss
</span>    
    <span class="c1"># 标准 SFT loss
</span>    <span class="n">loss_std</span> <span class="o">=</span> <span class="n">sft_loss</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span> <span class="n">ignore_index</span><span class="o">=-</span><span class="mi">100</span><span class="p">,</span> <span class="n">label_smoothing</span><span class="o">=</span><span class="mf">0.0</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Standard CE loss: </span><span class="si">{</span><span class="n">loss_std</span><span class="p">.</span><span class="n">item</span><span class="p">()</span><span class="si">:</span><span class="p">.</span><span class="mi">4</span><span class="n">f</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    
    <span class="c1"># Label smoothing SFT loss
</span>    <span class="n">loss_smooth</span> <span class="o">=</span> <span class="n">sft_loss</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span> <span class="n">ignore_index</span><span class="o">=-</span><span class="mi">100</span><span class="p">,</span> <span class="n">label_smoothing</span><span class="o">=</span><span class="mf">0.1</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Label smoothing loss (ε=0.1): </span><span class="si">{</span><span class="n">loss_smooth</span><span class="p">.</span><span class="n">item</span><span class="p">()</span><span class="si">:</span><span class="p">.</span><span class="mi">4</span><span class="n">f</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    
    <span class="c1"># 验证：label smoothing loss 应略高于标准 CE（因为给非目标 token 也分配了概率）
</span>    
    <span class="c1"># 验证 ignore_index 效果
</span>    <span class="c1"># 全部 label 都计算 vs 只有 response 部分
</span>    <span class="n">labels_all</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randint</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">V</span><span class="p">,</span> <span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">T</span><span class="p">))</span>  <span class="c1"># 没有 -100
</span>    <span class="n">loss_all</span> <span class="o">=</span> <span class="n">sft_loss</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">labels_all</span><span class="p">,</span> <span class="n">ignore_index</span><span class="o">=-</span><span class="mi">100</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Loss (all tokens): </span><span class="si">{</span><span class="n">loss_all</span><span class="p">.</span><span class="n">item</span><span class="p">()</span><span class="si">:</span><span class="p">.</span><span class="mi">4</span><span class="n">f</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Loss (only response): </span><span class="si">{</span><span class="n">loss_std</span><span class="p">.</span><span class="n">item</span><span class="p">()</span><span class="si">:</span><span class="p">.</span><span class="mi">4</span><span class="n">f</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    <span class="c1"># 只算 response 的 loss 应更高（因为分母更小：10 vs 16 tokens）
</span>    
    <span class="c1"># 对比 PyTorch 内置实现
</span>    <span class="n">loss_builtin</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">cross_entropy</span><span class="p">(</span><span class="n">logits</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">V</span><span class="p">),</span> <span class="n">labels</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">),</span> <span class="n">ignore_index</span><span class="o">=-</span><span class="mi">100</span><span class="p">,</span> <span class="n">reduction</span><span class="o">=</span><span class="s">"mean"</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"PyTorch builtin CE: </span><span class="si">{</span><span class="n">loss_builtin</span><span class="p">.</span><span class="n">item</span><span class="p">()</span><span class="si">:</span><span class="p">.</span><span class="mi">4</span><span class="n">f</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Our implementation: </span><span class="si">{</span><span class="n">loss_std</span><span class="p">.</span><span class="n">item</span><span class="p">()</span><span class="si">:</span><span class="p">.</span><span class="mi">4</span><span class="n">f</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    <span class="c1"># 应完全一致
</span></code></pre></div></div>

<p><strong>关键点解析：</strong></p>

<ol>
  <li>
    <p><strong>ignore_index 的作用</strong>：SFT 数据格式为 <code class="language-plaintext highlighter-rouge">prompt + response</code> → loss 只在 response 部分（自回归地预测下一个 token）。prompt 部分的 token 标为 -100 → 不计算 loss → 不更新梯度。</p>
  </li>
  <li>
    <p><strong>Shift 关系</strong>：SFT 的 labels 做 shift → <code class="language-plaintext highlighter-rouge">labels[i] = input_ids[i+1]</code>（预测下一个 token）。通常数据预处理已完成 shift → 代码中假设 labels 已 shift。</p>
  </li>
  <li>
    <p><strong>Label smoothing 实现</strong>：<code class="language-plaintext highlighter-rouge">smooth_labels</code> 构造方式——先用 $\epsilon/V$ 填充所有位置 → 再用 <code class="language-plaintext highlighter-rouge">scatter_</code> 在 target 位置设 $(1-\epsilon + \epsilon/V)$。这样做避免了手动循环 → 更高效。</p>
  </li>
  <li>
    <p><strong>F.cross_entropy 的局限</strong>：内置了 <code class="language-plaintext highlighter-rouge">ignore_index</code> 和 <code class="language-plaintext highlighter-rouge">label_smoothing</code> 参数，但两者<strong>不能同时使用</strong>（PyTorch &lt;2.3 的版本）→ 需要手动实现 label smoothing 版本。</p>
  </li>
</ol>

<p>⚠️ 常见坑：忘记做 shift → labels 和 logits 位置不对齐 → loss 计算完全错误。典型 SFT 数据：<code class="language-plaintext highlighter-rouge">input_ids = [prompt_tokens, response_tokens]</code>, <code class="language-plaintext highlighter-rouge">labels = [-100, -100, ..., response_tokens]</code>（shift 后）。</p>

<p>⚠️ 另一个坑：label smoothing ε=0.1 → 以为会让模型”更不确定” → 实际上 LS 的主要效果是<strong>防止模型过度自信</strong>（log prob 极高）→ 提升泛化性。ε=0 在训练集上可能效果更好但泛化差。</p>

<p>⚠️ 另一个常见错误：<code class="language-plaintext highlighter-rouge">reduction="mean"</code> 的分母是<strong>所有 token 数</strong>而非<strong>有效 token 数</strong> → PyTorch 的 F.cross_entropy 用<strong>有效 token 数</strong>做分母 → 注意区别。</p>

<p><strong>复杂度分析：</strong></p>
<ul>
  <li>时间：$O(B \cdot T \cdot V)$（log_softmax 是瓶颈 → 需要对 V 维做 softmax）</li>
  <li>显存：$O(B \cdot T \cdot V)$（log_probs 存储）</li>
  <li>实际中 $V$ 通常很大（32K-128K）→ log_softmax 是 SFT 训练的主要计算瓶颈</li>
</ul>

<p><strong>面试官常见追问：</strong></p>
<ul>
  <li>为什么 SFT loss 只在 response 部分计算？（prompt 是已知的 → 模型不需要学习预测 prompt → 若计算 prompt 的 loss → 模型会花精力学 prompt 的模式 → 浪费）</li>
  <li>label smoothing 在 SFT 中有用吗？（有用但效果有限：ε=0.1 在 LLaMA-2 的 SFT 中约提升 0.5-1% 泛化。但ε太大（&gt;0.2）→ 模型太不确定 → 辧出质量下降）</li>
  <li>如何处理多轮对话的 SFT loss？（只在最后一轮 response 计算 loss → 前面轮次的 response 标为 -100。或全部 response 都计算 loss → 效果通常更好）</li>
</ul>

<p><strong>参考资料：</strong></p>
<ul>
  <li><a href="https://arxiv.org/abs/2203.02155">Training language models to follow instructions with human feedback (InstructGPT)</a></li>
  <li><a href="https://arxiv.org/abs/1512.00567">Label Smoothing</a> (Szegedy et al., 2016)</li>
</ul>

<hr />]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><summary type="html"><![CDATA[LLM / MLLM 后训练面试题库（60 题精选）]]></summary></entry><entry><title type="html">Diffusion 模型的条件注入演进史：从通道拼接到单流 DiT</title><link href="https://liyongzhi.xyz/posts/2026/05/diffusion-condition-injection/" rel="alternate" type="text/html" title="Diffusion 模型的条件注入演进史：从通道拼接到单流 DiT" /><published>2026-05-09T00:00:00+08:00</published><updated>2026-05-09T00:00:00+08:00</updated><id>https://liyongzhi.xyz/posts/2026/05/blog-post-diffusion-condition-injection</id><content type="html" xml:base="https://liyongzhi.xyz/posts/2026/05/diffusion-condition-injection/"><![CDATA[<p>如果你看过 Stable Diffusion、ControlNet、IP-Adapter，又听说过最近的 Qwen-Image 和 Z-Image，可能会有一个共同的疑问：</p>

<blockquote>
  <p>这些模型架构看上去差别很大，但它们要解决的问题其实都是同一个：<strong>怎么把”用户想要什么”这件事告诉模型？</strong></p>
</blockquote>

<p>文本 prompt、参考图、姿态骨架、深度图、音频、mask……每一种”控制信号”进入网络的方式都不一样。这篇文章想做的事情是：</p>

<blockquote>
  <p>把过去几年 Diffusion 里<strong>主流的条件注入方式</strong>串成一条线，讲清楚每一步的动机、做法、优劣，以及它们最后是怎么汇聚到今天的 DiT 体系里的。</p>
</blockquote>

<p>读完之后，你应该能：</p>

<ul>
  <li>理解为什么 inpainting 用 channel concat，而 text-to-image 用 cross-attention；</li>
  <li>知道 ControlNet 和 IP-Adapter 各自解决的是什么 prompt 解决不了的问题；</li>
  <li>看懂 adaLN-Zero 为什么是 DiT 的”默认调制器”；</li>
  <li>理解 Qwen-Image 的多流 MMDiT 和 Z-Image 的单流 S3-DiT 在条件控制上到底差在哪里；</li>
  <li>对未来 Diffusion 架构的演进方向有一个直觉性的判断。</li>
</ul>

<hr />

<h2 id="1-起点diffusion-模型为什么需要条件">1. 起点：Diffusion 模型为什么需要”条件”</h2>

<p>在最朴素的 DDPM 里，模型学的是数据分布 <code class="language-plaintext highlighter-rouge">p(x)</code>：</p>

<blockquote>
  <p>什么样的样本看起来像”真实世界中的自然图像”。</p>
</blockquote>

<p>它知道怎么把高斯噪声慢慢拉回自然图像流形，但它<strong>不知道你想要什么</strong>。给它噪声，它能生成一张图，但你没法控制是猫是狗、是写实还是油画、是白天还是夜晚。</p>

<p>所以条件 Diffusion 真正学的是：</p>

\[p(x \mid c)\]

<p>也就是”给定条件 <code class="language-plaintext highlighter-rouge">c</code> 的情况下，样本 <code class="language-plaintext highlighter-rouge">x</code> 长什么样”。这里的 <code class="language-plaintext highlighter-rouge">c</code> 可以是文本、参考图、姿态骨架、音频、深度图……几乎任何能数字化表示的”用户意图”。</p>

<p>问题来了：<strong><code class="language-plaintext highlighter-rouge">c</code> 怎么进入模型？</strong></p>

<p>这看似是个工程细节，但它直接决定了模型能不能用这个条件、能用得多准、计算成本多大。过去几年，Diffusion 社区在这个问题上其实经历了好几代演化。我们一个一个看。</p>

<hr />

<h2 id="2-第一代通道拼接-channel-concat">2. 第一代：通道拼接 (Channel Concat)</h2>

<h3 id="21-动机inpainting-是最早的条件生成">2.1 动机：inpainting 是最早的”条件生成”</h3>

<p>最早的”有条件”扩散模型其实就是 <strong>inpainting</strong> —— 给一张图、一个 mask，让模型把 mask 区域补全。这种任务的特点是：</p>

<ul>
  <li>条件本身是<strong>和图像同一空间结构</strong>的（mask 是 H×W 的图，masked image 也是 H×W 的图）；</li>
  <li>用户的意图就是”按照原图的结构和上下文，把这块补好”。</li>
</ul>

<p>那最简单的做法是什么？</p>

<p><strong>直接把 mask 和 masked image 当成额外通道，拼到 noisy latent 上一起送进 U-Net。</strong></p>

<p>Stable Diffusion Inpainting 就是这么干的。原本的 U-Net 输入是 <code class="language-plaintext highlighter-rouge">[B, 4, H, W]</code>（4 个 latent channel），inpainting 版本变成 <code class="language-plaintext highlighter-rouge">[B, 9, H, W]</code> —— 多出来的 5 个通道分别是 1 个 mask 和 4 个 masked latent。</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>unet_input = concat[
    noisy_latent,        # [B, 4, H, W]   被去噪的对象
    mask,                # [B, 1, H, W]   告诉模型哪里要改
    masked_latent,       # [B, 4, H, W]   告诉模型其它地方长什么样
]
</code></pre></div></div>

<h3 id="22-latentsync把这个思路用到-lip-sync">2.2 LatentSync：把这个思路用到 lip-sync</h3>

<p>最近字节开源的 <a href="https://github.com/bytedance/LatentSync">LatentSync</a> 把这个思路扩展到了视频口型同步。它的 U-Net 输入是 13 通道：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>unet_input = concat[
    noisy_gt_latents,    # 4 channels  当前 diffusion step 下的 noisy target
    masks,               # 1 channel   嘴部 mask
    masked_latents,      # 4 channels  嘴部被遮住的当前帧 latent
    ref_latents,         # 4 channels  同一视频里另一段参考帧 latent
]                        # total: 13
</code></pre></div></div>

<p>可以看出来，channel concat 适合的是<strong>和图像在空间上对齐的低层视觉条件</strong>：mask、被遮挡的图像、参考帧。它的优点是简单直接，VAE encoder 的输出可以直接拼上去；缺点是它没法处理<strong>变长</strong>或<strong>异构模态</strong>的条件，比如一段文字、一段音频。</p>

<h3 id="23-局限">2.3 局限</h3>

<p>如果用户的 prompt 是 “a red sports car running on the highway”，你怎么把它”拼”到一张图上？文字根本不是 H×W 的张量。</p>

<p>这就引出了下一代的方案。</p>

<hr />

<h2 id="3-第二代交叉注意力-cross-attention">3. 第二代：交叉注意力 (Cross-Attention)</h2>

<h3 id="31-动机stable-diffusion-让文本成为-prompt">3.1 动机：Stable Diffusion 让文本成为 prompt</h3>

<p>2022 年 Stable Diffusion 一炮走红的关键，不是它的 VAE，也不是它的 U-Net 结构，而是它把<strong>文本→图像</strong>这件事做成了 cross-attention：</p>

<ul>
  <li>文本经过 CLIP text encoder 变成 <code class="language-plaintext highlighter-rouge">[N_text, D]</code> 的 token 序列；</li>
  <li>U-Net 中的图像 feature 作为 query，文本 token 作为 key/value；</li>
  <li>在每一层 attention 里，图像 feature 会去”问”文本 token：你想让我画什么？</li>
</ul>

<p>数学上：</p>

\[\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d}}\right)V\]

<ul>
  <li>$Q = W_q \cdot \text{image_feature}$</li>
  <li>$K = W_k \cdot \text{text_token}$</li>
  <li>$V = W_v \cdot \text{text_token}$</li>
</ul>

<p>这种设计天然适合<strong>变长、异构</strong>的条件：文本可以是 5 个词也可以是 50 个词，attention 自适应处理。</p>

<h3 id="32-不只是文本音频video-embedding-都能这么干">3.2 不只是文本：音频、video embedding 都能这么干</h3>

<p>LatentSync 用 Whisper 提取音频 embedding，然后通过 cross-attention 注入 U-Net；Wan2.1 用 T5 编码文本、CLIP 编码图像，两路条件都通过 cross-attention 进入 DiT。这套机制已经成了”高层语义条件”的事实标准。</p>

<h3 id="33-局限">3.3 局限</h3>

<p>cross-attention 不便宜。原始 DiT 论文实验过几种条件注入方案，发现 cross-attention 比 adaLN 多大约 <strong>15% 的 FLOPs</strong><sup id="fnref:1" role="doc-noteref"><a href="#fn:1" class="footnote" rel="footnote">1</a></sup>。当模型规模上到几十亿参数时，这个开销不容忽视。</p>

<p>更深层的问题是：cross-attention 的条件信息<strong>只在 attention 层生效</strong>，对 LayerNorm、MLP 这些非 attention 层是”透明”的。如果条件本身是一个全局信号（比如”现在是 timestep=500”、”这是猫这一类”），用 cross-attention 杀鸡用牛刀。</p>

<hr />

<h2 id="4-第三代自适应归一化-film--adaln">4. 第三代：自适应归一化 (FiLM / adaLN)</h2>

<h3 id="41-动机有些条件其实是全局调制信号">4.1 动机：有些条件其实是”全局调制信号”</h3>

<p>看一下 cross-attention 在做什么：它让每个图像 token 去关注一些条件 token。但如果条件本身就是一个<strong>全局向量</strong>（比如时间步 <code class="language-plaintext highlighter-rouge">t</code>、类别 <code class="language-plaintext highlighter-rouge">c</code>、说话人 id），让每个图像 token 都去 attend 它，纯粹是浪费算力。</p>

<p>更高效的做法是 <strong>FiLM (Feature-wise Linear Modulation)</strong>：</p>

\[\text{FiLM}(x, c) = \gamma(c) \odot x + \beta(c)\]

<p>也就是用条件 <code class="language-plaintext highlighter-rouge">c</code> 算出一组缩放和偏移参数，直接对 feature 做线性调制。</p>

<h3 id="42-adaln把-film-用到-layernorm-上">4.2 adaLN：把 FiLM 用到 LayerNorm 上</h3>

<p>普通的 LayerNorm 是这样：</p>

\[\text{LN}(x) = \gamma \cdot \text{normalize}(x) + \beta\]

<p>这里的 <code class="language-plaintext highlighter-rouge">γ, β</code> 是学习出来的固定参数。<strong>adaLN (Adaptive LayerNorm)</strong> 把它换成条件的函数：</p>

\[\text{adaLN}(x, c) = \gamma(c) \cdot \text{normalize}(x) + \beta(c)\]

<p>其中 <code class="language-plaintext highlighter-rouge">γ(c), β(c) = MLP(c)</code>。</p>

<p>DiT 的论文系统比较了几种条件注入方案——in-context conditioning、cross-attention、adaLN——最终发现 <strong>adaLN 是 FLOPs 最低、效果最好的方案</strong><sup id="fnref:2" role="doc-noteref"><a href="#fn:2" class="footnote" rel="footnote">2</a></sup>。</p>

<h3 id="43-adaln-zero让深层-transformer-训练得更稳">4.3 adaLN-Zero：让深层 Transformer 训练得更稳</h3>

<p>DiT 在 adaLN 上又加了一个改动，叫 <strong>adaLN-Zero</strong>：除了 scale 和 shift，还预测一个 <strong>gate</strong>，并把它的初始值设为 0：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">shift</span><span class="p">,</span> <span class="n">scale</span><span class="p">,</span> <span class="n">gate</span> <span class="o">=</span> <span class="n">MLP</span><span class="p">(</span><span class="n">c</span><span class="p">).</span><span class="n">chunk</span><span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>

<span class="n">y</span> <span class="o">=</span> <span class="n">LN</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="n">y</span> <span class="o">=</span> <span class="n">y</span> <span class="o">*</span> <span class="p">(</span><span class="mi">1</span> <span class="o">+</span> <span class="n">scale</span><span class="p">)</span> <span class="o">+</span> <span class="n">shift</span>
<span class="n">x</span> <span class="o">=</span> <span class="n">x</span> <span class="o">+</span> <span class="n">gate</span> <span class="o">*</span> <span class="n">Attention</span><span class="p">(</span><span class="n">y</span><span class="p">)</span>   <span class="c1"># gate 初始为 0，所以这一项一开始是 0
</span></code></pre></div></div>

<p>这意味着：<strong>训练初始时，每个 block 都近似 identity function</strong>，新加进来的网络层不会立即破坏预训练的表征。这个技巧让深层 DiT 训练得非常稳，几乎成了今天所有 Transformer-based diffusion 的标配。</p>

<h3 id="44-局限">4.4 局限</h3>

<p>adaLN 的本质是<strong>全局调制</strong>：它对所有 token 施加同一组 <code class="language-plaintext highlighter-rouge">(γ, β, gate)</code>。这意味着它擅长做的事情是：</p>

<ul>
  <li>时间步注入（<code class="language-plaintext highlighter-rouge">t</code> 是一个标量）</li>
  <li>类别注入（class id 是一个 one-hot）</li>
  <li>全局风格控制（style embedding）</li>
</ul>

<p>它不擅长的事情是：</p>

<ul>
  <li>告诉模型”左手在这里”</li>
  <li>告诉模型”这条边缘要保留”</li>
</ul>

<p>要做空间结构控制，还需要更专门的机制。</p>

<hr />

<h2 id="5-第四代controlnet--给-u-net-外挂一支控制分支">5. 第四代：ControlNet —— 给 U-Net 外挂一支控制分支</h2>

<h3 id="51-动机文本说不清楚姿态">5.1 动机：文本说不清楚”姿态”</h3>

<p>文字 prompt 有一个根本局限：<strong>它没法精确描述空间结构</strong>。你说 “a person standing with left hand raised”，模型可能给你一个右手抬起来的人，或者两只手都抬着的人。</p>

<p>但如果你能给模型一张姿态骨架图、一张边缘图、一张深度图，事情就完全不一样了 —— 这些条件本身就是<strong>空间对齐</strong>的，每个像素都精确地告诉模型”这个位置应该是什么”。</p>

<h3 id="52-做法复制一份-encoder--零卷积">5.2 做法：复制一份 encoder + 零卷积</h3>

<p>ControlNet 的设计很巧妙<sup id="fnref:3" role="doc-noteref"><a href="#fn:3" class="footnote" rel="footnote">3</a></sup>：</p>

<ol>
  <li><strong>冻结</strong>原始 Stable Diffusion U-Net，保留它所有的预训练能力；</li>
  <li><strong>复制</strong>一份 U-Net 的 encoder（包括 down blocks 和 mid block）作为 trainable branch；</li>
  <li>condition 图（pose / edge / depth）作为这条分支的输入；</li>
  <li>这条分支在每个尺度产生 residual feature，<strong>通过零卷积加到原 U-Net 的对应层</strong>。</li>
</ol>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>condition_map (pose/edge/depth)
    ↓
[ControlNet branch (trainable copy of encoder)]
    ↓
multi-scale residuals
    ↓
加到 frozen U-Net 的 down/mid blocks
</code></pre></div></div>

<h3 id="53-零卷积让训练从零干扰开始">5.3 零卷积：让训练从”零干扰”开始</h3>

<p>ControlNet 最关键的技巧是 <strong>zero convolution</strong>：连接 ControlNet 分支和原 U-Net 的卷积层，<strong>初始权重和 bias 都是 0</strong>。</p>

<p>这意味着：</p>

<ul>
  <li>训练第 0 步：ControlNet 输出全是 0，原 U-Net 行为完全不变；</li>
  <li>随着训练进行，零卷积逐渐学到非零参数，ControlNet 的影响逐渐增强；</li>
  <li><strong>不会有训练初期”噪声破坏预训练模型”的问题</strong><sup id="fnref:3:1" role="doc-noteref"><a href="#fn:3" class="footnote" rel="footnote">3</a></sup>。</li>
</ul>

<p>这个设计让 ControlNet 在很小的数据量上就能训练成功 —— 不像 fine-tuning 那样需要担心 catastrophic forgetting。</p>

<h3 id="54-局限">5.4 局限</h3>

<ul>
  <li><strong>每种 condition 要训一个 ControlNet</strong>：pose、edge、depth、normal、segmentation 各自一个分支；</li>
  <li><strong>额外参数量不小</strong>：复制了一半 U-Net，参数量大约是原模型的 50%；</li>
  <li><strong>只控制结构，不控制语义/身份</strong>：ControlNet 告诉模型”这个位置有边缘”，但没法告诉它”这个人长得像谁”。</li>
</ul>

<p>最后这一点引出了下一个机制。</p>

<hr />

<h2 id="6-第五代ip-adapter--让图片成为-prompt">6. 第五代：IP-Adapter —— 让图片成为 prompt</h2>

<h3 id="61-动机有些东西文字真的描述不了">6.1 动机：有些东西文字真的描述不了</h3>

<p>试着用文字描述一个具体的人长什么样：</p>

<blockquote>
  <p>“a woman with long brown hair, brown eyes, oval face, slightly upturned nose…”</p>
</blockquote>

<p>写得再细，模型也只能给你一个<strong>满足这些属性的随机面孔</strong>，不是你脑子里那个<strong>具体的人</strong>。</p>

<p>类似地：</p>

<ul>
  <li>一个 logo 的精确视觉风格；</li>
  <li>一件衣服的具体花纹；</li>
  <li>一种独特的画风。</li>
</ul>

<p>这些”视觉概念”用文字 prompt 几乎不可能精确传达。但如果你能直接给模型一张<strong>参考图</strong>，事情就简单多了。</p>

<h3 id="62-朴素做法把图像和文字-token-拼起来">6.2 朴素做法：把图像和文字 token 拼起来</h3>

<p>最直观的想法是：用 CLIP image encoder 编码参考图，然后把图像 token 和文本 token 拼到一起，喂给同一个 cross-attention。</p>

<p>但 IP-Adapter 论文指出，<strong>这种朴素拼接会导致图像信息被文本特征”覆盖”</strong><sup id="fnref:4" role="doc-noteref"><a href="#fn:4" class="footnote" rel="footnote">4</a></sup>。原因是：原模型的 cross-attention <code class="language-plaintext highlighter-rouge">K, V</code> projection 是针对文本特征训练的，强行让图像特征通过同一组 projection 进入 attention，会损失大量图像信息。</p>

<h3 id="63-解耦交叉注意力-decoupled-cross-attention">6.3 解耦交叉注意力 (Decoupled Cross-Attention)</h3>

<p>IP-Adapter 的核心创新是 <strong>decoupled cross-attention</strong>：<strong>给图像单独开一套 cross-attention</strong>，而不是和文本共享。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">ip_adapter_attention</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">text_tokens</span><span class="p">,</span> <span class="n">image_tokens</span><span class="p">):</span>
    <span class="n">q</span> <span class="o">=</span> <span class="n">Wq</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>

    <span class="c1"># 原本的 text cross-attention（保持不变）
</span>    <span class="n">k_text</span> <span class="o">=</span> <span class="n">Wk_text</span><span class="p">(</span><span class="n">text_tokens</span><span class="p">)</span>
    <span class="n">v_text</span> <span class="o">=</span> <span class="n">Wv_text</span><span class="p">(</span><span class="n">text_tokens</span><span class="p">)</span>
    <span class="n">out_text</span> <span class="o">=</span> <span class="n">attention</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">k_text</span><span class="p">,</span> <span class="n">v_text</span><span class="p">)</span>

    <span class="c1"># 新增的 image cross-attention（新训练的）
</span>    <span class="n">k_img</span> <span class="o">=</span> <span class="n">Wk_img</span><span class="p">(</span><span class="n">image_tokens</span><span class="p">)</span>
    <span class="n">v_img</span> <span class="o">=</span> <span class="n">Wv_img</span><span class="p">(</span><span class="n">image_tokens</span><span class="p">)</span>
    <span class="n">out_img</span> <span class="o">=</span> <span class="n">attention</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">k_img</span><span class="p">,</span> <span class="n">v_img</span><span class="p">)</span>

    <span class="k">return</span> <span class="n">out_text</span> <span class="o">+</span> <span class="n">scale</span> <span class="o">*</span> <span class="n">out_img</span>
</code></pre></div></div>

<p>注意几个关键设计：</p>

<ol>
  <li><strong>复用 query</strong>：图像 attention 和文本 attention 共享同一个 <code class="language-plaintext highlighter-rouge">Q</code>，因为 query 来自图像 latent；</li>
  <li><strong>独立的 K/V projection</strong>：<code class="language-plaintext highlighter-rouge">Wk_img, Wv_img</code> 是新训练的，专门为图像特征设计；</li>
  <li><strong>加和融合</strong>：两路 attention 的输出直接相加，可以用 <code class="language-plaintext highlighter-rouge">scale</code> 控制图像 prompt 的强度。</li>
</ol>

<h3 id="64-优势极小的可训练参数">6.4 优势：极小的可训练参数</h3>

<p>IP-Adapter 论文报告，<strong>只需要约 22M 可训练参数</strong>，就能在冻结的 Stable Diffusion 上实现强大的图像 prompt 能力，并且和文本 prompt、ControlNet 等工具完全兼容<sup id="fnref:4:1" role="doc-noteref"><a href="#fn:4" class="footnote" rel="footnote">4</a></sup>。</p>

<h3 id="65-ip-adapter-vs-ref_latent-latentsync">6.5 IP-Adapter vs ref_latent (LatentSync)</h3>

<p>值得对比一下，因为它们看起来都是”用参考图做条件”：</p>

<table>
  <thead>
    <tr>
      <th>方式</th>
      <th>信息类型</th>
      <th>注入方式</th>
      <th>优点</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>LatentSync 的 ref_latents</strong></td>
      <td>低层像素/纹理/位置</td>
      <td>VAE latent → channel concat</td>
      <td>空间对齐强，重建保真度高</td>
    </tr>
    <tr>
      <td><strong>IP-Adapter 的 image tokens</strong></td>
      <td>高层语义/身份/风格</td>
      <td>CLIP image encoder → cross-attention</td>
      <td>泛化强，可以”风格迁移”</td>
    </tr>
  </tbody>
</table>

<p>简单来说：</p>
<ul>
  <li>想保留<strong>精确的像素结构</strong>（同一个人、同一个场景的不同角度）→ ref latent concat；</li>
  <li>想保留<strong>风格和身份语义</strong>（这个人的”长相”，但姿势可以不一样）→ IP-Adapter。</li>
</ul>

<hr />

<h2 id="7-一张表把前五代串起来">7. 一张表把前五代串起来</h2>

<p>到这里我们已经看了五种主流的条件注入方式。它们其实是<strong>互补</strong>而不是替代关系。一个现代的 Diffusion 系统通常会同时用到好几种：</p>

<table>
  <thead>
    <tr>
      <th>机制</th>
      <th>数学形式</th>
      <th>条件类型</th>
      <th>空间对齐</th>
      <th>参数量</th>
      <th>典型用途</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>Channel Concat</strong></td>
      <td><code class="language-plaintext highlighter-rouge">concat([z_t, c], dim=1)</code></td>
      <td>空间对齐的低层视觉信息</td>
      <td>强</td>
      <td>几乎为 0</td>
      <td>inpainting, ref frame, mask</td>
    </tr>
    <tr>
      <td><strong>Cross-Attention</strong></td>
      <td><code class="language-plaintext highlighter-rouge">Attn(Q=z_t, K=c, V=c)</code></td>
      <td>变长高层语义</td>
      <td>弱</td>
      <td>中</td>
      <td>text, audio, video tokens</td>
    </tr>
    <tr>
      <td><strong>adaLN-Zero</strong></td>
      <td><code class="language-plaintext highlighter-rouge">x = x + gate(c) · f(scale(c)·LN(x)+shift(c))</code></td>
      <td>全局信号</td>
      <td>无</td>
      <td>极小</td>
      <td>timestep, class, style</td>
    </tr>
    <tr>
      <td><strong>ControlNet</strong></td>
      <td><code class="language-plaintext highlighter-rouge">h_i ← h_i + ControlNet_i(cond)</code></td>
      <td>空间结构图</td>
      <td>强</td>
      <td>大（~50% U-Net）</td>
      <td>pose, edge, depth, seg</td>
    </tr>
    <tr>
      <td><strong>IP-Adapter</strong></td>
      <td><code class="language-plaintext highlighter-rouge">Attn_text + λ · Attn_image</code></td>
      <td>参考图语义/身份</td>
      <td>中</td>
      <td>极小（~22M）</td>
      <td>reference image, style</td>
    </tr>
  </tbody>
</table>

<p>注意一个有意思的事实：<strong>这五种机制里，真正被 diffusion 去噪的只有 <code class="language-plaintext highlighter-rouge">z_t</code> 一个</strong>。其他所有的”条件” —— 不管是 mask、文本 token、ControlNet residual、image token —— 都只是<strong>改变 U-Net 对 <code class="language-plaintext highlighter-rouge">z_t</code> 噪声预测方向的指引</strong>，它们自己不会被更新。</p>

<p>理解这一点，对接下来看 DiT 的演进非常重要。</p>

<hr />

<h2 id="8-第六代从-u-net-到-dit条件注入也跟着进化">8. 第六代：从 U-Net 到 DiT，条件注入也跟着进化</h2>

<h3 id="81-为什么大家开始用-transformer-替代-u-net">8.1 为什么大家开始用 Transformer 替代 U-Net</h3>

<p>U-Net 的核心是卷积 + skip connection，它的归纳偏置很适合处理图像，但它有几个问题：</p>

<ol>
  <li><strong>难以 scale</strong>：当模型规模上到 10B 以上，U-Net 的训练不如 Transformer 稳定；</li>
  <li><strong>跨模态融合受限</strong>：U-Net 主要靠 cross-attention 注入文本/图像条件，深度有限；</li>
  <li><strong>不适合统一多模态</strong>：当条件不只是文本、还包括图像、音频、视频时，U-Net 的结构不够灵活。</li>
</ol>

<p>DiT (Diffusion Transformer) 解决了前两个问题：把 U-Net 换成 ViT 风格的 Transformer，把图像 patchify 成 token 序列，然后用 Transformer block 去噪<sup id="fnref:2:1" role="doc-noteref"><a href="#fn:2" class="footnote" rel="footnote">2</a></sup>。</p>

<p>但<strong>条件怎么注入到 DiT 里</strong>，又出现了不同的设计哲学。这就引出了今天最热的两条路线：<strong>多流 MMDiT</strong> vs <strong>单流 S3-DiT</strong>。</p>

<h3 id="82-多流-mmdit文本和图像各走一条流">8.2 多流 MMDiT：文本和图像各走一条流</h3>

<p>代表作是阿里的 <strong>Qwen-Image</strong><sup id="fnref:5" role="doc-noteref"><a href="#fn:5" class="footnote" rel="footnote">5</a></sup>。它是一个 20B 参数的多模态 Diffusion Transformer，核心架构包括三个组件：</p>

<ol>
  <li><strong>冻结的 Qwen2.5-VL</strong>（视觉语言模型）：负责文本和图像的语义对齐；</li>
  <li><strong>VAE encoder/decoder</strong>：负责图像潜变量的压缩和重建；</li>
  <li><strong>MMDiT diffusion backbone</strong>：负责在 latent 空间去噪。</li>
</ol>

<p>它的”多流”体现在：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>text prompt ─→ Qwen2.5-VL ─→ semantic tokens ─┐
                                              ├─→ MMDiT (cross-modal fusion) ─→ noise
input image ─→ VAE ───────→ reconstruction tokens ┐
              + Qwen2.5-VL → semantic tokens   ──┘
              ↑
              这就是 "dual encoding" 机制
</code></pre></div></div>

<p><strong>Dual encoding 是 Qwen-Image 的关键创新</strong>：同一张输入图同时被 Qwen2.5-VL 和 VAE 编码，前者提供语义信息，后者提供重建信息，两者在 MMDiT 里融合。这种设计在图像编辑任务上特别有用 —— 编辑既要”理解你想改什么”（语义），也要”保持原图其他地方不变”（重建保真度）<sup id="fnref:5:1" role="doc-noteref"><a href="#fn:5" class="footnote" rel="footnote">5</a></sup>。</p>

<h3 id="83-单流-s3-dit所有-token-拼成一条序列">8.3 单流 S3-DiT：所有 token 拼成一条序列</h3>

<p>代表作是 <strong>Z-Image</strong>，由 Tongyi Lab 提出的 6B 参数 Scalable Single-Stream Diffusion Transformer (<strong>S3-DiT</strong>)<sup id="fnref:6" role="doc-noteref"><a href="#fn:6" class="footnote" rel="footnote">6</a></sup>。</p>

<p>它的设计哲学和 MMDiT 完全相反：<strong>所有模态的 token 都拼到一条序列里，共享同一个 Transformer 处理</strong>：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[text tokens | visual semantic tokens | noisy image VAE tokens]
                         ↓
                Single Transformer (S3-DiT)
                         ↓
             只取 image token 部分预测 noise
</code></pre></div></div>

<p>S3-DiT 的几个关键技术点：</p>

<ol>
  <li><strong>3D RoPE</strong>：把文本、空间、通道位置都编码到同一空间；</li>
  <li><strong>轻量的模态 stem</strong>：每种模态有自己的小 MLP，把它映射到共享的 hidden space；</li>
  <li><strong>FiLM 式条件适配器</strong>：timestep 和 global condition 通过 FiLM-like scale/shift 注入；</li>
  <li><strong>流匹配 + 蒸馏</strong>：用 flow matching loss 训练，并通过蒸馏得到 Z-Image-Turbo，可以在消费级 GPU 上做亚秒级推理<sup id="fnref:6:1" role="doc-noteref"><a href="#fn:6" class="footnote" rel="footnote">6</a></sup>。</li>
</ol>

<p><strong>为什么 6B 单流能挑战 20B 多流？</strong> 因为单流架构的参数效率更高 —— 文本和图像共享同一套 Transformer 权重，避免了双流之间的冗余表征。</p>

<h3 id="84-两种路线的对比">8.4 两种路线的对比</h3>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>多流 MMDiT (Qwen-Image)</th>
      <th>单流 S3-DiT (Z-Image)</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>token 组织</strong></td>
      <td>文本流 + 图像流，分开处理后融合</td>
      <td>所有 token 拼成统一序列</td>
    </tr>
    <tr>
      <td><strong>条件注入</strong></td>
      <td>cross-modal fusion + dual encoding</td>
      <td>unified self-attention + FiLM</td>
    </tr>
    <tr>
      <td><strong>控制风格</strong></td>
      <td>显式、模块化、可分解</td>
      <td>隐式、统一、端到端</td>
    </tr>
    <tr>
      <td><strong>优势</strong></td>
      <td>强语义、强编辑、复杂条件稳定</td>
      <td>参数效率高、结构简洁、推理友好</td>
    </tr>
    <tr>
      <td><strong>劣势</strong></td>
      <td>参数巨大、计算重、系统复杂</td>
      <td>长序列 attention 成本高，强空间控制需额外设计</td>
    </tr>
    <tr>
      <td><strong>典型场景</strong></td>
      <td>专业图像编辑、文字渲染、多条件控制</td>
      <td>高效 T2I、统一多模态生成、轻量部署</td>
    </tr>
  </tbody>
</table>

<h3 id="85-直观比喻">8.5 直观比喻</h3>

<p>如果把 Diffusion 模型比作一个画师：</p>

<ul>
  <li><strong>多流 MMDiT</strong> 像是一个团队：一个文本理解专家、一个图像理解专家、一个绘画师傅，三人开会沟通后画师傅落笔。每个人都是该领域的专家，分工明确，但维护成本高。</li>
  <li><strong>单流 S3-DiT</strong> 像是一个全能选手：他自己读 prompt、看参考图、想空间结构、动笔作画，全在一个脑子里完成。沟通成本低、效率高，但需要训练数据足够多样才能学会所有这些技能。</li>
</ul>

<hr />

<h2 id="9-一个例子图像编辑任务下两种路线怎么处理">9. 一个例子：图像编辑任务下两种路线怎么处理</h2>

<p>假设任务是：</p>

<blockquote>
  <p>“输入一张人像，把衣服改成红色西装，脸和背景保持不变。”</p>
</blockquote>

<p><strong>多流 MMDiT 的做法</strong>：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>原图          ─→ Qwen2.5-VL ─→ "这是一个穿白衬衫的男人" (语义)
              ─→ VAE        ─→ pixel-level 重建特征
文本指令       ─→ Qwen2.5-VL ─→ "把衣服改成红色西装" (编辑指令)
noisy target  ─→ MMDiT denoising stream

MMDiT 通过显式的 cross-modal fusion：
- 语义层面理解"衣服→红色西装"；
- 重建层面知道"脸和背景保持原样"；
- denoising 层面生成最终结果。
</code></pre></div></div>

<p><strong>单流 S3-DiT 的做法</strong>：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[ text instruction tokens
| input image semantic tokens
| input image VAE tokens
| noisy target tokens ]
                ↓
        Single Transformer
                ↓
        从 target tokens 输出 noise
</code></pre></div></div>

<p>模型自己通过 attention 学习”哪些 token 对哪些区域重要”。这种端到端的学习理论上更灵活，但需要大量精心标注的编辑数据才能训稳。</p>

<hr />

<h2 id="10-未来趋势混合路线--多模态统一">10. 未来趋势：混合路线 + 多模态统一</h2>

<p>看完前面这一切，我对未来 Diffusion 架构的演进有几个判断。</p>

<h3 id="101-单流会成为基础架构但不会纯单流">10.1 单流会成为基础架构，但不会”纯单流”</h3>

<p>单流的优势是简洁和效率，但它在<strong>强空间控制</strong>（pose、depth、edge）和<strong>复杂编辑</strong>上还有差距。最实用的系统大概率是 hybrid：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Single-stream DiT backbone (主干)
    + Control branch (空间结构控制，类似 ControlNet)
    + Reference adapter (参考图，类似 IP-Adapter)
    + Mask/edit branch (区域级编辑)
    + Layout/typography expert (排版、文字渲染)
</code></pre></div></div>

<p>也就是说，<strong>单流解决”统一和效率”，专门分支解决”精确和可控”</strong>。</p>

<h3 id="102-控制粒度从-prompt-level-走向-region-level">10.2 控制粒度从 prompt-level 走向 region-level</h3>

<p>现在的主流是”一句 prompt 控制整张图”。但实际场景中，用户更需要：</p>

<ul>
  <li>这个区域保持不变；</li>
  <li>这个物体换材质；</li>
  <li>这个人保持身份；</li>
  <li>这几个字必须准确显示；</li>
  <li>这个 logo 不能变形。</li>
</ul>

<p>这要求模型支持<strong>区域绑定、对象绑定、文字位置绑定、多参考图绑定</strong>。Qwen-Image 强调的”复杂文字渲染”和”精确图像编辑”已经在往这个方向走。</p>

<h3 id="103-生成与编辑会统一为一个模型">10.3 生成与编辑会统一为一个模型</h3>

<p>以前模型经常分得很细：T2I model、inpainting model、editing model、ControlNet model……。未来会逐渐合并：</p>

<blockquote>
  <p>一个模型同时支持 T2I、image editing、multi-image composition、style transfer、layout generation、text rendering、object replacement。</p>
</blockquote>

<p>Qwen-Image 的 multi-task training 已经包括 T2I、TI2I (text+image→image)、I2I reconstruction<sup id="fnref:5:2" role="doc-noteref"><a href="#fn:5" class="footnote" rel="footnote">5</a></sup>。这说明主流方向已经是<strong>统一训练</strong>，而不是每个任务单独一个模型。</p>

<h3 id="104-高效化会变成核心竞争力">10.4 高效化会变成核心竞争力</h3>

<p>Z-Image 的 6B 单流路线说明：<strong>不是只能靠堆参数</strong>。数据质量、架构设计、蒸馏和推理优化同样重要。未来会越来越重视：</p>

<ul>
  <li>few-step generation（4 步、2 步、甚至 1 步采样）；</li>
  <li>distillation（把大模型的能力蒸馏到小模型）；</li>
  <li>FP8 / INT8 / NF4 quantization；</li>
  <li>KV cache / feature cache；</li>
  <li>MoE DiT（稀疏激活的 Diffusion Transformer）；</li>
  <li>consumer GPU fine-tuning。</li>
</ul>

<h3 id="105-跨模态联合学习将成为标配">10.5 跨模态联合学习将成为标配</h3>

<p>未来的”条件”不会只有文本和图像。视频、音频、3D、甚至 IMU/sensor 数据都会一起训练。模型需要更复杂的条件注入策略来处理多模态信息 —— adaLN、cross-attention、ControlNet-like branch、IP-Adapter-like adapter 都会在不同位置发挥作用。</p>

<hr />

<h2 id="11-总结">11. 总结</h2>

<p>回到一开始那个问题：<strong>怎么把”用户想要什么”告诉模型？</strong></p>

<p>过去几年，Diffusion 社区给出的答案大致是这样一条主线：</p>

<ol>
  <li><strong>Channel Concat</strong>：条件和图像在同一空间结构 → 直接拼通道。简单粗暴，适合 inpainting。</li>
  <li><strong>Cross-Attention</strong>：条件是变长语义 → 让图像 token 去 attend 它。文本/音频条件的标配。</li>
  <li><strong>adaLN-Zero</strong>：条件是全局信号 → 用它生成 LayerNorm 的 scale/shift/gate。极低开销，DiT 标配。</li>
  <li><strong>ControlNet</strong>：条件是空间结构图 → 复制一份 encoder + 零卷积。强空间控制，但每种条件要训一个分支。</li>
  <li><strong>IP-Adapter</strong>：条件是参考图 → 解耦 cross-attention，给图像单开一套。极小参数实现图像 prompt。</li>
  <li><strong>多流 MMDiT (Qwen-Image)</strong>：文本和图像各走一条流，通过 cross-modal fusion 融合。强编辑、强语义。</li>
  <li><strong>单流 S3-DiT (Z-Image)</strong>：所有 token 拼成一条序列共享 Transformer。参数高效、推理友好。</li>
</ol>

<p><strong>核心洞察</strong>：这些机制不是替代关系，而是互补的。一个现代 Diffusion 系统往往同时用到多种 —— DiT 主干用 adaLN-Zero 做时间步调制，cross-attention 接文本，ControlNet 做空间控制，IP-Adapter 做参考图。理解每种机制的”擅长什么、不擅长什么”，比死记某一个架构更有用。</p>

<p>未来的方向是<strong>单流主干 + 多种专门分支的 hybrid 架构</strong>，再叠加蒸馏、量化和高效推理。从条件注入的角度看，这场演化远没有结束。</p>

<hr />

<h2 id="sources">Sources</h2>

<ul>
  <li><a href="https://arxiv.org/abs/2006.11239">DDPM (Ho et al., 2020)</a></li>
  <li><a href="https://github.com/bytedance/LatentSync">LatentSync (字节, 2024)</a></li>
  <li><a href="https://arxiv.org/abs/2112.10752">Stable Diffusion / Latent Diffusion (Rombach et al., 2022)</a></li>
  <li><a href="https://arxiv.org/abs/2212.09748">DiT: Scalable Diffusion Models with Transformers (Peebles &amp; Xie, 2022)</a></li>
  <li><a href="https://arxiv.org/abs/2302.05543">ControlNet (Zhang et al., 2023)</a></li>
  <li><a href="https://arxiv.org/abs/2308.06721">IP-Adapter (Ye et al., 2023)</a></li>
  <li><a href="https://github.com/QwenLM/Qwen-Image">Qwen-Image 官方仓库</a></li>
  <li><a href="https://arxiv.org/abs/2508.02324">Qwen-Image Technical Report</a></li>
  <li><a href="https://github.com/Tongyi-MAI/Z-Image">Z-Image / S3-DiT</a></li>
  <li><a href="https://arxiv.org/abs/1709.07871">FiLM (Perez et al., 2017)</a></li>
  <li><a href="https://arxiv.org/abs/2210.02747">Flow Matching for Generative Modeling (Lipman et al., 2023)</a></li>
</ul>

<div class="footnotes" role="doc-endnotes">
  <ol>
    <li id="fn:1" role="doc-endnote">
      <p>DiT 论文中比较了 in-context conditioning、cross-attention、adaLN 三种条件注入方式，cross-attention 比 adaLN 多约 15% Gflops。 <a href="#fnref:1" class="reversefootnote" role="doc-backlink">&#8617;</a></p>
    </li>
    <li id="fn:2" role="doc-endnote">
      <p>Peebles &amp; Xie. <em>Scalable Diffusion Models with Transformers</em>. ICCV 2023. adaLN-Zero 把 residual block 中调制参数初始化为零，使 block 初始接近 identity function。 <a href="#fnref:2" class="reversefootnote" role="doc-backlink">&#8617;</a> <a href="#fnref:2:1" class="reversefootnote" role="doc-backlink">&#8617;<sup>2</sup></a></p>
    </li>
    <li id="fn:3" role="doc-endnote">
      <p>Zhang et al. <em>Adding Conditional Control to Text-to-Image Diffusion Models</em>. ICCV 2023. ControlNet 通过 zero-initialized convolution 让参数从零逐渐增长，避免训练初期破坏预训练模型。 <a href="#fnref:3" class="reversefootnote" role="doc-backlink">&#8617;</a> <a href="#fnref:3:1" class="reversefootnote" role="doc-backlink">&#8617;<sup>2</sup></a></p>
    </li>
    <li id="fn:4" role="doc-endnote">
      <p>Ye et al. <em>IP-Adapter: Text Compatible Image Prompt Adapter for Text-to-Image Diffusion Models</em>. 2023. 关键设计是 decoupled cross-attention，把 text 和 image 的 cross-attention 分开。 <a href="#fnref:4" class="reversefootnote" role="doc-backlink">&#8617;</a> <a href="#fnref:4:1" class="reversefootnote" role="doc-backlink">&#8617;<sup>2</sup></a></p>
    </li>
    <li id="fn:5" role="doc-endnote">
      <p>Qwen-Image 是 20B MMDiT 模型，使用 Qwen2.5-VL + VAE 双编码机制，支持 T2I、TI2I、I2I 多任务训练，强调复杂文字渲染和图像编辑能力。 <a href="#fnref:5" class="reversefootnote" role="doc-backlink">&#8617;</a> <a href="#fnref:5:1" class="reversefootnote" role="doc-backlink">&#8617;<sup>2</sup></a> <a href="#fnref:5:2" class="reversefootnote" role="doc-backlink">&#8617;<sup>3</sup></a></p>
    </li>
    <li id="fn:6" role="doc-endnote">
      <p>Z-Image 是 6B 参数的 Scalable Single-Stream Diffusion Transformer (S3-DiT)，把 text、visual semantic tokens、image VAE tokens 拼成统一输入流，通过 3D RoPE 和 FiLM 适配器注入条件，支持亚秒级推理。 <a href="#fnref:6" class="reversefootnote" role="doc-backlink">&#8617;</a> <a href="#fnref:6:1" class="reversefootnote" role="doc-backlink">&#8617;<sup>2</sup></a></p>
    </li>
  </ol>
</div>]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><category term="blog" /><category term="diffusion" /><category term="generative model" /><category term="DiT" /><category term="ControlNet" /><category term="IP-Adapter" /><category term="multimodal" /><summary type="html"><![CDATA[如果你看过 Stable Diffusion、ControlNet、IP-Adapter，又听说过最近的 Qwen-Image 和 Z-Image，可能会有一个共同的疑问：]]></summary></entry><entry><title type="html">Agentic RL 训练全景：环境、信号、分布与系统的协同闭环</title><link href="https://liyongzhi.xyz/posts/2026/04/agentic-rl/" rel="alternate" type="text/html" title="Agentic RL 训练全景：环境、信号、分布与系统的协同闭环" /><published>2026-04-28T00:00:00+08:00</published><updated>2026-04-28T00:00:00+08:00</updated><id>https://liyongzhi.xyz/posts/2026/04/blog-post-agentic-rl</id><content type="html" xml:base="https://liyongzhi.xyz/posts/2026/04/agentic-rl/"><![CDATA[<blockquote>
  <p>过去一年，各家大模型公司公开的技术报告透出的最重要信号，不是又出现了一个更好的 PPO/GRPO 变体，而是<strong>真正有效的 Agentic RL 已经从”单轮文本优化”转向了”在长上下文、工具调用、部分可观测、异步执行环境中的系统性策略学习”</strong>。</p>

  <p>Kimi K1.5[1] 把长上下文 RL、partial rollout 重用和 mirror-descent 风格的 policy optimization 拉到了台前；Kimi K2[2]/K2.5[3] 又把 agentic 数据合成、多模态 RL、token-level clipping、GRM rubric、Toggle、PARL / Agent Swarm 这些关键部件公开；MiniMax 把另一个事实讲得更彻底：<strong>当 rollout 时长从秒级扩到小时级，训练瓶颈就不再是 loss design，而是吞吐、稳定性与 agent 灵活性之间的三难权衡</strong>；GLM 则强调分阶段 RL：Reasoning RL、Agentic RL、General RL 不是混在一起一次训完，而是通过顺序化 pipeline 逐步推进，并借助异步 RL 基础设施与跨阶段蒸馏来兼顾长时程 agent 学习与能力保持。</p>

  <p>Agentic RL 的核心问题，已经从”怎么更新参数”扩展为”<strong>怎么在真实 Agent 环境里持续制造可用的学习信号，并用在线交互的轨迹数据驱动优化</strong>“。</p>
</blockquote>

<hr />

<h2 id="一为什么-agentic-rl-与传统-rlhf--rlvr-不同">一、为什么 Agentic RL 与传统 RLHF / RLVR 不同</h2>

<p>Agentic RL 的训练对象不再是”给定一个 prompt，输出一个答案”的单轮文本映射，而是<strong>一个在环境中交互的策略</strong>。这个策略要处理：状态更新、工具调用、外部观察、上下文整理、子任务委派、终止条件判断，以及成本 / 时延 / 安全约束。换句话说，agentic RL 更像是在做一类带有长时间尺度、部分可观测性和结构化动作空间的策略学习，而不是简单地对文本续写概率做后验重排。</p>

<p>这直接带来四个训练上的变化：</p>

<ol>
  <li><strong>状态不再只由用户输入决定</strong>：它由历史轨迹、工具返回、环境回馈、记忆摘要和当前上下文共同构成。</li>
  <li><strong>动作也不再只是下一个 token</strong>：它可能是”选哪个工具、填什么参数、要不要压缩上下文、是否并行分派子任务”。</li>
  <li><strong>奖励更延迟、更稀疏、更复合</strong>：既要看结果对不对，也要看过程是否准确、是否高效、是否节省 token 和单位时间有效训练效率。</li>
  <li><strong>Rollout 时间高度不均匀</strong>：同步训练代价高，异步训练又引入分布偏移。</li>
</ol>

<p>因此，agentic RL 的本质<strong>不是把 GRPO / PPO 套到更长的输出上</strong>，而是把环境、奖励、采样、调度、缓存、优化器和评测接到同一个闭环里。</p>

<hr />

<h2 id="二理解-agentic-rl-的三个不变量">二、理解 Agentic RL 的三个不变量</h2>

<p>如果把 Agentic RL 理解成一个”在真实环境里持续交互、持续采样、持续更新”的策略学习系统，那么真正重要的就不再是”这一步用哪种 RL 算法”，而是<strong>训练闭环能否长期守住三个更底层的条件</strong>。</p>

<p>这里的”不变量”不是指某个量在数学上严格恒定，而是指它们虽然会天然漂移，却必须在整个训练过程中被不断拉回到一个仍然可学习、可优化的区间里。前两个是<strong>不应跌破的下限</strong>，第三个是<strong>不应越过的上限</strong>。</p>

<h3 id="1第一不变量策略的可探索空间不能过早塌缩">1）第一不变量：策略的可探索空间不能过早塌缩</h3>

<p>第一不变量<strong>不是</strong>要求输出更随机、token 熵更高，而是：<strong>模型在给定状态下，仍然保有一组彼此可区分、语义上不同、并且真实可行的行为路径</strong>。</p>

<p>对 Agentic RL 来说，这个探索空间不只是”不同措辞”，而是：</p>

<ul>
  <li>不同的<strong>任务分解</strong>方式</li>
  <li>不同的<strong>工具调用顺序</strong></li>
  <li>不同的<strong>记忆读写</strong>策略</li>
  <li>不同的<strong>上下文整理</strong>方式</li>
  <li>不同的<strong>停止条件</strong>与<strong>自我修正</strong>路径</li>
</ul>

<p>它之所以会塌缩，是因为训练天然会把概率质量压向”少数当前最占优的模式”。只要训练目标主要奖励”更短、更像标准流程、更容易被 verifier 识别”的行为，模型就会把其他原本也可能成功的路径边缘化。在 agent 场景下，这种压缩比单轮问答更严重——工具接口、scaffold、上下文模板和终止逻辑本身就会<strong>暗中偏好</strong>某类固定 workflow。</p>

<p><strong>保持这一不变量的意义</strong>：它决定了后续 RL 是否还有真正的搜索空间。RL 的价值不是把已知最好答案重复推高概率，而是让模型在交互中持续发现”此前还没被放大的高回报行为”。如果可探索空间已经提前塌缩，后面的采样大多只是对同一种套路做表面扰动，reward spread 越来越小，训练看似还在继续，实际上只是在一个已经缩水的空间里做局部扰动。</p>

<h3 id="2第二不变量学习信号必须持续非退化">2）第二不变量：学习信号必须持续非退化</h3>

<p>即使模型仍然保有多种可行路径，这些路径也不一定会被<strong>学到</strong>。参数更新依赖的不是”存在别的可能性”，而是<strong>不同轨迹之间的差异能否稳定地转成非零、方向明确、尺度合理的梯度</strong>。</p>

<p>Agentic RL 的奖励结构天然容易让信号塌缩：真实任务奖励延迟、结果稀疏、过程很长，最终常常只有成败标签、粗粒度 rubric，或少数高层质量分。于是同一组采样很容易出现两种退化情形——</p>

<ul>
  <li><strong>简单任务几乎全对</strong>（模型已在该局部饱和）</li>
  <li><strong>困难任务几乎全错</strong>（模型尚未进入可学习区域）</li>
</ul>

<p>但对梯度而言，这两类样本都会导向同一个结果：<strong>组内没有足够差异，优势接近消失，更新方向随之退化</strong>。再叠加长轨迹的信用分配、部分可观测性带来的归因模糊、工具噪声和 verifier 噪声对比较关系的污染——系统表面上在大量收集交互数据，实际上却在不断生产”不可学样本”。</p>

<p><strong>这里有一个关键观察</strong>：学习信号的质量，不取决于奖励项有多少，而取决于<strong>比较是否可学</strong>。奖励可以很复杂，但如果它无法在模型当前边界附近稳定区分”略好”与”略差”的轨迹，它仍然会产生退化梯度。反过来，一个看上去更简单的反馈，只要能持续打开轨迹间的有效差异，也能成为高质量学习信号。</p>

<p><strong>第二不变量真正要求保持不变的，不是奖励总量，而是可比较性与可更新性。</strong></p>

<h3 id="3第三不变量训练--更新--部署三者的分布偏移必须可控">3）第三不变量：训练 / 更新 / 部署三者的分布偏移必须可控</h3>

<p>前两个不变量解决”还有没有别的路径”和”这些路径能不能变成梯度”，第三个不变量解决<strong>这些梯度是不是作用在了正确的分布上</strong>。</p>

<p>在 Agentic RL 中，有三个天然不一致的分布：</p>

<ul>
  <li>策略模型<strong>采样</strong>出的 rollout 分布</li>
  <li>learner 真正<strong>拿来更新</strong>的样本分布</li>
  <li>最终<strong>部署执行</strong>的策略分布</li>
</ul>

<p>Agent 训练会持续制造分布漂移：</p>

<ul>
  <li>轨迹长短差异极大，严格同步的 on-policy 不现实，异步采样、缓存、续跑、复用、过滤都会让”生成样本时的策略”和”更新参数时的策略”发生时间错位；</li>
  <li>Agent 状态由工具返回、环境反馈、上下文裁剪、记忆摘要、调度决策共同构成，只要其中任何一层在 rollout / training / serving 三阶段的表示不完全一致，模型学到的可能就不是同一个动作语义；</li>
  <li>训练和部署脚手架常常并不完全相同：解码设置、context packing、tool schema、tokenizer/engine、middleware、日志序列化方式都会改变模型真正面对的决策问题。</li>
</ul>

<p>结果是：被优化的不再是一个干净统一的策略分布，而是<strong>多个相似但不相同的分布拼接而成的近似对象</strong>。</p>

<p>对长轨迹 Agent，这一点尤其致命——轨迹越长，前面每一点小的偏移都会沿着后续状态转移不断累积，最终把策略推向”在训练里看起来合理、在真实环境里却不可执行”的方向。</p>

<p><strong>Agentic RL 里的分布偏移，不只是外部环境变化带来的，它在很大程度上是系统自己制造出来的</strong>。这也是为什么第三不变量不是单纯的算法修正问题，而是一个系统级的一致性问题。</p>

<h3 id="4为什么这三个不变量要放在一起理解">4）为什么这三个不变量要放在一起理解</h3>

<p>它们不是彼此独立的三条要素，而是<strong>同一个训练系统的三个耦合边界</strong>：</p>

<table>
  <thead>
    <tr>
      <th>不变量</th>
      <th>本质问题</th>
      <th>失守后果</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>第一</td>
      <td>策略空间是否还够宽</td>
      <td>没有可探索的新路径</td>
    </tr>
    <tr>
      <td>第二</td>
      <td>空间里的差异能否转成有效梯度</td>
      <td>有路径但学不到</td>
    </tr>
    <tr>
      <td>第三</td>
      <td>梯度是否作用在正确分布上</td>
      <td>学到的行为在部署时失真</td>
    </tr>
  </tbody>
</table>

<ul>
  <li>只有<strong>探索</strong>没有<strong>信号</strong>：训练变成高噪声试错；</li>
  <li>只有<strong>信号</strong>没有<strong>探索</strong>：训练迅速收缩到狭窄局部最优；</li>
  <li>探索和信号都有，但<strong>分布偏移失控</strong>：学到的也不是部署时真正需要的行为。</li>
</ul>

<p>它们彼此之间还天然存在张力：探索更强，会让比较更稀、分布偏移更难控；过度追求稳定更新，又容易压平探索空间；为了制造更锋利的信号把 verifier 设计得过于严格，又会让模型朝少数投机模式收缩。</p>

<p><strong>Agentic RL 真正要解决的，不是把某个 loss 降得更低，而是在一个持续变化、持续异步、持续与外部环境交互的系统里，始终把探索、信号和分布维持在同一个可学习区间内。</strong></p>

<hr />

<h2 id="三agentic-rl-的九个关键维度">三、Agentic RL 的九个关键维度</h2>

<p>三个不变量是”要守住什么”；下面九个维度是”在哪些具体位置守”。前八个维度对应训练系统的核心环节，第九个维度（评测与可观测性）回答的是一个更基础的问题——<strong>如果你连三个不变量是否正在被守住都测不出来，就根本谈不上管理它们</strong>。</p>

<h3 id="1-环境与接口建模先搞清楚环境允许-agent-做什么再谈正确答案是什么">1. 环境与接口建模：先搞清楚”环境允许 Agent 做什么”，再谈”正确答案是什么”</h3>

<p>Agentic RL 和普通”对一道题生成一个答案”的最大差别在于：模型不再只是从 prompt 里猜一个 completion，而是在一个<strong>可交互、可执行、带状态转移</strong>的世界里学 policy。</p>

<p>决定 Agentic RL 训练效果的<strong>第一个变量</strong>不是 reward model，而是环境和接口本身是否设计清楚：</p>

<ul>
  <li>每一步模型能看到哪些信息？</li>
  <li>能采取哪些动作、哪些工具调用是允许且有效的？</li>
  <li>任务在什么条件下结束、成功如何判断？</li>
  <li>训练时使用的工具接口和交互流程，是否和真实部署一致？</li>
</ul>

<p>当前几家的共识非常一致：</p>

<ul>
  <li>Kimi K2 把大规模 agentic 数据合成 + 真实/合成环境 RL 放进后训练主线；</li>
  <li>K2.5 把 Agentic RL 统一到 <strong>Gym-like 接口</strong>，并支持大规模异步任务管理；</li>
  <li>GLM-5[8] 把 agentic RL 扩展到<strong>超过 10K 个可验证的软件工程环境、terminal 环境和多跳搜索任务</strong>；</li>
  <li>Forge[9] 强调系统跨越了十万级 real-world scaffolds 与数千种工具调用格式。</li>
</ul>

<p>真正的 agentic capability 不是从静态数据里背下来的，而是从<strong>结构化、可验证、可迁移的环境</strong>里训练出来的。</p>

<p>环境建模的核心，不是把现实世界完整模拟出来，而是把真实工作转写成一个<strong>结构上不失真的可训练决策过程</strong>——重要的不是表面真实，而是 <strong>structural fidelity</strong>：动作空间、关键信息流、失败模式和成功判据，是否与真实部署保持一致。举一个典型例子：一个客服 agent 不必复现公司所有噪声，但必须保留库存状态、退款规则、权限边界、上下文记忆、工具接口、升级流程和最终评分 rubric；否则学到的只是”像在做客服”，而不是”真的能做客服”。</p>

<p><strong>环境覆盖度，是 Agentic RL 的第一条 scaling axis</strong>。但真实任务的难点往往不在 data scaling，而在 <strong>specification scaling</strong>：很多高价值任务之所以难进训练闭环，不是因为模型不够聪明，而是因为任务没有被写成机器可执行、机器可验证的规范。下一代 env scaling 更像三个”编译器”问题：</p>

<ul>
  <li><strong>task compiler</strong>：把模糊请求编译成初始状态、工具、约束和终止条件；</li>
  <li><strong>verifier compiler</strong>：把”做得好不好”编译成可执行检查、rubric 和必要时的人类审阅；</li>
  <li><strong>scaffold compiler</strong>：把同一能力放进不同 agent scaffold、tool schema 和 orchestration loop，避免模型只记住单一 workflow。</li>
</ul>

<p>Forge 强调跨大规模 scaffold 训练，本质上就在处理第三个问题。真实人类任务里最大的问题不是”任务太少”，而是 <strong>evaluator 太弱</strong>——一旦 verifier 失真，模型就会学会 hacking，而不是学会工作。SWE-Universe[10] 把环境构建、self-verification 和 hacking detection 自动化，说明大家已经开始把”防投机评测”当成环境的一部分。</p>

<h3 id="2-探索能力与多样性保持不是把-temperature-调高而是维护可探索行为的空间">2. 探索能力与多样性保持：不是把 temperature 调高，而是维护可探索行为的空间</h3>

<p>很多人一谈”探索”就想到：调高 temperature、多采几个 rollout、加 entropy regularization。但对 agentic RL，这些都只是表层现象。核心问题是：<strong>模型在训练的不同阶段，是否仍然保有一组彼此可区分、都可能成功、且在参数空间里真实可达的行为路径</strong>。</p>

<p>对 reasoning 模型，这个问题已经被直接观察到：随着 SFT 推进，Pass@1 可以继续上升，但 <strong>Pass@k 会快速恶化</strong>，而且后续 RL 往往也恢复不了；仅靠 token-level 的多样化解码，距离理论上的 oracle 上界仍有明显差距。<strong>真正塌缩的不是采样温度，而是模型权重层面的行为可探索空间。</strong></p>

<p>所以这一节最本质的思想是：<strong>探索本质上是一个 support management 问题</strong>。你要管理的不是 token 级噪声，而是模型是否还保有：</p>

<ul>
  <li>多种合法任务分解</li>
  <li>多种工具调用顺序</li>
  <li>多种上下文组织方式</li>
  <li>多种长度的 reasoning path</li>
  <li>在 agent 场景下的多种 memory / planning / action 组合</li>
</ul>

<p>只要这些分支在参数里还活着，后续 RL 才有可能通过 verifier 和 rollout 把它们放大；一旦在进入 RL 前就被压没，训练再稳定也只是在缩水的空间里做局部优化。</p>

<p><strong>预训练 / 基座阶段</strong>决定的是 <strong>reachable support</strong>——模型是否已经具备足够多的技能碎片、长上下文耐受性、工具使用先验和任务分解能力：</p>

<ul>
  <li>MiniMax-M1[11] 把额外 7.5T continual pretraining 直接称为 “<em>Foundation for RL Scaling</em>“；</li>
  <li>Kimi K2 用 diverse agents、tool combinations 和 rubric-guided tasks，把未来 agent 可能探索的 action space 和 task space 提前做宽；</li>
  <li>DeepSeek-R1-Zero[12] 提供了另一个很有代表性的例子：它在没有 SFT 冷启动的前提下直接 RL，模型会自然增加思考时长，并逐步长出更长推理和自我修正的行为——这说明对能力足够强的基础模型，<strong>RL 过程本身就可能激发并放大更长程的推理与自我修正行为</strong>。</li>
</ul>

<p><strong>冷启动 / SFT 阶段</strong>真正要解决的，不仅是”把模型教得更会答题”，而是<strong>不要在进入 RL 之前就把分布压塌</strong>：</p>

<ul>
  <li>GEM[4] 的重要性不在于又提出一个新的 SFT loss，而在于它把问题说透了：标准交叉熵 SFT 会压缩输出分布，抹掉很多 alternative plausible outputs，而在线 RL 恰恰需要这些行为分歧来形成探索空间；</li>
  <li>Getting Your LLMs Ready for RL[13] 进一步指出：最适合接 RL 的 checkpoint，<strong>往往不是 validation 上表现最好的那个</strong>——在传统过拟合发生之前，模型就可能已经出现 distributional forgetting，过度偏离 base distribution，从而损害后续 RL 的潜力。</li>
</ul>

<p><strong>到了在线 RL 阶段</strong>，探索问题又会表现成另一种形态：即便模型内部还保留着多种路径，如果 RL 目标只盯 correctness，训练仍然会把概率质量持续推向少数高回报模式：</p>

<ul>
  <li>DAPO[14] 把 Clip-Higher 明确写成 “<em>promotes diversity and avoids entropy collapse</em>“；</li>
  <li>Diversity-Aware Policy Optimization[15] 在 12 个 LLM 上给出更强的经验结论：solution diversity 与 Potential@k 存在强正相关，因此<strong>在 RL 目标中显式促进 token-level diversity</strong>，平均带来 3.5% 的数学推理提升。</li>
</ul>

<p>这里真正重要的不是某一个技巧，而是一个更深的转向：<strong>探索，第一次从”训练自然会保住的东西”，变成了需要被显式优化的对象</strong>。</p>

<p>这一维度今天仍有几个未解决的问题：</p>

<ol>
  <li>当前很多方法管理的仍然是 token entropy 或字符串级 diversity，但 agentic RL 真正需要保住的是<strong>语义层和策略层的多样性</strong>——不同工具顺序、不同 memory 操作、不同任务分解不一定表现为更高的 token entropy；</li>
  <li>很多系统的 verifier 偏 outcome-only，天然低估那些”短期看更绕、长期却更有价值”的探索路径；</li>
  <li>社区仍过度依赖 Pass@1，而对 Pass@k、Potential@k、解法簇数量、跨 scaffold 迁移这些更接近探索前沿的指标重视不够。</li>
</ol>

<h3 id="3-算力分配与学习信号整理谁拿到-rollout谁才真正有机会被学到">3. 算力分配与学习信号整理：谁拿到 rollout，谁才真正有机会被学到</h3>

<p>上一节讨论”多样化的采样路径是否存在”，这一节讨论<strong>在固定 rollout 预算下，这些路径里哪些会真正进入梯度</strong>。探索解决<strong>可达性</strong>，算力分配解决<strong>可学习性</strong>。</p>

<p>对 reasoning / agentic RL 来说，模型内部也许还保留着多种策略，但如果 rollout 总是平均分给”已经学会的简单题”和”暂时完全学不会的极难题”，训练既看不到组内差异，也形成不了有效梯度——在稀疏奖励和 group baseline 设置下，很多 prompt-group 会退化成全 0 或全 1，<strong>advantage energy 为 0，gate 关闭</strong>，这些组消耗了算力却没有产生 usable learning signal。</p>

<p>因此，真正该优化的目标不只是平均 reward，而是更接近训练动力学本身的量：</p>

<ul>
  <li>non-zero gradient ratio</li>
  <li>gate-open probability</li>
  <li>组内 reward spread</li>
  <li>单位训练时间内的有效样本率</li>
</ul>

<p><strong>算力分配是 credit assignment 的上游机制</strong>：谁拿到更多 rollout，谁就更有机会被比较、被区分、被学到。</p>

<p>主流做法可以分成三类——</p>

<p><strong>① 方差控制视角</strong>。既然不同 prompt 对梯度方差的贡献不同，那么 rollout 预算就应该优先投给那些最能减少估计方差、最可能恢复学习信号的 prompt：</p>

<ul>
  <li>GVM-RAFT[17] 从 acceptance rate 和 gradient noise 的角度做动态分配；</li>
  <li>VIP[18] 更系统，用轻量高斯过程预测 prompt 成功概率，再转成 gradient variance 估计，并在固定预算约束下解一个 rollout allocation 优化问题。VIP 明确把目标写成 <em>minimize the expected gradient variance of the policy update</em>，而不是机械拉高 pass rate。</li>
</ul>

<p>这标志着 rollout allocation 开始从经验 heuristics 变成 <strong>policy optimization 的一部分</strong>。</p>

<p><strong>② 学习价值—成本权衡视角</strong>。Knapsack RL[6] 把每个任务的探索看成”具有不同 value 和 cost 的 item”，由此推出自适应资源分配规则——把预算从已经学饱和的题转移到更可能产出信号的题。预算分配不是为了省钱，而是<strong>避免把大量算力烧在注定不会更新参数的地方</strong>。</p>

<p><strong>③ 主动恢复信号视角</strong>。Reinforce-Ada[19] 认为很多”所谓难 prompt”没法学，其实是 undersampling 造成的统计假象，而不是模型真没潜力。于是它不再用固定小组、统一采样被动等待 mixed outcomes，而是<strong>根据 prompt 难度动态增加推理预算</strong>，主动去找出那些本来会被 uniform GRPO 漏掉的信号。</p>

<p>这个话题还有不少未解问题：</p>

<ul>
  <li>现有 allocator 主要依赖 pass rate、variance proxy 或近期 rollout 统计，但这些量不等于<strong>长期训练价值</strong>——一个 prompt 今天方差大，不代表明天最值得更多预算；</li>
  <li>现有方法仍把单条 prompt 作为分配单位，但 agentic RL 的训练难度更多取决于<strong>交互结构和执行状态</strong>（scaffold、工具链、历史记忆、任务阶段），而不只是 prompt 文本；</li>
  <li>大多数分配器优化的是局部训练效率，还没有把预算分配、reward 结构、hinting、off-policy freshness、长时程 credit assignment <strong>联合起来</strong>。</li>
</ul>

<p>下一步真正值得做的是：把 semantic difficulty、uncertainty、verifier sharpness、历史 learning gain、scaffold transfer 价值、甚至 hinting 后的 gate-open probability，<strong>一起纳入 allocation policy</strong>；把 prompt-level allocation 推广到 trajectory segment、tool-call branch、memory operation 这类更细的 agent 单位。到那时，算力分配才会真正从”更高效的训练技巧”变成 <strong>agentic RL 的核心算法层</strong>。</p>

<h3 id="4-目标函数与策略优化不先问用哪种-rl先问现在到底坏在哪">4. 目标函数与策略优化：不先问”用哪种 RL”，先问”现在到底坏在哪”</h3>

<p>这一部分重点不是 PPO、GRPO、REINFORCE 的技术细节，而是<strong>Agentic RL 的优化器究竟在控制什么</strong>。更本质地说，它在回答两个问题：</p>

<ol>
  <li>高回报轨迹要以多大力度被推回当前策略？</li>
  <li>rollout 分布、learner 更新分布、deployment 执行动作之间允许多大偏移？</li>
</ol>

<p>这里有一条<strong>常被忽略的基本事实</strong>：PPO 那一整套 value network machinery 未必是必要的。ReMax[5] 提醒我们，在文本生成这种快仿真、近似确定性转移、轨迹级奖励的设定下，REINFORCE 路线也可以既简单又稳定。Kimi K1.5 则把长 CoT RL 明确写成 <strong>relative-entropy regularized 的 online mirror descent</strong> 问题。</p>

<p>到了 K2.5、MiniMax-M1 和 GLM-5，问题进一步从”如何估 advantage”转成”如何控制长轨迹、异步 rollout、训练 / 推理 mismatch 下的 off-policy drift”，于是出现了这些看起来很细但实际上很关键的设计：</p>

<ul>
  <li><strong>K2.5 的 token-level clipping</strong>：处理 train-inference framework 差异放大的 off-policy divergence；</li>
  <li><strong>M1 的 CISPO</strong>：裁 importance weights 而不是裁 token updates，在保留更多 token 级梯度的同时控制比值爆炸；</li>
  <li><strong>GLM-5 的 TITO + 双边重要性采样</strong>：确保被优化的动作尽可能还是当时真正被采样的动作。</li>
</ul>

<p><strong>未来真正有价值的优化研究，不是继续修改 PPO 或 GRPO 的公式</strong>，而是先诊断：训练当前究竟受限于哪一类瓶颈——</p>

<ul>
  <li>梯度噪声过大？</li>
  <li>策略漂移过快？</li>
  <li>训练目标与真实任务不匹配？</li>
</ul>

<p>只有先定位清楚，才能决定是改进优势估计、采样方式、更新约束，还是训练调度策略。</p>

<h3 id="5-rollout-采样异步并行与调度调度策略本身就是算法的一部分">5. Rollout 采样、异步并行与调度：调度策略本身就是算法的一部分</h3>

<p>在真实 agent 场景，理想化的同步 on-policy RL 很难被满足：不同 rollout 完成时间差异极大，短的几秒，长的可能几十分钟甚至更久。<strong>坚持严格同步会被 straggler 拖死；完全贪心异步又把训练拖入过重的 off-policy 偏移</strong>。</p>

<p>各家给出的折中方案非常有代表性：</p>

<ul>
  <li><strong>Kimi K1.5 的 partial rollout</strong>：长轨迹切段，未完成部分进 replay buffer，下一轮继续，只有当前段要求 on-policy；</li>
  <li><strong>K2.5</strong>：每个 agent task 都当作独立异步 coroutine，通过专门的 Rollout Manager 支持高并发；</li>
  <li><strong>MiniMax 的 Windowed FIFO</strong>：在”严格 FIFO（稳但慢）”和”完全异步（快但漂移大）”之间做折中——不要求全局严格排队，只在有限窗口内保持大致顺序，让窗口里的已完成任务可以灵活先训练；</li>
  <li><strong>GLM-5</strong>：直接把采样和训练分开，一边持续并行生成轨迹，另一边独立消费数据，再用 TITO + 双边重要性采样 + 陈旧样本过滤来控制异步训练中不可避免的 off-policy 偏移。</li>
</ul>

<p><strong>很多人把 queueing、resume、tail-latency、staleness 当成工程问题，但在 agentic RL 里，调度实际上会改写训练分布</strong>。K1.5 的 partial rollout 意味着一条长轨迹由新旧段拼接而成；MiniMax 的 Windowed FIFO 直接控制了”允许新鲜样本先于更早提交的样本进入训练”的程度；GLM-5 的异步 Agent RL 更是明确承认”不现实去追踪所有历史行为策略，必须在可接受的偏差内做近似校正”。</p>

<p>Agentic RL 的核心不是”如何保持纯 on-policy”，而是<strong>如何在不可避免的异步与陈旧性下，让偏移保持在仍然有学习价值的范围内</strong>。这就是为什么 rollout system 不是承载算法的底座——<strong>它本身就是算法的一部分</strong>。</p>

<h3 id="6-奖励验证器与效率约束reward-定义的不只是答对而是怎样工作才算好">6. 奖励、验证器与效率约束：Reward 定义的不只是”答对”，而是”怎样工作才算好”</h3>

<p>很多关于 agentic RL 的讨论会说”verifier 就够了”。这对真实 Agent 任务其实不成立：agent 的成功不只体现在 final correctness 上，还体现在动作是否合理、工具调用是否合适、是否浪费上下文、是否无意义过度思考、是否拖慢总完成时间、以及输出是否符合更高层的质量和交互要求。</p>

<p>几家的具体做法非常有参考价值：</p>

<ul>
  <li><strong>K2.5</strong>：可验证任务用 rule-based outcome reward，token 成本用 budget-control reward，开放任务用多 rubric GRM，并通过 Toggle 在”尽量做对”和”尽量省 token”之间交替优化；</li>
  <li><strong>MiniMax-M1</strong>：verifiable 与 unverifiable 任务分开处理，用 GenRM 处理不能靠规则验证的任务，并<strong>特别讨论了长 CoT 下 GenRM 的 length bias</strong>——奖励模型偏好更长但未必更好的回答，会直接诱发 reward hacking；</li>
  <li><strong>GLM-5</strong>：把 rule-based reward、ORM、GRM 组合成 hybrid reward system，并明确写出三者权衡——规则奖励精确但窄，ORM 低方差但容易被 exploit，GRM 更灵活但方差更高；</li>
  <li><strong>Forge</strong>：进一步把中间过程质量和<strong>任务完成时间</strong>都纳入 agent RL——真实用户需要的不是”最终做对但过程低效、等待很久”的系统，而是”既能做对、又能较快完成”的 agent。</li>
</ul>

<p><strong>对 reward 正确的理解是”工作方式的规范化”，而不只是”答案质量的评分器”</strong>：</p>

<ul>
  <li>K2.5 用多个 GRM rubric，是因为单一偏好信号太容易被过拟合；</li>
  <li>M1 专门处理 GenRM 长度偏置，是因为 reward model 一旦系统性偏向 verbose response，整个 RL 就会被带偏；</li>
  <li>Forge 引入完成时间相关奖励，是因为真实部署中 agent 的效用不只由正确率决定，还取决于实际耗时。</li>
</ul>

<p>Reward design 的关键<strong>不是给模型更多分数</strong>，而是把 correctness、quality、efficiency、robustness 拆开，再决定哪些可以硬验证、哪些要用模型判断、哪些必须通过对抗测试和 OOD transfer 来防止被投机。</p>

<h3 id="7-记忆层级与并行-agent被训练的对象已经不只是-token-policy而是-operating-policy">7. 记忆、层级与并行 Agent：被训练的对象已经不只是 Token Policy，而是 Operating Policy</h3>

<p>很多人一谈 long-context agent 就想”把 context window 做大一点”。但<strong>长上下文不等于记忆，更不等于好的 agent</strong>。核心问题是：当交互历史越来越长、工具观察越来越多时，模型如何决定什么该保留、什么该丢弃、什么该压缩、什么时候拆任务、什么时候并行多个子 agent？</p>

<ul>
  <li><strong>MiniMax Forge 的 Context Rot</strong>：即使没有触到绝对 context window 上限，长轮次交互中累积的中间推理和冗余 observation 也会造成 attention dilution，让模型失焦。于是 Forge 直接把 <strong>Context Management 纳入 RL 交互回路</strong>，把它当作一种显式 action，让 context transition 本身成为环境状态转移的一部分；</li>
  <li><strong>GLM-5</strong> 在搜索 agent 上也观察到极长上下文会明显伤害性能，因此使用 <strong>keep-recent-k 与 discard-all 的层级式 context management</strong>；</li>
  <li><strong>K2.5 的 Agent Swarm 与 PARL</strong>：当单 agent 顺序执行的延迟变得不可接受时，让 orchestrator 学会<strong>动态任务分解、子 agent 创建和并行调度</strong>。训练时只更新 orchestrator、冻结 sub-agent，以规避最难的 credit ambiguity 与训练不稳定。</li>
</ul>

<p><strong>被优化的对象已经从”token 级生成策略”扩展成”操作系统级策略”</strong>——模型不再只决定下一个 token，而是在决定：</p>

<ul>
  <li>算力怎么花</li>
  <li>上下文怎么管</li>
  <li>任务怎么拆</li>
  <li>子 agent 怎么协作</li>
</ul>

<p>K2.5 的一个关键 insight：<strong>真正的并行 agent 不是把同一个模型复制几份并发运行</strong>，而是让 orchestrator 学会”什么时候值得并行、如何分配子任务、如何在最终汇总时保持全局一致性”。Forge 则强调：记忆管理如果只在 inference 端手工加规则、训练时没见过这种状态转移，最终会形成严重的 inference-training mismatch。</p>

<p>未来 agentic RL 的 frontier，未必是让模型”再想更久”，而是<strong>把 memory editing、hierarchical decomposition 和 agent orchestration 一起纳入训练目标</strong>。</p>

<h3 id="8-infra-基础设施它不是承载算法的底座而是在塑造训练分布">8. Infra 基础设施：它不是承载算法的底座，而是在塑造训练分布</h3>

<p>如果说 RLHF 是在一个相对规整的 prompt → completion → reward → update 闭环上做优化，那么 Agentic RL 面对的是<strong>长短极不均匀、工具调用密集、环境反馈异步、动作语义复杂</strong>的真实交互轨迹。</p>

<p>在这种设定下，基础设施直接决定：</p>

<ul>
  <li>rollout 以什么顺序完成</li>
  <li>哪些样本因过时被丢弃</li>
  <li>哪些前缀能够复用</li>
  <li>训练端和推理端看到的是否还是同一个动作空间</li>
</ul>

<p>这里有三层 infra：</p>

<p><strong>① 塑造训练分布的 rollout / learner 基础设施</strong>。由于任务完成时间可能从秒级跨到小时级，同步 on-policy 几乎不可能，系统必须处理 actor–learner 解耦、队列调度、buffer freshness、checkpoint staleness、partial rollout reuse、stale sample filtering。MiniMax Forge 把 strict FIFO / greedy async / Windowed FIFO 的权衡直接写成”吞吐与分布稳定之间的核心矛盾”；GLM-5 通过异步 generation-training 解耦 + TITO + double-sided importance sampling 控制偏移；K1.5 的 partial rollout reuse 说明<strong>长轨迹能否被复用，本身就是训练 recipe 的一部分</strong>。这一层 infra 直接塑造了”模型真正看到的训练分布”。</p>

<p><strong>② 提升吞吐与成本效率的规模化训练 / 推理 infra</strong>。包括训练 / 推理解耦、数据池缓存、KV / prefix 复用、动态 batching、各种并行化和异构资源调度策略。它们解决的核心问题不是”单点算法是否成立”，而是”这些方法能否在现实成本下真正跑到足够规模”。对 agent workload 来说，模型生成、环境执行、工具调用、verifier 计算、日志存储的资源瓶颈完全不同，基础设施<strong>必须是解耦和分层的</strong>，不能继续沿用单一、同步、同构的训练范式。</p>

<p><strong>③ 保证数值一致性和训练—推理一致性的 serving infra</strong>。最容易被低估但其实最关键：Agentic RL 优化的不是抽象文本，而是<strong>具有明确执行语义的动作序列</strong>——训练时、采样时、部署时对动作的表示或接口稍有错位，模型学到的策略就可能在上线时部分失效。GLM-5 的 TITO 之所以重要，不只是为了省一次 re-tokenization，而是为了精确保持 sampled action 与 optimized action 的对应；MiniMax Forge 的 gateway 与 middleware 设计本质上也在做 action interface standardization。因此，tokenizer / engine 对齐、tool schema 标准化、trajectory serialization、metadata logging、train-serving alignment——都不再只是工程细节，而是在决定<strong>训练时被优化的动作，是否真的是部署时会执行的那个动作</strong>。</p>

<h3 id="9-评测与可观测性测不出来的不变量就守不住">9. 评测与可观测性：测不出来的不变量，就守不住</h3>

<p>前面八节讲了”在哪里守不变量”，但有一个被大多数文章忽视的基础问题：<strong>如果你连三个不变量是否正在被守住都测不出来，就根本无从管理它们</strong>。</p>

<p>Agentic RL 的 evaluation 不能只看 Pass@1 或 final reward，至少需要三类互补的观测维度：</p>

<p><strong>① 探索健康度（对应第一不变量）</strong>：</p>

<ul>
  <li>Pass@k、Potential@k、解法簇数量（semantic cluster count）</li>
  <li>行为路径的 scaffold 迁移率（同一能力在不同 scaffold 下的成功率）</li>
  <li>长期 entropy trajectory 与 action-level 多样性（而不仅是 token-level）</li>
</ul>

<p><strong>② 学习信号健康度（对应第二不变量）</strong>：</p>

<ul>
  <li><strong>non-zero advantage ratio</strong>：一个 batch 内多少 group 产生了非零梯度</li>
  <li><strong>gate-open probability</strong>：group-based 方法中 advantage 有效的样本比例</li>
  <li><strong>组内 reward spread</strong> 和 <strong>gradient SNR</strong></li>
  <li><strong>单位训练时间的有效样本率</strong>（effective tokens per GPU-hour）</li>
</ul>

<p>这些量往往比 loss 曲线更能解释”为什么训练看起来还在跑，但能力没长”。</p>

<p><strong>③ 分布一致性（对应第三不变量）</strong>：</p>

<ul>
  <li><strong>training–serving KL</strong>：相同 prompt 下训练 checkpoint 与部署 checkpoint 的输出分布差异</li>
  <li><strong>rollout staleness 分布</strong>：样本被生成时的策略与被学习时的策略相隔多少步</li>
  <li><strong>tokenizer / tool schema mismatch 率</strong>：训练端与部署端接口一致性的硬指标</li>
  <li><strong>长轨迹误差累积曲线</strong>：模型表现随交互步数退化的速度</li>
</ul>

<p>在更高层次上，还需要一套<strong>对抗性评测</strong>：verifier-hacking 检测、reward-model OOD 探针、scaffold 替换测试、工具噪声注入测试——这些不是”锦上添花的 benchmark”，而是<strong>第一不变量和第二不变量是否被守住的直接证据</strong>。</p>

<p>SWE-Universe 把 hacking detection 自动化进环境，本质上就是在承认：<strong>评测已经不是 pipeline 的末端，而是训练系统的一部分</strong>。没有这层观测，所谓”调参”就只是在黑箱里做随机扰动。</p>

<hr />

<h2 id="四结语agentic-rl-的真正竞争不在单点算法">四、结语：Agentic RL 的真正竞争，不在单点算法</h2>

<p>回到开头那句话——<strong>Agentic RL 的核心问题，已经从”怎么更新参数”扩展为”怎么在真实 Agent 环境里持续制造可用的学习信号”</strong>。</p>

<p>把三条技术路线放在一起看，信号非常清楚：</p>

<ul>
  <li><strong>Kimi 路线</strong>告诉我们：(1) 长上下文本身是一条 RL scaling axis；(2) 复杂的 value function / MCTS / process RM 不是唯一道路，简洁但分布一致的 policy optimization 也能跑出很强的长链能力；(3) 当 agent 工作流变复杂后，奖励模型、token-level clipping、token efficiency 控制和 learned parallel orchestration 会越来越重要。K1.5 → K2.5 的演进，本质上是从”把长 reasoning RL 跑通”走向”把多步 agentic / multimodal RL 规模化”；</li>
  <li><strong>MiniMax 路线</strong>说明：长时程 agent RL 一进到真实环境，首要问题很快就从”模型能不能推理”转向”<strong>系统能不能稳定地持续学习</strong>“。M1 的 CISPO 的价值在于修复长轨迹 RL 的 off-policy 和梯度裁剪副作用；Forge 进一步证明，异步调度、上下文管理、完成时间奖励、跨任务联合训练、前缀树合并这类”看起来很工程”的东西，<strong>实际上决定了你最终能否在大规模真实环境里把 RL 跑起来</strong>；</li>
  <li><strong>GLM 路线</strong>强调：后训练不应该一股脑混在一起，而要按能力类型分阶段组织，并借助蒸馏机制保护已有能力。Reasoning RL → Agentic RL → General RL 的顺序<strong>不只是训练日程安排，而是一种能力编排方式</strong>。GLM-5 对异步 RL 基础设施、TITO、double-sided importance sampling 的强调，也再次说明：<strong>训练系统与策略优化之间已经没有清晰边界</strong>。</li>
</ul>

<p>综合这些路线，一个清晰的结论是：</p>

<blockquote>
  <p>Agentic RL 不只是”更大模型 × 更多数据 × 更多 token”，而是：</p>

  <ul>
    <li>更丰富的<strong>环境覆盖</strong></li>
    <li>更高密度的<strong>有效学习信号</strong></li>
    <li>更一致的 <strong>rollout / update / serving 分布</strong></li>
    <li>更高的<strong>单位时间有效训练效率</strong></li>
    <li>以及能让你<strong>确认前四者正在发生</strong>的评测与可观测性</li>
  </ul>
</blockquote>

<p>在完善高效的 infra 支持下，谁在这五个维度上同时做得更好，谁就更可能真正把 agent 训出来。</p>

<hr />

<h2 id="参考文献">参考文献</h2>

<p>[1] Kimi Team. <em>Kimi k1.5: Scaling Reinforcement Learning with LLMs</em>. arXiv:2501.12599, 2025. <a href="https://arxiv.org/abs/2501.12599">https://arxiv.org/abs/2501.12599</a></p>

<p>[2] Kimi Team. <em>Kimi K2: Open Agentic Intelligence</em>. arXiv:2507.20534, 2025. <a href="https://arxiv.org/abs/2507.20534">https://arxiv.org/abs/2507.20534</a></p>

<p>[3] Kimi Team. <em>Kimi K2.5: Visual Agentic Intelligence</em>. arXiv:2602.02276, 2026. <a href="https://arxiv.org/abs/2602.02276">https://arxiv.org/abs/2602.02276</a></p>

<p>[4] Ziniu Li et al. <em>Preserving Diversity in Supervised Fine-Tuning of Large Language Models</em>. arXiv:2408.16673, 2024. <a href="https://arxiv.org/abs/2408.16673">https://arxiv.org/abs/2408.16673</a></p>

<p>[5] Ziniu Li et al. <em>ReMax: A Simple, Effective, and Efficient Reinforcement Learning Method for Aligning Large Language Models</em>. arXiv:2310.10505, 2023. <a href="https://arxiv.org/abs/2310.10505">https://arxiv.org/abs/2310.10505</a></p>

<p>[6] Ziniu Li et al. <em>Knapsack RL: Unlocking Exploration of LLMs via Optimizing Budget Allocation</em>. arXiv:2509.25849, 2025. <a href="https://arxiv.org/abs/2509.25849">https://arxiv.org/abs/2509.25849</a></p>

<p>[7] Hanze Dong. <em>Curate the Learning Signal for Reinforcement Learning: Variance Minimization, Adaptive Sampling, and Self-Hinting</em>. Blog post, 2026. <a href="https://hendrydong.github.io/blogs/pages/rl-ada.html">https://hendrydong.github.io/blogs/pages/rl-ada.html</a></p>

<p>[8] GLM-5 Team. <em>GLM-5: from Vibe Coding to Agentic Engineering</em>. arXiv:2602.15763, 2026. <a href="https://arxiv.org/abs/2602.15763">https://arxiv.org/abs/2602.15763</a></p>

<p>[9] MiniMax. <em>Forge: Scalable Agent RL Framework and Algorithm</em>. MiniMax News, 2026. <a href="https://www.minimax.io/news/forge-scalable-agent-rl-framework-and-algorithm">https://www.minimax.io/news/forge-scalable-agent-rl-framework-and-algorithm</a></p>

<p>[10] Mouxiang Chen et al. <em>SWE-Universe: Scale Real-World Verifiable Environments to Millions</em>. arXiv:2602.02361, 2026. <a href="https://arxiv.org/abs/2602.02361">https://arxiv.org/abs/2602.02361</a></p>

<p>[11] MiniMax. <em>MiniMax-M1: Scaling Test-Time Compute Efficiently with Lightning Attention</em>. arXiv:2506.13585, 2025. <a href="https://arxiv.org/abs/2506.13585">https://arxiv.org/abs/2506.13585</a></p>

<p>[12] DeepSeek-AI. <em>DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via Reinforcement Learning</em>. arXiv:2501.12948, 2025. <a href="https://arxiv.org/abs/2501.12948">https://arxiv.org/abs/2501.12948</a></p>

<p>[13] Xinran Li et al. <em>Getting Your LLMs Ready for Reinforcement Learning with Lightweight SFT</em>. OpenReview / ICLR 2026. <a href="https://openreview.net/forum?id=yezWGJmODg">https://openreview.net/forum?id=yezWGJmODg</a></p>

<p>[14] Qiyuan Yu et al. <em>DAPO: An Open-Source LLM Reinforcement Learning System at Scale</em>. arXiv:2503.14476, 2025. <a href="https://arxiv.org/abs/2503.14476">https://arxiv.org/abs/2503.14476</a></p>

<p>[15] Jian Yao et al. <em>Diversity-Aware Policy Optimization for Large Language Model Reasoning</em>. arXiv:2505.23433, 2025. <a href="https://arxiv.org/abs/2505.23433">https://arxiv.org/abs/2505.23433</a></p>

<p>[16] Xingyu Dang et al. <em>Assessing Diversity Collapse in Reasoning</em>. OpenReview, 2025. <a href="https://openreview.net/forum?id=AMiKsHLjQh">https://openreview.net/forum?id=AMiKsHLjQh</a></p>

<p>[17] Jiarui Yao et al. <em>Optimizing Chain-of-Thought Reasoners via Gradient Variance Minimization in Rejection Sampling and RL</em>. arXiv:2505.02391, 2025. <a href="https://arxiv.org/abs/2505.02391">https://arxiv.org/abs/2505.02391</a></p>

<p>[18] Hieu Trung Nguyen et al. <em>Adaptive Rollout Allocation for Online Reinforcement Learning with Verifiable Rewards</em>. arXiv:2602.01601, 2026. <a href="https://arxiv.org/abs/2602.01601">https://arxiv.org/abs/2602.01601</a></p>

<p>[19] Wei Xiong et al. <em>Reinforce-Ada: An Adaptive Sampling Framework for Reinforce-Style LLM Training</em>. arXiv:2510.04996, 2025. <a href="https://arxiv.org/abs/2510.04996">https://arxiv.org/abs/2510.04996</a></p>]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><category term="blog" /><category term="reinforcement learning" /><category term="agentic" /><category term="llm" /><category term="post-training" /><category term="infrastructure" /><summary type="html"><![CDATA[过去一年，各家大模型公司公开的技术报告透出的最重要信号，不是又出现了一个更好的 PPO/GRPO 变体，而是真正有效的 Agentic RL 已经从”单轮文本优化”转向了”在长上下文、工具调用、部分可观测、异步执行环境中的系统性策略学习”。 Kimi K1.5[1] 把长上下文 RL、partial rollout 重用和 mirror-descent 风格的 policy optimization 拉到了台前；Kimi K2[2]/K2.5[3] 又把 agentic 数据合成、多模态 RL、token-level clipping、GRM rubric、Toggle、PARL / Agent Swarm 这些关键部件公开；MiniMax 把另一个事实讲得更彻底：当 rollout 时长从秒级扩到小时级，训练瓶颈就不再是 loss design，而是吞吐、稳定性与 agent 灵活性之间的三难权衡；GLM 则强调分阶段 RL：Reasoning RL、Agentic RL、General RL 不是混在一起一次训完，而是通过顺序化 pipeline 逐步推进，并借助异步 RL 基础设施与跨阶段蒸馏来兼顾长时程 agent 学习与能力保持。 Agentic RL 的核心问题，已经从”怎么更新参数”扩展为”怎么在真实 Agent 环境里持续制造可用的学习信号，并用在线交互的轨迹数据驱动优化“。]]></summary></entry><entry><title type="html">图解 Wan2.1 I2V：从一张图到一段视频，模型到底发生了什么</title><link href="https://liyongzhi.xyz/posts/2026/04/wan21-i2v-explained/" rel="alternate" type="text/html" title="图解 Wan2.1 I2V：从一张图到一段视频，模型到底发生了什么" /><published>2026-04-24T00:00:00+08:00</published><updated>2026-04-24T00:00:00+08:00</updated><id>https://liyongzhi.xyz/posts/2026/04/blog-post-wan21-i2v-explained</id><content type="html" xml:base="https://liyongzhi.xyz/posts/2026/04/wan21-i2v-explained/"><![CDATA[<p>最近视频生成模型卷得很快，<code class="language-plaintext highlighter-rouge">Wan2.1</code> 是阿里 Wan 团队开源的那一套。它最常用的场景之一就是 <strong>I2V（Image-to-Video）</strong>：给一张参考图加一句文字 prompt，模型给你生成一段几秒的视频，首帧基本还是那张图，后续的镜头就按你写的文字去演。</p>

<p>这篇文章想做的事情是：</p>

<blockquote>
  <p>把 Wan2.1 I2V 里<strong>每一步数据发生了什么</strong>讲清楚，让从没接触过视频生成的人也能看懂。</p>
</blockquote>

<p>我们会从最外层的”图像 + 文字 → 视频”讲起，一路剥开壳子：<br />
VAE 到底在压缩什么、CLIP 和 T5 各自管什么、DiT 内部是怎么把图像信息和文字信息混进去的、采样循环为什么要跑那么多步、以及为什么首帧会这么”像”你给的那张图。</p>

<p><br />
<img align="center" width="1000" src="https://liyongzhi.xyz/images/posts/wan21-i2v-overview.svg" alt="Wan2.1 I2V overall architecture" />
<br /></p>

<p>这张图是全文的总地图。下面的每一节都是在放大它的某一块。</p>

<hr />

<h2 id="1-先做一次外行翻译i2v-到底在做什么">1. 先做一次”外行翻译”：I2V 到底在做什么</h2>

<p>如果用一句日常语言来描述 I2V，其实是：</p>

<blockquote>
  <p>我们有一张图（<code class="language-plaintext highlighter-rouge">3 × H × W</code>，RGB 像素），想把它”续写”成一段视频（<code class="language-plaintext highlighter-rouge">3 × F × H × W</code>，F 帧），而且这段视频的内容要符合文字 prompt。</p>
</blockquote>

<p>朴素想法是直接训练一个”图 + 文字 → 视频”的网络。问题有二：</p>

<ol>
  <li>视频的体积太大。即便是 480p × 24fps × 4 秒，也已经是 1.1 亿像素级别，直接建模太贵。</li>
  <li>我们希望生成过程是<strong>可控的</strong>——能调 guidance，能控制风格，能多步修正——而不是一次性跑完一个巨大网络就结束。</li>
</ol>

<p>Diffusion 模型的套路恰好能解决这两件事：</p>

<ul>
  <li><strong>压缩</strong>：用 VAE 把视频压到一个小很多的 latent 空间，之后所有运算都在 latent 上做。</li>
  <li><strong>迭代</strong>：扩散模型天然是多步的，每一步都在”把更接近噪声的视频”往”更清晰的视频”方向推一点。</li>
</ul>

<p>所以 Wan2.1 I2V 的骨架分成两大块：</p>

<table>
  <thead>
    <tr>
      <th>模块</th>
      <th>角色</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>Wan-VAE</strong></td>
      <td>像素 ⇄ latent 的翻译员</td>
    </tr>
    <tr>
      <td><strong>DiT</strong></td>
      <td>在 latent 空间里”去噪”的大脑</td>
    </tr>
  </tbody>
</table>

<p>外加两个<strong>条件编码器</strong>：</p>

<table>
  <thead>
    <tr>
      <th>模块</th>
      <th>角色</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>CLIP ViT-H/14</strong></td>
      <td>把参考图变成”这张图看起来讲了什么”的高层语义向量</td>
    </tr>
    <tr>
      <td><strong>umT5</strong></td>
      <td>把文字 prompt 编码成一串 token embedding</td>
    </tr>
  </tbody>
</table>

<p>接下来我们分别看每一块。</p>

<hr />

<h2 id="2-wan-vae把视频压缩-256-倍后再还原">2. Wan-VAE：把视频压缩 256 倍后再还原</h2>

<p><code class="language-plaintext highlighter-rouge">Wan-VAE</code> 是一个 <strong>3D Causal VAE</strong>。它做的事很朴素：</p>

<ul>
  <li>输入：<code class="language-plaintext highlighter-rouge">[3, F, H, W]</code> 的视频（或单张图当作 <code class="language-plaintext highlighter-rouge">F=1</code>）</li>
  <li>输出：<code class="language-plaintext highlighter-rouge">[16, F/4, H/8, W/8]</code> 的 latent</li>
</ul>

<p>换句话说：</p>

<ul>
  <li>空间下采样 <code class="language-plaintext highlighter-rouge">8×8</code> 倍</li>
  <li>时间下采样 <code class="language-plaintext highlighter-rouge">4</code> 倍</li>
  <li>通道数从 <code class="language-plaintext highlighter-rouge">3</code> 变成 <code class="language-plaintext highlighter-rouge">16</code>（表达能力变强）</li>
</ul>

<p>总体积约压缩 <strong>256 倍</strong>（<code class="language-plaintext highlighter-rouge">8·8·4 / (16/3) ≈ 24×24 / ...</code>，算下来大约 48× 的”信息体积”，但浮点数要少 200+ 倍）。</p>

<blockquote>
  <p><strong>为什么叫 Causal？</strong> 指的是它的时间卷积只看”过去”不看”未来”，这样可以支持变长视频、流式推理，和后续滚动生成新帧。</p>
</blockquote>

<p>一个关键点是 <strong>I2V 里 VAE 会被用两次</strong>：</p>

<ol>
  <li><strong>编码参考图</strong>：把那张图当成一个 <code class="language-plaintext highlighter-rouge">F=1</code> 的视频编码，得到它的 latent。</li>
  <li><strong>解码最终视频</strong>：DiT 输出 latent，扔给 VAE 解码回像素视频。</li>
</ol>

<p>其中第一次编码的结果被塞进 DiT 作为”低层像素/结构”条件——这是后面讲 I2V 双路条件时的关键一环。</p>

<hr />

<h2 id="3-两路文字图像条件编码clip-和-t5-各自做什么">3. 两路文字/图像条件编码：CLIP 和 T5 各自做什么</h2>

<p>这两个模型很多人容易搞混，但它们在 Wan2.1 里分工很清晰。</p>

<h3 id="31-umt5把文字变成-512--4096-的-token-序列">3.1 umT5：把文字变成 512 × 4096 的 token 序列</h3>

<p><code class="language-plaintext highlighter-rouge">umT5</code> 是 T5 的多语言版。输入是你的 prompt，输出是：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[seq_len, 4096]   # 每个 token 一个 4096 维向量
</code></pre></div></div>

<p>Wan2.1 统一把这个序列 padding / truncate 到 <strong>512 个 token</strong>，所以文本总是 <code class="language-plaintext highlighter-rouge">[512, 4096]</code>。</p>

<blockquote>
  <p>T5 是一个纯文本的大模型，它的向量很”语言化”，擅长表达语义、句法关系。</p>
</blockquote>

<h3 id="32-clip-vit-h14把图像变成-257--1280-的-token-序列">3.2 CLIP ViT-H/14：把图像变成 257 × 1280 的 token 序列</h3>

<p><code class="language-plaintext highlighter-rouge">CLIP</code> 是一个<strong>跨模态</strong>模型（图像 + 文本对齐训练的），这里我们只用它的<strong>图像编码器</strong>（ViT-H/14）。</p>

<p>它吃一张 <code class="language-plaintext highlighter-rouge">224 × 224</code> 的图，输出：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[257, 1280]
</code></pre></div></div>

<p><strong>257 从哪来？</strong> 这是一个很常见的数字：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">ViT-H/14</code> 把 224×224 切成 <code class="language-plaintext highlighter-rouge">14×14</code> 的 patch</li>
  <li><code class="language-plaintext highlighter-rouge">224 / 14 = 16</code>，所以一张图变成 <code class="language-plaintext highlighter-rouge">16 × 16 = 256</code> 个 patch token</li>
  <li>再加一个 <code class="language-plaintext highlighter-rouge">CLS</code> token，总共 <strong>257</strong> 个</li>
</ul>

<p>每个 token 的通道数是 <code class="language-plaintext highlighter-rouge">1280</code>（ViT-H 的隐藏维度）。</p>

<blockquote>
  <p>CLIP 给出的是<strong>图像的高层语义</strong>：它知道这张图里是”一只猫”、”傍晚的海边”、”油画风格”之类的语义抽象，但几乎不保留像素级的精细结构。</p>
</blockquote>

<h3 id="33-clip-vs-t5为什么两个都要">3.3 CLIP vs T5：为什么两个都要？</h3>

<p>这是 I2V 非常关键的一点。两者的”关注点”不一样：</p>

<table>
  <thead>
    <tr>
      <th> </th>
      <th>擅长</th>
      <th>不擅长</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><strong>T5</strong></td>
      <td>文字描述的动作、意图、场景</td>
      <td>图像具体长什么样</td>
    </tr>
    <tr>
      <td><strong>CLIP</strong></td>
      <td>参考图的整体风格、主体</td>
      <td>精确的像素/空间结构</td>
    </tr>
  </tbody>
</table>

<p>所以两者是<strong>互补</strong>的——都给 DiT 看一遍，DiT 再自己挑。这也是为什么后面会看到 cross-attention 是”双流”的。</p>

<hr />

<h2 id="4-把参考图的像素也塞进模型条件-latent-y">4. 把”参考图的像素”也塞进模型：条件 latent <code class="language-plaintext highlighter-rouge">y</code></h2>

<p>到这里我们已经有了两条图像通路：CLIP（语义）和 T5（文字）。但对 I2V 来说，仅靠 CLIP 的语义是不够的——生成的第一帧如果不能”长得非常像”输入图，用户立刻会觉得不对。</p>

<p>于是 Wan2.1 加了第三条通路：<strong>把参考图用 VAE 编码后，直接在通道维度拼到噪声 latent 上</strong>。</p>

<h3 id="41-构造-y">4.1 构造 <code class="language-plaintext highlighter-rouge">y</code></h3>

<p>假设目标视频是 <code class="language-plaintext highlighter-rouge">F</code> 帧，latent 形状 <code class="language-plaintext highlighter-rouge">[16, F/4, H/8, W/8]</code>。我们把 <code class="language-plaintext highlighter-rouge">T_latent = F/4</code>。</p>

<p><strong>第 1 步：把参考图放到第 0 帧，其余帧置零。</strong></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">video_clip</span> <span class="o">=</span> <span class="n">concat</span><span class="p">([</span>
    <span class="n">img_resized</span><span class="p">,</span>                <span class="c1"># [3, 1, H, W]  ← 第 0 帧 = 参考图
</span>    <span class="n">zeros</span><span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="n">F</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">H</span><span class="p">,</span> <span class="n">W</span><span class="p">)</span>         <span class="c1"># 其余帧为 0
</span><span class="p">],</span> <span class="n">dim</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>                       <span class="c1"># → [3, F, H, W]
</span></code></pre></div></div>

<p><strong>第 2 步：VAE 编码。</strong></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">y_latent</span> <span class="o">=</span> <span class="n">VAE</span><span class="p">.</span><span class="n">encode</span><span class="p">(</span><span class="n">video_clip</span><span class="p">)</span>   <span class="c1"># → [16, T_latent, H/8, W/8]
</span></code></pre></div></div>

<p><strong>第 3 步：构造时间 mask，标记”哪些帧是已知的”。</strong></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">msk</span> <span class="o">=</span> <span class="n">ones</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">F</span><span class="p">,</span> <span class="n">H_lat</span><span class="p">,</span> <span class="n">W_lat</span><span class="p">)</span>
<span class="n">msk</span><span class="p">[:,</span> <span class="mi">1</span><span class="p">:]</span> <span class="o">=</span> <span class="mi">0</span>                         <span class="c1"># 只有第 0 帧 = 1
# 把 msk[:, 0:1] 沿时间 repeat 4 次，和 msk[:, 1:] 拼接
# 再 reshape 成 [4, T_latent, H_lat, W_lat]
</span></code></pre></div></div>

<p>这里的 <code class="language-plaintext highlighter-rouge">4</code> 是 VAE 的时间 stride——我们需要让 mask 通道数足够”表达”被 VAE 压缩掉的时间细节。</p>

<p><strong>第 4 步：mask 和 VAE latent 通道拼接，得到 <code class="language-plaintext highlighter-rouge">y</code>。</strong></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">y</span> <span class="o">=</span> <span class="n">concat</span><span class="p">([</span><span class="n">msk</span><span class="p">,</span> <span class="n">y_latent</span><span class="p">],</span> <span class="n">dim</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>    <span class="c1"># [4 + 16 = 20, T_latent, H_lat, W_lat]
</span></code></pre></div></div>

<blockquote>
  <p>把 <code class="language-plaintext highlighter-rouge">y</code> 想成一块”透明纸”：第 0 帧那一层写满了”你要照着这张图画”，其它帧那一层是空白，同时还有一层专门标注”哪里非空白”。</p>
</blockquote>

<h3 id="42-y-怎么进-dit">4.2 <code class="language-plaintext highlighter-rouge">y</code> 怎么进 DiT</h3>

<p>DiT 的输入是噪声 latent <code class="language-plaintext highlighter-rouge">x_t: [16, T_latent, H_lat, W_lat]</code>。进网络前做一次通道拼接：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x = concat(x_t, y) = [16 + 20, T, H, W] = [36, T, H, W]
</code></pre></div></div>

<p>所以 I2V 的 DiT 输入通道是 <strong>36</strong>（T2V 是 16）。也正因为这个差别，I2V checkpoint 的 <code class="language-plaintext highlighter-rouge">patch_embedding</code> 卷积权重和 T2V 不是一回事。</p>

<hr />

<h2 id="5-dit-内部一个时间步里到底跑了什么">5. DiT 内部：一个时间步里到底跑了什么</h2>

<p>接下来进入最核心的部分。我们放一下 DiT 单层的结构图：</p>

<p><br />
<img align="center" width="1100" src="https://liyongzhi.xyz/images/posts/wan21-i2v-block.svg" alt="Wan2.1 DiT block internal" />
<br /></p>

<p>整体看，DiT 是一个典型的 Transformer 栈，但有三个重要定制：</p>

<ol>
  <li><strong>时空 3D RoPE</strong>（self-attention 里的位置编码）</li>
  <li><strong>双流 cross-attention</strong>（image KV + text KV）</li>
  <li><strong>AdaLN-Zero 风格的 timestep 调制</strong></li>
</ol>

<p>下面一条一条讲。</p>

<h3 id="51-patchify把视频-latent-变成-transformer-的-token-序列">5.1 Patchify：把视频 latent 变成 Transformer 的 token 序列</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="bp">self</span><span class="p">.</span><span class="n">patch_embedding</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Conv3d</span><span class="p">(</span><span class="mi">36</span><span class="p">,</span> <span class="n">dim</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span><span class="mi">2</span><span class="p">,</span><span class="mi">2</span><span class="p">),</span> <span class="n">stride</span><span class="o">=</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span><span class="mi">2</span><span class="p">,</span><span class="mi">2</span><span class="p">))</span>
</code></pre></div></div>

<p>这是一个”<strong>3D patchify</strong>“：用 <code class="language-plaintext highlighter-rouge">Conv3d</code> 把每个 <code class="language-plaintext highlighter-rouge">1×2×2</code> 的时空小块打成一个 token。</p>

<ul>
  <li>时间方向 kernel=1，意味着<strong>时间维度不被合并</strong>（每一个 latent 帧仍然是独立的一层 token）。</li>
  <li>空间方向 kernel=2，把 <code class="language-plaintext highlighter-rouge">H_lat × W_lat</code> 的网格再进一步压 <code class="language-plaintext highlighter-rouge">2×2</code>，得到 <code class="language-plaintext highlighter-rouge">H_lat/2 × W_lat/2</code> 个 token。</li>
</ul>

<p>最终序列长度是：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>S = T_latent × (H_lat/2) × (W_lat/2)
</code></pre></div></div>

<p>每个 token 是 <code class="language-plaintext highlighter-rouge">dim</code> 维向量（1.3B 版本里 <code class="language-plaintext highlighter-rouge">dim=2048</code>）。</p>

<h3 id="52-timestep-embedding让每层都知道现在在第几步">5.2 Timestep embedding：让每层都知道”现在在第几步”</h3>

<p>扩散模型的一个关键差别是每一步的处理方式不一样。T=T_max 时几乎全是噪声，T=0 时已经是完整视频，所以模型在不同 step 应该”轻重不一”。</p>

<p>Wan2.1 的做法是 <strong>AdaLN-Zero</strong>（DiT 论文里的那一套）：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">e</span> <span class="o">=</span> <span class="n">sinusoidal_embedding_1d</span><span class="p">(</span><span class="mi">256</span><span class="p">,</span> <span class="n">t</span><span class="p">)</span>          <span class="c1"># 标量 t → 256 维向量
</span><span class="n">e</span> <span class="o">=</span> <span class="n">time_embedding</span><span class="p">(</span><span class="n">e</span><span class="p">)</span>                        <span class="c1"># MLP 投到 dim
</span><span class="n">e0</span> <span class="o">=</span> <span class="n">time_projection</span><span class="p">(</span><span class="n">e</span><span class="p">).</span><span class="n">unflatten</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="p">(</span><span class="mi">6</span><span class="p">,</span><span class="n">dim</span><span class="p">))</span>  <span class="c1"># 再投成 [B, 6, dim]
</span></code></pre></div></div>

<p>然后把这 6 份向量分发给每个 block，块内再加上自己可学习的 <code class="language-plaintext highlighter-rouge">modulation</code> 参数，切成 6 组：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>(shift1, scale1, gate1,  shift2, scale2, gate2)
</code></pre></div></div>

<ul>
  <li><code class="language-plaintext highlighter-rouge">shift, scale</code> 用在 LayerNorm 之后：<code class="language-plaintext highlighter-rouge">x' = norm(x) · (1 + scale) + shift</code></li>
  <li><code class="language-plaintext highlighter-rouge">gate</code> 用在残差分支：<code class="language-plaintext highlighter-rouge">x = x + gate · f(x')</code></li>
</ul>

<blockquote>
  <p><strong>“Zero” 的含义</strong>：<code class="language-plaintext highlighter-rouge">gate</code> 初始化为 0，使得模型训练开始时每个 block 都是恒等映射——DiT 从一个干净的起点开始学。</p>
</blockquote>

<p>注意：cross-attention <strong>不被 AdaLN 调制</strong>，只有 self-attention 和 FFN 被调制。</p>

<h3 id="53-self-attention3d-全局注意力--分解式-rope">5.3 Self-Attention：3D 全局注意力 + 分解式 RoPE</h3>

<p>这一步做的事很简单：<strong>视频 token 之间互相看</strong>。</p>

<p>代码上是标准的 QKV flash attention，但有两处定制：</p>

<p><strong>① QK 做 RMSNorm</strong>。这是稳定训练用的技巧：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">q</span> <span class="o">=</span> <span class="n">RMSNorm</span><span class="p">(</span><span class="n">Linear_q</span><span class="p">(</span><span class="n">x</span><span class="p">))</span>
<span class="n">k</span> <span class="o">=</span> <span class="n">RMSNorm</span><span class="p">(</span><span class="n">Linear_k</span><span class="p">(</span><span class="n">x</span><span class="p">))</span>
<span class="n">v</span> <span class="o">=</span> <span class="n">Linear_v</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
</code></pre></div></div>

<p><strong>② 3D 分解式 RoPE 作用在 Q/K 上</strong>（不作用于 V）。</p>

<p>视频 token 有三个坐标：<code class="language-plaintext highlighter-rouge">(frame, height, width)</code>。Wan2.1 把每个 head 的维度 <code class="language-plaintext highlighter-rouge">d</code> 切成三段：</p>

<table>
  <thead>
    <tr>
      <th>段</th>
      <th>通道数</th>
      <th>编码的是</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>时间</td>
      <td><code class="language-plaintext highlighter-rouge">d − 4·(d/6)</code></td>
      <td>帧索引 <code class="language-plaintext highlighter-rouge">f</code></td>
    </tr>
    <tr>
      <td>高</td>
      <td><code class="language-plaintext highlighter-rouge">2·(d/6)</code></td>
      <td>行索引 <code class="language-plaintext highlighter-rouge">h</code></td>
    </tr>
    <tr>
      <td>宽</td>
      <td><code class="language-plaintext highlighter-rouge">2·(d/6)</code></td>
      <td>列索引 <code class="language-plaintext highlighter-rouge">w</code></td>
    </tr>
  </tbody>
</table>

<p>三段分别应用一维 RoPE（复数旋转），然后沿通道拼回一起：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>q_rot = RoPE_T(q[:, :, :dT], f) ⊕ RoPE_H(q[:, :, dT:dT+dH], h) ⊕ RoPE_W(q[:, :, dT+dH:], w)
</code></pre></div></div>

<p><strong>为什么这样设计？</strong></p>

<ul>
  <li>可以支持<strong>任意分辨率和帧数</strong>，因为 RoPE 是外推良好的位置编码。</li>
  <li>时间/空间的频率独立，模型可以各自学合适的”时间尺度”和”空间尺度”。</li>
  <li>相比绝对位置嵌入，训练时可以在一个尺度下训，推理时换尺度不会崩。</li>
</ul>

<p><strong>注意力的范围是”全 3D 全局”</strong>——所有视频 token 互相能看。这就是为什么视频生成模型这么贵：序列长度是 <code class="language-plaintext highlighter-rouge">T × H × W</code>，attention 是 O(S²)。</p>

<h3 id="54-cross-attention双流融合图像--文字">5.4 Cross-Attention：双流融合图像 + 文字</h3>

<p>到了 I2V 最有意思的设计。先回忆一下 context 长什么样：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>context = [ CLIP_257 ∥ T5_512 ]         # shape [769, dim]
         ─── 前 257 ───   ── 后 512 ──
             (图像)           (文本)
</code></pre></div></div>

<p>如果用朴素 cross-attention，你会一次算 <code class="language-plaintext highlighter-rouge">attn(q, K=k_all, V=v_all)</code>，让视频 token 对这 769 个 token 做 softmax。问题是图像和文本的分布差距很大，softmax 会把注意力偏到一侧。</p>

<p>Wan2.1 的做法是<strong>双流独立</strong>：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 共享 Query
</span><span class="n">q</span> <span class="o">=</span> <span class="n">Linear_q</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>

<span class="c1"># 图像分支（独立的 k_img, v_img）
</span><span class="n">k_img</span> <span class="o">=</span> <span class="n">Linear_k_img</span><span class="p">(</span><span class="n">context</span><span class="p">[:,</span> <span class="p">:</span><span class="mi">257</span><span class="p">])</span>
<span class="n">v_img</span> <span class="o">=</span> <span class="n">Linear_v_img</span><span class="p">(</span><span class="n">context</span><span class="p">[:,</span> <span class="p">:</span><span class="mi">257</span><span class="p">])</span>

<span class="c1"># 文本分支（共享 T2V 的 k, v）
</span><span class="n">k_txt</span> <span class="o">=</span> <span class="n">Linear_k</span><span class="p">(</span><span class="n">context</span><span class="p">[:,</span> <span class="mi">257</span><span class="p">:])</span>
<span class="n">v_txt</span> <span class="o">=</span> <span class="n">Linear_v</span><span class="p">(</span><span class="n">context</span><span class="p">[:,</span> <span class="mi">257</span><span class="p">:])</span>

<span class="n">out_img</span> <span class="o">=</span> <span class="n">flash_attn</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">k_img</span><span class="p">,</span> <span class="n">v_img</span><span class="p">)</span>
<span class="n">out_txt</span> <span class="o">=</span> <span class="n">flash_attn</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">k_txt</span><span class="p">,</span> <span class="n">v_txt</span><span class="p">)</span>

<span class="n">out</span> <span class="o">=</span> <span class="n">Linear_o</span><span class="p">(</span><span class="n">out_img</span> <span class="o">+</span> <span class="n">out_txt</span><span class="p">)</span>   <span class="c1"># 逐元素相加再过输出投影
</span></code></pre></div></div>

<p>几个关键设计点：</p>

<ol>
  <li><strong>独立的 K/V 投影</strong>：图像用 <code class="language-plaintext highlighter-rouge">k_img, v_img</code>，文本用 <code class="language-plaintext highlighter-rouge">k, v</code>。每一模态在自己的几何空间里算 attention，不会互相挤压 softmax。</li>
  <li><strong>两次独立 attention 再相加</strong>：相当于两种信号<strong>分别</strong>给每个视频 token 打了一次分，再叠加作为新的残差。</li>
  <li><strong>Q 共享</strong>：视频 token 只有一份”问题”，问图和文字同一个问题：”你们谁和我相关？”</li>
  <li><strong>无 RoPE</strong>：cross-attn 中的 K/V 是外部序列，不需要视频的时空位置编码。</li>
</ol>

<blockquote>
  <p>直观理解：<strong>image 分支管”我希望长什么样”，text 分支管”我希望怎么演”，两个加在一起就是视频 token 的条件梯度</strong>。</p>
</blockquote>

<h3 id="55-ffn标准-mlp再来一次-adaln-门控">5.5 FFN：标准 MLP，再来一次 AdaLN 门控</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">y</span> <span class="o">=</span> <span class="n">ffn</span><span class="p">(</span><span class="n">norm2</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="err">·</span> <span class="p">(</span><span class="mi">1</span><span class="o">+</span><span class="n">scale2</span><span class="p">)</span> <span class="o">+</span> <span class="n">shift2</span><span class="p">)</span>
<span class="n">x</span> <span class="o">=</span> <span class="n">x</span> <span class="o">+</span> <span class="n">gate2</span> <span class="err">·</span> <span class="n">y</span>
</code></pre></div></div>

<p><code class="language-plaintext highlighter-rouge">ffn</code> 就是常规的 <code class="language-plaintext highlighter-rouge">Linear → GELU → Linear</code>，中间维度是 <code class="language-plaintext highlighter-rouge">4 × dim</code>（比如 dim=2048 时 ffn_dim=8192）。</p>

<p>到这里一个 block 就结束了。把这个 block 叠 32 层（1.3B 版）或 40 层（14B 版），最后过一个 <code class="language-plaintext highlighter-rouge">Head</code>（也带 AdaLN 和 unpatchify），就能把 <code class="language-plaintext highlighter-rouge">[B, S, dim]</code> 变回 <code class="language-plaintext highlighter-rouge">[16, T, H, W]</code>——也就是模型对<strong>当前时间步的”速度场” <code class="language-plaintext highlighter-rouge">v</code></strong> 的预测。</p>

<hr />

<h2 id="6-训练目标为什么叫-flow-matching不再叫预测噪声">6. 训练目标：为什么叫 Flow Matching，不再叫”预测噪声”</h2>

<p>DDPM 早年是让模型预测”这张图里的噪声 <code class="language-plaintext highlighter-rouge">ε</code>“。Wan2.1 用的是 <strong>Flow Matching / Rectified Flow</strong> 的范式——本质上是把扩散过程理解成<strong>一条从噪声到数据的直线路径</strong>，模型学的是这条路径上每一点的”速度”。</p>

<p>具体来说，定义一条插值：</p>

\[x_t = (1 - t) \cdot x_0 + t \cdot \epsilon, \quad t \in [0, 1], \quad \epsilon \sim \mathcal{N}(0, I)\]

<p>那么真值速度就是：</p>

\[v^* = \frac{d x_t}{d t} = \epsilon - x_0\]

<p>训练目标：</p>

\[\mathcal{L}_{\text{FM}} = \mathbb{E}_{x_0, \epsilon, t} \left\| v_\theta(x_t, t, c) - (\epsilon - x_0) \right\|^2\]

<p>其中 <code class="language-plaintext highlighter-rouge">c = {y, CLIP_fea, T5_text}</code> 是所有条件的合集。</p>

<p><strong>Flow Matching 相比预测 ε 有什么好处？</strong></p>

<ul>
  <li>训练 loss 更稳定，对 <code class="language-plaintext highlighter-rouge">t</code> 的依赖更平滑。</li>
  <li>采样时可以用更少的步数。典型配置 <strong>25–50 步</strong>即可出不错结果（早期 DDPM 需要 1000 步）。</li>
  <li>路径”直”这件事意味着模型不容易陷入局部的噪声拟合。</li>
</ul>

<hr />

<h2 id="7-一次完整的推理25-步里到底发生了什么">7. 一次完整的推理：25 步里到底发生了什么</h2>

<p>现在把所有东西串起来。假设你给了一张 <code class="language-plaintext highlighter-rouge">H × W</code> 的图、一句 prompt，让模型生成 F 帧的视频：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>─── 推理前准备（只做 1 次）────────────────────────────────
1. t5_ctx  = umT5(prompt)                      # [512, 4096] → MLP → [512, dim]
2. clip_fea = CLIP.visual(image)               # [257, 1280] → MLPProj → [257, dim]
3. img_lat = VAE.encode([image, zeros, ...])   # [16, T, H_lat, W_lat]
4. msk     = build_mask(first_frame=1)         # [4, T, H_lat, W_lat]
5. y       = concat(msk, img_lat)              # [20, T, H_lat, W_lat]
6. context = concat(clip_fea, t5_ctx)          # [769, dim]
7. x_T     ~ N(0, I)                           # [16, T, H_lat, W_lat]

─── 采样循环（跑 25~50 次）─────────────────────────────
for t in schedule:  # e.g. [1.0, 0.96, ..., 0.0]
    # (可选) CFG: 各跑一次有/无条件
    x_in = concat(x_t, y)          # [36, T, H_lat, W_lat]
    v_cond = DiT(x_in, t, context, clip_fea, y)
    # v_uncond = DiT(x_in, t, empty_context, ...)
    # v = v_uncond + s · (v_cond - v_uncond)
    v = v_cond
    
    x_{t-Δt} = x_t - v · Δt        # flow matching 欧拉步

─── 解码 ──────────────────────────────────────────
video_latent = x_0                  # [16, T, H_lat, W_lat]
video        = VAE.decode(video_latent)  # [3, F, H, W]
</code></pre></div></div>

<p>几个细节：</p>

<ul>
  <li><strong><code class="language-plaintext highlighter-rouge">y</code> 只构造一次</strong>，在整个 25 步里都用同一份。因为参考图是不变的。</li>
  <li><strong>CFG（Classifier-Free Guidance）</strong>：Wan2.1 训练时会随机丢弃条件，所以推理时可以通过 <code class="language-plaintext highlighter-rouge">v = v_u + s·(v_c - v_u)</code> 放大条件信号（典型 <code class="language-plaintext highlighter-rouge">s=5~7.5</code>）。每步需要跑两遍 DiT。</li>
  <li><strong>首帧为什么保真？</strong>：因为第 0 帧的 <code class="language-plaintext highlighter-rouge">mask=1</code> 和 <code class="language-plaintext highlighter-rouge">VAE(img)</code> 一直被塞进输入，DiT 每步都在”被提醒”首帧应该长什么样。随着 t 变小，模型越来越相信这个约束。</li>
</ul>

<hr />

<h2 id="8-几个关键数字一张表带走">8. 几个关键数字一张表带走</h2>

<table>
  <thead>
    <tr>
      <th>参数</th>
      <th>值</th>
      <th>解释</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>VAE 空间 stride</td>
      <td>8</td>
      <td>H/W 方向下采样倍率</td>
    </tr>
    <tr>
      <td>VAE 时间 stride</td>
      <td>4</td>
      <td>F 方向下采样倍率</td>
    </tr>
    <tr>
      <td>VAE latent 通道</td>
      <td>16</td>
      <td>压缩后的通道数</td>
    </tr>
    <tr>
      <td>I2V <code class="language-plaintext highlighter-rouge">y</code> 通道</td>
      <td><strong>20</strong></td>
      <td>4 (mask) + 16 (VAE latent)</td>
    </tr>
    <tr>
      <td>DiT 输入通道</td>
      <td><strong>36</strong></td>
      <td>16 (noise) + 20 (y)</td>
    </tr>
    <tr>
      <td>Patch size</td>
      <td>(1, 2, 2)</td>
      <td>时间不并、空间 2×2</td>
    </tr>
    <tr>
      <td>文本 token 数</td>
      <td>512</td>
      <td>umT5 输出 padded</td>
    </tr>
    <tr>
      <td>CLIP token 数</td>
      <td><strong>257</strong></td>
      <td>1 CLS + 16×16 patches</td>
    </tr>
    <tr>
      <td>CLIP 维度</td>
      <td>1280</td>
      <td>ViT-H 的 hidden</td>
    </tr>
    <tr>
      <td>DiT hidden</td>
      <td>2048 (1.3B) / 更大 (14B)</td>
      <td> </td>
    </tr>
    <tr>
      <td>DiT 层数</td>
      <td>32 / 40</td>
      <td> </td>
    </tr>
    <tr>
      <td>注意力头</td>
      <td>16</td>
      <td>head_dim=128</td>
    </tr>
    <tr>
      <td>Sampling 步数</td>
      <td>25–50</td>
      <td>Flow Matching 下</td>
    </tr>
  </tbody>
</table>

<hr />

<h2 id="9-t2v-vs-i2v到底改了哪里">9. T2V vs I2V：到底改了哪里</h2>

<p>最后来一张对比表，帮你一眼看清两种模型的差别：</p>

<table>
  <thead>
    <tr>
      <th>方面</th>
      <th>T2V</th>
      <th>I2V</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>输入条件</td>
      <td>只有文本</td>
      <td>文本 + 参考图</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">patch_embedding</code> in_channels</td>
      <td>16</td>
      <td><strong>36</strong></td>
    </tr>
    <tr>
      <td>Cross-Attention 类型</td>
      <td>单流（只有文本 K/V）</td>
      <td><strong>双流</strong>（image K_img/V_img + text K/V）</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">img_emb</code> (CLIP → dim MLP)</td>
      <td>❌ 无</td>
      <td>✅ 有</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">y</code>（mask + image latent）</td>
      <td>❌ 无</td>
      <td>✅ 有，通道拼接到 <code class="language-plaintext highlighter-rouge">x</code></td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">clip_fea</code></td>
      <td>❌ 无</td>
      <td>✅ 前置到 context</td>
    </tr>
    <tr>
      <td>采样过程</td>
      <td>一样（flow matching）</td>
      <td>一样</td>
    </tr>
  </tbody>
</table>

<p>所以 T2V → I2V 的改造量其实并不大：<strong>多了两条图像通路（CLIP 语义 + VAE 像素），外加一组额外的 cross-attention K/V 权重</strong>，其它骨架完全一致。这也是为什么很多团队能从 T2V checkpoint 微调出 I2V 版本。</p>

<hr />

<h2 id="10-常见疑问答疑">10. 常见疑问答疑</h2>

<p><strong>Q1：为什么不只用 CLIP、不要 VAE latent？</strong><br />
只用 CLIP 的话，模型知道”这是一只猫”，但不知道”这只猫在图里具体长什么样、坐在什么位置、毛色分布怎样”。CLIP 太高层。VAE latent 保留了像素级结构，所以首帧能做到”几乎像素级一致”。</p>

<p><strong>Q2：为什么不只用 VAE latent、不要 CLIP？</strong><br />
VAE latent 是”为了重建像素而设计的压缩特征”，它缺乏跨模态语义。CLIP 的语义向量能让模型在后续帧里理解”这张图在讲什么”，从而和 prompt 对齐得更好。两者是语义和像素的两极，缺一不可。</p>

<p><strong>Q3：mask 通道为什么要 4 维，不能是 1 维？</strong><br />
因为 VAE 的时间 stride = 4，一个 latent 帧对应 4 个像素帧。4 通道的 mask 让每个 latent 帧能独立标记”这 4 帧里各自是不是已知”。这样在滚动生成或多帧条件 I2V 里能无缝扩展。</p>

<p><strong>Q4：为什么 cross-attn 不做 RoPE？</strong><br />
RoPE 是为 query/key 在同一个坐标系下的相对距离准备的。cross-attn 的 key 来自外部序列（文本/图像 token），没有和视频 token 共享的”时空坐标”，用 RoPE 反而有害。</p>

<p><strong>Q5：CFG 在 I2V 里到底丢的是什么？</strong><br />
Wan2.1 做 CFG 时通常<strong>只丢文本</strong>（把 <code class="language-plaintext highlighter-rouge">t5_ctx</code> 置空），保留 CLIP 和 VAE latent。因为 I2V 的核心约束是参考图，不能丢；被用来”放大信号”的是文本 prompt。有些实现也会同时丢 CLIP，做”image guidance”。</p>

<p><strong>Q6：能不能做多图 / 多首尾帧条件？</strong><br />
可以。<code class="language-plaintext highlighter-rouge">y</code> 的结构天然支持——只需要把对应帧位置的 mask 设为 1、在 VAE 输入里把那些帧填真实图像即可。这就是社区里各种”首尾帧控制”、”关键帧插值”玩法的实现基础。</p>

<hr />

<h2 id="11-总结">11. 总结</h2>

<p>回到开头那张大图，现在你应该能一眼看懂每一块发生了什么：</p>

<ul>
  <li><strong>VAE</strong> 负责压缩像素和还原像素；</li>
  <li><strong>T5</strong> 负责理解文字；</li>
  <li><strong>CLIP</strong> 负责理解图像的”长相和风格”；</li>
  <li><strong>DiT</strong> 在一个压缩的 latent 空间里，一步一步把噪声拉回视频，拉的方向由前三个模块的条件决定；</li>
  <li><strong>I2V</strong> 的所有”魔法”就是把参考图的信息<strong>同时</strong>从两条通路（像素 / 语义）塞给 DiT，再用 cross-attention 双流、AdaLN 门控把它们融合进每个视频 token。</li>
</ul>

<p>一旦把这张图想清楚，你去读 Wan2.1 源码、甚至去扩展它（做首尾帧、多图参考、风格迁移），都会容易很多。</p>

<hr />

<h2 id="sources">Sources</h2>

<ul>
  <li><a href="https://github.com/Wan-Video/Wan2.1">Wan2.1 官方仓库</a></li>
  <li><a href="https://arxiv.org/abs/2503.20314">Wan 技术报告 arXiv:2503.20314</a></li>
  <li><a href="https://huggingface.co/spaces/2chch/Wan2.1/blob/2a07a689c8837aac720a73915b783ea23b371927/wan/modules/model.py">Wan2.1 model.py 源码（HF 镜像）</a></li>
  <li><a href="https://github.com/Wan-Video/Wan2.1/blob/main/wan/image2video.py">Wan2.1 image2video.py</a></li>
  <li><a href="https://arxiv.org/abs/2212.09748">DiT: Scalable Diffusion Models with Transformers (Peebles &amp; Xie, 2022)</a></li>
  <li><a href="https://arxiv.org/abs/2210.02747">Flow Matching for Generative Modeling (Lipman et al., 2023)</a></li>
  <li><a href="https://arxiv.org/abs/2209.03003">Rectified Flow (Liu et al., 2022)</a></li>
  <li><a href="https://arxiv.org/abs/2103.00020">CLIP (Radford et al., 2021)</a></li>
  <li><a href="https://arxiv.org/abs/2104.09864">RoFormer: RoPE (Su et al., 2021)</a></li>
</ul>]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><category term="blog" /><category term="diffusion" /><category term="video generation" /><category term="image-to-video" /><category term="DiT" /><category term="multimodal" /><summary type="html"><![CDATA[最近视频生成模型卷得很快，Wan2.1 是阿里 Wan 团队开源的那一套。它最常用的场景之一就是 I2V（Image-to-Video）：给一张参考图加一句文字 prompt，模型给你生成一段几秒的视频，首帧基本还是那张图，后续的镜头就按你写的文字去演。]]></summary></entry><entry><title type="html">大模型面试手撕题全攻略：Attention、Transformer、归一化与损失函数</title><link href="https://liyongzhi.xyz/posts/2026/04/llm-interview-implementations/" rel="alternate" type="text/html" title="大模型面试手撕题全攻略：Attention、Transformer、归一化与损失函数" /><published>2026-04-22T00:00:00+08:00</published><updated>2026-04-22T00:00:00+08:00</updated><id>https://liyongzhi.xyz/posts/2026/04/blog-post-llm-interview-implementations</id><content type="html" xml:base="https://liyongzhi.xyz/posts/2026/04/llm-interview-implementations/"><![CDATA[<blockquote>
  <p>大模型算法岗面试中，手撕代码是几乎绕不过去的一环。面试官会盯着你从零实现 Attention、MHA、GQA、LayerNorm、RMSNorm、SafeSoftmax、Cross-Entropy 等模块，既考察你对原理的理解，也考察你是否能在紧张的环境下把数值稳定性、维度对齐、broadcasting 这些细节处理干净。</p>

  <p>这篇文章把这些高频手撕题系统梳理一遍：每一节都给出<strong>核心原理 → 数学公式 → 从零手写的 PyTorch 实现 → 面试容易追问的点</strong>，读完之后这一类题你应该都能在白板上 10 分钟内写出来。</p>
</blockquote>

<hr />

<h2 id="1-self-attention所有-transformer-的起点">1. Self-Attention：所有 Transformer 的起点</h2>

<h3 id="11-核心思想">1.1 核心思想</h3>

<p>Self-Attention 要回答的问题非常简单：</p>

<blockquote>
  <p>给定一个序列里的每个 token，它应该从其它 token 里”抄”多少信息过来？</p>
</blockquote>

<p>它的三件套是 <code class="language-plaintext highlighter-rouge">Query</code>、<code class="language-plaintext highlighter-rouge">Key</code>、<code class="language-plaintext highlighter-rouge">Value</code>：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">Query</code>：当前 token 想”找什么”</li>
  <li><code class="language-plaintext highlighter-rouge">Key</code>：其它 token”能提供什么”（用来被检索）</li>
  <li><code class="language-plaintext highlighter-rouge">Value</code>：其它 token”真正要传递的内容”</li>
</ul>

<p>计算流程就一句话：<strong>Query 和 Key 做点积得到相似度，softmax 归一化后再加权求和 Value</strong>。</p>

<h3 id="12-数学公式">1.2 数学公式</h3>

<p>标准的 Scaled Dot-Product Attention：</p>

\[\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V\]

<p>其中：</p>

<ul>
  <li>$Q \in \mathbb{R}^{n \times d_k}$，$K \in \mathbb{R}^{m \times d_k}$，$V \in \mathbb{R}^{m \times d_v}$</li>
  <li>$n$ 是 query 序列长度，$m$ 是 key/value 序列长度</li>
  <li>$d_k$ 是每个头的维度</li>
</ul>

<p><strong>为什么要除以 $\sqrt{d_k}$？</strong></p>

<p>当 $d_k$ 很大时，$QK^\top$ 的方差会随着 $d_k$ 线性增长。点积数值过大会让 softmax 落到极端区域——梯度趋近 0，模型训不动。除以 $\sqrt{d_k}$ 可以把方差拉回 $O(1)$ 的量级。</p>

<p>简单推导：假设 $q$ 和 $k$ 每个分量均值 0、方差 1 且独立，那么</p>

\[\text{Var}(q \cdot k) = \text{Var}\left(\sum_{i=1}^{d_k} q_i k_i\right) = d_k\]

<p>所以除以 $\sqrt{d_k}$ 后方差变回 1。</p>

<h3 id="13-从零手撕实现">1.3 从零手撕实现</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>
<span class="kn">import</span> <span class="nn">math</span>

<span class="k">def</span> <span class="nf">scaled_dot_product_attention</span><span class="p">(</span><span class="n">Q</span><span class="p">,</span> <span class="n">K</span><span class="p">,</span> <span class="n">V</span><span class="p">,</span> <span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
    <span class="s">"""
    Q: [B, n, d_k]
    K: [B, m, d_k]
    V: [B, m, d_v]
    mask: [B, n, m]  True 的位置会被屏蔽
    """</span>
    <span class="n">d_k</span> <span class="o">=</span> <span class="n">Q</span><span class="p">.</span><span class="n">size</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span>
    <span class="c1"># [B, n, m]
</span>    <span class="n">scores</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">Q</span><span class="p">,</span> <span class="n">K</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">))</span> <span class="o">/</span> <span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">d_k</span><span class="p">)</span>

    <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
        <span class="n">scores</span> <span class="o">=</span> <span class="n">scores</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>

    <span class="n">attn</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">scores</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>       <span class="c1"># [B, n, m]
</span>    <span class="n">out</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">attn</span><span class="p">,</span> <span class="n">V</span><span class="p">)</span>            <span class="c1"># [B, n, d_v]
</span>    <span class="k">return</span> <span class="n">out</span><span class="p">,</span> <span class="n">attn</span>
</code></pre></div></div>

<h3 id="14-面试常见追问">1.4 面试常见追问</h3>

<ul>
  <li>
    <p><strong>Q：为什么用点积而不是加性注意力（Bahdanau）？</strong>
点积可以用矩阵乘法高效实现，GPU 友好；加性注意力多一层线性变换，不利于大规模并行。</p>
  </li>
  <li>
    <p><strong>Q：mask 怎么做？</strong>
Decoder 的 causal mask 是一个上三角为 True 的矩阵；Padding mask 是把 padding 位置置为 True。二者常合并使用。</p>
  </li>
  <li>
    <p><strong>Q：为什么要用 softmax 而不是别的归一化？</strong>
softmax 保证权重非负且和为 1，符合”加权平均”的语义；同时它是可导的。</p>
  </li>
</ul>

<hr />

<h2 id="2-multi-head-attention-mha">2. Multi-Head Attention (MHA)</h2>

<h3 id="21-为什么要多头">2.1 为什么要多头？</h3>

<p>单头注意力只能学到一种”相似度”模式。多头允许模型在<strong>不同的子空间里关注不同类型的关系</strong>——比如一个头学句法，一个头学语义，一个头学远距离依赖。</p>

<h3 id="22-数学公式">2.2 数学公式</h3>

\[\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \dots, \text{head}_h) W^O\]

<p>其中每个头独立计算：</p>

\[\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)\]

<p>维度上：</p>

<ul>
  <li>输入维度 $d_{\text{model}}$，头数 $h$，每头维度 $d_k = d_{\text{model}} / h$</li>
  <li>$W_i^Q, W_i^K \in \mathbb{R}^{d_{\text{model}} \times d_k}$</li>
  <li>$W^O \in \mathbb{R}^{d_{\text{model}} \times d_{\text{model}}}$</li>
</ul>

<p>注意：<strong>参数总量不变</strong>——切成 $h$ 个头之后每个头更”瘦”，但拼起来总维度还是 $d_{\text{model}}$。</p>

<h3 id="23-从零手撕实现">2.3 从零手撕实现</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">MultiHeadAttention</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">num_heads</span><span class="p">,</span> <span class="n">dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="k">assert</span> <span class="n">d_model</span> <span class="o">%</span> <span class="n">num_heads</span> <span class="o">==</span> <span class="mi">0</span><span class="p">,</span> <span class="s">"d_model 必须能被 num_heads 整除"</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">d_model</span> <span class="o">=</span> <span class="n">d_model</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span> <span class="o">=</span> <span class="n">num_heads</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span> <span class="o">=</span> <span class="n">d_model</span> <span class="o">//</span> <span class="n">num_heads</span>

        <span class="c1"># 一次性投影，效率更高
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">W_q</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">W_k</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">W_v</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">W_o</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>

        <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">q</span><span class="p">,</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span><span class="p">,</span> <span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
        <span class="s">"""
        q: [B, n, d_model]
        k, v: [B, m, d_model]
        mask: [B, 1, n, m] 或 [B, n, m]
        """</span>
        <span class="n">B</span><span class="p">,</span> <span class="n">n</span><span class="p">,</span> <span class="n">_</span> <span class="o">=</span> <span class="n">q</span><span class="p">.</span><span class="n">shape</span>
        <span class="n">m</span> <span class="o">=</span> <span class="n">k</span><span class="p">.</span><span class="n">size</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>

        <span class="c1"># 1. 线性投影
</span>        <span class="n">Q</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">W_q</span><span class="p">(</span><span class="n">q</span><span class="p">)</span>  <span class="c1"># [B, n, d_model]
</span>        <span class="n">K</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">W_k</span><span class="p">(</span><span class="n">k</span><span class="p">)</span>
        <span class="n">V</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">W_v</span><span class="p">(</span><span class="n">v</span><span class="p">)</span>

        <span class="c1"># 2. 切分成多头: [B, n, d_model] -&gt; [B, h, n, d_k]
</span>        <span class="n">Q</span> <span class="o">=</span> <span class="n">Q</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">n</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">K</span> <span class="o">=</span> <span class="n">K</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">m</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span> <span class="c1"># [B, h, m, d_k]
</span>        <span class="n">V</span> <span class="o">=</span> <span class="n">V</span><span class="p">.</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">m</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span> <span class="c1"># [B, h, m, d_k]
</span>
        <span class="c1"># 3. Scaled Dot-Product Attention
</span>        <span class="n">scores</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">Q</span><span class="p">,</span> <span class="n">K</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">))</span> <span class="o">/</span> <span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">)</span>
        <span class="c1"># scores: [B, h, n, m]
</span>        <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="k">if</span> <span class="n">mask</span><span class="p">.</span><span class="n">dim</span><span class="p">()</span> <span class="o">==</span> <span class="mi">3</span><span class="p">:</span>
                <span class="n">mask</span> <span class="o">=</span> <span class="n">mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>   <span class="c1"># 广播到 head 维
</span>            <span class="n">scores</span> <span class="o">=</span> <span class="n">scores</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>

        <span class="n">attn</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">scores</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span> <span class="c1"># [B,h,n,m]
</span>        <span class="n">attn</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span><span class="p">(</span><span class="n">attn</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">attn</span><span class="p">,</span> <span class="n">V</span><span class="p">)</span>        <span class="c1"># attn [B, h, n, m] V: [B, h, m, d_k] -&gt; [B, h, n, d_k]
</span>
        <span class="c1"># 4. 拼回来: [B, h, n, d_k] -&gt; [B, n, d_model]
</span>        <span class="n">out</span> <span class="o">=</span> <span class="n">out</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">).</span><span class="n">contiguous</span><span class="p">().</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">n</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_model</span><span class="p">)</span>

        <span class="c1"># 5. 输出投影
</span>        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">W_o</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
</code></pre></div></div>

<h3 id="24-面试常见追问">2.4 面试常见追问</h3>

<ul>
  <li>
    <p><strong>Q：头数 h 越多越好吗？</strong>
不一定。头数过多会让每个头的 $d_k$ 过小，表达能力反而下降；同时 KV Cache 也会膨胀。实际工程通常 32~64 个头。</p>
  </li>
  <li>
    <p><strong>Q：MHA 的时间复杂度？</strong>
$O(n^2 \cdot d_{\text{model}})$，序列长度是瓶颈。这也是 Flash Attention、线性 Attention 等工作的优化对象。</p>
  </li>
</ul>

<hr />

<h2 id="3-grouped-query-attention-gqa">3. Grouped-Query Attention (GQA)</h2>

<h3 id="31-为什么要-gqa">3.1 为什么要 GQA？</h3>

<p>推理阶段的显存瓶颈主要是 <strong>KV Cache</strong>——每一步解码都要存下所有历史 token 的 K 和 V。</p>

<ul>
  <li>MHA：每个 Q 头都有独立的 K、V 头，KV Cache = $2 \cdot n \cdot h \cdot d_k$</li>
  <li>MQA（Multi-Query Attention）：所有 Q 头共享一组 K、V，KV Cache 缩小 $h$ 倍，但质量下降</li>
  <li><strong>GQA</strong>：折中方案——Q 头分成 $g$ 组，每组共享一组 K、V</li>
</ul>

<p>当 $g = h$ 就是 MHA，当 $g = 1$ 就是 MQA。LLaMA-2/3、Mixtral 等主流开源模型都用的是 GQA。</p>

<p><br />
<img align="center" width="800" src="https://liyongzhi.xyz/images/posts/gqa-comparison.svg" alt="MHA vs GQA vs MQA comparison" />
<br /></p>

<h3 id="32-数学形式">3.2 数学形式</h3>

<p>设 Q 头数为 $h$，KV 头组数为 $g$，每组包含 $h / g$ 个 Q 头共享同一对 K、V：</p>

\[\text{head}_i = \text{Attention}(Q_i, K_{\lfloor i / (h/g) \rfloor}, V_{\lfloor i / (h/g) \rfloor})\]

<h3 id="33-从零手撕实现">3.3 从零手撕实现</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">GroupedQueryAttention</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">num_heads</span><span class="p">,</span> <span class="n">num_kv_groups</span><span class="p">,</span> <span class="n">dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="k">assert</span> <span class="n">d_model</span> <span class="o">%</span> <span class="n">num_heads</span> <span class="o">==</span> <span class="mi">0</span>
        <span class="k">assert</span> <span class="n">num_heads</span> <span class="o">%</span> <span class="n">num_kv_groups</span> <span class="o">==</span> <span class="mi">0</span><span class="p">,</span> <span class="s">"头数必须能被组数整除"</span>

        <span class="bp">self</span><span class="p">.</span><span class="n">d_model</span> <span class="o">=</span> <span class="n">d_model</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span> <span class="o">=</span> <span class="n">num_heads</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">num_kv_groups</span> <span class="o">=</span> <span class="n">num_kv_groups</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span> <span class="o">=</span> <span class="n">d_model</span> <span class="o">//</span> <span class="n">num_heads</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">group_size</span> <span class="o">=</span> <span class="n">num_heads</span> <span class="o">//</span> <span class="n">num_kv_groups</span>

        <span class="bp">self</span><span class="p">.</span><span class="n">W_q</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">num_heads</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="c1"># K, V 只用组数个头
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">W_k</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">num_kv_groups</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">W_v</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">num_kv_groups</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">W_o</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>

        <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">q</span><span class="p">,</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span><span class="p">,</span> <span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
        <span class="n">B</span><span class="p">,</span> <span class="n">n</span><span class="p">,</span> <span class="n">_</span> <span class="o">=</span> <span class="n">q</span><span class="p">.</span><span class="n">shape</span>
        <span class="n">m</span> <span class="o">=</span> <span class="n">k</span><span class="p">.</span><span class="n">size</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>

        <span class="c1"># Q: [B, h,  n, d_k]
</span>        <span class="c1"># K,V: [B, g, m, d_k]
</span>        <span class="n">Q</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">W_q</span><span class="p">(</span><span class="n">q</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">n</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">num_heads</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">K</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">W_k</span><span class="p">(</span><span class="n">k</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">m</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">num_kv_groups</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">V</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">W_v</span><span class="p">(</span><span class="n">v</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">m</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">num_kv_groups</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>

        <span class="c1"># 把 K, V 在头维度上"重复 group_size 次"对齐到 Q
</span>        <span class="c1"># [B, g, m, d_k] -&gt; [B, h, m, d_k]
</span>        <span class="n">K</span> <span class="o">=</span> <span class="n">K</span><span class="p">.</span><span class="n">repeat_interleave</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">group_size</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
        <span class="n">V</span> <span class="o">=</span> <span class="n">V</span><span class="p">.</span><span class="n">repeat_interleave</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">group_size</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>

        <span class="n">scores</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">Q</span><span class="p">,</span> <span class="n">K</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">))</span> <span class="o">/</span> <span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">)</span>
        <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="k">if</span> <span class="n">mask</span><span class="p">.</span><span class="n">dim</span><span class="p">()</span> <span class="o">==</span> <span class="mi">3</span><span class="p">:</span>
                <span class="n">mask</span> <span class="o">=</span> <span class="n">mask</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>
            <span class="n">scores</span> <span class="o">=</span> <span class="n">scores</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>

        <span class="n">attn</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">scores</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>
        <span class="n">attn</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span><span class="p">(</span><span class="n">attn</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">attn</span><span class="p">,</span> <span class="n">V</span><span class="p">)</span>

        <span class="n">out</span> <span class="o">=</span> <span class="n">out</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">).</span><span class="n">contiguous</span><span class="p">().</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">n</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_model</span><span class="p">)</span>
        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">W_o</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
</code></pre></div></div>

<p><strong>实现细节</strong>：<code class="language-plaintext highlighter-rouge">repeat_interleave</code> 在工程上其实可以避免——直接用 einsum 或在 attention 计算里广播更省显存。但面试场景下写 <code class="language-plaintext highlighter-rouge">repeat_interleave</code> 更直观易读。</p>

<h3 id="34-面试常见追问">3.4 面试常见追问</h3>

<ul>
  <li>
    <p><strong>Q：GQA 为什么比 MQA 效果好？</strong>
MQA 所有 Q 头共享一套 K、V，表达能力瓶颈明显；GQA 保留了多组 K、V，能在”显存”和”表达”之间做权衡。</p>
  </li>
  <li>
    <p><strong>Q：KV Cache 具体省多少？</strong>
对 LLaMA-2-70B，MHA KV Cache 每层是 $2 \cdot 64 \cdot d_k$，GQA 是 $2 \cdot 8 \cdot d_k$，直接省 8 倍。</p>
  </li>
</ul>

<hr />

<h2 id="4-transformer-encoder-模块">4. Transformer Encoder 模块</h2>

<h3 id="41-结构图">4.1 结构图</h3>

<p>一个完整的 Transformer Encoder Block 由以下部分组成：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x ──▶ LayerNorm ──▶ MHA ──▶ + ──▶ LayerNorm ──▶ FFN ──▶ + ──▶ out
  │                          ▲                           ▲
  └──────────────────────────┘                           │
                │                                         │
                └─────────────────────────────────────────┘
</code></pre></div></div>

<p>这是 Pre-Norm 结构（GPT、LLaMA 都用这种）。原版 Transformer 用 Post-Norm，训练稳定性差一些，现在主流都切到了 Pre-Norm。</p>

<h3 id="42-公式">4.2 公式</h3>

\[\begin{aligned}
z &amp;= x + \text{MHA}(\text{LN}(x)) \\
y &amp;= z + \text{FFN}(\text{LN}(z))
\end{aligned}\]

<p>其中 FFN 通常是：</p>

\[\text{FFN}(x) = W_2 \cdot \text{GELU}(W_1 x + b_1) + b_2\]

<p>中间维度一般取 $4 \times d_{\text{model}}$。</p>

<h3 id="43-从零手撕实现">4.3 从零手撕实现</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">FeedForward</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">,</span> <span class="n">dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">fc1</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">fc2</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_ff</span><span class="p">,</span> <span class="n">d_model</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">fc2</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">dropout</span><span class="p">(</span><span class="n">F</span><span class="p">.</span><span class="n">gelu</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">fc1</span><span class="p">(</span><span class="n">x</span><span class="p">))))</span>


<span class="k">class</span> <span class="nc">TransformerEncoderBlock</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">num_heads</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">,</span> <span class="n">dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">ln1</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">LayerNorm</span><span class="p">(</span><span class="n">d_model</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">attn</span> <span class="o">=</span> <span class="n">MultiHeadAttention</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">num_heads</span><span class="p">,</span> <span class="n">dropout</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">ln2</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">LayerNorm</span><span class="p">(</span><span class="n">d_model</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">ffn</span> <span class="o">=</span> <span class="n">FeedForward</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">,</span> <span class="n">dropout</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">dropout</span><span class="p">)</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
        <span class="c1"># Pre-Norm + Residual
</span>        <span class="n">h</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">ln1</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="n">x</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">attn</span><span class="p">(</span><span class="n">h</span><span class="p">,</span> <span class="n">h</span><span class="p">,</span> <span class="n">h</span><span class="p">,</span> <span class="n">mask</span><span class="o">=</span><span class="n">mask</span><span class="p">))</span>

        <span class="n">h</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">ln2</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="n">x</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">ffn</span><span class="p">(</span><span class="n">h</span><span class="p">))</span>
        <span class="k">return</span> <span class="n">x</span>
</code></pre></div></div>

<h3 id="44-面试常见追问">4.4 面试常见追问</h3>

<ul>
  <li>
    <p><strong>Q：Pre-Norm 和 Post-Norm 的区别？</strong>
Post-Norm（原版）：$y = \text{LN}(x + \text{Sublayer}(x))$，残差链路被 LN 截断，深层训练不稳定。
Pre-Norm：$y = x + \text{Sublayer}(\text{LN}(x))$，残差是”干净”的恒等映射，梯度更稳，但可能轻微损失表达能力。</p>
  </li>
  <li>
    <p><strong>Q：FFN 中间维度为什么是 4 倍？</strong>
经验值。它提供了模型的主要参数量（约占 2/3），也是存储世界知识的主要场所。</p>
  </li>
  <li>
    <p><strong>Q：GELU 和 ReLU 的区别？</strong>
GELU 是 $x \cdot \Phi(x)$，平滑版 ReLU，负半轴有小幅激活，在语言模型上效果更好。LLaMA 进一步用 SwiGLU。</p>
  </li>
</ul>

<hr />

<h2 id="5-layernorm-vs-batchnorm">5. LayerNorm vs BatchNorm</h2>

<h3 id="51-区别一句话">5.1 区别一句话</h3>

<ul>
  <li><strong>BatchNorm</strong>：对每个<strong>特征维度</strong>，在 <strong>batch 维度</strong>上算均值和方差</li>
  <li><strong>LayerNorm</strong>：对每个<strong>样本</strong>，在<strong>特征维度</strong>上算均值和方差</li>
</ul>

<h3 id="52-公式">5.2 公式</h3>

<p>设输入 $x \in \mathbb{R}^{B \times L \times D}$（Batch × SeqLen × Dim）。</p>

<p><strong>BatchNorm</strong>（对 NLP 几乎不用）：</p>

\[\mu_d = \frac{1}{B \cdot L} \sum_{b, l} x_{b, l, d}, \quad
\sigma_d^2 = \frac{1}{B \cdot L} \sum_{b, l} (x_{b, l, d} - \mu_d)^2\]

\[\hat{x}_{b, l, d} = \frac{x_{b, l, d} - \mu_d}{\sqrt{\sigma_d^2 + \epsilon}}, \quad
y = \gamma \hat{x} + \beta\]

<p><strong>LayerNorm</strong>：</p>

\[\mu_{b, l} = \frac{1}{D} \sum_d x_{b, l, d}, \quad
\sigma_{b, l}^2 = \frac{1}{D} \sum_d (x_{b, l, d} - \mu_{b, l})^2\]

\[\hat{x}_{b, l, d} = \frac{x_{b, l, d} - \mu_{b, l}}{\sqrt{\sigma_{b, l}^2 + \epsilon}}, \quad
y = \gamma \hat{x} + \beta\]

<p>注意 LayerNorm 的 $\gamma, \beta$ 是 $D$ 维向量，不依赖 batch 和 seq。</p>

<h3 id="53-为什么-nlp-用-layernorm-而不是-batchnorm">5.3 为什么 NLP 用 LayerNorm 而不是 BatchNorm？</h3>

<ol>
  <li><strong>变长序列</strong>：NLP 输入有大量 padding，padding 位置参与 BN 统计会污染结果。</li>
  <li><strong>小 batch</strong>：语言模型 batch 常常不大（长序列更吃显存），BN 在小 batch 上统计量不稳。</li>
  <li><strong>训练/推理一致</strong>：BN 推理时用 running mean/var，语言模型分布漂移敏感；LN 训推完全一致。</li>
  <li><strong>每个 token 独立归一化</strong>更贴合语言模型”逐 token 建模”的直觉。</li>
</ol>

<h3 id="54-从零手撕实现">5.4 从零手撕实现</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">LayerNorm</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">dim</span><span class="p">,</span> <span class="n">eps</span><span class="o">=</span><span class="mf">1e-5</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">gamma</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">dim</span><span class="p">))</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">beta</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">dim</span><span class="p">))</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">eps</span> <span class="o">=</span> <span class="n">eps</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="c1"># x: [..., dim]
</span>        <span class="n">mean</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">mean</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
        <span class="n">var</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">var</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">unbiased</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="n">x_hat</span> <span class="o">=</span> <span class="p">(</span><span class="n">x</span> <span class="o">-</span> <span class="n">mean</span><span class="p">)</span> <span class="o">/</span> <span class="n">torch</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">var</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">eps</span><span class="p">)</span>
        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">gamma</span> <span class="o">*</span> <span class="n">x_hat</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">beta</span>


<span class="k">class</span> <span class="nc">BatchNorm1d</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">num_features</span><span class="p">,</span> <span class="n">eps</span><span class="o">=</span><span class="mf">1e-5</span><span class="p">,</span> <span class="n">momentum</span><span class="o">=</span><span class="mf">0.1</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">gamma</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">num_features</span><span class="p">))</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">beta</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">num_features</span><span class="p">))</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">register_buffer</span><span class="p">(</span><span class="s">'running_mean'</span><span class="p">,</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">num_features</span><span class="p">))</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">register_buffer</span><span class="p">(</span><span class="s">'running_var'</span><span class="p">,</span> <span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">num_features</span><span class="p">))</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">eps</span> <span class="o">=</span> <span class="n">eps</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">momentum</span> <span class="o">=</span> <span class="n">momentum</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="c1"># x: [B, D]
</span>        <span class="k">if</span> <span class="bp">self</span><span class="p">.</span><span class="n">training</span><span class="p">:</span>
            <span class="n">mean</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">mean</span><span class="p">(</span><span class="n">dim</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
            <span class="n">var</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">var</span><span class="p">(</span><span class="n">dim</span><span class="o">=</span><span class="mi">0</span><span class="p">,</span> <span class="n">unbiased</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
            <span class="c1"># 更新 running 统计量
</span>            <span class="bp">self</span><span class="p">.</span><span class="n">running_mean</span> <span class="o">=</span> <span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="bp">self</span><span class="p">.</span><span class="n">momentum</span><span class="p">)</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">running_mean</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">momentum</span> <span class="o">*</span> <span class="n">mean</span><span class="p">.</span><span class="n">detach</span><span class="p">()</span>
            <span class="bp">self</span><span class="p">.</span><span class="n">running_var</span> <span class="o">=</span> <span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="bp">self</span><span class="p">.</span><span class="n">momentum</span><span class="p">)</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">running_var</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">momentum</span> <span class="o">*</span> <span class="n">var</span><span class="p">.</span><span class="n">detach</span><span class="p">()</span>
        <span class="k">else</span><span class="p">:</span>
            <span class="n">mean</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">running_mean</span>
            <span class="n">var</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">running_var</span>

        <span class="n">x_hat</span> <span class="o">=</span> <span class="p">(</span><span class="n">x</span> <span class="o">-</span> <span class="n">mean</span><span class="p">)</span> <span class="o">/</span> <span class="n">torch</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">var</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">eps</span><span class="p">)</span>
        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">gamma</span> <span class="o">*</span> <span class="n">x_hat</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">beta</span>
</code></pre></div></div>

<p><strong>注意</strong>：<code class="language-plaintext highlighter-rouge">unbiased=False</code> 表示用 $1/N$ 而不是 $1/(N-1)$，这是神经网络里常用的有偏估计，和 PyTorch 默认实现一致。</p>

<hr />

<h2 id="6-rmsnormllama-时代的归一化标配">6. RMSNorm：LLaMA 时代的归一化标配</h2>

<h3 id="61-动机">6.1 动机</h3>

<p>LayerNorm 做了两件事：<strong>减均值</strong>（中心化） + <strong>除标准差</strong>（缩放）。
但有研究发现，<strong>减均值这一步对性能的贡献非常小</strong>——真正起作用的是”缩放”。</p>

<p>于是 RMSNorm 提出：直接去掉减均值，只做缩放，用 RMS（Root Mean Square）替代标准差：</p>

<ul>
  <li>省掉一次均值计算和减法</li>
  <li>实测 7%~64% 的速度提升</li>
  <li>效果和 LayerNorm 持平甚至更好</li>
</ul>

<p>LLaMA、LLaMA-2/3、Mistral、Qwen 等主流模型全都用 RMSNorm。</p>

<h3 id="62-公式">6.2 公式</h3>

\[\text{RMS}(x) = \sqrt{\frac{1}{D} \sum_{d=1}^{D} x_d^2}\]

\[\text{RMSNorm}(x) = \frac{x}{\text{RMS}(x) + \epsilon} \odot \gamma\]

<p>注意：</p>

<ul>
  <li>没有 $\beta$（平移项）</li>
  <li>分母里是 RMS 而不是 std</li>
  <li>$\gamma$ 是可学习的缩放向量</li>
</ul>

<h3 id="63-和-layernorm-的关系">6.3 和 LayerNorm 的关系</h3>

<p>如果 $x$ 的均值恰好为 0，那么 $\text{RMS}(x) = \text{std}(x)$，RMSNorm 就退化为没有 bias 的 LayerNorm。</p>

<p>换句话说，<strong>RMSNorm = LayerNorm 扔掉均值平移</strong>。</p>

<h3 id="64-从零手撕实现">6.4 从零手撕实现</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">RMSNorm</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">dim</span><span class="p">,</span> <span class="n">eps</span><span class="o">=</span><span class="mf">1e-6</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">gamma</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">dim</span><span class="p">))</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">eps</span> <span class="o">=</span> <span class="n">eps</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="c1"># x: [..., dim]
</span>        <span class="c1"># 注意用 float32 算 rms 避免 fp16 下溢
</span>        <span class="n">rms</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="nb">float</span><span class="p">().</span><span class="nb">pow</span><span class="p">(</span><span class="mi">2</span><span class="p">).</span><span class="n">mean</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">).</span><span class="n">add</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">eps</span><span class="p">).</span><span class="n">rsqrt</span><span class="p">()</span>
        <span class="k">return</span> <span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="nb">float</span><span class="p">()</span> <span class="o">*</span> <span class="n">rms</span><span class="p">).</span><span class="n">to</span><span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="n">dtype</span><span class="p">)</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">gamma</span>
</code></pre></div></div>

<p><strong>工程细节</strong>：</p>
<ul>
  <li><code class="language-plaintext highlighter-rouge">rsqrt()</code> 比 <code class="language-plaintext highlighter-rouge">1 / sqrt()</code> 更快</li>
  <li>混合精度下先 cast 到 fp32 计算，结果再 cast 回原 dtype，可以避免数值不稳定</li>
  <li>eps 通常取 $10^{-6}$ 或 $10^{-5}$</li>
</ul>

<h3 id="65-面试常见追问">6.5 面试常见追问</h3>

<ul>
  <li>
    <p><strong>Q：RMSNorm 为什么比 LayerNorm 快？</strong>
少了一次”求均值 + 做减法”的操作，访存和计算都减半。</p>
  </li>
  <li>
    <p><strong>Q：去掉 $\beta$ 不会损失表达能力吗？</strong>
理论上会，但实验发现对语言模型几乎无影响。可能原因是残差连接本身已经提供了足够的 bias 能力。</p>
  </li>
</ul>

<hr />

<h2 id="7-safe-softmax数值稳定性的必考点">7. Safe Softmax：数值稳定性的必考点</h2>

<h3 id="71-朴素-softmax-的问题">7.1 朴素 Softmax 的问题</h3>

\[\text{softmax}(x_i) = \frac{e^{x_i}}{\sum_j e^{x_j}}\]

<p>当 $x_i$ 很大时（比如 1000），$e^{x_i}$ 直接 <strong>上溢为 inf</strong>；当 $x_i$ 很小时，$e^{x_i}$ 下溢为 0。fp16/bf16 下这个问题尤其严重（fp16 的最大值只有 65504）。</p>

<h3 id="72-safe-softmax-技巧">7.2 Safe Softmax 技巧</h3>

<p>利用 softmax 的<strong>平移不变性</strong>：</p>

\[\text{softmax}(x_i) = \text{softmax}(x_i - c)\]

<p>因为分子分母同时乘 $e^{-c}$ 会约掉。所以我们可以取 $c = \max_j x_j$：</p>

\[\text{softmax}(x_i) = \frac{e^{x_i - \max_j x_j}}{\sum_j e^{x_j - \max_j x_j}}\]

<p>这样：</p>

<ul>
  <li>指数的最大值变成 $e^0 = 1$，永远不会上溢</li>
  <li>分母至少有一项是 1，不会下溢为 0</li>
</ul>

<h3 id="73-从零手撕实现">7.3 从零手撕实现</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">safe_softmax</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">):</span>
    <span class="s">"""
    数值稳定的 softmax
    x: 任意 shape 的 tensor
    dim: 在哪个维度做 softmax
    """</span>
    <span class="c1"># 减去最大值防溢出（注意 detach/不 detach 都不影响梯度，因为是平移不变的）
</span>    <span class="n">x_max</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">dim</span><span class="o">=</span><span class="n">dim</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">).</span><span class="n">values</span>
    <span class="n">x_shifted</span> <span class="o">=</span> <span class="n">x</span> <span class="o">-</span> <span class="n">x_max</span>

    <span class="n">exp_x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">exp</span><span class="p">(</span><span class="n">x_shifted</span><span class="p">)</span>
    <span class="n">sum_exp</span> <span class="o">=</span> <span class="n">exp_x</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">dim</span><span class="o">=</span><span class="n">dim</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">exp_x</span> <span class="o">/</span> <span class="n">sum_exp</span>


<span class="k">def</span> <span class="nf">safe_log_softmax</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">):</span>
    <span class="s">"""
    数值稳定的 log_softmax，后面算交叉熵要用到
    log_softmax(x) = x - max - log(sum(exp(x - max)))
    """</span>
    <span class="n">x_max</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">dim</span><span class="o">=</span><span class="n">dim</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">).</span><span class="n">values</span>
    <span class="n">x_shifted</span> <span class="o">=</span> <span class="n">x</span> <span class="o">-</span> <span class="n">x_max</span>
    <span class="c1"># log-sum-exp
</span>    <span class="n">log_sum_exp</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">log</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">exp</span><span class="p">(</span><span class="n">x_shifted</span><span class="p">).</span><span class="nb">sum</span><span class="p">(</span><span class="n">dim</span><span class="o">=</span><span class="n">dim</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">))</span>
    <span class="k">return</span> <span class="n">x_shifted</span> <span class="o">-</span> <span class="n">log_sum_exp</span>
</code></pre></div></div>

<h3 id="74-面试常见追问">7.4 面试常见追问</h3>

<ul>
  <li>
    <p><strong>Q：减最大值这个操作会影响梯度吗？</strong>
不会。softmax 对平移不变，数学上完全等价，梯度也完全等价。</p>
  </li>
  <li>
    <p><strong>Q：如果所有输入都是 -inf（比如全被 mask）怎么办？</strong>
会出现 NaN（0/0）。实际实现里对 attention 会保证至少有一个非 mask 的位置，或者在 softmax 之后把整行置零。</p>
  </li>
  <li>
    <p><strong>Q：Flash Attention 里的 online softmax 是什么？</strong>
Flash Attention 在 tiling 过程中不能一次看到所有的 logits，它用 <strong>递推公式</strong> 维护当前已见的最大值和分母，逐块更新。这是 safe softmax 的分块在线版本。</p>
  </li>
</ul>

<hr />

<h2 id="8-cross-entropy-loss分类任务的灵魂">8. Cross-Entropy Loss：分类任务的灵魂</h2>

<h3 id="81-公式">8.1 公式</h3>

<p>对于 $C$ 分类问题，模型输出 logits $z \in \mathbb{R}^C$，真实标签 $y \in {0, 1, \dots, C-1}$。</p>

<p>先 softmax 得到概率：</p>

\[p_i = \frac{e^{z_i}}{\sum_j e^{z_j}}\]

<p>交叉熵损失：</p>

\[\mathcal{L} = -\log p_y = -\log \frac{e^{z_y}}{\sum_j e^{z_j}} = -z_y + \log \sum_j e^{z_j}\]

<p>最右边的形式就是我们常说的 <strong>LogSumExp</strong> 形式，非常适合数值稳定实现。</p>

<p><strong>批量版</strong>：</p>

\[\mathcal{L} = -\frac{1}{N} \sum_{n=1}^{N} \log p_{n, y_n}\]

<h3 id="82-为什么要用交叉熵而不是-mse">8.2 为什么要用交叉熵而不是 MSE？</h3>

<ol>
  <li><strong>梯度性质好</strong>：对 logits 求导，$\frac{\partial \mathcal{L}}{\partial z_i} = p_i - \mathbb{1}[i = y]$，形式简洁且不会出现”梯度消失”。</li>
  <li><strong>概率语义契合</strong>：交叉熵衡量两个分布的距离，和 softmax 输出的”概率”天然配对。</li>
  <li><strong>MSE 配 softmax 会梯度饱和</strong>：预测很离谱时梯度反而很小，训练慢。</li>
</ol>

<h3 id="83-从零手撕实现">8.3 从零手撕实现</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">cross_entropy_loss</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">targets</span><span class="p">,</span> <span class="n">reduction</span><span class="o">=</span><span class="s">'mean'</span><span class="p">):</span>
    <span class="s">"""
    logits: [N, C]  未经 softmax 的原始输出
    targets: [N]    每个样本的真实标签 index
    """</span>
    <span class="c1"># 用 log_softmax 的形式，避免 log(0)
</span>    <span class="n">log_probs</span> <span class="o">=</span> <span class="n">safe_log_softmax</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>  <span class="c1"># [N, C]
</span>
    <span class="c1"># gather 出真实标签对应的 log_prob
</span>    <span class="c1"># log_probs.gather(1, targets.unsqueeze(1)).squeeze(1): [N]
</span>    <span class="n">nll</span> <span class="o">=</span> <span class="o">-</span><span class="n">log_probs</span><span class="p">.</span><span class="n">gather</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">targets</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">)).</span><span class="n">squeeze</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>

    <span class="k">if</span> <span class="n">reduction</span> <span class="o">==</span> <span class="s">'mean'</span><span class="p">:</span>
        <span class="k">return</span> <span class="n">nll</span><span class="p">.</span><span class="n">mean</span><span class="p">()</span>
    <span class="k">elif</span> <span class="n">reduction</span> <span class="o">==</span> <span class="s">'sum'</span><span class="p">:</span>
        <span class="k">return</span> <span class="n">nll</span><span class="p">.</span><span class="nb">sum</span><span class="p">()</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="k">return</span> <span class="n">nll</span>
</code></pre></div></div>

<h3 id="84-带-label-smoothing-的版本">8.4 带 Label Smoothing 的版本</h3>

<p>大模型训练里经常用到 Label Smoothing：不把真实标签当成 one-hot，而是留一点点给别的类别，防止模型”过分自信”。</p>

\[\tilde{y}_i = \begin{cases}
1 - \epsilon + \epsilon / C &amp; i = y \\
\epsilon / C &amp; i \neq y
\end{cases}\]

\[\mathcal{L} = -\sum_{i} \tilde{y}_i \log p_i\]

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">cross_entropy_with_label_smoothing</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">targets</span><span class="p">,</span> <span class="n">smoothing</span><span class="o">=</span><span class="mf">0.1</span><span class="p">):</span>
    <span class="n">N</span><span class="p">,</span> <span class="n">C</span> <span class="o">=</span> <span class="n">logits</span><span class="p">.</span><span class="n">shape</span>
    <span class="n">log_probs</span> <span class="o">=</span> <span class="n">safe_log_softmax</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>

    <span class="c1"># 构造平滑后的标签分布
</span>    <span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
        <span class="n">true_dist</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">full_like</span><span class="p">(</span><span class="n">log_probs</span><span class="p">,</span> <span class="n">smoothing</span> <span class="o">/</span> <span class="n">C</span><span class="p">)</span>
        <span class="n">true_dist</span><span class="p">.</span><span class="n">scatter_</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">targets</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="mi">1</span> <span class="o">-</span> <span class="n">smoothing</span> <span class="o">+</span> <span class="n">smoothing</span> <span class="o">/</span> <span class="n">C</span><span class="p">)</span>

    <span class="c1"># -(true_dist * log_probs).sum(dim=-1).mean()
</span>    <span class="k">return</span> <span class="o">-</span><span class="p">(</span><span class="n">true_dist</span> <span class="o">*</span> <span class="n">log_probs</span><span class="p">).</span><span class="nb">sum</span><span class="p">(</span><span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">).</span><span class="n">mean</span><span class="p">()</span>
</code></pre></div></div>

<h3 id="85-面试常见追问">8.5 面试常见追问</h3>

<ul>
  <li>
    <p><strong>Q：PyTorch 的 <code class="language-plaintext highlighter-rouge">nn.CrossEntropyLoss</code> 输入是 logits 还是概率？</strong>
logits！它内部融合了 log_softmax + nll_loss，比手写分开两步更稳定、更快。</p>
  </li>
  <li>
    <p><strong>Q：对 logits 求导的结果推一下？</strong>
$\frac{\partial \mathcal{L}}{\partial z_i} = p_i - \mathbb{1}[i = y]$，这就是为什么反向传播时只需要”预测概率减去 one-hot”。</p>
  </li>
  <li>
    <p><strong>Q：语言模型训练时怎么处理 padding 的 loss？</strong>
用 <code class="language-plaintext highlighter-rouge">ignore_index</code>（PyTorch 原生支持），或者构造 mask 在 reduction 之前把 padding 位置的 loss 置零。</p>
  </li>
</ul>

<hr />

<h2 id="9-一个完整的白板样板">9. 一个完整的”白板样板”</h2>

<p>最后给一个浓缩版的”应急套路”——如果面试官让你 5 分钟手撕一个精简 Transformer Block，就按下面这个最小实现来：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>
<span class="kn">import</span> <span class="nn">math</span>


<span class="k">class</span> <span class="nc">RMSNorm</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">dim</span><span class="p">,</span> <span class="n">eps</span><span class="o">=</span><span class="mf">1e-6</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">weight</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Parameter</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">ones</span><span class="p">(</span><span class="n">dim</span><span class="p">))</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">eps</span> <span class="o">=</span> <span class="n">eps</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="n">rms</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="nb">pow</span><span class="p">(</span><span class="mi">2</span><span class="p">).</span><span class="n">mean</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">).</span><span class="n">add</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">eps</span><span class="p">).</span><span class="n">rsqrt</span><span class="p">()</span>
        <span class="k">return</span> <span class="n">x</span> <span class="o">*</span> <span class="n">rms</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">weight</span>


<span class="k">class</span> <span class="nc">MHA</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">h</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span> <span class="o">=</span> <span class="n">n_heads</span><span class="p">,</span> <span class="n">d_model</span> <span class="o">//</span> <span class="n">n_heads</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">wq</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">wk</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">wv</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">wo</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
        <span class="n">B</span><span class="p">,</span> <span class="n">N</span><span class="p">,</span> <span class="n">D</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">shape</span>
        <span class="n">q</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">wq</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">N</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">h</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">k</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">wk</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">N</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">h</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">v</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">wv</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">view</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">N</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">h</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>

        <span class="n">scores</span> <span class="o">=</span> <span class="n">q</span> <span class="o">@</span> <span class="n">k</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">2</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">)</span> <span class="o">/</span> <span class="n">math</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">d_k</span><span class="p">)</span>
        <span class="k">if</span> <span class="n">mask</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">scores</span> <span class="o">=</span> <span class="n">scores</span><span class="p">.</span><span class="n">masked_fill</span><span class="p">(</span><span class="n">mask</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="s">'-inf'</span><span class="p">))</span>
        <span class="n">attn</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">scores</span><span class="p">,</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="p">(</span><span class="n">attn</span> <span class="o">@</span> <span class="n">v</span><span class="p">).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">).</span><span class="n">reshape</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">N</span><span class="p">,</span> <span class="n">D</span><span class="p">)</span>
        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">wo</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>


<span class="k">class</span> <span class="nc">Block</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">n1</span> <span class="o">=</span> <span class="n">RMSNorm</span><span class="p">(</span><span class="n">d_model</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">attn</span> <span class="o">=</span> <span class="n">MHA</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">n_heads</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">n2</span> <span class="o">=</span> <span class="n">RMSNorm</span><span class="p">(</span><span class="n">d_model</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">ffn</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Sequential</span><span class="p">(</span>
            <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_model</span><span class="p">,</span> <span class="n">d_ff</span><span class="p">),</span> <span class="n">nn</span><span class="p">.</span><span class="n">GELU</span><span class="p">(),</span> <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">d_ff</span><span class="p">,</span> <span class="n">d_model</span><span class="p">)</span>
        <span class="p">)</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">mask</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
        <span class="n">x</span> <span class="o">=</span> <span class="n">x</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">attn</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">n1</span><span class="p">(</span><span class="n">x</span><span class="p">),</span> <span class="n">mask</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="n">x</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">ffn</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">n2</span><span class="p">(</span><span class="n">x</span><span class="p">))</span>
        <span class="k">return</span> <span class="n">x</span>
</code></pre></div></div>

<p>把这套板子背熟，再根据面试官追问填 GQA、KV Cache、RoPE 等变种即可。</p>

<hr />

<h2 id="10-总结">10. 总结</h2>

<p>本文把大模型算法岗最常考的一组手撕题串起来：</p>

<table>
  <thead>
    <tr>
      <th>模块</th>
      <th>一句话总结</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>Self-Attention</td>
      <td>Q·Kᵀ / √d_k 再 softmax 加权 V</td>
    </tr>
    <tr>
      <td>MHA</td>
      <td>切成 h 个头并行算 Attention，拼回再投影</td>
    </tr>
    <tr>
      <td>GQA</td>
      <td>Q 独立 h 个头，KV 只有 g 组头，省 KV Cache</td>
    </tr>
    <tr>
      <td>Transformer Encoder</td>
      <td>Pre-Norm + MHA + Pre-Norm + FFN，两段残差</td>
    </tr>
    <tr>
      <td>LayerNorm</td>
      <td>每个样本在特征维上做归一化</td>
    </tr>
    <tr>
      <td>BatchNorm</td>
      <td>每个特征在 batch 维上做归一化（NLP 不用）</td>
    </tr>
    <tr>
      <td>RMSNorm</td>
      <td>LayerNorm 去掉均值平移，只保留 RMS 缩放</td>
    </tr>
    <tr>
      <td>Safe Softmax</td>
      <td>减去最大值再做 exp，避免上溢</td>
    </tr>
    <tr>
      <td>Cross-Entropy</td>
      <td>$-\log p_y$，实战用 log_softmax + nll_loss</td>
    </tr>
  </tbody>
</table>

<p>建议的学习路径：先把每一节的公式自己手推一遍，再白板默写实现，然后用 <code class="language-plaintext highlighter-rouge">torch.allclose</code> 和 <code class="language-plaintext highlighter-rouge">nn</code> 自带模块对拍一下数值，最后在纸上做时空复杂度分析。真正把这一整套走完之后，这类题目你都能在面试里淡定应付了。</p>

<hr />

<h2 id="参考资料">参考资料</h2>

<ol>
  <li>Vaswani et al., <em>Attention is All You Need</em>, 2017</li>
  <li>Ba et al., <em>Layer Normalization</em>, 2016</li>
  <li>Zhang &amp; Sennrich, <em>Root Mean Square Layer Normalization</em>, 2019</li>
  <li>Ainslie et al., <em>GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints</em>, 2023</li>
  <li>Touvron et al., <em>LLaMA: Open and Efficient Foundation Language Models</em>, 2023</li>
</ol>]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><category term="blog" /><category term="llm" /><category term="transformer" /><category term="attention" /><category term="interview" /><category term="machine learning" /><summary type="html"><![CDATA[大模型算法岗面试中，手撕代码是几乎绕不过去的一环。面试官会盯着你从零实现 Attention、MHA、GQA、LayerNorm、RMSNorm、SafeSoftmax、Cross-Entropy 等模块，既考察你对原理的理解，也考察你是否能在紧张的环境下把数值稳定性、维度对齐、broadcasting 这些细节处理干净。 这篇文章把这些高频手撕题系统梳理一遍：每一节都给出核心原理 → 数学公式 → 从零手写的 PyTorch 实现 → 面试容易追问的点，读完之后这一类题你应该都能在白板上 10 分钟内写出来。]]></summary></entry><entry><title type="html">从 Classifier Guidance 到 Classifier-Free Guidance：一文讲清 Diffusion 里的 CFG</title><link href="https://liyongzhi.xyz/posts/2026/04/diffusion-cfg/" rel="alternate" type="text/html" title="从 Classifier Guidance 到 Classifier-Free Guidance：一文讲清 Diffusion 里的 CFG" /><published>2026-04-20T00:00:00+08:00</published><updated>2026-04-20T00:00:00+08:00</updated><id>https://liyongzhi.xyz/posts/2026/04/blog-post-diffusion-cfg</id><content type="html" xml:base="https://liyongzhi.xyz/posts/2026/04/diffusion-cfg/"><![CDATA[<p>Diffusion 模型发展到今天，<code class="language-plaintext highlighter-rouge">CFG</code> 几乎已经成了文本生成图像系统里的“默认组件”。<br />
但很多人第一次看到它时都会困惑：</p>

<ul>
  <li>为什么一个模型要做“有条件前向”一次、“无条件前向”一次？</li>
  <li>为什么不能直接训练一个很强的条件模型？</li>
  <li><code class="language-plaintext highlighter-rouge">Classifier Guidance</code> 和 <code class="language-plaintext highlighter-rouge">Classifier-Free Guidance</code> 到底是什么关系？</li>
  <li>为什么后来的蒸馏模型又常常说“已经不需要 CFG 了”？</li>
</ul>

<p>这篇文章想做的事情，就是把这条知识脉络从头理顺：<br />
从无条件 Diffusion 的起点出发，讲到 <code class="language-plaintext highlighter-rouge">Classifier Guidance</code>，再讲到今天真正主流的 <code class="language-plaintext highlighter-rouge">Classifier-Free Guidance (CFG)</code>，最后补上 few-step / distillation 路线和近两年的一些延伸工作。</p>

<p><br />
<img align="center" width="1000" src="https://liyongzhi.xyz/images/posts/diffusion-cfg-evolution.svg" alt="Diffusion guidance evolution timeline" />
<br /></p>

<p>上图可以先当作全文导航来读：<br />
最左边是无条件 diffusion 的起点，中间是两代 guidance，最右边是把 guidance 效果蒸进 few-step 模型的后续路线。</p>

<hr />

<h2 id="1-起点无条件-diffusion-学的只是-px">1. 起点：无条件 Diffusion 学的只是 <code class="language-plaintext highlighter-rouge">p(x)</code></h2>

<p>先从最原始的 DDPM 视角看问题。<br />
一个<strong>无条件</strong> diffusion 模型学的是数据分布 <code class="language-plaintext highlighter-rouge">p(x)</code>，也就是：</p>

<blockquote>
  <p>什么样的样本看起来像“真实世界中的自然图像”。</p>
</blockquote>

<p>它并不知道你想生成什么。<br />
所以如果你只给它高斯噪声，它学会的是“把噪声慢慢拉回自然图像流形上”，而不是“把噪声拉成一只猫”或者“拉成一辆红色跑车”。</p>

<p>在 DDPM 的常见参数化里，前向加噪可以写成：</p>

\[x_t = \sqrt{\bar{\alpha}_t}x_0 + \sqrt{1-\bar{\alpha}_t}\epsilon,\quad \epsilon \sim \mathcal{N}(0, I)\]

<p>训练时模型通常预测噪声：</p>

\[\epsilon_\theta(x_t, t)\]

<p>Ho 等人在 DDPM 里说明了这种噪声预测与 denoising score matching 存在紧密联系，因此我们常把 diffusion 模型理解为在不同噪声水平下学习一个 score function 的近似器。[1]</p>

<p>更直观一点说：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">x_t</code> 是某个时刻的噪声图</li>
  <li>模型要回答的问题是：“这张图里哪一部分是噪声？应该往哪个方向去噪？”</li>
  <li>如果模型只学了 <code class="language-plaintext highlighter-rouge">p(x)</code>，它最多只能说“往更像自然图像的方向去”</li>
</ul>

<p>但条件生成想做的是：</p>

\[p(x|c)\]

<p>也就是：</p>

<blockquote>
  <p>在满足条件 <code class="language-plaintext highlighter-rouge">c</code> 的前提下，什么样的图像是合理的？</p>
</blockquote>

<p>这里的 <code class="language-plaintext highlighter-rouge">c</code> 可以是类别标签、文本、另一张图像、语音 embedding，甚至更复杂的多模态条件。</p>

<p>问题于是变成了：</p>

<blockquote>
  <p>怎么把条件信号注入去噪过程？</p>
</blockquote>

<p>这条路上，Diffusion 社区先后给出了两代代表性答案。</p>

<hr />

<h2 id="2-第一代方案classifier-guidance">2. 第一代方案：Classifier Guidance</h2>

<p><code class="language-plaintext highlighter-rouge">Classifier Guidance</code> 的代表工作是 Dhariwal 和 Nichol 在 2021 年的 <em>Diffusion Models Beat GANs on Image Synthesis</em>。[2]</p>

<p>它的核心想法非常优雅：<br />
先保留一个无条件 diffusion 模型来提供“自然图像先验”，然后再额外训练一个分类器，告诉模型“怎样更像某个类别”。</p>

<h3 id="21-贝叶斯公式是整个故事的起点">2.1 贝叶斯公式是整个故事的起点</h3>

<p>对于类别条件 <code class="language-plaintext highlighter-rouge">c</code>，有：</p>

\[p(x_t|c) \propto p(c|x_t)\,p(x_t)\]

<p>两边取对数并对 <code class="language-plaintext highlighter-rouge">x_t</code> 求梯度，可得：</p>

\[\nabla_{x_t}\log p(x_t|c)=\nabla_{x_t}\log p(x_t)+\nabla_{x_t}\log p(c|x_t)\]

<p>这行式子的含义特别重要：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">\nabla_{x_t}\log p(x_t)</code>：让样本更像“自然图像”</li>
  <li><code class="language-plaintext highlighter-rouge">\nabla_{x_t}\log p(c|x_t)</code>：让样本更像“类别 c”</li>
</ul>

<p>于是条件生成的 score 可以拆成：</p>

<blockquote>
  <p>无条件生成能力 + 分类器给的条件方向</p>
</blockquote>

<h3 id="22-它在采样时是怎么工作的">2.2 它在采样时是怎么工作的</h3>

<p>直觉化地写，Classifier Guidance 的采样过程可以理解为：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Step 1: diffusion 模型在当前 x_t 上做一次前向
        得到无条件噪声预测或 score

Step 2: 分类器在同一个 x_t 上判断“这有多像类别 c”
        然后对 x_t 求梯度

Step 3: 用这个梯度去修正 diffusion 模型的去噪方向

Step 4: 按修正后的方向走一步，得到 x_{t-1}
</code></pre></div></div>

<p>在 Dhariwal &amp; Nichol 的论文里，这个修正以“对采样均值加上分类器梯度项”的形式写进算法；如果改写到更常见的 <code class="language-plaintext highlighter-rouge">\epsilon</code>-prediction 记号中，会看到大家熟悉的“沿着分类器梯度方向做引导”。<br />
需要注意的是：<strong>不同采样器、不同参数化下常数项会略有差异</strong>，所以教程里出现的公式看起来可能不完全一样，但核心思想是一致的。[2]</p>

<h3 id="23-这个方法为什么当时很重要">2.3 这个方法为什么当时很重要</h3>

<p>因为它第一次清楚地证明了：</p>

<blockquote>
  <p>diffusion 模型也可以像 GAN 一样，通过引导在“样本保真度”和“分布覆盖度”之间做可控权衡。</p>
</blockquote>

<p>Dhariwal &amp; Nichol 在 ImageNet 上展示了很强的结果：通过 classifier guidance，他们把 conditional diffusion 的质量显著拉高，并把 diffusion 真正推到了“能和当时顶级 GAN 正面竞争”的阶段。[2]</p>

<h3 id="24-但它有一个很重的代价分类器必须在噪声图上工作">2.4 但它有一个很重的代价：分类器必须在噪声图上工作</h3>

<p>这是 Classifier Guidance 最大的工程痛点。</p>

<p>普通分类器只见过干净图像 <code class="language-plaintext highlighter-rouge">x_0</code>，但 diffusion 采样时给它的是各种噪声水平下的 <code class="language-plaintext highlighter-rouge">x_t</code>。<br />
因此你不能直接拿一个普通 ImageNet 分类器来引导 diffusion，而必须训练一个<strong>噪声感知分类器</strong>：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">train_noisy_classifier</span><span class="p">(</span><span class="n">x0</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
    <span class="n">t</span> <span class="o">=</span> <span class="n">sample_timestep</span><span class="p">()</span>
    <span class="n">noise</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn_like</span><span class="p">(</span><span class="n">x0</span><span class="p">)</span>
    <span class="n">x_t</span> <span class="o">=</span> <span class="n">add_noise</span><span class="p">(</span><span class="n">x0</span><span class="p">,</span> <span class="n">noise</span><span class="p">,</span> <span class="n">t</span><span class="p">)</span>
    <span class="n">logits</span> <span class="o">=</span> <span class="n">classifier</span><span class="p">(</span><span class="n">x_t</span><span class="p">,</span> <span class="n">t</span><span class="p">)</span>
    <span class="n">loss</span> <span class="o">=</span> <span class="n">cross_entropy</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">loss</span>
</code></pre></div></div>

<p>也就是说，整个系统变成了：</p>

<ol>
  <li>一个 diffusion 模型</li>
  <li>一个额外分类器</li>
  <li>分类器还要在所有噪声水平上都能稳定工作</li>
</ol>

<p>而且推理时为了拿到 <code class="language-plaintext highlighter-rouge">\nabla_{x_t}\log p(c|x_t)</code>，还需要对分类器做反向传播，这会带来额外显存和速度开销。</p>

<hr />

<h2 id="3-第二代方案classifier-free-guidance">3. 第二代方案：Classifier-Free Guidance</h2>

<p>真正把 diffusion 条件生成推成主流工业范式的，是 Ho 和 Salimans 在 2021 年提出的 <em>Classifier-Free Diffusion Guidance</em>。[3]</p>

<p>这篇工作的标题其实就已经说明了本质：</p>

<blockquote>
  <p>classifier guidance without a classifier</p>
</blockquote>

<p>也就是：</p>

<blockquote>
  <p>我还想要 guidance 的效果，但我不想再训练一个分类器。</p>
</blockquote>

<h3 id="31-关键替换把分类器梯度改写成两个-score-的差">3.1 关键替换：把“分类器梯度”改写成两个 score 的差</h3>

<p>从贝叶斯公式出发：</p>

\[p(c|x_t)=\frac{p(x_t|c)p(c)}{p(x_t)}\]

<p>对数化后对 <code class="language-plaintext highlighter-rouge">x_t</code> 求梯度：</p>

\[\nabla_{x_t}\log p(c|x_t)
=
\nabla_{x_t}\log p(x_t|c)
-\nabla_{x_t}\log p(x_t)\]

<p>这里 <code class="language-plaintext highlighter-rouge">p(c)</code> 对 <code class="language-plaintext highlighter-rouge">x_t</code> 来说是常数，所以梯度为 0。</p>

<p>这一步非常关键，因为它说明：</p>

<blockquote>
  <p>分类器梯度，本质上可以看成“条件 score”和“无条件 score”的差值。</p>
</blockquote>

<p>而 diffusion 模型本来就在学 score 的近似。<br />
所以如果同一个模型既能输出条件版本，又能输出无条件版本，那么分类器的作用就可以被“模型自己的两次前向”替代。</p>

<h3 id="32-从-score-形式到大家熟悉的-cfg-公式">3.2 从 score 形式到大家熟悉的 CFG 公式</h3>

<p>在常见 VP diffusion / <code class="language-plaintext highlighter-rouge">\epsilon</code>-prediction 的记号下，可以把这件事写成：</p>

\[\hat{\epsilon}
=
\epsilon_\theta(x_t,t,\varnothing)
+ w\cdot\left[\epsilon_\theta(x_t,t,c)-\epsilon_\theta(x_t,t,\varnothing)\right]\]

<p>其中：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">\epsilon_\theta(x_t,t,c)</code>：有条件噪声预测</li>
  <li><code class="language-plaintext highlighter-rouge">\epsilon_\theta(x_t,t,\varnothing)</code>：无条件噪声预测</li>
  <li><code class="language-plaintext highlighter-rouge">w</code>：guidance scale</li>
</ul>

<p>有时你也会看到等价写法：</p>

\[\hat{\epsilon}
=
(1+w)\epsilon_\theta(x_t,t,c)-w\epsilon_\theta(x_t,t,\varnothing)\]

<p>两者完全等价，只是展开方式不同。</p>

<h3 id="33-为什么-score-能改写成噪声预测">3.3 为什么 score 能改写成噪声预测</h3>

<p>上面这一步里，很多人最容易卡住的地方是：</p>

<blockquote>
  <p>前面推的是 score，为什么后面突然变成了噪声预测 <code class="language-plaintext highlighter-rouge">\epsilon_\theta</code>？</p>
</blockquote>

<p>关键在于 VP diffusion 的前向分布：</p>

\[q(x_t|x_0)=\mathcal{N}\left(\sqrt{\bar{\alpha}_t}x_0,\,(1-\bar{\alpha}_t)I\right)\]

<p>对 <code class="language-plaintext highlighter-rouge">x_t</code> 求对数梯度：</p>

\[\nabla_{x_t}\log q(x_t|x_0)
=
-\frac{x_t-\sqrt{\bar{\alpha}_t}x_0}{1-\bar{\alpha}_t}
=
-\frac{\epsilon}{\sqrt{1-\bar{\alpha}_t}}\]

<p>因为：</p>

\[x_t-\sqrt{\bar{\alpha}_t}x_0=\sqrt{1-\bar{\alpha}_t}\,\epsilon\]

<p>所以在这个参数化下，score 和噪声只差一个与时间步有关的缩放因子。<br />
这也是为什么 diffusion 文献里常把：</p>

\[\nabla_{x_t}\log p_t(x_t|c)\]

<p>近似写成：</p>

\[-\frac{1}{\sqrt{1-\bar{\alpha}_t}}\epsilon_\theta(x_t,t,c)\]

<p>严格地说，这里对应的是扰动后边缘分布的 score 近似；不同参数化和不同 sampler 下写法会有差别，但对理解 CFG 来说，这个关系已经足够用了。</p>

<h3 id="34-直觉上它到底在干什么">3.4 直觉上它到底在干什么</h3>

<p>把上面的公式拆开看：</p>

\[\epsilon_\theta(x_t,t,c)-\epsilon_\theta(x_t,t,\varnothing)\]

<p>这一项可以理解为：</p>

<blockquote>
  <p>条件 <code class="language-plaintext highlighter-rouge">c</code> 相对于“什么都不说”所带来的额外方向</p>
</blockquote>

<p>所以 CFG 其实是在做一件非常朴素的事：</p>

<ol>
  <li>先得到“自然去噪方向”</li>
  <li>再提取出“条件带来的额外偏移”</li>
  <li>把这部分偏移乘上一个更大的系数</li>
</ol>

<p>这就是为什么很多人把 CFG 理解成“方向放大器”。<br />
它不是凭空发明一个新方向，而是在<strong>有条件</strong>和<strong>无条件</strong>之间做对比，把“条件真正贡献的那一小段方向”放大。</p>

<hr />

<h2 id="4-训练阶段为什么一个模型能同时学会有条件和无条件">4. 训练阶段：为什么一个模型能同时学会有条件和无条件</h2>

<p>CFG 成立的前提是：</p>

<blockquote>
  <p>同一个模型既会做 <code class="language-plaintext highlighter-rouge">\epsilon(x_t,t,c)</code>，也会做 <code class="language-plaintext highlighter-rouge">\epsilon(x_t,t,\varnothing)</code>。</p>
</blockquote>

<p>Ho 和 Salimans 的做法很简单：<strong>训练时随机把条件丢掉</strong>。[3]</p>

<p>伪代码如下：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="n">x0</span><span class="p">,</span> <span class="n">c</span><span class="p">):</span>
    <span class="n">t</span> <span class="o">=</span> <span class="n">sample_timestep</span><span class="p">()</span>
    <span class="n">noise</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">randn_like</span><span class="p">(</span><span class="n">x0</span><span class="p">)</span>
    <span class="n">x_t</span> <span class="o">=</span> <span class="n">add_noise</span><span class="p">(</span><span class="n">x0</span><span class="p">,</span> <span class="n">noise</span><span class="p">,</span> <span class="n">t</span><span class="p">)</span>

    <span class="k">if</span> <span class="n">random</span><span class="p">()</span> <span class="o">&lt;</span> <span class="n">p_drop</span><span class="p">:</span>
        <span class="n">c</span> <span class="o">=</span> <span class="n">null_condition</span>

    <span class="n">eps_pred</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">x_t</span><span class="p">,</span> <span class="n">t</span><span class="p">,</span> <span class="n">c</span><span class="p">)</span>
    <span class="n">loss</span> <span class="o">=</span> <span class="n">mse</span><span class="p">(</span><span class="n">eps_pred</span><span class="p">,</span> <span class="n">noise</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">loss</span>
</code></pre></div></div>

<p>其中 <code class="language-plaintext highlighter-rouge">p_drop</code> 通常设成 10% 到 20% 左右。</p>

<p>这意味着训练过程中模型会见到两类样本：</p>

<ul>
  <li>大多数时候：正常条件训练</li>
  <li>少数时候：条件被替换为空条件</li>
</ul>

<p>于是模型自然学会了两种行为模式：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>输入真实条件 c   -&gt; 输出条件去噪结果
输入空条件 ∅     -&gt; 输出无条件去噪结果
</code></pre></div></div>

<p>所以推理时只要把同一个 <code class="language-plaintext highlighter-rouge">x_t</code> 喂给模型两次：</p>

<ul>
  <li>一次给真实条件</li>
  <li>一次给空条件</li>
</ul>

<p>就能得到 CFG 所需的两个分支。</p>

<hr />

<h2 id="5-classifier-guidance-和-cfg-到底差在哪里">5. Classifier Guidance 和 CFG 到底差在哪里</h2>

<p>把两代方法摆在一起，对比会非常清楚。</p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>Classifier Guidance</th>
      <th>Classifier-Free Guidance</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>需要几个模型</td>
      <td>2 个</td>
      <td>1 个</td>
    </tr>
    <tr>
      <td>是否需要额外分类器</td>
      <td>需要</td>
      <td>不需要</td>
    </tr>
    <tr>
      <td>分类器是否要适配噪声图</td>
      <td>需要</td>
      <td>不需要</td>
    </tr>
    <tr>
      <td>训练方式</td>
      <td>diffusion 无条件训练 + 分类器单独训练</td>
      <td>条件训练 + condition dropout</td>
    </tr>
    <tr>
      <td>推理开销</td>
      <td>diffusion 前向 + 分类器反向</td>
      <td>两次 diffusion 前向</td>
    </tr>
    <tr>
      <td>条件类型</td>
      <td>更适合离散类别</td>
      <td>任意 embedding 条件</td>
    </tr>
    <tr>
      <td>工程复杂度</td>
      <td>高</td>
      <td>低</td>
    </tr>
    <tr>
      <td>当前使用情况</td>
      <td>历史重要，但已非主流</td>
      <td>现代主流</td>
    </tr>
  </tbody>
</table>

<p>如果只看数学推导，这两者像是“同一家族的两个版本”。<br />
但如果从工程实现看，差距非常大。</p>

<p><br />
<img align="center" width="1000" src="https://liyongzhi.xyz/images/posts/diffusion-cfg-compare.svg" alt="Classifier Guidance versus Classifier-Free Guidance" />
<br /></p>

<p>如果你只想先抓住最核心的工程差异，那么看这张图就够了：<br />
<code class="language-plaintext highlighter-rouge">Classifier Guidance</code> 是“diffusion 模型 + 外部分类器”，<code class="language-plaintext highlighter-rouge">CFG</code> 则是“同一个 diffusion 模型跑两次”。</p>

<h3 id="51-为什么-classifier-guidance-会被边缘化">5.1 为什么 Classifier Guidance 会被边缘化</h3>

<p>主要有四个原因。</p>

<h4 id="原因-1分类器负担太重">原因 1：分类器负担太重</h4>

<p>你不只是多训练了一个网络，而是多训练了一个<strong>噪声鲁棒</strong>分类器。<br />
这件事本身就不轻，而且每换一种条件形式都要重来。</p>

<h4 id="原因-2推理要反向传播">原因 2：推理要反向传播</h4>

<p>Classifier Guidance 采样时需要拿分类器对 <code class="language-plaintext highlighter-rouge">x_t</code> 的梯度。<br />
这意味着：</p>

<ul>
  <li>需要保留更多中间激活</li>
  <li>显存压力更大</li>
  <li>速度通常不如“两次前向”的 CFG</li>
</ul>

<h4 id="原因-3条件类型太受限">原因 3：条件类型太受限</h4>

<p>分类器天然适合“猫 / 狗 / 车”这种离散标签。<br />
但现代生成任务的条件往往是：</p>

<ul>
  <li>长文本 prompt</li>
  <li>图像条件</li>
  <li>layout / depth / segmentation</li>
  <li>多模态混合 embedding</li>
</ul>

<p>这时候“训练一个噪声图上的条件分类器”就变得很不自然。</p>

<h4 id="原因-4高噪声阶段梯度不稳定">原因 4：高噪声阶段梯度不稳定</h4>

<p>当 <code class="language-plaintext highlighter-rouge">x_t</code> 很接近纯噪声时，分类器给出的判断很容易不可靠。<br />
这会导致引导方向 noisy、过激，甚至带来类似对抗样本的伪影。</p>

<hr />

<h2 id="6-为什么直接做条件生成还不够偏偏要有无条件分支">6. 为什么“直接做条件生成”还不够，偏偏要有无条件分支</h2>

<p>这往往是大家第一次学 CFG 时最困惑的点。</p>

<p>一个很自然的问题是：</p>

<blockquote>
  <p>如果模型本来就是条件模型，为什么不直接用 <code class="language-plaintext highlighter-rouge">\epsilon_\theta(x_t,t,c)</code> 去采样？<br />
为什么还要再算一次无条件分支？</p>
</blockquote>

<p>答案是：<strong>因为 CFG 不只是“让模型有条件”，而是在推理阶段对条件方向进行再加权。</strong>
这并不是说纯条件模型不能工作，而是说它少了一个可以在推理时显式放大条件信号的控制手柄。</p>

<h3 id="61-纯条件模型只是在做条件均值意义下的去噪">6.1 纯条件模型只是在做“条件均值意义下的去噪”</h3>

<p>如果 prompt 很宽泛，比如“a cat”，那么满足这个条件的图像分布其实很宽：</p>

<ul>
  <li>橘猫</li>
  <li>黑猫</li>
  <li>正脸</li>
  <li>侧脸</li>
  <li>写实</li>
  <li>插画</li>
</ul>

<p>单纯条件模型学到的是这个条件分布下的平均去噪规律。<br />
而 MSE 型训练目标天然倾向于学习“平均意义上最稳妥的预测”。</p>

<p>结果往往是：</p>

<ul>
  <li>条件满足了</li>
  <li>但语义不够“尖锐”</li>
  <li>细节和风格不够坚定</li>
</ul>

<p>CFG 做的事情则更像是在说：</p>

<blockquote>
  <p>我不仅要满足条件，我还要更坚定地朝条件特征前进。</p>
</blockquote>

<h3 id="62-无条件分支提供了一个基线">6.2 无条件分支提供了一个“基线”</h3>

<p>这一点非常重要。</p>

<p>有了无条件预测，你才能问出下面这个问题：</p>

<blockquote>
  <p>相比于“什么都不说”时的去噪方向，这个条件到底额外改变了什么？</p>
</blockquote>

<p>也就是这项：</p>

\[\epsilon_\theta(x_t,t,c)-\epsilon_\theta(x_t,t,\varnothing)\]

<p>这就是条件的“纯增量”。</p>

<p>没有无条件分支，你只有 <code class="language-plaintext highlighter-rouge">\epsilon_\theta(x_t,t,c)</code>，却不知道里面有多少是：</p>

<ul>
  <li>来自数据分布本身的通用图像先验</li>
  <li>来自条件 <code class="language-plaintext highlighter-rouge">c</code> 的额外要求</li>
</ul>

<p>CFG 恰恰把这两部分显式分离开了。</p>

<h3 id="63-从分布角度看cfg-可以理解为后验锐化">6.3 从分布角度看，CFG 可以理解为“后验锐化”</h3>

<p>很多教程会把 CFG 理解成对条件分布的 sharpen。<br />
直觉上它对应于：</p>

\[\tilde{p}(x|c)\propto \frac{p(x|c)^w}{p(x)^{w-1}}\]

<p>也就是在对数空间里放大条件分布相对于无条件分布的优势方向。</p>

<p>这个视角对于理解“为什么更贴 prompt，但多样性下降”非常有帮助：<br />
<code class="language-plaintext highlighter-rouge">w</code> 越大，分布越尖，语义通常越强，但 mode coverage 往往会下降。</p>

<p>这里要加一个重要注记：</p>

<blockquote>
  <p>上面这个“锐化分布”视角是一个非常有用的直觉，但并不是对所有离散采样器、所有有限步采样过程都严格成立的完整结论。</p>
</blockquote>

<p>近年的理论工作开始更仔细地讨论 CFG 的有限步行为，以及它为什么会出现“过饱和、模式坍缩、编辑不可逆”等副作用；例如 CFG++ 就把部分问题解释为 off-manifold 现象。[7]</p>

<hr />

<h2 id="7-cfg-的实际效果为什么它能长期统治文本生成图像">7. CFG 的实际效果，为什么它能长期统治文本生成图像</h2>

<p>CFG 能成为主流，不只是因为“省了一个分类器”，更因为它刚好踩中了现代大模型系统的需求。</p>

<h3 id="71-它天然支持任意条件-embedding">7.1 它天然支持任意条件 embedding</h3>

<p>只要条件能被编码成向量，CFG 就能工作：</p>

<ul>
  <li>class embedding</li>
  <li>text encoder 输出</li>
  <li>image encoder 输出</li>
  <li>audio embedding</li>
  <li>layout / depth / pose control signal</li>
</ul>

<p>这跟文本生成图像时代的需求几乎完美匹配。</p>

<h3 id="72-它给了推理阶段一个可调节旋钮">7.2 它给了推理阶段一个可调节旋钮</h3>

<p><code class="language-plaintext highlighter-rouge">guidance scale</code> 是一个极其实用的控制参数。</p>

<ul>
  <li>小一些：更多样，但可能没那么贴 prompt</li>
  <li>大一些：更贴 prompt，但更容易失真、过饱和、重复</li>
</ul>

<p>这让同一个基础模型能覆盖很多场景，而不必为每种“对齐强度”重新训练一份。</p>

<h3 id="73-它非常契合-latent-diffusion--stable-diffusion-这类架构">7.3 它非常契合 latent diffusion / Stable Diffusion 这类架构</h3>

<p>现代文本生成图像系统通常采用：</p>

<ol>
  <li>文本编码器把 prompt 编成 embedding</li>
  <li>latent diffusion / UNet 在潜空间做去噪</li>
  <li>采样时同时跑 conditional 和 unconditional 分支</li>
  <li>用 CFG 线性组合</li>
</ol>

<p>这套接口很简单，也很模块化。<br />
所以从 Stable Diffusion 到许多后来的文本生成图像系统，CFG 都成为了默认推理机制。</p>

<hr />

<h2 id="8-cfg-的副作用为什么-guidance-scale-不能无限加大">8. CFG 的副作用：为什么 guidance scale 不能无限加大</h2>

<p>CFG 不是越大越好。</p>

<p>如果把 <code class="language-plaintext highlighter-rouge">w</code> 开得很高，常见问题包括：</p>

<ul>
  <li>图像过饱和</li>
  <li>纹理僵硬</li>
  <li>构图重复</li>
  <li>多样性下降</li>
  <li>细节出现“被强行往 prompt 上扯”的伪影</li>
</ul>

<p>这背后的直觉并不难理解：</p>

<blockquote>
  <p>你在不断放大“条件增量方向”，但这个方向本来只是一个局部修正。<br />
放大过头，就会从“更对齐”变成“过度纠偏”。</p>
</blockquote>

<p>很多用户在 Stable Diffusion 里都有很直观的经验：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">scale</code> 太小，图像“没听话”</li>
  <li><code class="language-plaintext highlighter-rouge">scale</code> 太大，图像“过于用力”</li>
</ul>

<p>从理论和经验上看，CFG 一直在做一件 trade-off：</p>

<blockquote>
  <p>对齐更强，通常意味着多样性更弱。</p>
</blockquote>

<p>这也是后来很多改进工作的出发点。</p>

<hr />

<h2 id="9-后续发展为什么蒸馏模型常常不再需要-cfg">9. 后续发展：为什么蒸馏模型常常“不再需要 CFG”</h2>

<p>理解了 CFG 之后，再看 few-step 模型就顺了。</p>

<p>很多人第一次看到 LCM、SDXL Turbo 这类模型时会觉得奇怪：</p>

<blockquote>
  <p>为什么原始模型要几十步、还要 CFG；<br />
蒸馏后的模型却几步就行，甚至不再依赖传统 CFG？</p>
</blockquote>

<p>答案是：</p>

<blockquote>
  <p>因为 CFG 的效果可以在训练时被“蒸”进学生模型里。</p>
</blockquote>

<h3 id="91-progressive-distillation先解决步数太多">9.1 Progressive Distillation：先解决“步数太多”</h3>

<p>Salimans 和 Ho 在 2022 年提出 <em>Progressive Distillation for Fast Sampling of Diffusion Models</em>，核心思想是把一个多步 deterministic sampler 逐轮蒸成更少步数的模型，每一轮把步数减半。[4]</p>

<p>它解决的是：</p>

<blockquote>
  <p>diffusion 太慢，如何把 8192 步、1024 步慢慢蒸成 4 步？</p>
</blockquote>

<p>这一步不一定直接针对 CFG，但为后面 few-step 生成打了基础。</p>

<h3 id="92-consistency-models直接学习从噪声到数据的快速映射">9.2 Consistency Models：直接学习从噪声到数据的快速映射</h3>

<p>Song 等人在 2023 年提出 <em>Consistency Models</em>，把目标推进到“一步或少步生成”。[5]</p>

<p>它的关键思想是让模型学会不同时间点之间的一致映射，从而绕开传统 diffusion 逐步积分的高成本。</p>

<h3 id="93-latent-consistency-models把-cfg-蒸进-latent-diffusion-体系">9.3 Latent Consistency Models：把 CFG 蒸进 latent diffusion 体系</h3>

<p>真正和现代文本生成图像工作流贴得很近的是 2023 年的 <em>Latent Consistency Models (LCM)</em>。[6]</p>

<p>LCM 很关键的一句话是：</p>

<blockquote>
  <p>它是从<strong>预训练的 classifier-free guided diffusion models</strong> 高效蒸馏出来的。</p>
</blockquote>

<p>换句话说，teacher 本身就带着 CFG 的行为。<br />
学生模型学习的是：</p>

<blockquote>
  <p>teacher 做完 guidance 之后的结果</p>
</blockquote>

<p>于是推理阶段就不必再显式执行：</p>

<ol>
  <li>一次 conditional 前向</li>
  <li>一次 unconditional 前向</li>
  <li>线性组合</li>
</ol>

<p>学生模型已经把“有 guidance 的好处”折进自己参数里了。</p>

<h3 id="94-adversarial-diffusion-distillation把-few-step-做到更激进">9.4 Adversarial Diffusion Distillation：把 few-step 做到更激进</h3>

<p>2023 年的 <em>Adversarial Diffusion Distillation (ADD)</em> 更进一步，把 few-step / one-step 的质量继续往上推。[8]</p>

<p>它利用预训练 diffusion 模型作为 teacher signal，再加上 adversarial loss，目标是在极少步数下依然维持高质量图像。</p>

<p>所以如果把这一整条线串起来，你会得到一个很清晰的演化逻辑：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>先有多步 diffusion
-&gt; 再有 CFG，让条件生成更强
-&gt; 再有蒸馏，把“多步 + CFG 的能力”压缩进少步学生模型
-&gt; 最后出现 few-step / one-step 的实用系统
</code></pre></div></div>

<p>这也是为什么今天很多快速模型虽然表面上“不再跑传统 CFG”，但它们背后的 teacher 往往仍然深受 CFG 体系影响。</p>

<hr />

<h2 id="10-最近两年的一些延伸大家在改进-cfg-的什么">10. 最近两年的一些延伸：大家在改进 CFG 的什么</h2>

<p>如果说 2021 到 2023 年的主线是“把 CFG 变成标准件”，那么 2024 到 2025 年的很多工作则是在回答：</p>

<blockquote>
  <p>CFG 很有用，但它到底哪里还不够好？</p>
</blockquote>

<h3 id="101-self-attention-guidance不用额外条件也能做训练自由的引导">10.1 Self-Attention Guidance：不用额外条件，也能做训练自由的引导</h3>

<p>Hong 等人在 <em>Self-Attention Guidance</em> 中提出，除了 classifier guidance 和 CFG 之外，还可以利用模型内部 self-attention 信息来做 guidance。[9]</p>

<p>这个方向的重要意义在于：</p>

<ul>
  <li>guidance 不一定非得来自外部分类器</li>
  <li>guidance 也不一定非得来自条件 dropout 训练</li>
  <li>可以从模型内部结构本身提取“纠偏信号”</li>
</ul>

<h3 id="102-pag把-guidance-扩展到无条件与下游任务场景">10.2 PAG：把 guidance 扩展到无条件与下游任务场景</h3>

<p>2024 年的 <em>Perturbed-Attention Guidance (PAG)</em> 则进一步展示：<br />
即使在 unconditional generation 或某些 CFG 不方便使用的任务里，也可以通过扰动 attention 构造 guidance 信号。[10]</p>

<p>这说明一个更大的趋势：</p>

<blockquote>
  <p>“guidance” 已经不再只是一条固定公式，而是在演化成一个更宽的推理控制框架。</p>
</blockquote>

<h3 id="103-cfg讨论-vanilla-cfg-的-off-manifold-问题">10.3 CFG++：讨论 vanilla CFG 的 off-manifold 问题</h3>

<p>2025 年的 <em>CFG++</em> 指出，传统 CFG 的一些副作用并不一定是 diffusion 本身的问题，而可能和 CFG 把采样轨迹推离数据流形有关。[7]</p>

<p>这类工作之所以值得关注，是因为它们开始从“经验调 scale”走向：</p>

<ul>
  <li>更系统地理解 CFG 为什么有效</li>
  <li>更具体地解释 CFG 为什么会失真</li>
  <li>更有针对性地修复它的缺点</li>
</ul>

<hr />

<h2 id="11-一张总图把整条知识脉络串起来">11. 一张总图，把整条知识脉络串起来</h2>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DDPM / score-based diffusion
  先学会从噪声中恢复“自然样本”
  核心对象是 p(x) 或其 score

        |
        v

Classifier Guidance (2021)
  用无条件 diffusion 提供 p(x)
  再训练噪声分类器提供 ∇ log p(c|x_t)
  条件 score = 无条件 score + 分类器梯度

        |
        v

Classifier-Free Guidance (2021/2022)
  不再训练分类器
  通过条件 dropout 让同一模型同时学会：
    ε(x_t, t, c)
    ε(x_t, t, ∅)
  采样时做：
    ε_cfg = ε_∅ + w (ε_c - ε_∅)

        |
        v

现代文本生成图像系统
  CFG 成为默认推理机制
  通过 guidance scale 控制 prompt 对齐和多样性

        |
        v

Few-step / Distillation 路线
  Progressive Distillation
  Consistency Models
  LCM
  ADD

  把“多步 + CFG”的效果蒸进学生模型

        |
        v

近期改进
  SAG / PAG / CFG++
  从训练自由 guidance、attention guidance、
  以及理论修正等角度继续优化 CFG
</code></pre></div></div>

<hr />

<h2 id="12-一个最常见的误区">12. 一个最常见的误区</h2>

<p>很多教程会把 CFG 说成：</p>

<blockquote>
  <p>“就是条件减去无条件，再乘一个 scale。”</p>
</blockquote>

<p>这当然没错，但如果只停在这一步，会漏掉最关键的理解：</p>

<blockquote>
  <p><code class="language-plaintext highlighter-rouge">ε_cond - ε_uncond</code> 不是一个拍脑袋的 engineering trick，<br />
它来自贝叶斯公式下“分类器梯度 = 条件 score - 无条件 score”的推导。</p>
</blockquote>

<p>也就是说，CFG 不是“经验上好用的 hack”，而是一个有明确 probabilistic 来历、又极度工程友好的近似方案。</p>

<p>它真正厉害的地方在于：</p>

<ul>
  <li>数学上和 classifier guidance 是一脉相承的</li>
  <li>工程上却省掉了最麻烦的那个分类器</li>
  <li>还顺手把条件接口扩展成了任意 embedding</li>
</ul>

<p>这就是为什么它几乎成为现代 diffusion 条件生成的默认答案。</p>

<hr />

<h2 id="13-总结">13. 总结</h2>

<p>如果只用一句话总结整条脉络，我会写成：</p>

<blockquote>
  <p><code class="language-plaintext highlighter-rouge">Classifier Guidance</code> 证明了 diffusion 可以被“引导”；<br />
<code class="language-plaintext highlighter-rouge">Classifier-Free Guidance</code> 则把这种引导从一个昂贵、受限的两模型系统，变成了一个几乎所有条件 diffusion 都能直接使用的标准模块。</p>
</blockquote>

<p>更具体地说：</p>

<ol>
  <li>无条件 diffusion 只学 <code class="language-plaintext highlighter-rouge">p(x)</code>，不知道用户要什么。</li>
  <li>Classifier Guidance 用额外分类器给出 <code class="language-plaintext highlighter-rouge">\nabla \log p(c|x_t)</code>，第一次把条件引导明确做进 diffusion 采样。</li>
  <li>CFG 发现这个分类器梯度可以由“条件 score - 无条件 score”替代，于是只靠一个模型、两次前向就能完成引导。</li>
  <li>CFG 因为简单、通用、兼容文本条件，最终成为文本生成图像时代的主流。</li>
  <li>后续的蒸馏与 consistency 路线，又把“多步 + CFG”的能力进一步压缩进 few-step 模型。</li>
</ol>

<p>所以从历史上看，CFG 不是 diffusion 里的一个小技巧，它几乎就是现代条件 diffusion 能真正大规模落地的关键转折点之一。</p>

<hr />

<h2 id="参考资料">参考资料</h2>

<ol>
  <li>
    <p>Ho, Jain, Abbeel. <em>Denoising Diffusion Probabilistic Models</em>. NeurIPS 2020.<br />
<a href="https://proceedings.neurips.cc/paper/2020/file/4c5bcfec8584af0d967f1ab10179ca4b-Paper.pdf">Paper</a></p>
  </li>
  <li>
    <p>Dhariwal, Nichol. <em>Diffusion Models Beat GANs on Image Synthesis</em>. NeurIPS 2021.<br />
<a href="https://proceedings.neurips.cc/paper/2021/file/49ad23d1ec9fa4bd8d77d02681df5cfa-Paper.pdf">Paper</a></p>
  </li>
  <li>
    <p>Ho, Salimans. <em>Classifier-Free Diffusion Guidance</em>. NeurIPS 2021 Workshop / OpenReview.<br />
<a href="https://openreview.net/forum?id=qw8AKxfYbI">OpenReview</a></p>
  </li>
  <li>
    <p>Salimans, Ho. <em>Progressive Distillation for Fast Sampling of Diffusion Models</em>. ICLR 2022.<br />
<a href="https://arxiv.org/abs/2202.00512">arXiv</a></p>
  </li>
  <li>
    <p>Song, Dhariwal, Chen, Sutskever. <em>Consistency Models</em>. ICML 2023.<br />
<a href="https://proceedings.mlr.press/v202/song23a.html">PMLR</a></p>
  </li>
  <li>
    <p>Luo, Tan, Huang, Li, Zhao. <em>Latent Consistency Models: Synthesizing High-Resolution Images with Few-Step Inference</em>. 2023.<br />
<a href="https://arxiv.org/abs/2310.04378">arXiv</a></p>
  </li>
  <li>
    <p>Chung, Kim, Park, Nam, Ye. <em>CFG++: Manifold-constrained Classifier Free Guidance for Diffusion Models</em>. ICLR 2025.<br />
<a href="https://openreview.net/forum?id=E77uvbOTtp">OpenReview</a></p>
  </li>
  <li>
    <p>Sauer, Lorenz, Blattmann, Rombach. <em>Adversarial Diffusion Distillation</em>. 2023.<br />
<a href="https://arxiv.org/abs/2311.17042">arXiv</a></p>
  </li>
  <li>
    <p>Hong, Lee, Jang, Kim. <em>Improving Sample Quality of Diffusion Models Using Self-Attention Guidance</em>. ICCV 2023.<br />
<a href="https://openaccess.thecvf.com/content/ICCV2023/html/Hong_Improving_Sample_Quality_of_Diffusion_Models_Using_Self-Attention_Guidance_ICCV_2023_paper.html">Paper</a></p>
  </li>
  <li>
    <p>Ahn et al. <em>Self-Rectifying Diffusion Sampling with Perturbed-Attention Guidance</em>. ECCV 2024 / arXiv.<br />
<a href="https://arxiv.org/abs/2403.17377">arXiv</a></p>
  </li>
</ol>]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><category term="blog" /><category term="diffusion" /><category term="generative model" /><category term="machine learning" /><category term="computer vision" /><summary type="html"><![CDATA[Diffusion 模型发展到今天，CFG 几乎已经成了文本生成图像系统里的“默认组件”。 但很多人第一次看到它时都会困惑：]]></summary></entry><entry><title type="html">从 DDPO 到 Flow-GRPO：一文看懂 Diffusion 模型的强化学习过程与发展脉络</title><link href="https://liyongzhi.xyz/posts/2026/04/diffusion-rl/" rel="alternate" type="text/html" title="从 DDPO 到 Flow-GRPO：一文看懂 Diffusion 模型的强化学习过程与发展脉络" /><published>2026-04-20T00:00:00+08:00</published><updated>2026-04-20T00:00:00+08:00</updated><id>https://liyongzhi.xyz/posts/2026/04/blog-post-diffusion-rl</id><content type="html" xml:base="https://liyongzhi.xyz/posts/2026/04/diffusion-rl/"><![CDATA[<p>Diffusion 模型最初是按“去噪 MSE / 似然近似”来训练的，但真正上线时，我们更关心的往往不是似然，而是：</p>

<ul>
  <li>人类是否更喜欢这张图</li>
  <li>图像和 prompt 是否更对齐</li>
  <li>视频动作是否更连贯</li>
  <li>输出是否更安全</li>
  <li>结果是否更符合物理或任务约束</li>
</ul>

<p>这类目标通常没有一个漂亮、统一、可微、稳定的监督损失。<br />
于是 2023 年开始，一批工作把 Diffusion 的采样过程重新解释成<strong>多步决策</strong>，再用 RL、偏好优化、reward backprop 等方法对它做后训练。</p>

<p>这篇文章想讲清两件事：</p>

<ol>
  <li>Diffusion 模型为什么能被看成一个 RL 问题，以及一次 RL fine-tuning 到底在做什么。</li>
  <li>这条线是如何从 <code class="language-plaintext highlighter-rouge">DDPO / DPOK</code>，发展到 <code class="language-plaintext highlighter-rouge">Diffusion-DPO / D3PO</code>、<code class="language-plaintext highlighter-rouge">DRaFT / AlignProp</code>、视频对齐，再走到 <code class="language-plaintext highlighter-rouge">Flow-GRPO</code> 和 2026 年更高效的方法。</li>
</ol>

<p><br />
<img align="center" width="1000" src="https://liyongzhi.xyz/images/posts/diffusion-rl-evolution.svg" alt="Diffusion RL evolution timeline" />
<br /></p>

<p>上图可以先当全文导航：<br />
2023 年是“把去噪过程变成决策过程”的起点；2023 年下半年到 2024 年，方法开始沿着不同反馈类型分叉；2025 年进入 Flow Matching 时代；到 2026 年，研究重点明显转向了<strong>效率、credit assignment 和对 reverse process likelihood 的替代</strong>。</p>

<hr />

<h2 id="1-为什么要用-rl-优化-diffusion">1. 为什么要用 RL 优化 Diffusion</h2>

<p>标准 Diffusion 训练，学到的是一个“如何把噪声慢慢拉回数据分布”的模型。<br />
在 DDPM 记号下，它通常优化的是噪声预测误差：</p>

\[\mathcal{L}_{\text{diffusion}} = \mathbb{E}\left[\|\epsilon - \epsilon_\theta(x_t, t, c)\|^2\right]\]

<p>这个目标和“生成结果更符合人类偏好”之间，并不是一回事。</p>

<p>更准确地说，预训练阶段优化的是：</p>

<blockquote>
  <p>生成样本要像训练分布中的样本。</p>
</blockquote>

<p>而后训练阶段往往优化的是：</p>

<blockquote>
  <p>在不要严重偏离预训练分布的前提下，让样本在某个外部指标上拿更高分。</p>
</blockquote>

<p>外部指标可以是：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">ImageReward</code>、<code class="language-plaintext highlighter-rouge">HPSv2</code> 这类偏好或美学 reward model</li>
  <li><code class="language-plaintext highlighter-rouge">CLIP</code>、VLM 或 OCR-based 的对齐指标</li>
  <li>视频里的时序一致性、运动平滑性</li>
  <li>分子、蛋白、材料里的任务分数</li>
  <li>人类偏好对本身</li>
</ul>

<p>RL 的价值就在这里：<br />
你不再要求目标必须长得像一个监督学习 loss，而是只要求它能为最终样本给出一个分数，或者至少给出偏好关系。</p>

<hr />

<h2 id="2-关键重写去噪过程其实是一个-mdp">2. 关键重写：去噪过程其实是一个 MDP</h2>

<p>Diffusion RL 的真正起点，不是某个具体算法，而是这个重写：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>state      s_t = (x_t, t, c)
action     a_t = sample or predict the next denoising move
transition x_t -&gt; x_{t-1}
reward     r(x_0, c) or preference over final samples
policy     the diffusion / flow model itself
</code></pre></div></div>

<p>如果用 DDPM 风格的随机反向过程来看，模型每一步都定义了一个条件高斯分布：</p>

\[p_\theta(x_{t-1}\mid x_t, c)=\mathcal{N}(\mu_\theta(x_t,t,c), \sigma_t^2 I)\]

<p>这时“动作”可以理解为：<br />
<strong>在状态 <code class="language-plaintext highlighter-rouge">x_t</code> 下，策略选择了一个从高斯反向转移里采样出来的 <code class="language-plaintext highlighter-rouge">x_{t-1}</code>。</strong></p>

<p>于是整条采样轨迹就像：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x_T -&gt; x_{T-1} -&gt; ... -&gt; x_1 -&gt; x_0
</code></pre></div></div>

<p>最后只在 <code class="language-plaintext highlighter-rouge">x_0</code> 上拿到一个终止 reward。<br />
这正是 RL 里最经典、也最麻烦的一类问题：</p>

<ul>
  <li>奖励是稀疏的</li>
  <li>credit assignment 跨越很多步</li>
  <li>你还不希望模型把预训练学到的图像先验彻底破坏掉</li>
</ul>

<h3 id="21-compute_log_prob-到底在算什么">2.1 <code class="language-plaintext highlighter-rouge">compute_log_prob</code> 到底在算什么</h3>

<p>很多人第一次看 DDPO 代码时，最困惑的是 <code class="language-plaintext highlighter-rouge">compute_log_prob</code>。<br />
它算的其实非常朴素：</p>

<blockquote>
  <p>在当前策略下，模型把 <code class="language-plaintext highlighter-rouge">x_t</code> 变成这一步实际采样到的 <code class="language-plaintext highlighter-rouge">x_{t-1}</code>，这件事的对数概率有多大。</p>
</blockquote>

<p>因为反向过程是高斯分布，所以：</p>

\[\log p_\theta(x_{t-1}\mid x_t,c)
\propto
-\frac{1}{2\sigma_t^2}\|x_{t-1}-\mu_\theta(x_t,t,c)\|^2\]

<p>这件事之所以重要，是因为 policy gradient 需要的正是：</p>

\[\nabla_\theta \log \pi_\theta(a_t\mid s_t)\]

<p>在 Diffusion 里，它就变成了每一步反向转移的 log-prob gradient。</p>

<h3 id="22-为什么-ddpm-比-ddim-更容易做-policy-gradient">2.2 为什么 DDPM 比 DDIM 更容易做 policy gradient</h3>

<p>如果你用的是 DDPM 风格采样，每一步天然有随机性，高斯转移和 log-prob 都是良定义的。<br />
但 DDIM 和很多 Flow Matching 采样器本质上更接近确定性 ODE，这会带来两个麻烦：</p>

<ol>
  <li>exploration 不够自然</li>
  <li>likelihood / log-prob 不容易直接写出来</li>
</ol>

<p>这也是为什么后面 <code class="language-plaintext highlighter-rouge">Flow-GRPO</code> 需要专门做 <strong>ODE-to-SDE conversion</strong>，而 2026 年一些方法开始尝试<strong>不再直接依赖 reverse-process likelihood</strong>。[11][13]</p>

<hr />

<h2 id="3-一次-diffusion-rl-训练迭代到底发生了什么">3. 一次 Diffusion RL 训练迭代到底发生了什么</h2>

<p>先不急着分论文，先看一次最典型的 <code class="language-plaintext highlighter-rouge">DDPO / PPO</code> 风格训练循环。<br />
把细节都压缩掉，它本质上只有五步：</p>

<h3 id="31-第一步用旧策略生成完整轨迹">3.1 第一步：用旧策略生成完整轨迹</h3>

<p>给一批 prompt，用当前模型 <code class="language-plaintext highlighter-rouge">\theta_{\text{old}}</code> 从噪声开始生成图像，并记录整条去噪轨迹：</p>

<ul>
  <li>每个时间步的 <code class="language-plaintext highlighter-rouge">x_t</code></li>
  <li>实际采样得到的 <code class="language-plaintext highlighter-rouge">x_{t-1}</code></li>
  <li>旧策略下对应的 <code class="language-plaintext highlighter-rouge">old_log_prob</code></li>
</ul>

<h3 id="32-第二步对最终样本打分">3.2 第二步：对最终样本打分</h3>

<p>只在最后生成出的 <code class="language-plaintext highlighter-rouge">x_0</code> 上调用 reward：</p>

\[r = r(x_0, c)\]

<p>这个 reward 可以是：</p>

<ul>
  <li>美学分数</li>
  <li>prompt 对齐分数</li>
  <li>安全分数</li>
  <li>视频时序 reward</li>
  <li>人类偏好数据训练出来的 reward model</li>
</ul>

<h3 id="33-第三步把-reward-变成-advantage">3.3 第三步：把 reward 变成 advantage</h3>

<p>最简单时，整条轨迹共享同一个终止 reward；更稳一些时会做 batch normalization、baseline 或组内归一化：</p>

\[A = \frac{r - \mathrm{mean}(r)}{\mathrm{std}(r) + \epsilon}\]

<h3 id="34-第四步用新参数重算每一步-log-prob">3.4 第四步：用新参数重算每一步 log prob</h3>

<p>对同一条轨迹，用当前正在更新的参数重新计算：</p>

\[\rho_t = \exp\left(\log p_\theta(x_{t-1}\mid x_t,c) - \log p_{\theta_{\text{old}}}(x_{t-1}\mid x_t,c)\right)\]

<p>然后套进 PPO clip 或带 KL 的目标里。</p>

<h3 id="35-第五步让高-reward-轨迹更可能低-reward-轨迹更不可能">3.5 第五步：让高 reward 轨迹更可能、低 reward 轨迹更不可能</h3>

<p>如果一张图最终分数高，就提升这条去噪轨迹上各步动作的概率；<br />
如果一张图分数低，就降低这些动作的概率。</p>

<p>直觉上，Diffusion RL 学到的不是某个“神秘奖励魔法”，而是：</p>

<blockquote>
  <p>哪些去噪路径更容易通向人类真正想要的结果。</p>
</blockquote>

<hr />

<h2 id="4-发展脉络这条线是怎么长出来的">4. 发展脉络：这条线是怎么长出来的</h2>

<p>到这里已经能理解“过程”了，接下来再看历史就会顺很多。<br />
我更推荐按<strong>反馈类型</strong>和<strong>建模约束</strong>来看，而不是只按年份背论文名。</p>

<h3 id="41-2023从-reward-model-到-online-rl">4.1 2023：从 reward model 到 online RL</h3>

<p>2023 年最重要的转折，是社区开始承认：</p>

<blockquote>
  <p>Diffusion 的预训练目标和下游目标不一致，所以需要单独的 post-training。</p>
</blockquote>

<p>这一年最早的标志性工作之一是 <code class="language-plaintext highlighter-rouge">ImageReward</code>。它不仅提出了一个通用文本生成图像 reward model，还给出了 <code class="language-plaintext highlighter-rouge">ReFL</code>，把 reward feedback 直接用于模型调优。[1]</p>

<p>紧接着，<code class="language-plaintext highlighter-rouge">DDPO</code> 在 2023 年 5 月把 denoising 明确建模成多步决策过程，并系统引入 policy gradient / PPO 风格更新。[2]</p>

<p>几天后，<code class="language-plaintext highlighter-rouge">DPOK</code> 进一步强调了 <strong>KL regularization</strong> 的重要性：<br />
你不只是要优化 reward，还要约束模型不要偏离预训练分布太远，否则很快就会 reward hacking、图像质量塌陷。[3]</p>

<p>这一阶段的核心思想可以压缩成一句话：</p>

<blockquote>
  <p>先承认 reward 存在，再把 Diffusion 当策略来优化。</p>
</blockquote>

<h3 id="42-2023-下半年到-2024按反馈类型开始分叉">4.2 2023 下半年到 2024：按反馈类型开始分叉</h3>

<p>当“Diffusion 可以做 RL 后训练”这个大门打开后，下一步的问题自然变成：</p>

<blockquote>
  <p>你手里到底有什么反馈信号？</p>
</blockquote>

<p>如果答案不同，方法也会不同。</p>

<h4 id="路线-a你有黑盒-scalar-reward">路线 A：你有黑盒 scalar reward</h4>

<p>那就最适合 <code class="language-plaintext highlighter-rouge">DDPO / DPOK</code> 这种 policy gradient 路线：</p>

<ul>
  <li>reward 可黑盒</li>
  <li>不要求可微</li>
  <li>但方差高，采样贵</li>
</ul>

<h4 id="路线-b你有偏好对没有-reward-model">路线 B：你有偏好对，没有 reward model</h4>

<p>那就更接近 <code class="language-plaintext highlighter-rouge">DPO</code> 家族。</p>

<p><code class="language-plaintext highlighter-rouge">Diffusion-DPO</code> 在 2023 年 11 月把 LLM 里的 DPO 思路迁到 text-to-image diffusion：不先训 reward model，而是直接用人类偏好对优化模型相对偏好概率。[6]</p>

<p>同月提交、后被 CVPR 2024 接收的 <code class="language-plaintext highlighter-rouge">D3PO</code>（<em>Using Human Feedback to Fine-tune Diffusion Models without Any Reward Model</em>）则进一步把“无 reward model 的直接偏好优化”扩展到多步 denoising MDP 视角。[7]</p>

<p>这一路的核心不是“直接做 RL”，而是：</p>

<blockquote>
  <p>如果你已经拿到了 winner / loser 对，就没必要再绕一圈学一个 reward model。</p>
</blockquote>

<h4 id="路线-c你的-reward-本身可微">路线 C：你的 reward 本身可微</h4>

<p>那就完全没必要忍受 REINFORCE / PPO 的高方差。</p>

<p><code class="language-plaintext highlighter-rouge">DRaFT</code> 在 2023 年 9 月提出，直接把 reward 梯度穿过采样过程反传回来，并进一步给出：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">DRaFT-K</code>：只反传最后 <code class="language-plaintext highlighter-rouge">K</code> 步</li>
  <li><code class="language-plaintext highlighter-rouge">DRaFT-LV</code>：在 <code class="language-plaintext highlighter-rouge">K=1</code> 时进一步降方差</li>
</ul>

<p>它的关键 trade-off 很直白：</p>

<ul>
  <li>梯度更准</li>
  <li>样本效率更高</li>
  <li>但显存和反传成本更高</li>
</ul>

<p><code class="language-plaintext highlighter-rouge">AlignProp</code> 也属于这条 reward-backprop 路线，并通过 <code class="language-plaintext highlighter-rouge">LoRA + gradient checkpointing</code> 让直接反传更实用。[5]</p>

<p>这里有一个需要澄清的点：<br />
<code class="language-plaintext highlighter-rouge">AlignProp</code> 的 arXiv 条目后来被作者撤回，并在页面上注明内容被后续工作吸收；但“通过 reward gradient 直接调 diffusion”的路线本身没有消失，反而继续扩展到了视频。[5][9]</p>

<hr />

<h2 id="5-三条主技术路线到底该怎么区分">5. 三条主技术路线，到底该怎么区分</h2>

<p>如果把 2023 到 2026 的方法压缩成一个表，最有用的不是“谁早谁晚”，而是下面这张图。</p>

<p><br />
<img align="center" width="1000" src="https://liyongzhi.xyz/images/posts/diffusion-rl-selector.svg" alt="How to choose among diffusion RL methods" />
<br /></p>

<p>再配合这张表会更直观：</p>

<table>
  <thead>
    <tr>
      <th>路线</th>
      <th>代表方法</th>
      <th>你需要什么反馈</th>
      <th>优点</th>
      <th>代价</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>Policy Gradient</td>
      <td>DDPO, DPOK, Flow-GRPO</td>
      <td>黑盒 scalar reward</td>
      <td>通用，reward 不必可微</td>
      <td>方差高，采样贵</td>
    </tr>
    <tr>
      <td>Preference Optimization</td>
      <td>Diffusion-DPO, D3PO</td>
      <td>winner / loser 偏好对</td>
      <td>不必先训 reward model</td>
      <td>需要成对偏好数据，likelihood 近似更复杂</td>
    </tr>
    <tr>
      <td>Direct Backprop</td>
      <td>DRaFT, AlignProp, Video Reward Gradients</td>
      <td>可微 reward</td>
      <td>梯度低方差，样本效率高</td>
      <td>显存大，reward 必须可微</td>
    </tr>
  </tbody>
</table>

<p>这里最值得记住的一点是：</p>

<blockquote>
  <p>这三条线不是互相否定，而是在回答不同的问题。</p>
</blockquote>

<p>不是所有场景都该上 PPO，也不是所有场景都该上 DPO。<br />
真正的分界线其实是：<strong>你的反馈是什么形式、能不能反传、采样预算有多贵。</strong></p>

<hr />

<h2 id="6-视频生成把问题又抬高了一个量级">6. 视频生成把问题又抬高了一个量级</h2>

<p>图像里已经够难了，视频里问题会立刻放大。</p>

<p>原因很简单。<br />
视频 reward 往往不是一个“单帧好不好看”的问题，而是至少同时包含：</p>

<ul>
  <li>单帧质量</li>
  <li>文本对齐</li>
  <li>时序一致性</li>
  <li>运动合理性</li>
  <li>镜头与物理连续性</li>
</ul>

<p>而视频采样又比图像慢得多，显存开销也高得多。</p>

<h3 id="61-instructvideo把视频-reward-fine-tuning-重新表述成-editing">6.1 InstructVideo：把视频 reward fine-tuning 重新表述成 editing</h3>

<p><code class="language-plaintext highlighter-rouge">InstructVideo</code> 在 2023 年 12 月提交、后被 CVPR 2024 接收。<br />
它的做法很典型：不是傻乎乎每次都把完整 DDIM 采样链跑到底，而是把 fine-tuning 重写成 editing，从而减少 full-chain sampling 成本。[8]</p>

<p>更重要的是，它面对了视频路线最现实的问题之一：</p>

<blockquote>
  <p>当时并没有一个像 ImageReward 那样成熟的视频偏好 reward model。</p>
</blockquote>

<p>所以它把图像 reward model 重用于视频，并提出：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">Segmental Video Reward</code></li>
  <li><code class="language-plaintext highlighter-rouge">Temporally Attenuated Reward</code></li>
</ul>

<p>本质上是在说：<br />
视频 reward 不能只看“最后整段视频的一个总分”，而要想办法把 reward 更稳定地分配到片段和时序结构上。</p>

<h3 id="62-video-diffusion-alignment-via-reward-gradients把-direct-backprop-扩展到视频">6.2 Video Diffusion Alignment via Reward Gradients：把 direct backprop 扩展到视频</h3>

<p>2024 年 7 月的 <code class="language-plaintext highlighter-rouge">Video Diffusion Alignment via Reward Gradients</code> 则把 reward-backprop 路线明确推进到视频 Diffusion：<br />
既然 reward model 对 RGB 像素有稠密梯度，那就把这个梯度直接反传回视频生成过程。[9]</p>

<p>这条线的意义在于：</p>

<ul>
  <li>它说明 direct backprop 不只是图像 trick</li>
  <li>在视频这种搜索空间更大、采样更贵的场景里，低方差梯度反而更有价值</li>
</ul>

<p>所以如果你关心的是“视频 Diffusion 的 RL 怎么做”，真正要抓住的不是某个单独 paper 名，而是这两个事实：</p>

<ol>
  <li>视频 reward 一定要显式考虑时序结构。</li>
  <li>由于采样成本太高，视频场景通常更偏爱 editing、局部反传、LoRA 和稀疏 reward 设计。</li>
</ol>

<hr />

<h2 id="7-flow-matching-时代为什么-flow-grpo-是新的转折点">7. Flow Matching 时代：为什么 <code class="language-plaintext highlighter-rouge">Flow-GRPO</code> 是新的转折点</h2>

<p>到了 2025 年，主流大模型里已经不全是传统 DDPM/DDIM 了。<br />
像 SD3、FLUX 这类系统更接近 <strong>Flow Matching / Rectified Flow</strong> 范式。</p>

<p>这时旧问题又回来了：</p>

<blockquote>
  <p>RL 需要随机性和可处理的 log-prob；<br />
但 Flow Matching 的采样通常是确定性 ODE。</p>
</blockquote>

<p><code class="language-plaintext highlighter-rouge">Flow-GRPO</code> 在 2025 年 NeurIPS 上给出的回答是两个关键设计：[11]</p>

<h3 id="71-ode-to-sde-conversion">7.1 ODE-to-SDE conversion</h3>

<p>它把原本的确定性 ODE 改写成与原边际分布一致的 SDE，于是：</p>

<ul>
  <li>采样过程重新拥有随机性</li>
  <li>exploration 成立</li>
  <li>每一步转移又能写成统计上可处理的形式</li>
</ul>

<p>这一步非常重要，因为它不是在“给 Flow 模型硬套 DDPO”，而是在修补：</p>

<blockquote>
  <p>Flow 模型天然不适合直接做 reverse-process policy gradient</p>
</blockquote>

<h3 id="72-denoising-reduction">7.2 Denoising Reduction</h3>

<p><code class="language-plaintext highlighter-rouge">Flow-GRPO</code> 的第二个关键点是：<br />
训练时减少 denoising steps，推理时保留原本的高质量步数。</p>

<p>这说明一个很现实的经验：</p>

<blockquote>
  <p>RL 后训练并不一定需要最完美的生成样本，只需要足够稳定、足够可区分的 reward 信号。</p>
</blockquote>

<p>从工程角度看，这让 Flow 模型的 RL 后训练第一次真正变得可用。</p>

<hr />

<h2 id="8-截至-2026-年-4-月这条线又在往哪里走">8. 截至 2026 年 4 月，这条线又在往哪里走</h2>

<p>如果只看 2023 到 2025，你会觉得主线大概是：</p>

<p><code class="language-plaintext highlighter-rouge">DDPO -&gt; DPO / direct backprop -&gt; Flow-GRPO</code></p>

<p>但到 2026 年，研究重点已经明显转向了两个更细的问题。</p>

<h3 id="81-更好的-credit-assignment-和-rollout-复用">8.1 更好的 credit assignment 和 rollout 复用</h3>

<p><code class="language-plaintext highlighter-rouge">TreeGRPO</code>（ICLR 2026）把 denoising 过程重写成一棵搜索树，通过共享轨迹前缀来提升样本效率，并尝试解决“终止 reward 被粗暴平均分给所有时间步”的问题。[12]</p>

<p>这件事说明社区已经意识到：<br />
<strong>uniform terminal reward assignment</strong> 其实是老一代 DDPO / GRPO 风格方法的一个核心瓶颈。</p>

<h3 id="82-不再执着于-reverse-process-likelihood">8.2 不再执着于 reverse-process likelihood</h3>

<p><code class="language-plaintext highlighter-rouge">DiffusionNFT</code>（ICLR 2026 Oral）更进一步，直接质疑了“必须在 reverse sampling 里估 log-prob 才能做 online RL”这个前提。[13]</p>

<p>它的出发点很清楚：</p>

<ul>
  <li>reverse likelihood 受 solver 选择限制</li>
  <li>和 CFG 的兼容性复杂</li>
  <li>trajectory 级策略优化的成本太高</li>
</ul>

<p>所以它改走了 forward-process / flow-matching 风格目标，并声称在效率上显著超过 <code class="language-plaintext highlighter-rouge">Flow-GRPO</code>。[13]</p>

<p>这说明截至 <strong>2026 年 4 月 20 日</strong>，这个领域已经不再只是“把 PPO 套到去噪轨迹上”，而是在认真重构：</p>

<ul>
  <li>policy objective 到底该写成什么</li>
  <li>log-prob 是不是必须的</li>
  <li>terminal reward 该如何分配到各步</li>
  <li>采样器、CFG、solver 和 RL 目标能不能更自然地统一</li>
</ul>

<hr />

<h2 id="9-实践里最难的其实不是公式而是-reward-design">9. 实践里最难的其实不是公式，而是 reward design</h2>

<p>从工程上看，Diffusion RL 最大的敌人一直都不是“推不出梯度”，而是 <strong>reward hacking</strong>。</p>

<p>典型例子包括：</p>

<ul>
  <li>只优化美学分，模型会学会“糖水色”和过饱和</li>
  <li>只优化 CLIP 对齐，模型可能学会骗 VLM</li>
  <li>只优化时序一致性，视频可能变成几乎静止</li>
</ul>

<p>所以大多数可用系统都会同时做四件事：</p>

<h3 id="91-用-kl-约束守住预训练分布">9.1 用 KL 约束守住预训练分布</h3>

<p><code class="language-plaintext highlighter-rouge">DPOK</code> 之后，这几乎已经是标配。<br />
你可以把它理解为一个护栏：</p>

<blockquote>
  <p>允许模型变好，但不允许它为了 reward 快速偏航。</p>
</blockquote>

<h3 id="92-用多维-reward-而不是单分数独裁">9.2 用多维 reward 而不是单分数独裁</h3>

<p>实际系统里更常见的是加权组合：</p>

<ul>
  <li>质量</li>
  <li>对齐</li>
  <li>安全</li>
  <li>时序</li>
  <li>OCR / counting / composition</li>
</ul>

<p>单一 reward 往往最容易被钻空子。</p>

<h3 id="93-用-lora低步数反传和局部更新控制成本">9.3 用 LoRA、低步数反传和局部更新控制成本</h3>

<p>这是 <code class="language-plaintext highlighter-rouge">DRaFT</code>、<code class="language-plaintext highlighter-rouge">AlignProp</code>、视频 alignment 工作都反复验证过的经验：</p>

<ul>
  <li>不是所有参数都要动</li>
  <li>不是所有时间步都要反传</li>
  <li>不是每次都要完整链路采样</li>
</ul>

<h3 id="94-先问清楚你手里的反馈是什么">9.4 先问清楚“你手里的反馈是什么”</h3>

<p>这是最重要的实操建议：</p>

<ul>
  <li>如果 reward 可微，优先考虑 direct backprop</li>
  <li>如果只有偏好对，优先考虑 DPO 类方法</li>
  <li>如果 reward 是黑盒且可调用，再考虑 policy gradient</li>
  <li>如果模型已经是 Flow Matching，要特别注意随机性和 likelihood 的定义问题</li>
</ul>

<hr />

<h2 id="10-应该怎么选方法">10. 应该怎么选方法</h2>

<p>如果只给一个实用版判断树，我会这样用：</p>

<ol>
  <li>
    <p>reward 可微吗？
可微：优先 <code class="language-plaintext highlighter-rouge">DRaFT / AlignProp / Video Reward Gradients</code> 这类直接反传路线。</p>
  </li>
  <li>
    <p>reward 不可微，但有偏好对吗？
有：优先 <code class="language-plaintext highlighter-rouge">Diffusion-DPO / D3PO</code>。</p>
  </li>
  <li>
    <p>既没有可微 reward，也没有偏好对，只有黑盒打分器？
图像 DDPM 类：<code class="language-plaintext highlighter-rouge">DDPO / DPOK</code>。<br />
Flow Matching 类：先看 <code class="language-plaintext highlighter-rouge">Flow-GRPO</code>，再关注 <code class="language-plaintext highlighter-rouge">DiffusionNFT</code> 这类新范式。</p>
  </li>
  <li>
    <p>是视频吗？
默认把“reward 设计”和“采样成本”放在首位，再决定是 editing 式 fine-tuning、direct backprop，还是更传统的 RL 更新。</p>
  </li>
</ol>

<hr />

<h2 id="11-总结">11. 总结</h2>

<p>如果只用一句话总结这条脉络，我会写成：</p>

<blockquote>
  <p>Diffusion RL 的本质，不是把 PPO 生搬硬套到生成模型上；<br />
而是把“逐步去噪”重新看成一个可优化的决策过程，再根据你手里的反馈形式，选择 policy gradient、偏好优化或 reward backprop 这三类不同工具。</p>
</blockquote>

<p>更具体地说：</p>

<ol>
  <li><code class="language-plaintext highlighter-rouge">DDPO / DPOK</code> 解决的是“黑盒 reward 怎么直接优化”。</li>
  <li><code class="language-plaintext highlighter-rouge">Diffusion-DPO / D3PO</code> 解决的是“只有偏好对时，能不能跳过 reward model”。</li>
  <li><code class="language-plaintext highlighter-rouge">DRaFT / AlignProp</code> 解决的是“如果 reward 可微，为什么还要忍受高方差 RL”。</li>
  <li><code class="language-plaintext highlighter-rouge">InstructVideo</code> 和后续视频工作说明，视频不是图像方法的简单复制，而是一个 reward design 和效率问题都更尖锐的场景。</li>
  <li><code class="language-plaintext highlighter-rouge">Flow-GRPO</code> 则标志着这条线正式进入 Flow Matching 时代。</li>
  <li>到 2026 年，研究已经开始进一步追问：<code class="language-plaintext highlighter-rouge">log-prob</code> 是否必须、credit assignment 能否更细、采样器与 RL 目标能否更统一。</li>
</ol>

<p>所以从历史上看，Diffusion 模型的“强化学习过程”并不是一套固定算法，而是一条不断重新定义<strong>策略、反馈、轨迹和约束</strong>的演化路线。</p>

<hr />

<h2 id="参考资料">参考资料</h2>

<ol>
  <li>
    <p>Xu, Liu, Wu, Tong, Li, Ding, Tang, Dong. <em>ImageReward: Learning and Evaluating Human Preferences for Text-to-Image Generation</em>. arXiv 2023.<br />
<a href="https://arxiv.org/abs/2304.05977">arXiv</a></p>
  </li>
  <li>
    <p>Black, Janner, Du, Kostrikov, Levine. <em>Training Diffusion Models with Reinforcement Learning</em>. arXiv 2023.<br />
<a href="https://arxiv.org/abs/2305.13301">arXiv</a> | <a href="https://rl-diffusion.github.io/">Project</a></p>
  </li>
  <li>
    <p>Fan, Watkins, Du, Liu, Ryu, Boutilier, Abbeel, Ghavamzadeh, K. Lee, K. Lee. <em>DPOK: Reinforcement Learning for Fine-tuning Text-to-Image Diffusion Models</em>. arXiv 2023.<br />
<a href="https://arxiv.org/abs/2305.16381">arXiv</a></p>
  </li>
  <li>
    <p>Clark, Vicol, Swersky, Fleet. <em>Directly Fine-Tuning Diffusion Models on Differentiable Rewards</em>. arXiv 2023 / ICLR 2024.<br />
<a href="https://arxiv.org/abs/2309.17400">arXiv</a></p>
  </li>
  <li>
    <p>Prabhudesai, Goyal, Pathak, Fragkiadaki. <em>Aligning Text-to-Image Diffusion Models with Reward Backpropagation</em>. arXiv 2023.<br />
<a href="https://arxiv.org/abs/2310.03739">arXiv</a> | <a href="https://align-prop.github.io/">Project</a></p>
  </li>
  <li>
    <p>Wallace, Dang, Rafailov, Zhou, Lou, Purushwalkam, Ermon, Xiong, Joty, Naik. <em>Diffusion Model Alignment Using Direct Preference Optimization</em>. arXiv 2023.<br />
<a href="https://arxiv.org/abs/2311.12908">arXiv</a></p>
  </li>
  <li>
    <p>Yang, Tao, Lyu, Ge, Chen, Li, Shen, Zhu, Li. <em>Using Human Feedback to Fine-tune Diffusion Models without Any Reward Model</em>. arXiv 2023 / CVPR 2024.<br />
<a href="https://arxiv.org/abs/2311.13231">arXiv</a></p>
  </li>
  <li>
    <p>Yuan, Zhang, Wang, Wei, Feng, Pan, Zhang, Liu, Albanie, Ni. <em>InstructVideo: Instructing Video Diffusion Models with Human Feedback</em>. arXiv 2023 / CVPR 2024.<br />
<a href="https://arxiv.org/abs/2312.12490">arXiv</a></p>
  </li>
  <li>
    <p>Prabhudesai, Mendonca, Qin, Fragkiadaki, Pathak. <em>Video Diffusion Alignment via Reward Gradients</em>. arXiv 2024.<br />
<a href="https://arxiv.org/abs/2407.08737">arXiv</a></p>
  </li>
  <li>
    <p>Uehara, Zhao, Biancalani, Levine. <em>Understanding Reinforcement Learning-Based Fine-Tuning of Diffusion Models: A Tutorial and Review</em>. arXiv 2024.<br />
<a href="https://arxiv.org/abs/2407.13734">arXiv</a></p>
  </li>
  <li>
    <p>Liu, Liu, Liang, Li, Liu, Wang, Wan, Zhang, Ouyang. <em>Flow-GRPO: Training Flow Matching Models via Online RL</em>. NeurIPS 2025.<br />
<a href="https://openreview.net/forum?id=oCBKGw5HNf">OpenReview</a></p>
  </li>
  <li>
    <p>Ding, Ye. <em>TreeGRPO: Tree-Advantage GRPO for Online RL Post-Training of Diffusion Models</em>. ICLR 2026.<br />
<a href="https://openreview.net/forum?id=3rZdp4TmUb">OpenReview</a></p>
  </li>
  <li>
    <p>Zheng, Chen, Ye, Wang, Zhang, Jiang, Su, Ermon, Zhu, Liu. <em>DiffusionNFT: Online Diffusion Reinforcement with Forward Process</em>. ICLR 2026 Oral.<br />
<a href="https://openreview.net/forum?id=VJZ477R89F">OpenReview</a></p>
  </li>
</ol>]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><category term="blog" /><category term="diffusion" /><category term="reinforcement learning" /><category term="generative model" /><category term="machine learning" /><category term="computer vision" /><summary type="html"><![CDATA[Diffusion 模型最初是按“去噪 MSE / 似然近似”来训练的，但真正上线时，我们更关心的往往不是似然，而是：]]></summary></entry><entry><title type="html">让大模型快 8 倍：从投机解码到 DDTree 的完整原理</title><link href="https://liyongzhi.xyz/posts/2026/04/speculative-decoding/" rel="alternate" type="text/html" title="让大模型快 8 倍：从投机解码到 DDTree 的完整原理" /><published>2026-04-20T00:00:00+08:00</published><updated>2026-04-20T00:00:00+08:00</updated><id>https://liyongzhi.xyz/posts/2026/04/blog-post-speculative-decoding</id><content type="html" xml:base="https://liyongzhi.xyz/posts/2026/04/speculative-decoding/"><![CDATA[<blockquote>
  <p>本文从零开始，带你理解 LLM 推理加速的核心思路，读完之后你会明白：大模型为什么慢、投机解码如何加速、为什么加速后输出质量完全不变，以及 DDTree 这篇 2026 年的新论文究竟做了什么创新。</p>
</blockquote>

<video controls="" style="width: 100%; max-width: 960px; display: block; margin: 2rem auto;">
  <source src="/videos/spec-decoding.mp4" type="video/mp4" />
  你的浏览器不支持 HTML5 视频播放。
</video>

<hr />

<h2 id="1-大模型推理为什么慢">1. 大模型推理为什么慢？</h2>

<p>你每次跟 ChatGPT 或 Claude 对话，它生成文字的方式其实非常朴素——<strong>每次只生成一个 token（词片段），然后把这个 token 加进上下文，再生成下一个，如此往复</strong>。</p>

<p>这种方式叫做<strong>自回归解码（Autoregressive Decoding）</strong>。它的问题在于天然的串行性：第 $n$ 个 token 必须等第 $n-1$ 个 token 生成完才能开始，完全无法并行。</p>

<p>更麻烦的是，现代大模型动辄数百亿参数，每生成一个 token 就要做一次完整的前向传播，把所有参数都过一遍。而 GPU 最擅长的恰恰是<strong>大批量并行计算</strong>——生成单个 token 这件事，对 GPU 来说几乎是一种浪费，大量算力处于闲置状态。</p>

<p>用一个类比：这就好比你雇了一个有 100 条生产线的工厂，却每次只让它生产一个零件，然后等零件出来之后再决定下一个生产什么。工厂的产能大部分都在空转。</p>

<p>那有没有办法，让大模型一次性并行生成多个 token，又不影响输出质量？</p>

<p>这就是<strong>投机解码（Speculative Decoding）</strong>的出发点。</p>

<hr />

<h2 id="2-投机解码用小模型猜用大模型验">2. 投机解码：用小模型”猜”，用大模型”验”</h2>

<p>投机解码的核心思想简单到令人惊喜：</p>

<blockquote>
  <p>用一个<strong>轻量的草稿模型（Draft Model）</strong> 先快速生成多个候选 token，再让<strong>大的目标模型（Target Model）</strong> 一次性并行验证这些候选 token，接受正确的，拒绝错误的。</p>
</blockquote>

<p>为什么这能加速？关键在于一个不对称性：<strong>大模型验证多个 token 和验证单个 token，所需要的计算量几乎相同</strong>。这是因为验证本质上是一次 prefill（并行处理整个序列），而不是逐个自回归生成。</p>

<p>具体流程是这样的：</p>

<ol>
  <li><strong>草稿阶段</strong>：草稿模型连续生成 $k$ 个 token（比如 4 个），速度很快</li>
  <li><strong>验证阶段</strong>：目标模型把这 $k$ 个候选 token 一次性并行处理，输出对每个位置的概率判断</li>
  <li><strong>接受/拒绝</strong>：从左到右逐个判断，接受的 token 保留，遇到第一个被拒绝的 token 就停止，后面的全部丢弃</li>
  <li><strong>下一轮</strong>：从被拒绝的位置重新开始草稿</li>
</ol>

<p>如果草稿模型猜对了 3 个 token，那这一轮相当于大模型用一次前向传播的时间，产出了 3 个高质量 token，而不是通常的 1 个。加速比可以非常显著。</p>

<h3 id="草稿模型的-qx-究竟是什么">草稿模型的 $q(x)$ 究竟是什么？</h3>

<p>一个常见的细节困惑：草稿模型给每个 token 的概率 $q(x)$，到底指什么？</p>

<ul>
  <li><strong>Temperature = 0（贪心）</strong>：草稿直接取 argmax 那个 token，$q(x)$ 就是该 token 在 softmax 后对应的那个最大概率值</li>
  <li><strong>Temperature &gt; 0（采样）</strong>：按分布采样一个 token，$q(x)$ 就是采到的那个 token 对应的概率值</li>
  <li><strong>DFlash 这类块扩散草稿</strong>：一次 forward pass 直接输出每个位置的完整边际分布，$q(x)$ 就是该分布在对应 token 上的概率值</li>
</ul>

<p>无论哪种情况，$q(x)$ 都是”草稿模型对它自己选出的那个 token 所赋予的概率”，而不是整个分布。这一点在后面理解接受/拒绝规则时非常重要。</p>

<hr />

<h2 id="3-最关键的问题输出质量真的不变吗">3. 最关键的问题：输出质量真的不变吗？</h2>

<p>这是整个方法最精妙的地方，也是很多人最困惑的地方。</p>

<p><strong>答案是：严格不变，数学上可以证明。</strong></p>

<p>关键在于接受/拒绝的规则设计。对于草稿模型提出的某个 token $x$：</p>

<ul>
  <li>草稿模型给它的概率是 $q(x)$</li>
  <li>目标模型给它的概率是 $p(x)$</li>
</ul>

<p><strong>以概率 $\min!\left(1,\, \frac{p(x)}{q(x)}\right)$ 接受这个 token。</strong></p>

<p>直觉很简单：如果目标模型认为这个 token 的概率至少和草稿模型一样高（$p \geq q$），就无条件接受；如果目标模型认为它概率更低（$p &lt; q$），就按比例拒绝，避免这个 token 被过度采样。</p>

<p><strong>被拒绝时怎么办？</strong> 不是直接放弃，而是从一个”残差分布”里重新采样：</p>

\[p'(x) = \frac{\max(0,\; p(x) - q(x))}{1 - \beta}, \quad \text{其中 } \beta = \sum_y \min(p(y), q(y))\]

<p>这个残差分布的含义是：草稿模型过度提名了某些 token（$q &gt; p$），那些 token 残差为 0，不再有机会；而那些被草稿欠缺代表的 token（$p &gt; q$），按亏欠量分配权重，作为补偿。</p>

<h3 id="一个具体的数字例子">一个具体的数字例子</h3>

<p>抽象公式不好懂，换一个只有 5 个 token 的小词表来算一遍。假设目标分布 $p$ 和草稿分布 $q$ 如下，草稿模型抽出来的 token 是 “cat”：</p>

<table>
  <thead>
    <tr>
      <th>Token</th>
      <th>$p(x)$</th>
      <th>$q(x)$</th>
      <th>$\min(p,q)$</th>
      <th>$\max(0,\, p-q)$</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>cat</td>
      <td>0.50</td>
      <td>0.70</td>
      <td>0.50</td>
      <td>0</td>
    </tr>
    <tr>
      <td>dog</td>
      <td>0.20</td>
      <td>0.10</td>
      <td>0.10</td>
      <td>0.10</td>
    </tr>
    <tr>
      <td>sat</td>
      <td>0.15</td>
      <td>0.05</td>
      <td>0.05</td>
      <td>0.10</td>
    </tr>
    <tr>
      <td>on</td>
      <td>0.10</td>
      <td>0.10</td>
      <td>0.10</td>
      <td>0</td>
    </tr>
    <tr>
      <td>mat</td>
      <td>0.05</td>
      <td>0.05</td>
      <td>0.05</td>
      <td>0</td>
    </tr>
  </tbody>
</table>

<p><strong>第一步：判断是否接受 “cat”</strong>。因为 $p(\text{cat})/q(\text{cat}) = 0.50/0.70 \approx 0.71 &lt; 1$，以 71% 的概率接受，29% 的概率拒绝。</p>

<p><strong>第二步：一旦被拒绝，从残差分布 $p’$ 重新采样</strong>。先算 $\beta = \sum \min(p, q) = 0.50 + 0.10 + 0.05 + 0.10 + 0.05 = 0.80$，因此 $1 - \beta = 0.20$。把 $\max(0,\, p-q)$ 这一列每个元素除以 $0.20$：</p>

<table>
  <thead>
    <tr>
      <th>Token</th>
      <th>残差概率 $p’(x)$</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>cat</td>
      <td>0 / 0.20 = <strong>0</strong></td>
    </tr>
    <tr>
      <td>dog</td>
      <td>0.10 / 0.20 = <strong>0.50</strong></td>
    </tr>
    <tr>
      <td>sat</td>
      <td>0.10 / 0.20 = <strong>0.50</strong></td>
    </tr>
    <tr>
      <td>on</td>
      <td>0 / 0.20 = <strong>0</strong></td>
    </tr>
    <tr>
      <td>mat</td>
      <td>0 / 0.20 = <strong>0</strong></td>
    </tr>
  </tbody>
</table>

<p>于是：被拒绝后，以 50/50 的概率从 “dog” 和 “sat” 中补采一个。</p>

<h3 id="为什么残差分布不是把-cat-去掉再归一化">为什么残差分布不是”把 cat 去掉再归一化”？</h3>

<p>这里最容易踩的坑：很多人第一反应是”既然 cat 被拒了，就把 cat 从 $p$ 里划掉，剩下的 ${0.20, 0.15, 0.10, 0.05}$ 归一化一下（和是 0.50），按这个分布采一个就好”。</p>

<p>但那样算出来 “on” 会得到 $0.10/0.50 = 0.20$ 的概率。而残差分布给 “on” 的概率是 <strong>0</strong>。差别在哪？</p>

<p>关键是 “on” 满足 $p(\text{on}) = q(\text{on}) = 0.10$——草稿模型已经给 “on” 分配了完全正确比例的概率质量，不多不少。再在补采里给它机会，就会让 “on” 在最终输出分布里超出 0.10。</p>

<p><strong>残差分布只补偿那些被草稿欠缺代表的 token（$p &gt; q$）</strong>，而不是盲目地把剩下的 $p$ 归一化。这才是最终边际分布严格等于 $p$ 的关键。</p>

<h3 id="为什么整体输出分布仍然等于-p">为什么整体输出分布仍然等于 $p$？</h3>

<p>把直接接受和补采的贡献加起来：</p>

\[P_{\text{out}}(x) = \underbrace{\min(p(x),\, q(x))}_{\text{直接接受贡献}} + \underbrace{\max(0,\, p(x) - q(x))}_{\text{补采贡献}} = p(x)\]

<p>这个等式对任意 $p, q$ 都严格成立，无需任何假设。草稿模型质量只影响速度（草稿越准，$\beta$ 越大，补采越少），<strong>永远不影响输出正确性</strong>。</p>

<p>这种技术叫做<strong>拒绝采样（Rejection Sampling）</strong>，是概率论里的经典工具，被投机解码巧妙地应用到了 LLM 推理加速上。</p>

<hr />

<h2 id="4-级联丢弃为什么第一个错了后面全废">4. 级联丢弃：为什么第一个错了，后面全废</h2>

<p>这里有个重要细节需要理解。</p>

<p>草稿模型生成 token 序列时，每个 token 都是条件于前面的 token 生成的。比如草稿模型生成的是 “The cat sat on the mat”，其中 “sat” 是在”已知前两个 token 是 The, cat”的条件下预测的；”on” 是在”已知前三个是 The, cat, sat”的条件下预测的。</p>

<p>一旦 “cat” 被目标模型拒绝，目标模型会在那个位置补采一个不同的 token，比如 “dog”。此时，原草稿里的 “sat” 是基于错误前缀 “The, cat” 生成的 $q(\text{sat} \mid \text{The, cat})$，但真实需要的条件概率是 $q(\text{sat} \mid \text{The, dog})$——这是两个完全不同的分布，原来的 $q$ 值对新上下文毫无参考价值。</p>

<p>因此：<strong>第 $i$ 个位置被拒绝，位置 $i+1$ 及之后的所有草稿 token 必须全部丢弃</strong>，不再验证。</p>

<p>这也意味着：</p>

<ul>
  <li>草稿模型的质量极其重要——第一个 token 就被拒绝，整轮草稿全部白费</li>
  <li>优化的核心指标是<strong>期望接受长度</strong>（expected accepted length）：拒绝发生得越晚，这一轮能白捡的 token 越多</li>
  <li>从 $\beta = \sum_x \min(p(x), q(x))$ 也能看出：当 $q \equiv p$ 时 $\beta = 1$，每个 token 都被接受；当 $q$ 与 $p$ 完全不重叠时 $\beta \to 0$，几乎每个 token 都被拒，速度反而比纯自回归还慢（多跑了一次草稿模型）</li>
</ul>

<p>正是因此，专门为投机解码训练的草稿模型（比如 EAGLE、DFlash）要比通用小模型效果好得多：它们被明确训练成”输出分布尽可能像目标模型”，而不仅仅是”在自己的训练集上 loss 最小”。</p>

<hr />

<h2 id="5-dflash块扩散草稿模型">5. DFlash：块扩散草稿模型</h2>

<p>传统的草稿模型和目标模型一样，也是自回归的——要生成 4 个草稿 token，就要跑 4 次前向传播。自回归草稿的问题是：虽然模型小，但”串行”这个结构性瓶颈一点没变。EAGLE、EAGLE-2、EAGLE-3 这类主流草稿模型都是自回归派系的优化（更轻的 head、共享 hidden state 等），它们在单步上做得很快，但仍然要跑 $k$ 次 forward 才能出 $k$ 个草稿。</p>

<p><strong>DFlash</strong> 走了完全不同的一条路：<strong>块扩散（Block Diffusion）</strong>。它基于 Arriola 等人在 ICLR 2025 提出的 <code class="language-plaintext highlighter-rouge">BD3-LM</code>（Block Discrete Denoising Diffusion Language Model）范式，并针对投机解码场景做了专门适配。</p>

<h3 id="51-核心范式块之间自回归块内扩散">5.1 核心范式：块之间自回归，块内扩散</h3>

<p>BD3-LM 的一句话总结：<strong>块与块之间保持自回归，块内部用离散扩散一次性并行预测所有位置</strong>。</p>

<p>具体来说：</p>

<ul>
  <li>把要生成的长序列切成若干固定长度 $L_b$ 的块（例如 $L_b = 4, 8, 16$）</li>
  <li>块与块之间是严格的因果顺序：第 $j$ 块的生成必须在第 $j-1$ 块完成之后进行</li>
  <li>块内部的 $L_b$ 个位置通过<strong>离散扩散</strong>一次性联合预测</li>
</ul>

<p>这种混合结构解决了纯 diffusion 语言模型的两大痛点：<strong>固定长度限制</strong>和<strong>无法做 KV cache</strong>。因为块之间还是自回归，已经生成好的块就可以像普通 LLM 一样把 K/V 缓存下来，后续块只需要在缓存之上再做 attention——这和目标模型的 KV cache 使用完全兼容。</p>

<h3 id="52-训练目标块内的-masked-diffusion-elbo">5.2 训练目标：块内的 masked diffusion ELBO</h3>

<p>训练的时候，对每一个训练样本：</p>

<ol>
  <li>采一个噪声水平 $t \in [0, 1]$</li>
  <li>在<strong>当前块</strong>的 $L_b$ 个位置上，独立地以概率 $t$ 把每个 token 替换成特殊的 <code class="language-plaintext highlighter-rouge">[MASK]</code> token</li>
  <li>保留前面所有完整的块（上下文），只在当前块的位置上加噪</li>
  <li>模型的任务是从被 mask 的块中恢复原始 token</li>
</ol>

<p>损失函数本质上是离散 diffusion 的 ELBO，形式上可以理解成一个加权的交叉熵：</p>

\[\mathcal{L}_{\text{BD}} = \mathbb{E}_{t, x}\left[\sum_{i \in \text{当前块}} \mathbb{1}[x_i^{(t)} = \text{MASK}] \cdot \frac{1}{t} \cdot \log p_\theta(x_i \mid x^{(t)}, x_{&lt;\text{block}})\right]\]

<p>几个要点：</p>

<ul>
  <li>只在被 mask 的位置计算 loss（未被 mask 的位置模型只是看到了答案）</li>
  <li>$1/t$ 是重要性权重，把不同噪声水平的贡献归一化</li>
  <li>模型同时看到了”前面已完成的块”和”当前块的部分观测”，需要学会利用这两者一起做预测</li>
</ul>

<h3 id="53-架构块因果-attentionblock-causal-attention">5.3 架构：块因果 attention（Block-Causal Attention）</h3>

<p>BD3-LM 用的还是标准 Transformer，但 attention mask 变了：</p>

<ul>
  <li><strong>同一块内</strong>：全 attention（每个位置可以看到块内所有位置，包括被 mask 的）</li>
  <li><strong>跨块</strong>：因果 attention（当前块可以看前面所有块，反之不行）</li>
</ul>

<p>画出来像这样：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>            块1  块2  块3
       块1 [■ ■ ■][□ □ □][□ □ □]
       块2 [■ ■ ■][■ ■ ■][□ □ □]
       块3 [■ ■ ■][■ ■ ■][■ ■ ■]
</code></pre></div></div>

<p>（■ 可见，□ 不可见；每块 3 个位置）</p>

<p>这个结构让 KV cache 成为可能：块 1 算完之后，它的 K/V 就固定了；块 2 生成时复用块 1 的缓存，只需要新计算块 2 这 $L_b$ 个位置的 K/V；块 3 再叠加。跟自回归 LLM 的 KV cache 行为一致。</p>

<h3 id="54-采样t-步去噪一次得到一个块">5.4 采样：T 步去噪一次得到一个块</h3>

<p>给定前面的块做上下文，生成当前块的流程是：</p>

<ol>
  <li>把当前块的所有 $L_b$ 个位置初始化为 <code class="language-plaintext highlighter-rouge">[MASK]</code></li>
  <li>进行 $T$ 步去噪：每一步，模型看一眼当前的”部分 mask 状态”，对每个仍然是 <code class="language-plaintext highlighter-rouge">[MASK]</code> 的位置输出一个 softmax 分布，按某种策略（如置信度 top-k 或随机）选一部分位置把 mask 替换成采样到的 token</li>
  <li>直到所有位置都被填上，当前块完成</li>
</ol>

<p>关键：<strong>最后一步去噪得到的 $L_b$ 个 softmax 分布，就是 DDTree 要用的”每个位置的概率分布”</strong>。也就是说，DFlash 在块生成的最终一步 forward pass 里，同时暴露了块内每个位置的完整后验分布——这正是后面 DDTree 能用它构树的原因。</p>

<p>BD3-LM 原论文在一般生成时用 $T = 5000$ 步（质量优先），但在 draft 场景下完全不需要这么多——DFlash 把 $T$ 降到很小（论文级别通常 $T \in {1, 2, 4}$ 甚至 single-step），用”一次去噪直接输出块”的方式保证草稿延迟最低。</p>

<h3 id="55-块大小的权衡">5.5 块大小的权衡</h3>

<p>块大小 $L_b$ 是 DFlash 最关键的超参数：</p>

<table>
  <thead>
    <tr>
      <th>块大小</th>
      <th>草稿模型调用次数</th>
      <th>单次输出 token 数</th>
      <th>被接受长度期望</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>小（$L_b=4$）</td>
      <td>多</td>
      <td>4</td>
      <td>短</td>
    </tr>
    <tr>
      <td>中（$L_b=8$）</td>
      <td>中</td>
      <td>8</td>
      <td>中</td>
    </tr>
    <tr>
      <td>大（$L_b=16$）</td>
      <td>少</td>
      <td>16</td>
      <td>可能长也可能第一个就被拒</td>
    </tr>
  </tbody>
</table>

<p>BD3-LM 论文显示，$L_b$ 越大，困惑度恶化越明显（diffusion 建模 16 个位置比建模 4 个位置难得多）；但 $L_b$ 越小，草稿模型要被调用更多次才能”推进”相同距离。投机解码里通常选 $L_b \in {4, 8}$，在单次出 token 数和接受长度之间取平衡。</p>

<h3 id="56-dflash-针对投机解码的特化">5.6 DFlash 针对投机解码的特化</h3>

<p>把 block diffusion 作为草稿模型，DFlash 相比原始 BD3-LM 做了几项适配：</p>

<ul>
  <li><strong>目标模型蒸馏</strong>：训练时不只最大化 block diffusion ELBO，还加一项”草稿输出分布尽可能贴近目标模型”的 KL 损失，让 $\beta = \sum \min(p, q)$ 尽可能大</li>
  <li><strong>极少去噪步数</strong>：从 5000 步降到 1–4 步，在 draft 阶段用”近似一步到位”换延迟</li>
  <li><strong>输出带温度的完整分布</strong>：不像 BD3-LM 需要最终离散 token，DFlash 保留最后一步 softmax 的完整分布以供 DDTree 这类下游使用</li>
  <li><strong>共享 tokenizer 和 context</strong>：和目标模型使用同一份 tokenizer，KV cache 可以部分对齐复用</li>
</ul>

<h3 id="57-一次-forward-pass-实际输出什么">5.7 一次 forward pass 实际输出什么？</h3>

<p>有了上面这些背景，具体看 DFlash 一次前向传播（假设 $L_b = 16$，$T = 1$）的输出：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>输入：已经生成的 prefix tokens + 16 个 [MASK] 占位符
       │
       ▼（一次 block-causal transformer forward pass）
       │
输出：16 个 softmax 分布（每个分布大小 = 词表大小）

位置 1 的分布：cat=0.70, dog=0.20, sat=0.08, ...
位置 2 的分布：sat=0.60, lay=0.30, on=0.08, ...
位置 3 的分布：on=0.65, at=0.20, the=0.12, ...
...
位置 16 的分布：...
</code></pre></div></div>

<p>然后 DFlash 从每个位置的分布里各采样一个 token，拼成草稿序列，交给目标模型验证。</p>

<p>这比自回归草稿快很多——不管块有多长，草稿生成只需要一次前向传播。DFlash 在实测中已经超越了 EAGLE-3 等强大的自回归草稿模型，成为投机解码的领先方案。</p>

<p>但 DFlash 有一个明显的浪费：<strong>每个位置的概率分布包含丰富的候选信息，却只用了最高概率的一个 token</strong>。</p>

<p>比如位置 1，cat 概率 0.70，dog 概率 0.20。DFlash 只采了 cat，dog 的可能性完全被忽略了。如果目标模型拒绝了 cat，这一轮就白费了，而其实 dog 也是一个很有希望的候选。</p>

<p>能不能把这些信息都利用起来？</p>

<hr />

<h2 id="6-ddtree从一条路到一棵树">6. DDTree：从”一条路”到”一棵树”</h2>

<p><strong>DDTree（Diffusion Draft Tree）</strong> 就是这个问题的答案。</p>

<p>核心思路极其直觉：既然 DFlash 一次 forward pass 已经给出了每个位置的完整分布，为什么不用这些分布构建<strong>多条候选路径</strong>，组成一棵树，让目标模型一次性验证整棵树？</p>

<h3 id="61-草稿树长什么样">6.1 草稿树长什么样？</h3>

<p>树的根节点是上一轮结尾的 token，然后从每个位置的分布里选出 top-K 个候选，展开成多条分支：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>           root（…The）
          /            \
      cat (0.70)      dog (0.20)
      /      \              \
  sat(0.42) lay(0.21)    sat(0.12)
    /
 on(0.27)
</code></pre></div></div>

<p>每条从根到叶子的路径，就是一个完整的草稿序列。目标模型可以同时验证所有路径，找出其中最长的被接受前缀。</p>

<h3 id="62-节点预算-b速度与收益的权衡">6.2 节点预算 B：速度与收益的权衡</h3>

<p>树里的节点越多，覆盖的候选路径越丰富，被接受的 token 越多。但同时，目标模型验证时需要处理更多节点，开销也更大。</p>

<p>DDTree 用一个<strong>节点预算 $B$</strong> 来控制树的大小——$B$ 就是整棵树里最多允许几个节点槽。在 $B$ 范围内尽可能选出最有价值的节点。论文实验表明，最优预算大约是 $B = 512$。</p>

<h3 id="63-怎么在预算内选出最优树best-first-heap">6.3 怎么在预算内选出最优树？Best-First Heap</h3>

<p>每条路径的联合概率 = 各位置边际概率之积（因为 DFlash 输出的是独立的按位置分布）：</p>

\[P(\text{cat} \to \text{sat} \to \text{on}) = 0.70 \times 0.60 \times 0.65 \approx 0.27\]

<p>DDTree 用一个<strong>最优先堆（Best-First Heap）</strong> 来贪心构建最优树：</p>

<ol>
  <li>初始把根节点的所有子节点候选放入堆，按概率排序</li>
  <li>每次弹出概率最高的叶子节点，将其子节点（下一位置的 top-K）推入堆</li>
  <li>重复，直到节点总数达到预算 $B$</li>
</ol>

<p>整个过程<strong>完全基于 DFlash 同一次 forward pass 的输出</strong>，不需要再调草稿模型，时间复杂度仅 $O(B \log(B \cdot K))$，极快。可以证明，这个算法在预算 $B$ 内构建出的树，能<strong>最大化期望被接受的 token 数量</strong>。</p>

<h3 id="64-ancestor-only-attention-mask一次并行验证整棵树">6.4 Ancestor-only Attention Mask：一次并行验证整棵树</h3>

<p>树构建好之后，把所有节点按 BFS/DFS 顺序拉平成一维序列，输入目标模型做 prefill。但需要一个特殊的 attention mask：<strong>每个节点只能 attend 到它的祖先节点</strong>，不能看到兄弟或表亲节点。</p>

<p>这保证了每条路径的验证是独立正确的——就好像每条路径单独 prefill，但通过共享公共祖先的 KV（兄弟路径共享同一份祖先的 KV，只存一份），实际上一次 forward pass 就完成了所有路径的验证，效率极高。</p>

<p><strong>kernel 实现的坑</strong>：这种任意稀疏的 attention mask，标准 FlashAttention 并不原生支持，而且树结构每一轮都在变化，无法提前编译。实践上通常用 <strong>Triton 实现的 FlashAttention 变体</strong>，先把 mask 预处理为 block 格式，然后在 kernel 里跳过全零 block 的计算——用稀疏性换效率。这是 DDTree 工程落地的关键之一。</p>

<h3 id="65-贪心走树一个具体例子">6.5 贪心走树：一个具体例子</h3>

<p>验证完成后，<strong>贪心沿最长接受路径走树</strong>：从根出发，每一步看当前节点的哪些孩子被目标模型接受了，如果接受了就往下走一层，遇到第一个被拒绝的节点就停止，把该路径上所有接受的 token 一次性提交。</p>

<p>举个例子，假设树如下：</p>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>         root
       /      \
    cat ✓    dog ✗
   /    \
sat ✓  lay ✗
  /
on ✗
</code></pre></div></div>

<p>走树过程：从 root 出发 → “cat” 被接受，进入 cat → “sat” 被接受，进入 sat → “on” 被拒绝，停止。本轮提交 <code class="language-plaintext highlighter-rouge">cat → sat</code>（2 个 token），并在 “on” 的位置让目标模型额外采一个 bonus token 作为下一轮草稿的起点。被拒绝节点的子树（比如 “dog” 下面挂着的那一整条分支）完全丢弃。</p>

<h3 id="66-kv-cache-怎么回收">6.6 KV Cache 怎么回收？</h3>

<p>这是树形投机解码最优雅的地方——<strong>根本不需要显式”删除”</strong>。</p>

<p>验证本质是一次 prefill，prefill 结束后只把接受路径上的节点 KV 追加到主序列的 KV Cache 末尾。拒绝节点的 KV 值虽然在 prefill 里算过，但<strong>从未被持久化</strong>到 KV Cache 里，它们只是 attention 计算的中间张量，离开 kernel 就自动废弃，无需任何显式释放。</p>

<h3 id="67-batch-推理的挑战">6.7 Batch 推理的挑战</h3>

<p>单请求的 DDTree 已经比较清楚，但工业落地必然要支持多请求并发。方案是把多个请求的树节点 concat 成一个长序列做一次 prefill，请求之间用 attention mask 完全隔离——A 请求的任何节点看不到 B 请求的任何节点。</p>

<p>DDTree 论文的实现聚焦单请求，<strong>真正生产级的多请求 batch 需要类似 vLLM 的调度层配合树形 attention kernel</strong>，包括不同请求之间 KV Cache 的 paged 管理、不同请求树预算的动态分配等。这是当前投机解码生产落地最主要的工程难点。</p>

<hr />

<h2 id="7-效果全面超越-dflash">7. 效果：全面超越 DFlash</h2>

<p>DDTree 在 60 个数据集-模型-温度组合上，全部优于原始 DFlash。以 Qwen3-8B 在 temperature=0 时为例：</p>

<table>
  <thead>
    <tr>
      <th>数据集</th>
      <th>DFlash</th>
      <th>DDTree</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>MATH-500</td>
      <td>5.56×</td>
      <td><strong>7.52×</strong></td>
    </tr>
    <tr>
      <td>HumanEval</td>
      <td>4.84×</td>
      <td><strong>6.90×</strong></td>
    </tr>
    <tr>
      <td>GSM8K</td>
      <td>4.78×</td>
      <td><strong>6.75×</strong></td>
    </tr>
    <tr>
      <td>AIME’24</td>
      <td>—</td>
      <td><strong>7.3×</strong></td>
    </tr>
    <tr>
      <td>AIME’25</td>
      <td>—</td>
      <td><strong>7.2×</strong></td>
    </tr>
  </tbody>
</table>

<p>更大的 30B 代码模型上，HumanEval 加速比达到 <strong>8.22×</strong>，接近自回归解码速度的 8 倍。</p>

<p>而且 DDTree 是<strong>完全无损的</strong>——所有 token 都经过目标模型验证，输出分布与原始自回归解码在数学上完全等价。论文代码基于 HuggingFace Transformers 开源实现，结果可复现。</p>

<hr />

<h2 id="8-总结一张图理解全套系统">8. 总结：一张图理解全套系统</h2>

<div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[自回归解码的困境]
  每次只生成 1 个 token → GPU 闲置 → 延迟高

[投机解码的思路]
  草稿模型快速提出候选 → 目标模型并行验证 → 数学上等价于原始分布

[接受/拒绝规则]
  以 min(1, p/q) 的概率接受草稿 token
  被拒绝时从残差分布 p'(x) 补采 → 无论如何输出都等于 p(x)

[DFlash 的创新]
  块扩散：一次 forward pass 输出整块的按位置分布
  比自回归草稿快，但每个位置只用了一个 token

[DDTree 的创新]
  利用 DFlash 的完整分布，在节点预算 B 内用 Best-First Heap 构建最优草稿树
  目标模型用 ancestor-only attention mask 一次验证整棵树
  贪心走树取最长接受路径 → 每轮接受更多 token → 速度提升 35%+
</code></pre></div></div>

<p>投机解码的整个故事，本质上是一次关于”如何把 GPU 的并行能力用足”的精妙设计。从朴素的串行解码，到草稿-验证的并行框架，再到树形结构对信息的充分利用，每一步都在回答同一个问题：<strong>怎样让大模型在不降低质量的前提下，尽可能快地生成文字</strong>。</p>

<hr />

<p><em>参考论文：Ringel &amp; Romano, “Accelerating Speculative Decoding with Block Diffusion Draft Trees”, arXiv 2604.12989, 2026.</em></p>]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><category term="blog" /><category term="llm" /><category term="speculative decoding" /><category term="inference acceleration" /><category term="diffusion" /><category term="machine learning" /><summary type="html"><![CDATA[本文从零开始，带你理解 LLM 推理加速的核心思路，读完之后你会明白：大模型为什么慢、投机解码如何加速、为什么加速后输出质量完全不变，以及 DDTree 这篇 2026 年的新论文究竟做了什么创新。]]></summary></entry><entry><title type="html">Python 利用selenium 控制浏览器自动提交表单</title><link href="https://liyongzhi.xyz/posts/2023/02/python-selenium-auto-form-submit/" rel="alternate" type="text/html" title="Python 利用selenium 控制浏览器自动提交表单" /><published>2023-02-07T00:00:00+08:00</published><updated>2023-02-07T00:00:00+08:00</updated><id>https://liyongzhi.xyz/posts/2023/02/blog</id><content type="html" xml:base="https://liyongzhi.xyz/posts/2023/02/python-selenium-auto-form-submit/"><![CDATA[<hr />

<hr />

<script async="" src="//busuanzi.ibruce.info/busuanzi/2.3/busuanzi.pure.mini.js"></script>

<div>
<div class="button01">
      <visited_a href="#" display:inline=""><span id="busuanzi_container_site_pv">你是第<span id="busuanzi_value_site_pv"></span>位访客~</span></visited_a>
      <visited_p class="top">٩(๑^o^๑)۶</visited_p>
      <visited_p class="bottom">Σ(っ °Д °;)っ被你发现了！</visited_p>
</div>
<img align="center" width="100" src="https://liyongzhi.xyz/images/static/take_me.gif" alt="" display:inline="" />
</div>

<hr />

<h2 id="python-利用selenium-自动控制浏览器提交表单">Python 利用selenium 自动控制浏览器提交表单</h2>

<h3 id="前期准备">前期准备</h3>

<ol>
  <li>下载安装chrome webdriver <code class="language-plaintext highlighter-rouge">https://sites.google.com/chromium.org/driver/downloads?authuser=0</code></li>
  <li>安装selenium <code class="language-plaintext highlighter-rouge">pip install seleuim</code></li>
</ol>

<h3 id="执行代码">执行代码</h3>

<ul>
  <li>如果有一些网站需要登录，可以执行以下命令启动一个常驻浏览器，并且将用户信息写到指定路径</li>
</ul>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c"># 针对macos 上的Chrome 浏览器</span>
<span class="nb">export </span><span class="nv">PATH</span><span class="o">=</span><span class="s2">"/Applications/Google Chrome.app/Contents/MacOS:</span><span class="nv">$PATH</span><span class="s2">"</span>
Google<span class="se">\ </span>Chrome <span class="nt">--remote-debugging-port</span><span class="o">=</span>9222 <span class="nt">--user-data-dir</span><span class="o">=</span><span class="s2">"~/ChromeProfile"</span>
</code></pre></div></div>

<ul>
  <li>经过以上操作就会启动一个浏览器，后续使用代码可以控制该浏览器上的行为。</li>
</ul>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="nn">selenium</span> <span class="kn">import</span> <span class="n">webdriver</span> <span class="c1"># selenium.__version__ = 4.8.0
</span><span class="kn">from</span> <span class="nn">selenium.webdriver.common.by</span> <span class="kn">import</span> <span class="n">By</span>
<span class="kn">from</span> <span class="nn">selenium.webdriver.common.action_chains</span> <span class="kn">import</span> <span class="n">ActionChains</span>
<span class="kn">import</span> <span class="nn">time</span>


<span class="s">'''
如果想要保持登录状态：
打开一个terminal执行以下命令：
export PATH="/Applications/Google Chrome.app/Contents/MacOS:$PATH"
Google\ Chrome --remote-debugging-port=9222 --user-data-dir="~/ChromeProfile"

启动chrome驻留之后再执行以下代码
'''</span>

<span class="c1"># 打开浏览器驱动
</span><span class="n">option</span> <span class="o">=</span> <span class="n">webdriver</span><span class="p">.</span><span class="n">ChromeOptions</span><span class="p">()</span>
<span class="n">option</span><span class="p">.</span><span class="n">add_experimental_option</span><span class="p">(</span><span class="s">"debuggerAddress"</span><span class="p">,</span> <span class="s">"127.0.0.1:9222"</span><span class="p">)</span>

<span class="c1"># 启动浏览器
</span><span class="n">driver</span> <span class="o">=</span> <span class="n">webdriver</span><span class="p">.</span><span class="n">Chrome</span><span class="p">(</span><span class="n">options</span> <span class="o">=</span> <span class="n">option</span><span class="p">)</span>
<span class="n">driver</span><span class="p">.</span><span class="n">implicitly_wait</span><span class="p">(</span><span class="mi">10</span><span class="p">)</span>

<span class="k">class</span> <span class="nc">ServiceConfig</span><span class="p">():</span>

    <span class="c1"># 定义prepareWork函数，做准备工作
</span>    <span class="k">def</span> <span class="nf">prepareWork</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span><span class="n">url</span><span class="p">):</span>
        <span class="c1"># 打开百度首页
</span>        <span class="n">driver</span><span class="p">.</span><span class="n">get</span><span class="p">(</span><span class="n">url</span><span class="p">)</span>

        <span class="c1"># 查找搜索框元素
</span>        <span class="n">search_input</span> <span class="o">=</span> <span class="n">driver</span><span class="p">.</span><span class="n">find_element</span><span class="p">(</span><span class="n">By</span><span class="p">.</span><span class="n">XPATH</span><span class="p">,</span><span class="s">'//*[@id="root"]/div[1]/div[2]/div/div[2]/div[1]/div/div/div/input'</span><span class="p">)</span>

        <span class="c1"># 在搜索框中输入文本
</span>        <span class="n">search_input</span><span class="p">.</span><span class="n">send_keys</span><span class="p">(</span><span class="s">"人体工学椅子"</span><span class="p">)</span>
        <span class="c1"># time.sleep(2)
</span>
        <span class="c1"># 点击搜索按钮
</span>        <span class="n">search_button</span> <span class="o">=</span> <span class="n">driver</span><span class="p">.</span><span class="n">find_element</span><span class="p">(</span><span class="n">By</span><span class="p">.</span><span class="n">XPATH</span><span class="p">,</span><span class="s">'//*[@id="root"]/div[1]/div[2]/div/div[2]/div[1]/div/button/span'</span><span class="p">)</span>
        <span class="n">ActionChains</span><span class="p">(</span><span class="n">driver</span><span class="p">).</span><span class="n">move_to_element</span><span class="p">(</span><span class="n">search_button</span><span class="p">).</span><span class="n">click</span><span class="p">().</span><span class="n">perform</span><span class="p">()</span>
        <span class="c1"># time.sleep(2)
</span>
        <span class="n">setting_button</span> <span class="o">=</span> <span class="n">driver</span><span class="p">.</span><span class="n">find_element</span><span class="p">(</span><span class="n">By</span><span class="p">.</span><span class="n">XPATH</span><span class="p">,</span><span class="s">'//*[@id="root"]/div[2]/div/div[3]/div/div/div/div[1]/div[2]/p[1]'</span><span class="p">)</span>
        <span class="n">ActionChains</span><span class="p">(</span><span class="n">driver</span><span class="p">).</span><span class="n">move_to_element</span><span class="p">(</span><span class="n">setting_button</span><span class="p">).</span><span class="n">click</span><span class="p">().</span><span class="n">perform</span><span class="p">()</span>
        <span class="c1"># time.sleep(2)
</span>        <span class="n">windows</span> <span class="o">=</span> <span class="n">driver</span><span class="p">.</span><span class="n">window_handles</span>
        <span class="n">driver</span><span class="p">.</span><span class="n">switch_to</span><span class="p">.</span><span class="n">window</span><span class="p">(</span><span class="n">windows</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">])</span>

        <span class="n">sousuo_setting</span> <span class="o">=</span> <span class="n">driver</span><span class="p">.</span><span class="n">find_element</span><span class="p">(</span><span class="n">By</span><span class="p">.</span><span class="n">XPATH</span><span class="p">,</span><span class="s">'//*[@id="root"]/div[2]/div/div[3]/div[2]/div[5]/span[1]/button[1]/span'</span><span class="p">)</span>
        <span class="n">ActionChains</span><span class="p">(</span><span class="n">driver</span><span class="p">).</span><span class="n">move_to_element</span><span class="p">(</span><span class="n">sousuo_setting</span><span class="p">).</span><span class="n">click</span><span class="p">().</span><span class="n">perform</span><span class="p">()</span>



<span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">'__main__'</span><span class="p">:</span>
    <span class="n">url</span> <span class="o">=</span> <span class="s">'https://www.byte-mall.cn/'</span>
    <span class="n">sc</span> <span class="o">=</span> <span class="n">ServiceConfig</span><span class="p">()</span>
    <span class="n">sc</span><span class="p">.</span><span class="n">prepareWork</span><span class="p">(</span><span class="n">url</span><span class="p">)</span>
    <span class="n">time</span><span class="p">.</span><span class="n">sleep</span><span class="p">(</span><span class="mi">10000</span><span class="p">)</span>
    

</code></pre></div></div>

<ul>
  <li>具体其他复杂操作可以通过组合鼠标和键盘操作事件来实现</li>
</ul>

<div data-hk-top-pages="5"> 

</div>]]></content><author><name>李勇志 (Yongzhi Li)</name><email>yongzhili@pku.edu.cn</email></author><category term="blog" /><category term="python, tips" /><summary type="html"><![CDATA[]]></summary></entry></feed>