<?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://wqh011128.github.io/feed.xml" rel="self" type="application/atom+xml" /><link href="https://wqh011128.github.io/" rel="alternate" type="text/html" /><updated>2026-06-11T01:29:19+00:00</updated><id>https://wqh011128.github.io/feed.xml</id><title type="html">Qihang Wu</title><subtitle>Training and inference framework engineer focused on LLM and VLM systems, with earlier research on domain adaptation for semantic segmentation and efficient attention mechanisms.</subtitle><entry><title type="html">Pallas TPU Kernel 写法与优化文档</title><link href="https://wqh011128.github.io/blog/2026/06/10/Pallas_Kernel%E5%AD%A6%E4%B9%A0%E6%96%87%E6%A1%A3.html" rel="alternate" type="text/html" title="Pallas TPU Kernel 写法与优化文档" /><published>2026-06-10T00:00:00+00:00</published><updated>2026-06-10T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/06/10/Pallas_Kernel%E5%AD%A6%E4%B9%A0%E6%96%87%E6%A1%A3</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/06/10/Pallas_Kernel%E5%AD%A6%E4%B9%A0%E6%96%87%E6%A1%A3.html"><![CDATA[<h1 id="pallas-tpu-kernel-写法与优化文档">Pallas TPU Kernel 写法与优化文档</h1>

<h2 id="0-适用范围">0. 适用范围</h2>

<p>本文总结 Pallas TPU kernel 的常用生产级写法、设计规范与优化思路，重点面向以下类型的 kernel：</p>

<ul>
  <li>Attention / FlashAttention / GQA / MQA；</li>
  <li>Matmul-heavy backward kernel；</li>
  <li>带 reduction 的 block kernel；</li>
  <li>需要在 TPU MXU 上获得较高利用率的 Pallas kernel；</li>
  <li>需要控制 HBM、VMEM、register、scratch、tile shape 的性能敏感 kernel。</li>
</ul>

<p>本文尤其关注：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 如何设计 grid
2. 如何设计 BlockSpec
3. 如何使用 scratch 做 partial accumulation
4. 如何选择 major/minor tile
5. 如何避免单个 program 内部 loop 太长
6. 如何提升 MXU utilization
7. 如何 benchmark 和调 block size
8. 如何写出更接近生产级的 Pallas kernel
</code></pre></div></div>

<hr />

<h1 id="1-pallas-kernel-的基本心智模型">1. Pallas Kernel 的基本心智模型</h1>

<p>写 Pallas kernel 时，建议始终从下面几个概念出发。</p>

<h2 id="11-program">1.1 Program</h2>

<p>一个 Pallas <code class="language-plaintext highlighter-rouge">program</code> 可以理解成一个小的 tile 程序实例。</p>

<p>它由 <code class="language-plaintext highlighter-rouge">grid</code> 决定数量，例如：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">grid</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">batch_blocks</span><span class="p">,</span>
    <span class="n">num_heads</span><span class="p">,</span>
    <span class="n">q_blocks</span><span class="p">,</span>
    <span class="n">kv_blocks</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>每个 program 会通过：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="n">axis</span><span class="p">)</span>
</code></pre></div></div>

<p>拿到当前自己负责的 tile index。</p>

<p>例如：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">q_block_index</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span>
<span class="n">kv_block_index</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">3</span><span class="p">)</span>
</code></pre></div></div>

<p>可以理解为：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>一个 program 负责一个局部 tile。
</code></pre></div></div>

<hr />

<h2 id="12-grid">1.2 Grid</h2>

<p><code class="language-plaintext highlighter-rouge">grid</code> 决定有多少个 program 被调度。</p>

<p>一个好的 grid 应该：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 暴露足够并行度；
2. 把重要的 block 维度显式暴露给 compiler；
3. 避免把过长的 reduction loop 藏在单个 program 里；
4. 让每个 program 的工作量适中；
5. 让输出 tile 和 reduction tile 的关系清楚。
</code></pre></div></div>

<p>对于 attention backward，常见 grid 设计是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DQ:
  grid = batch × head × q_block × kv_block

DKV:
  grid = batch × head_or_kv_head × kv_block × q_block
</code></pre></div></div>

<hr />

<h2 id="13-blockspec">1.3 BlockSpec</h2>

<p><code class="language-plaintext highlighter-rouge">pl.BlockSpec</code> 决定每个 program 看到输入/输出 tensor 的哪一个 tile。</p>

<p>典型形式：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">q_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">q_index_map</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>其中：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_shape:
  当前 program 看到的局部 tile shape

index_map:
  从 program_id 映射到全局 tensor block index
</code></pre></div></div>

<p>生产级原则：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>BlockSpec 应该尽量描述局部 tile，而不是 full sequence。
</code></pre></div></div>

<p>例如，尽量避免：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">q_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">q_seq_len</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">q_full_index_map</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>因为这意味着一个 program 对 Q 的 sequence 维几乎没有 tile 化。</p>

<p>更推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">q_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q_major</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">q_index_map</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<hr />

<h2 id="14-mxu">1.4 MXU</h2>

<p>TPU MXU 是主要做矩阵乘法的硬件单元。</p>

<p>Pallas kernel 想跑得快，通常要做到：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. dot_general / dot 的 shape 规整；
2. matmul tile 足够大；
3. matmul 之间不要被太多 elementwise / mask / exp 阻塞；
4. 数据能及时 load 到 VMEM；
5. 不要让 MXU 等待数据或等待复杂控制流；
6. 不要让单个 program 的 live range 太长。
</code></pre></div></div>

<p>如果 MXU utilization 很低，例如只有 13%，通常说明：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>MXU 大量时间没有被有效喂饱。
</code></pre></div></div>

<p>常见原因包括：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. grid 并行度不足；
2. 单个 program 内部 loop 太长；
3. temporary tensor 太大；
4. VMEM/register 压力过高；
5. BlockSpec 过大；
6. dot tile shape 不合适；
7. matmul 中间夹杂过多 exp/mask/elementwise；
8. causal 上三角无效 block 没有跳过。
</code></pre></div></div>

<hr />

<h1 id="2-写-kernel-前必须先做的设计">2. 写 Kernel 前必须先做的设计</h1>

<p>不要一上来就写 Pallas kernel。生产级 kernel 应先写清楚：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 数学公式；
2. 输入输出 shape；
3. 哪些维度是并行维；
4. 哪些维度是 reduction 维；
5. 是否需要跨 block 累加；
6. 是否需要 scratch；
7. 是否需要保存 residuals；
8. 是否需要 mask/bias/causal/segment；
9. 是否需要支持 padding；
10. 是否需要支持 GQA/MQA。
</code></pre></div></div>

<hr />

<h2 id="21-先写数学公式">2.1 先写数学公式</h2>

<p>以 attention backward 为例：</p>

<p>前向：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>S = QK^T
P = softmax(S)
O = P V
</code></pre></div></div>

<p>反向：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dV = P^T dO
dP = dO V^T
D  = sum(O * dO, axis=-1)
dS = P * (dP - D)
dQ = dS K
dK = dS^T Q
</code></pre></div></div>

<p>如果有 scale：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>S = QK^T * sm_scale
</code></pre></div></div>

<p>则：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dS_unscaled = dS_scaled * sm_scale
</code></pre></div></div>

<p>如果有 additive attention bias：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>S = QK^T + AB
</code></pre></div></div>

<p>则：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dAB = dS
</code></pre></div></div>

<hr />

<h2 id="22-明确输出决定-kernel-方向">2.2 明确输出决定 kernel 方向</h2>

<p>对于 backward，常见拆法是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DQ kernel:
  输出 dQ
  沿 KV 维 reduction

DKV kernel:
  输出 dK, dV
  沿 Q 维 reduction
</code></pre></div></div>

<p>因此：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DQ 的 grid 应该优先围绕 Q output tile 设计；
DKV 的 grid 应该优先围绕 KV output tile 设计。
</code></pre></div></div>

<hr />

<h1 id="3-grid-设计规范">3. Grid 设计规范</h1>

<h2 id="31-不要把长-reduction-全塞进单个-program">3.1 不要把长 reduction 全塞进单个 program</h2>

<p>不推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">grid</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">batch_blocks</span><span class="p">,</span>
    <span class="n">num_heads</span><span class="p">,</span>
    <span class="n">q_blocks</span><span class="p">,</span>
<span class="p">)</span>

<span class="k">for</span> <span class="n">kv_start</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">kv_seq_len</span><span class="p">,</span> <span class="n">block_kv</span><span class="p">):</span>
    <span class="p">...</span>
</code></pre></div></div>

<p>因为这意味着：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>一个 program 负责一个 q_block，
然后在 program 内部扫完整 kv_seq。
</code></pre></div></div>

<p>如果 <code class="language-plaintext highlighter-rouge">kv_seq_len</code> 很长，program 内部 loop 会很长。</p>

<p>更推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">grid</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">batch_blocks</span><span class="p">,</span>
    <span class="n">num_heads</span><span class="p">,</span>
    <span class="n">q_blocks</span><span class="p">,</span>
    <span class="n">kv_blocks</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>让 <code class="language-plaintext highlighter-rouge">kv_blocks</code> 成为 grid 维度。</p>

<hr />

<h2 id="32-dq-推荐-grid">3.2 DQ 推荐 grid</h2>

<p>DQ 的公式：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dQ = sum_over_KV dS @ K
</code></pre></div></div>

<p>推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">grid_dq</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">batch_blocks</span><span class="p">,</span>
    <span class="n">num_kv_heads</span><span class="p">,</span>
    <span class="n">q_seq_len</span> <span class="o">//</span> <span class="n">block_q_major_dq</span><span class="p">,</span>
    <span class="n">kv_seq_len</span> <span class="o">//</span> <span class="n">block_kv_major_dq</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>含义：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>axis 0: batch block
axis 1: kv head
axis 2: q major block
axis 3: kv major block
</code></pre></div></div>

<p>如果是 GQA：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>一个 kv_head 对应 group_size 个 q_heads。
</code></pre></div></div>

<p>所以 Q tile 可以是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[block_b, group_size, block_q_major_dq, head_dim]
</code></pre></div></div>

<hr />

<h2 id="33-dkv-推荐-grid">3.3 DKV 推荐 grid</h2>

<p>DKV 的公式：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dK = sum_over_Q dS^T @ Q
dV = sum_over_Q P^T @ dO
</code></pre></div></div>

<p>推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">grid_dkv</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">batch_blocks</span><span class="p">,</span>
    <span class="n">num_kv_heads</span><span class="p">,</span>
    <span class="n">kv_seq_len</span> <span class="o">//</span> <span class="n">block_kv_major_dkv</span><span class="p">,</span>
    <span class="n">q_seq_len</span> <span class="o">//</span> <span class="n">block_q_major_dkv</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>含义：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>axis 0: batch block
axis 1: kv head
axis 2: kv major block
axis 3: q major block
</code></pre></div></div>

<p>这样每个 program 负责：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>一个 KV output tile
一个 Q reduction tile
</code></pre></div></div>

<p>再通过 scratch 跨 Q block 累加到最终 dK/dV。</p>

<hr />

<h2 id="34-哪些维度放-grid哪些维度放-loop">3.4 哪些维度放 grid，哪些维度放 loop？</h2>

<p>经验规则：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>短 loop 可以放 program 内部；
长 reduction loop 尽量暴露到 grid；
需要跨 block 累加时，用 scratch；
如果 scratch/ordering 很难处理，再考虑内部 loop。
</code></pre></div></div>

<p>推荐：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>batch/head/q_block/kv_block:
  通常适合放 grid

minor tile:
  通常适合放 program 内部 loop

完整 seq_len:
  不建议放单个 program 内部扫完
</code></pre></div></div>

<hr />

<h1 id="4-scratch-设计规范">4. Scratch 设计规范</h1>

<h2 id="41-scratch-的作用">4.1 Scratch 的作用</h2>

<p>Scratch 不是为了“直接加速”，而是为了允许：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>把 reduction 维拆到 grid
+
在 VMEM 中保存 partial accumulation
+
避免 atomic
+
最后一个 reduction block 写出结果
</code></pre></div></div>

<p>如果没有 scratch，通常只能让一个 program 扫完整 reduction 维。</p>

<p>这容易造成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 单个 program 太重；
2. 内部 loop 太长；
3. compiler 调度困难；
4. live range 变长；
5. MXU 利用率低。
</code></pre></div></div>

<hr />

<h2 id="42-dq-scratch-模式">4.2 DQ scratch 模式</h2>

<p>DQ 需要沿 KV 维累加：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dQ = sum_over_KV dS @ K
</code></pre></div></div>

<p>所以 DQ scratch 可以设计为：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">dq_scratch</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">VMEM</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q_major_dq</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>逻辑：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">kv_block_index</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">3</span><span class="p">)</span>

<span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">kv_block_index</span> <span class="o">==</span> <span class="mi">0</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">init</span><span class="p">():</span>
    <span class="n">dq_scratch</span><span class="p">[...]</span> <span class="o">=</span> <span class="mi">0</span>

<span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">should_run</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">run</span><span class="p">():</span>
    <span class="n">dq_scratch</span><span class="p">[...]</span> <span class="o">+=</span> <span class="n">partial_dq</span>

<span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">kv_block_index</span> <span class="o">==</span> <span class="n">last_kv_block</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">store</span><span class="p">():</span>
    <span class="n">dq_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">dq_scratch</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">dq_ref</span><span class="p">.</span><span class="n">dtype</span><span class="p">)</span>
</code></pre></div></div>

<hr />

<h2 id="43-dkv-scratch-模式">4.3 DKV scratch 模式</h2>

<p>DKV 需要沿 Q 维累加：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dK = sum_over_Q dS^T @ Q
dV = sum_over_Q P^T @ dO
</code></pre></div></div>

<p>所以 DKV scratch 可以设计为：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">dk_scratch</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">VMEM</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">block_kv_major_dkv</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">dv_scratch</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">VMEM</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">block_kv_major_dkv</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>逻辑：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">q_block_index</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">3</span><span class="p">)</span>

<span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">q_block_index</span> <span class="o">==</span> <span class="mi">0</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">init</span><span class="p">():</span>
    <span class="n">dk_scratch</span><span class="p">[...]</span> <span class="o">=</span> <span class="mi">0</span>
    <span class="n">dv_scratch</span><span class="p">[...]</span> <span class="o">=</span> <span class="mi">0</span>

<span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">should_run</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">run</span><span class="p">():</span>
    <span class="n">dk_scratch</span><span class="p">[...]</span> <span class="o">+=</span> <span class="n">partial_dk</span>
    <span class="n">dv_scratch</span><span class="p">[...]</span> <span class="o">+=</span> <span class="n">partial_dv</span>

<span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">q_block_index</span> <span class="o">==</span> <span class="n">last_q_block</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">store</span><span class="p">():</span>
    <span class="n">dk_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">dk_scratch</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">dk_ref</span><span class="p">.</span><span class="n">dtype</span><span class="p">)</span>
    <span class="n">dv_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">dv_scratch</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">dv_ref</span><span class="p">.</span><span class="n">dtype</span><span class="p">)</span>
</code></pre></div></div>

<hr />

<h2 id="44-使用-scratch-的注意事项">4.4 使用 scratch 的注意事项</h2>

<p>使用 scratch 跨 grid 维累加时，要注意：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. reduction grid 维必须有明确执行顺序；
2. 不能在完全 parallel 的 grid 维上假设先后顺序；
3. 初始化、累加、写出条件必须严格正确；
4. scratch shape 不要过大；
5. scratch 通常用 fp32 保证累加精度；
6. 最终写出时再 cast 到输出 dtype。
</code></pre></div></div>

<p>如果依赖某个 grid 维顺序推进，需要在 TPU compiler params 中正确设置该维的语义，例如将 reduction 维设置为 <code class="language-plaintext highlighter-rouge">"arbitrary"</code>，而不是所有维都当成完全 parallel。</p>

<hr />

<h1 id="5-blockspec-设计规范">5. BlockSpec 设计规范</h1>

<h2 id="51-blockspec-应该-tile-化">5.1 BlockSpec 应该 tile 化</h2>

<p>推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">q_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q_major</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">q_index_map</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>不推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">q_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">q_seq_len</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">q_full_index_map</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>因为后者意味着：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>对 q_seq_len 这个维度几乎没有 tile 化。
</code></pre></div></div>

<p>这会让一个 program 看到完整 sequence，容易导致：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. input window 过大；
2. prefetch 粒度过粗；
3. compiler 难以优化；
4. live range 变长；
5. VMEM/register 压力变大；
6. MXU pipeline 不稳定。
</code></pre></div></div>

<hr />

<h2 id="52-dq-blockspec-示例">5.2 DQ BlockSpec 示例</h2>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">q_index_map</span><span class="p">(</span><span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">q_block_index</span><span class="p">,</span> <span class="n">kv_block_index</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">q_block_index</span><span class="p">,</span> <span class="mi">0</span>

<span class="k">def</span> <span class="nf">kv_index_map</span><span class="p">(</span><span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">q_block_index</span><span class="p">,</span> <span class="n">kv_block_index</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">kv_block_index</span><span class="p">,</span> <span class="mi">0</span>

<span class="n">q_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q_major_dq</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">q_index_map</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">kv_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="n">block_kv_major_dq</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">kv_index_map</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">lse_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q_major_dq</span><span class="p">,</span> <span class="n">MIN_BLOCK_SIZE</span><span class="p">),</span>
    <span class="n">q_index_map</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<hr />

<h2 id="53-dkv-blockspec-示例">5.3 DKV BlockSpec 示例</h2>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">q_index_map</span><span class="p">(</span><span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">kv_block_index</span><span class="p">,</span> <span class="n">q_block_index</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">q_block_index</span><span class="p">,</span> <span class="mi">0</span>

<span class="k">def</span> <span class="nf">kv_index_map</span><span class="p">(</span><span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">kv_block_index</span><span class="p">,</span> <span class="n">q_block_index</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">kv_block_index</span><span class="p">,</span> <span class="mi">0</span>

<span class="n">q_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q_major_dkv</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">q_index_map</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">kv_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="n">block_kv_major_dkv</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">kv_index_map</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">lse_spec</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q_major_dkv</span><span class="p">,</span> <span class="n">MIN_BLOCK_SIZE</span><span class="p">),</span>
    <span class="n">q_index_map</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<hr />

<h2 id="54-skipped-block-的-safe-index">5.4 Skipped block 的 safe index</h2>

<p>如果 causal 下某些 block 完全无效，Pallas 仍可能需要 index_map 返回合法 tile。</p>

<p>可以使用 safe index：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">kv_index_map</span><span class="p">(</span><span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">q_block_index</span><span class="p">,</span> <span class="n">kv_block_index</span><span class="p">):</span>
    <span class="n">should_run</span> <span class="o">=</span> <span class="n">below_or_on_diag</span><span class="p">(</span>
        <span class="n">q_block_index</span><span class="p">,</span>
        <span class="n">block_q_major</span><span class="p">,</span>
        <span class="n">kv_block_index</span><span class="p">,</span>
        <span class="n">block_kv_major</span><span class="p">,</span>
    <span class="p">)</span>
    <span class="n">safe_kv_index</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">select</span><span class="p">(</span><span class="n">should_run</span><span class="p">,</span> <span class="n">kv_block_index</span><span class="p">,</span> <span class="mi">0</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">batch_index</span><span class="p">,</span> <span class="n">kv_head_index</span><span class="p">,</span> <span class="n">safe_kv_index</span><span class="p">,</span> <span class="mi">0</span>
</code></pre></div></div>

<p>这样可以避免无效或越界 prefetch。</p>

<hr />

<h1 id="6-major--minor-tile-分层">6. Major / Minor Tile 分层</h1>

<h2 id="61-为什么要分-majorminor">6.1 为什么要分 major/minor？</h2>

<p>不要只用一层：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q
block_kv
</code></pre></div></div>

<p>更推荐两层：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>major block:
  grid 级别 tile

minor block:
  program 内部 dot tile
</code></pre></div></div>

<p>好处：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. grid 粒度和 dot 粒度分开调；
2. 更容易平衡并行度和数据复用；
3. 更容易控制 VMEM/register 压力；
4. 更容易找到 MXU 友好的 matmul shape；
5. 更容易做 block size sweep；
6. 更接近成熟 FlashAttention kernel 的写法。
</code></pre></div></div>

<hr />

<h2 id="62-dq-tile-参数">6.2 DQ tile 参数</h2>

<p>推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">block_q_major_dq</span>
<span class="n">block_kv_major_dq</span>
<span class="n">block_kv_dq</span>
</code></pre></div></div>

<p>含义：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q_major_dq:
  DQ kernel 一个 q grid block 覆盖的 Q 长度

block_kv_major_dq:
  DQ kernel 一个 kv grid block 覆盖的 KV 长度

block_kv_dq:
  DQ kernel 内部每次 dot 的 KV minor tile
</code></pre></div></div>

<p>示例：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">block_q_major_dq</span> <span class="o">=</span> <span class="mi">128</span>
<span class="n">block_kv_major_dq</span> <span class="o">=</span> <span class="mi">256</span>
<span class="n">block_kv_dq</span> <span class="o">=</span> <span class="mi">128</span>
</code></pre></div></div>

<p>表示：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>一个 DQ program 覆盖 128 个 Q token 和 256 个 KV token，
但内部每次只用 128 个 KV token 做一次 dot。
</code></pre></div></div>

<p>伪代码：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">for</span> <span class="n">kv_minor</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">block_kv_major_dq</span><span class="p">,</span> <span class="n">block_kv_dq</span><span class="p">):</span>
    <span class="n">k</span> <span class="o">=</span> <span class="n">k_tile</span><span class="p">[:,</span> <span class="p">:,</span> <span class="n">kv_minor</span> <span class="p">:</span> <span class="n">kv_minor</span> <span class="o">+</span> <span class="n">block_kv_dq</span><span class="p">,</span> <span class="p">:]</span>
    <span class="n">v</span> <span class="o">=</span> <span class="n">v_tile</span><span class="p">[:,</span> <span class="p">:,</span> <span class="n">kv_minor</span> <span class="p">:</span> <span class="n">kv_minor</span> <span class="o">+</span> <span class="n">block_kv_dq</span><span class="p">,</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">T</span>
    <span class="p">...</span>
    <span class="n">dq_scratch</span> <span class="o">+=</span> <span class="n">ds</span> <span class="o">@</span> <span class="n">k</span>
</code></pre></div></div>

<hr />

<h2 id="63-dkv-tile-参数">6.3 DKV tile 参数</h2>

<p>推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">block_q_major_dkv</span>
<span class="n">block_kv_major_dkv</span>
<span class="n">block_q_dkv</span>
<span class="n">block_kv_dkv</span>
</code></pre></div></div>

<p>含义：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q_major_dkv:
  DKV kernel 一个 q grid block 覆盖的 Q 长度

block_kv_major_dkv:
  DKV kernel 一个 kv grid block 覆盖的 KV 长度

block_q_dkv:
  DKV kernel 内部每次处理的 Q minor tile

block_kv_dkv:
  DKV kernel 内部每次处理的 KV minor tile
</code></pre></div></div>

<p>示例：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">block_q_major_dkv</span> <span class="o">=</span> <span class="mi">128</span>
<span class="n">block_kv_major_dkv</span> <span class="o">=</span> <span class="mi">256</span>
<span class="n">block_q_dkv</span> <span class="o">=</span> <span class="mi">128</span>
<span class="n">block_kv_dkv</span> <span class="o">=</span> <span class="mi">128</span>
</code></pre></div></div>

<p>表示：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>一个 DKV program 覆盖 128 个 Q token 和 256 个 KV token，
但内部把 KV 再拆成两个 128 的 minor tile。
</code></pre></div></div>

<p>伪代码：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">for</span> <span class="n">q_minor</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">block_q_major_dkv</span><span class="p">,</span> <span class="n">block_q_dkv</span><span class="p">):</span>
    <span class="k">for</span> <span class="n">kv_minor</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">block_kv_major_dkv</span><span class="p">,</span> <span class="n">block_kv_dkv</span><span class="p">):</span>
        <span class="n">q</span> <span class="o">=</span> <span class="n">q_tile</span><span class="p">[:,</span> <span class="p">:,</span> <span class="n">q_minor</span> <span class="p">:</span> <span class="n">q_minor</span> <span class="o">+</span> <span class="n">block_q_dkv</span><span class="p">,</span> <span class="p">:]</span>
        <span class="n">k</span> <span class="o">=</span> <span class="n">k_tile</span><span class="p">[:,</span> <span class="p">:,</span> <span class="n">kv_minor</span> <span class="p">:</span> <span class="n">kv_minor</span> <span class="o">+</span> <span class="n">block_kv_dkv</span><span class="p">,</span> <span class="p">:]</span>
        <span class="n">v</span> <span class="o">=</span> <span class="n">v_tile</span><span class="p">[:,</span> <span class="p">:,</span> <span class="n">kv_minor</span> <span class="p">:</span> <span class="n">kv_minor</span> <span class="o">+</span> <span class="n">block_kv_dkv</span><span class="p">,</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">T</span>
        <span class="p">...</span>
        <span class="n">dk_scratch</span><span class="p">[</span><span class="n">kv_minor</span><span class="p">]</span> <span class="o">+=</span> <span class="n">ds</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">q</span>
        <span class="n">dv_scratch</span><span class="p">[</span><span class="n">kv_minor</span><span class="p">]</span> <span class="o">+=</span> <span class="n">p</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">do</span>
</code></pre></div></div>

<hr />

<h1 id="7-causal-mask-优化规范">7. Causal Mask 优化规范</h1>

<h2 id="71-区分-block-level-mask-和-element-level-mask">7.1 区分 block-level mask 和 element-level mask</h2>

<p>Element-level causal mask：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">scores</span> <span class="o">=</span> <span class="n">scores</span> <span class="o">+</span> <span class="n">jnp</span><span class="p">.</span><span class="n">where</span><span class="p">(</span><span class="n">mask</span><span class="p">,</span> <span class="mf">0.0</span><span class="p">,</span> <span class="n">DEFAULT_MASK_VALUE</span><span class="p">)</span>
</code></pre></div></div>

<p>这是必要的，因为对角线附近的 partial block 仍需要逐元素 mask。</p>

<p>但对于完全在 causal 上三角的 block，应该 block-level skip：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>如果 q_block_end &lt; kv_block_start，
则整个 block 无效，不需要做 QK^T。
</code></pre></div></div>

<hr />

<h2 id="72-block-level-skip">7.2 Block-level skip</h2>

<p>推荐函数：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">below_or_on_diag</span><span class="p">(</span><span class="n">q_block_index</span><span class="p">,</span> <span class="n">block_q</span><span class="p">,</span> <span class="n">kv_block_index</span><span class="p">,</span> <span class="n">block_kv</span><span class="p">):</span>
    <span class="n">q_block_end</span> <span class="o">=</span> <span class="p">(</span><span class="n">q_block_index</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)</span> <span class="o">*</span> <span class="n">block_q</span> <span class="o">-</span> <span class="mi">1</span>
    <span class="n">kv_block_start</span> <span class="o">=</span> <span class="n">kv_block_index</span> <span class="o">*</span> <span class="n">block_kv</span>
    <span class="k">return</span> <span class="n">q_block_end</span> <span class="o">&gt;=</span> <span class="n">kv_block_start</span>
</code></pre></div></div>

<p>DQ 中：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">should_run</span> <span class="o">=</span> <span class="n">below_or_on_diag</span><span class="p">(</span>
    <span class="n">q_block_index</span><span class="p">,</span>
    <span class="n">block_q_major_dq</span><span class="p">,</span>
    <span class="n">kv_block_index</span><span class="p">,</span>
    <span class="n">block_kv_major_dq</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>DKV 中：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">should_run</span> <span class="o">=</span> <span class="n">below_or_on_diag</span><span class="p">(</span>
    <span class="n">q_block_index</span><span class="p">,</span>
    <span class="n">block_q_major_dkv</span><span class="p">,</span>
    <span class="n">kv_block_index</span><span class="p">,</span>
    <span class="n">block_kv_major_dkv</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>然后：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">should_run</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">run</span><span class="p">():</span>
    <span class="p">...</span>
</code></pre></div></div>

<hr />

<h2 id="73-为什么-block-level-skip-很重要">7.3 为什么 block-level skip 很重要？</h2>

<p>对于 causal prefill，理论上上三角大约一半 attention block 无效。</p>

<p>如果不 skip，而是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><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">T</span>
<span class="n">scores</span> <span class="o">=</span> <span class="n">apply_causal_mask</span><span class="p">(</span><span class="n">scores</span><span class="p">)</span>
</code></pre></div></div>

<p>则无效 block 仍然消耗：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. QK^T matmul
2. exp
3. dp
4. ds
5. dQ/dK/dV partial accumulation
</code></pre></div></div>

<p>所以 block-level skip 可以直接降低 latency。</p>

<hr />

<h1 id="8-attention-backward-常用生产级模板">8. Attention Backward 常用生产级模板</h1>

<h2 id="81-dq-kernel-模板">8.1 DQ kernel 模板</h2>

<p>数学：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dQ = sum_over_KV dS @ K
</code></pre></div></div>

<p>结构：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">dq_kernel</span><span class="p">(</span>
    <span class="n">q_ref</span><span class="p">,</span>
    <span class="n">k_ref</span><span class="p">,</span>
    <span class="n">v_ref</span><span class="p">,</span>
    <span class="n">o_ref</span><span class="p">,</span>
    <span class="n">lse_ref</span><span class="p">,</span>
    <span class="n">do_ref</span><span class="p">,</span>
    <span class="n">dq_ref</span><span class="p">,</span>
    <span class="n">dq_scratch_ref</span><span class="p">,</span>
    <span class="o">*</span><span class="p">,</span>
    <span class="n">sm_scale</span><span class="p">,</span>
    <span class="n">block_kv_minor</span><span class="p">,</span>
    <span class="n">kv_seq_len</span><span class="p">,</span>
<span class="p">):</span>
    <span class="n">q_block_index</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span>
    <span class="n">kv_block_index</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">3</span><span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">kv_block_index</span> <span class="o">==</span> <span class="mi">0</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">init</span><span class="p">():</span>
        <span class="n">dq_scratch_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">zeros_like</span><span class="p">(</span><span class="n">dq_scratch_ref</span><span class="p">)</span>

    <span class="n">should_run</span> <span class="o">=</span> <span class="n">below_or_on_diag</span><span class="p">(</span>
        <span class="n">q_block_index</span><span class="p">,</span>
        <span class="n">block_q_major_dq</span><span class="p">,</span>
        <span class="n">kv_block_index</span><span class="p">,</span>
        <span class="n">block_kv_major_dq</span><span class="p">,</span>
    <span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">should_run</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">run</span><span class="p">():</span>
        <span class="n">q</span> <span class="o">=</span> <span class="n">q_ref</span><span class="p">[...]</span>
        <span class="n">o</span> <span class="o">=</span> <span class="n">o_ref</span><span class="p">[...]</span>
        <span class="n">do</span> <span class="o">=</span> <span class="n">do_ref</span><span class="p">[...]</span>
        <span class="n">lse</span> <span class="o">=</span> <span class="n">lse_ref</span><span class="p">[...]</span>

        <span class="n">di</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span>
            <span class="n">o</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span> <span class="o">*</span> <span class="n">do</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
            <span class="n">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span>
        <span class="p">)</span>

        <span class="k">for</span> <span class="n">kv_minor</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">block_kv_major_dq</span><span class="p">,</span> <span class="n">block_kv_minor</span><span class="p">):</span>
            <span class="n">k</span> <span class="o">=</span> <span class="n">k_ref</span><span class="p">[...,</span> <span class="n">kv_minor</span> <span class="p">:</span> <span class="n">kv_minor</span> <span class="o">+</span> <span class="n">block_kv_minor</span><span class="p">,</span> <span class="p">:]</span>
            <span class="n">v</span> <span class="o">=</span> <span class="n">v_ref</span><span class="p">[...,</span> <span class="n">kv_minor</span> <span class="p">:</span> <span class="n">kv_minor</span> <span class="o">+</span> <span class="n">block_kv_minor</span><span class="p">,</span> <span class="p">:]</span>

            <span class="n">scores</span> <span class="o">=</span> <span class="n">dot</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">T</span><span class="p">)</span>
            <span class="n">scores</span> <span class="o">*=</span> <span class="n">sm_scale</span>
            <span class="n">scores</span> <span class="o">=</span> <span class="n">apply_causal_mask_if_needed</span><span class="p">(</span><span class="n">scores</span><span class="p">)</span>

            <span class="n">p</span> <span class="o">=</span> <span class="n">exp</span><span class="p">(</span><span class="n">scores</span> <span class="o">-</span> <span class="n">lse</span><span class="p">)</span>
            <span class="n">dp</span> <span class="o">=</span> <span class="n">dot</span><span class="p">(</span><span class="n">do</span><span class="p">,</span> <span class="n">v</span><span class="p">.</span><span class="n">T</span><span class="p">)</span>
            <span class="n">ds</span> <span class="o">=</span> <span class="p">(</span><span class="n">dp</span> <span class="o">-</span> <span class="n">di</span><span class="p">[:,</span> <span class="bp">None</span><span class="p">])</span> <span class="o">*</span> <span class="n">p</span>
            <span class="n">ds</span> <span class="o">*=</span> <span class="n">sm_scale</span>

            <span class="n">dq_scratch_ref</span><span class="p">[...]</span> <span class="o">+=</span> <span class="n">dot</span><span class="p">(</span><span class="n">ds</span><span class="p">,</span> <span class="n">k</span><span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">kv_block_index</span> <span class="o">==</span> <span class="n">kv_seq_len</span> <span class="o">//</span> <span class="n">block_kv_major_dq</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">store</span><span class="p">():</span>
        <span class="n">dq_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">dq_scratch_ref</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">dq_ref</span><span class="p">.</span><span class="n">dtype</span><span class="p">)</span>
</code></pre></div></div>

<hr />

<h2 id="82-dkv-kernel-模板">8.2 DKV kernel 模板</h2>

<p>数学：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dK = sum_over_Q dS^T @ Q
dV = sum_over_Q P^T @ dO
</code></pre></div></div>

<p>结构：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">dkv_kernel</span><span class="p">(</span>
    <span class="n">q_ref</span><span class="p">,</span>
    <span class="n">k_ref</span><span class="p">,</span>
    <span class="n">v_ref</span><span class="p">,</span>
    <span class="n">o_ref</span><span class="p">,</span>
    <span class="n">lse_ref</span><span class="p">,</span>
    <span class="n">do_ref</span><span class="p">,</span>
    <span class="n">dk_ref</span><span class="p">,</span>
    <span class="n">dv_ref</span><span class="p">,</span>
    <span class="n">dk_scratch_ref</span><span class="p">,</span>
    <span class="n">dv_scratch_ref</span><span class="p">,</span>
    <span class="o">*</span><span class="p">,</span>
    <span class="n">sm_scale</span><span class="p">,</span>
    <span class="n">block_q_minor</span><span class="p">,</span>
    <span class="n">block_kv_minor</span><span class="p">,</span>
    <span class="n">q_seq_len</span><span class="p">,</span>
<span class="p">):</span>
    <span class="n">kv_block_index</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span>
    <span class="n">q_block_index</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">3</span><span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">q_block_index</span> <span class="o">==</span> <span class="mi">0</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">init</span><span class="p">():</span>
        <span class="n">dk_scratch_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">zeros_like</span><span class="p">(</span><span class="n">dk_scratch_ref</span><span class="p">)</span>
        <span class="n">dv_scratch_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">zeros_like</span><span class="p">(</span><span class="n">dv_scratch_ref</span><span class="p">)</span>

    <span class="n">should_run</span> <span class="o">=</span> <span class="n">below_or_on_diag</span><span class="p">(</span>
        <span class="n">q_block_index</span><span class="p">,</span>
        <span class="n">block_q_major_dkv</span><span class="p">,</span>
        <span class="n">kv_block_index</span><span class="p">,</span>
        <span class="n">block_kv_major_dkv</span><span class="p">,</span>
    <span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">should_run</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">run</span><span class="p">():</span>
        <span class="k">for</span> <span class="n">q_minor</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">block_q_major_dkv</span><span class="p">,</span> <span class="n">block_q_minor</span><span class="p">):</span>
            <span class="n">q</span> <span class="o">=</span> <span class="n">q_ref</span><span class="p">[...,</span> <span class="n">q_minor</span> <span class="p">:</span> <span class="n">q_minor</span> <span class="o">+</span> <span class="n">block_q_minor</span><span class="p">,</span> <span class="p">:]</span>
            <span class="n">o</span> <span class="o">=</span> <span class="n">o_ref</span><span class="p">[...,</span> <span class="n">q_minor</span> <span class="p">:</span> <span class="n">q_minor</span> <span class="o">+</span> <span class="n">block_q_minor</span><span class="p">,</span> <span class="p">:]</span>
            <span class="n">do</span> <span class="o">=</span> <span class="n">do_ref</span><span class="p">[...,</span> <span class="n">q_minor</span> <span class="p">:</span> <span class="n">q_minor</span> <span class="o">+</span> <span class="n">block_q_minor</span><span class="p">,</span> <span class="p">:]</span>
            <span class="n">lse</span> <span class="o">=</span> <span class="n">lse_ref</span><span class="p">[...,</span> <span class="n">q_minor</span> <span class="p">:</span> <span class="n">q_minor</span> <span class="o">+</span> <span class="n">block_q_minor</span><span class="p">,</span> <span class="p">:]</span>

            <span class="n">di</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span>
                <span class="n">o</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span> <span class="o">*</span> <span class="n">do</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
                <span class="n">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span>
            <span class="p">)</span>

            <span class="k">for</span> <span class="n">kv_minor</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">block_kv_major_dkv</span><span class="p">,</span> <span class="n">block_kv_minor</span><span class="p">):</span>
                <span class="n">k</span> <span class="o">=</span> <span class="n">k_ref</span><span class="p">[...,</span> <span class="n">kv_minor</span> <span class="p">:</span> <span class="n">kv_minor</span> <span class="o">+</span> <span class="n">block_kv_minor</span><span class="p">,</span> <span class="p">:]</span>
                <span class="n">v</span> <span class="o">=</span> <span class="n">v_ref</span><span class="p">[...,</span> <span class="n">kv_minor</span> <span class="p">:</span> <span class="n">kv_minor</span> <span class="o">+</span> <span class="n">block_kv_minor</span><span class="p">,</span> <span class="p">:]</span>

                <span class="n">scores</span> <span class="o">=</span> <span class="n">dot</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">T</span><span class="p">)</span>
                <span class="n">scores</span> <span class="o">*=</span> <span class="n">sm_scale</span>
                <span class="n">scores</span> <span class="o">=</span> <span class="n">apply_causal_mask_if_needed</span><span class="p">(</span><span class="n">scores</span><span class="p">)</span>

                <span class="n">p</span> <span class="o">=</span> <span class="n">exp</span><span class="p">(</span><span class="n">scores</span> <span class="o">-</span> <span class="n">lse</span><span class="p">)</span>
                <span class="n">dp</span> <span class="o">=</span> <span class="n">dot</span><span class="p">(</span><span class="n">do</span><span class="p">,</span> <span class="n">v</span><span class="p">.</span><span class="n">T</span><span class="p">)</span>
                <span class="n">ds</span> <span class="o">=</span> <span class="p">(</span><span class="n">dp</span> <span class="o">-</span> <span class="n">di</span><span class="p">[:,</span> <span class="bp">None</span><span class="p">])</span> <span class="o">*</span> <span class="n">p</span>
                <span class="n">ds</span> <span class="o">*=</span> <span class="n">sm_scale</span>

                <span class="n">dk_scratch_ref</span><span class="p">[...,</span> <span class="n">kv_minor</span><span class="p">,</span> <span class="p">:]</span> <span class="o">+=</span> <span class="n">dot</span><span class="p">(</span><span class="n">ds</span><span class="p">.</span><span class="n">T</span><span class="p">,</span> <span class="n">q</span><span class="p">)</span>
                <span class="n">dv_scratch_ref</span><span class="p">[...,</span> <span class="n">kv_minor</span><span class="p">,</span> <span class="p">:]</span> <span class="o">+=</span> <span class="n">dot</span><span class="p">(</span><span class="n">p</span><span class="p">.</span><span class="n">T</span><span class="p">,</span> <span class="n">do</span><span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">q_block_index</span> <span class="o">==</span> <span class="n">q_seq_len</span> <span class="o">//</span> <span class="n">block_q_major_dkv</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">store</span><span class="p">():</span>
        <span class="n">dk_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">dk_scratch_ref</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">dk_ref</span><span class="p">.</span><span class="n">dtype</span><span class="p">)</span>
        <span class="n">dv_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">dv_scratch_ref</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">dv_ref</span><span class="p">.</span><span class="n">dtype</span><span class="p">)</span>
</code></pre></div></div>

<hr />

<h1 id="9-lse--m--l-设计规范">9. LSE / m / l 设计规范</h1>

<h2 id="91-lse-是什么">9.1 LSE 是什么？</h2>

<p>LSE 是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>logsumexp(scores)
</code></pre></div></div>

<p>也就是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>lse_i = log(sum_j exp(scores_i,j))
</code></pre></div></div>

<p>backward 中可以用它重建 softmax：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">p</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">exp</span><span class="p">(</span><span class="n">scores</span> <span class="o">-</span> <span class="n">lse</span><span class="p">[:,</span> <span class="bp">None</span><span class="p">])</span>
</code></pre></div></div>

<hr />

<h2 id="92-ml-和-lse-的关系">9.2 m/l 和 LSE 的关系</h2>

<p>FlashAttention forward 有时保存：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>m = row max
l = sum(exp(scores - m))
</code></pre></div></div>

<p>那么：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>lse = m + log(l)
</code></pre></div></div>

<p>两种 backward 重建 P 的方式等价：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">p</span> <span class="o">=</span> <span class="n">exp</span><span class="p">(</span><span class="n">scores</span> <span class="o">-</span> <span class="n">lse</span><span class="p">)</span>
</code></pre></div></div>

<p>等价于：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">p</span> <span class="o">=</span> <span class="n">exp</span><span class="p">(</span><span class="n">scores</span> <span class="o">-</span> <span class="n">m</span><span class="p">)</span> <span class="o">/</span> <span class="n">l</span>
</code></pre></div></div>

<hr />

<h2 id="93-生产级建议">9.3 生产级建议</h2>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 如果已有 LSE，backward 用 exp(scores - lse) 最简单；
2. 如果 forward 是 streaming softmax，可保存 m/l；
3. residuals 应尽量比完整 P 小；
4. 不要保存完整 [B, H, Q, K] attention matrix；
5. LSE/m/l 通常用 fp32。
</code></pre></div></div>

<hr />

<h1 id="10-dtype-与数值稳定性规范">10. Dtype 与数值稳定性规范</h1>

<h2 id="101-matmul-accumulation">10.1 Matmul accumulation</h2>

<p>建议：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">preferred_element_type</span><span class="o">=</span><span class="n">jnp</span><span class="p">.</span><span class="n">float32</span>
</code></pre></div></div>

<p>尤其是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>QK^T
P @ V
dO @ V^T
dS @ K
dS^T @ Q
P^T @ dO
</code></pre></div></div>

<hr />

<h2 id="102-softmax-相关">10.2 Softmax 相关</h2>

<p>建议：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>scores / p / ds / accumulators 尽量使用 fp32
最终输出再 cast 到 q/k/v dtype
</code></pre></div></div>

<p>常见写法：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">scores</span> <span class="o">=</span> <span class="n">dot</span><span class="p">(...,</span> <span class="n">preferred_element_type</span><span class="o">=</span><span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
<span class="n">p</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">exp</span><span class="p">(</span><span class="n">scores</span> <span class="o">-</span> <span class="n">lse</span><span class="p">)</span>
<span class="n">ds</span> <span class="o">=</span> <span class="p">(</span><span class="n">dp</span> <span class="o">-</span> <span class="n">di</span><span class="p">[:,</span> <span class="bp">None</span><span class="p">])</span> <span class="o">*</span> <span class="n">p</span>
<span class="n">dq_acc</span> <span class="o">=</span> <span class="n">dq_acc</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
</code></pre></div></div>

<hr />

<h2 id="103-mask-value">10.3 mask value</h2>

<p>常见写法：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">DEFAULT_MASK_VALUE</span> <span class="o">=</span> <span class="o">-</span><span class="mf">0.7</span> <span class="o">*</span> <span class="nb">float</span><span class="p">(</span><span class="n">jnp</span><span class="p">.</span><span class="n">finfo</span><span class="p">(</span><span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">).</span><span class="nb">max</span><span class="p">)</span>
</code></pre></div></div>

<p>避免直接使用 <code class="language-plaintext highlighter-rouge">-inf</code> 导致某些平台上的数值或编译问题。</p>

<hr />

<h1 id="11-gqa--mqa-kernel-设计规范">11. GQA / MQA Kernel 设计规范</h1>

<h2 id="111-gqa-的核心">11.1 GQA 的核心</h2>

<p>GQA 中：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>num_q_heads &gt; num_kv_heads
group_size = num_q_heads // num_kv_heads
</code></pre></div></div>

<p>每个 KV head 被多个 Q heads 共享。</p>

<p>因此：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dK/dV 需要聚合来自同一 group 内多个 Q heads 的梯度。
</code></pre></div></div>

<hr />

<h2 id="112-推荐-layout">11.2 推荐 layout</h2>

<p>对于一个 KV head：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Q tile:
  [group_size, block_q, head_dim]

K/V tile:
  [1, block_kv, head_dim]
</code></pre></div></div>

<p>可以 reshape：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[group_size, block_q, head_dim]
-&gt;
[group_size * block_q, head_dim]
</code></pre></div></div>

<p>这样一次 matmul 覆盖整个 Q-head group。</p>

<hr />

<h2 id="113-group_size-的风险">11.3 group_size 的风险</h2>

<p>如果：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>group_size * block_q
</code></pre></div></div>

<p>太大，则：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>scores / p / dp / ds
</code></pre></div></div>

<p>会非常大。</p>

<p>例如：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>group_size = 8
block_q = 128
block_kv = 256
scores = [1024, 256]
</code></pre></div></div>

<p>这可能导致 VMEM/register 压力过大。</p>

<p>生产级策略：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>第一阶段:
  block_g = full group_size
  简化 dK/dV reduction

第二阶段:
  如果 group_size 很大且 VMEM 压力高，再引入 block_g
</code></pre></div></div>

<hr />

<h2 id="114-block_g-的代价">11.4 block_g 的代价</h2>

<p>如果引入：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">block_g</span> <span class="o">&lt;</span> <span class="n">group_size</span>
</code></pre></div></div>

<p>则 DKV 需要额外处理：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>不同 group block 对同一个 dK/dV 的累加。
</code></pre></div></div>

<p>这可能需要：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 更复杂的 scratch；
2. 额外 reduction kernel；
3. atomic-like accumulation 设计；
4. 更复杂的 correctness 测试。
</code></pre></div></div>

<p>所以不建议第一版就引入 <code class="language-plaintext highlighter-rouge">block_g</code>。</p>

<hr />

<h1 id="12-性能优化-checklist">12. 性能优化 Checklist</h1>

<h2 id="121-mxu-利用率低时优先检查">12.1 MXU 利用率低时优先检查</h2>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. dot_general shape 是否太小？
2. dot_general shape 是否不规整？
3. block size 是否 128 对齐？
4. 是否有过长 program 内部 loop？
5. reduction 维是否应该放进 grid？
6. 是否用了 full-seq BlockSpec？
7. scores/p/dp/ds 是否太大？
8. VMEM/register 是否压力太高？
9. 是否有 spill 或 compiler warning？
10. causal 上三角是否被 matmul 计算了？
11. matmul 中间是否夹杂太多 elementwise？
12. grid program 数是否太少？
13. batch/head 并行度是否足够？
14. block_b 是否过大或过小？
</code></pre></div></div>

<hr />

<h2 id="122-latency-高时优先检查">12.2 Latency 高时优先检查</h2>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 是否没有 block-level causal skip？
2. 是否物化了 [B, H, Q, K] 大 tensor？
3. 是否写出了 dAB/dS？
4. 是否 HBM 读写过多？
5. 是否重复读取 full Q/K/V？
6. scratch 是否太大？
7. block_kv 是否太大导致临时矩阵膨胀？
8. block_q 是否太大导致 group_size * block_q 过大？
9. 是否 padding 后多算太多 token？
10. DQ 和 DKV 哪一个更慢？
</code></pre></div></div>

<hr />

<h2 id="123-correctness-失败时优先检查">12.3 Correctness 失败时优先检查</h2>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. causal mask 的 row/col offset 是否正确？
2. q_start / kv_start 是否正确？
3. GQA head mapping 是否正确？
4. lse/m/l shape 是否正确 broadcast？
5. sm_scale 是否 forward/backward 一致？
6. ds 是否乘了 sm_scale？
7. di = sum(o * do) 是否按 row 计算？
8. dtype cast 是否过早？
9. padding 后是否正确 slice 回原 shape？
10. skipped block 是否误跳过了 diagonal block？
11. scratch init/store 条件是否正确？
</code></pre></div></div>

<hr />

<h1 id="13-benchmark-规范">13. Benchmark 规范</h1>

<h2 id="131-latency-是主指标">13.1 latency 是主指标</h2>

<p>当比较 Pallas kernel 和 JAX reference 时，不要直接混用 FLOP 口径。</p>

<p>常见 FLOP 口径包括：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. JAX compiler estimated FLOPs
2. manual Pallas matmul inventory
3. theoretical math FLOPs
4. actual executed FLOPs with causal skip
</code></pre></div></div>

<p>这些不是一回事。</p>

<p>速度比较以：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>latency
</code></pre></div></div>

<p>为主。</p>

<hr />

<h2 id="132-effective-tflops">13.2 effective TFLOP/s</h2>

<p>如果要计算：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>effective_TFLOP/s
</code></pre></div></div>

<p>必须固定 FLOP 口径。</p>

<p>建议使用：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>manual Pallas matmul inventory
</code></pre></div></div>

<p>公式：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>effective_TFLOP/s = manual_pallas_GFLOP / latency_ms
</code></pre></div></div>

<p>因为：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1 GFLOP / 1 ms = 1 TFLOP/s
</code></pre></div></div>

<hr />

<h2 id="133-必须分别测-dq-和-dkv">13.3 必须分别测 DQ 和 DKV</h2>

<p>不要只看 total。</p>

<p>建议记录：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DQ latency
DKV latency
Total latency
DQ MXU utilization
DKV MXU utilization
Total MXU utilization
</code></pre></div></div>

<p>原因：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DQ 和 DKV 的瓶颈可能完全不同。
</code></pre></div></div>

<p>例如：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DQ 可能卡在 KV reduction；
DKV 可能卡在 full Q BlockSpec、scratch、group reduction。
</code></pre></div></div>

<hr />

<h2 id="134-benchmark-矩阵">13.4 Benchmark 矩阵</h2>

<p>建议至少测：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>batch:
  1, 2, 4

seq_len:
  512, 1024, 2048, 4096, 8192

head_dim:
  64, 128

num_q_heads / num_kv_heads:
  16/4, 32/8, 32/4

group_size:
  2, 4, 8

dtype:
  bf16 input + fp32 accumulate
</code></pre></div></div>

<hr />

<h1 id="14-block-size-sweep-规范">14. Block Size Sweep 规范</h1>

<h2 id="141-forward-和-backward-的最优-block-size-可能不同">14.1 Forward 和 backward 的最优 block size 可能不同</h2>

<p>不要假设：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>forward 最优 block size = backward 最优 block size
</code></pre></div></div>

<p>原因：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>forward:
  QK^T + P@V + streaming softmax

backward:
  重建 P
  dV = P^T dO
  dP = dO V^T
  dS
  dK = dS^T Q
  dQ = dS K
</code></pre></div></div>

<p>backward 的 matmul 数量、reduction 方向、scratch 需求都不同。</p>

<hr />

<h2 id="142-dq-和-dkv-的最优-block-size-也可能不同">14.2 DQ 和 DKV 的最优 block size 也可能不同</h2>

<p>DQ：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>输出 dQ
沿 KV reduction
</code></pre></div></div>

<p>DKV：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>输出 dK/dV
沿 Q reduction
还要聚合 group_size 个 Q heads
</code></pre></div></div>

<p>所以：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q_major_dq
block_kv_major_dq
block_kv_dq
</code></pre></div></div>

<p>和：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q_major_dkv
block_kv_major_dkv
block_q_dkv
block_kv_dkv
</code></pre></div></div>

<p>应该分开调。</p>

<hr />

<h2 id="143-推荐初始-sweep">14.3 推荐初始 sweep</h2>

<p>DQ：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q_major_dq:
  64, 128

block_kv_major_dq:
  128, 256

block_kv_dq:
  128
</code></pre></div></div>

<p>DKV：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q_major_dkv:
  64, 128, 256

block_kv_major_dkv:
  128, 256

block_q_dkv:
  64, 128

block_kv_dkv:
  128
</code></pre></div></div>

<p>如果 group_size 较大：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q_major:
  优先尝试 64, 128

block_kv_major:
  优先尝试 128
</code></pre></div></div>

<hr />

<h1 id="15-生产级代码结构建议">15. 生产级代码结构建议</h1>

<h2 id="151-参数-dataclass">15.1 参数 dataclass</h2>

<p>推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="o">@</span><span class="n">dataclasses</span><span class="p">.</span><span class="n">dataclass</span><span class="p">(</span><span class="n">frozen</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="k">class</span> <span class="nc">BlockSizes</span><span class="p">:</span>
    <span class="n">block_b</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">1</span>

    <span class="n">block_q_major_dq</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">128</span>
    <span class="n">block_kv_major_dq</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">128</span>
    <span class="n">block_kv_dq</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">128</span>

    <span class="n">block_q_major_dkv</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">128</span>
    <span class="n">block_kv_major_dkv</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">128</span>
    <span class="n">block_q_dkv</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">128</span>
    <span class="n">block_kv_dkv</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">128</span>

    <span class="k">def</span> <span class="nf">__post_init__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="p">...</span>
</code></pre></div></div>

<p>校验：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 所有 block size &gt; 0；
2. minor &lt;= major；
3. major % minor == 0；
4. block_kv_* 是 128 的倍数；
5. q_seq_len / kv_seq_len 是否可整除，或是否需要 padding。
</code></pre></div></div>

<hr />

<h2 id="152-shape-validation">15.2 Shape validation</h2>

<p>在外层 API 做 shape 校验：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">validate_shapes</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">o</span><span class="p">,</span> <span class="n">lse</span><span class="p">,</span> <span class="n">do</span><span class="p">):</span>
    <span class="k">assert</span> <span class="n">q</span><span class="p">.</span><span class="n">shape</span> <span class="o">==</span> <span class="n">o</span><span class="p">.</span><span class="n">shape</span>
    <span class="k">assert</span> <span class="n">q</span><span class="p">.</span><span class="n">shape</span> <span class="o">==</span> <span class="n">do</span><span class="p">.</span><span class="n">shape</span>
    <span class="k">assert</span> <span class="n">k</span><span class="p">.</span><span class="n">shape</span> <span class="o">==</span> <span class="n">v</span><span class="p">.</span><span class="n">shape</span>
    <span class="k">assert</span> <span class="n">q</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</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">1</span><span class="p">]</span> <span class="o">==</span> <span class="mi">0</span>
    <span class="k">assert</span> <span class="n">q</span><span class="p">.</span><span class="n">shape</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">k</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">]</span>
</code></pre></div></div>

<p>不要把 shape 错误留到 kernel 内部才失败。</p>

<hr />

<h2 id="153-padding-和-slicing">15.3 Padding 和 slicing</h2>

<p>生产级 kernel 通常需要支持非整除 sequence length。</p>

<p>做法：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">q_pad</span> <span class="o">=</span> <span class="n">pad_axis_to_multiple</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="n">multiple</span><span class="o">=</span><span class="n">block_q</span><span class="p">)</span>
<span class="n">k_pad</span> <span class="o">=</span> <span class="n">pad_axis_to_multiple</span><span class="p">(</span><span class="n">k</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="n">multiple</span><span class="o">=</span><span class="n">block_kv</span><span class="p">)</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="p">:,</span> <span class="p">:</span><span class="n">original_seq_len</span><span class="p">,</span> <span class="p">:]</span>
</code></pre></div></div>

<p>注意：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>padding token 必须被 mask 或保证不影响结果。
</code></pre></div></div>

<hr />

<h2 id="154-named_scope">15.4 named_scope</h2>

<p>推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">name_scope</span> <span class="o">=</span> <span class="p">(</span>
    <span class="sa">f</span><span class="s">"gqa_bwd_dq_"</span>
    <span class="sa">f</span><span class="s">"</span><span class="si">{</span><span class="n">block_q_major_dq</span><span class="o">=</span><span class="si">}</span><span class="s">_"</span>
    <span class="sa">f</span><span class="s">"</span><span class="si">{</span><span class="n">block_kv_major_dq</span><span class="o">=</span><span class="si">}</span><span class="s">_"</span>
    <span class="sa">f</span><span class="s">"</span><span class="si">{</span><span class="n">block_kv_dq</span><span class="o">=</span><span class="si">}</span><span class="s">"</span>
<span class="p">)</span>

<span class="k">with</span> <span class="n">jax</span><span class="p">.</span><span class="n">named_scope</span><span class="p">(</span><span class="n">name_scope</span><span class="p">):</span>
    <span class="p">...</span>
</code></pre></div></div>

<p>好处：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. profile 更容易看；
2. HLO / trace 更容易定位；
3. benchmark 结果更清楚。
</code></pre></div></div>

<hr />

<h2 id="155-debug--interpret">15.5 debug / interpret</h2>

<p>建议外层 API 支持：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">debug</span><span class="p">:</span> <span class="nb">bool</span> <span class="o">=</span> <span class="bp">False</span>
<span class="n">interpret</span><span class="p">:</span> <span class="nb">bool</span> <span class="o">=</span> <span class="bp">False</span>
</code></pre></div></div>

<p>开发阶段：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>interpret=True 方便调试
</code></pre></div></div>

<p>性能测试阶段：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>interpret=False
debug=False
</code></pre></div></div>

<hr />

<h1 id="16-常见反模式">16. 常见反模式</h1>

<h2 id="161-一个-program-扫完整-seq_len">16.1 一个 program 扫完整 seq_len</h2>

<p>不推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">for</span> <span class="n">q_start</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">q_seq_len</span><span class="p">,</span> <span class="n">block_q</span><span class="p">):</span>
    <span class="p">...</span>
</code></pre></div></div>

<p>如果这个 loop 很长，应考虑把 <code class="language-plaintext highlighter-rouge">q_block</code> 放入 grid。</p>

<hr />

<h2 id="162-full-sequence-blockspec">16.2 full-sequence BlockSpec</h2>

<p>不推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">q_seq_len</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">index_map</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q_major</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">index_map</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<hr />

<h2 id="163-先-matmul-再-mask-掉整个无效-block">16.3 先 matmul 再 mask 掉整个无效 block</h2>

<p>不推荐：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><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">T</span>
<span class="n">scores</span> <span class="o">=</span> <span class="n">causal_mask</span><span class="p">(</span><span class="n">scores</span><span class="p">)</span>
</code></pre></div></div>

<p>对于完全无效 block，应先 skip：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">if</span> <span class="n">block_is_valid</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">T</span>
</code></pre></div></div>

<hr />

<h2 id="164-写出-ds--dab-大矩阵">16.4 写出 dS / dAB 大矩阵</h2>

<p>如果不是必须，不要写：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dAB = dS = [B, H, Q, K]
</code></pre></div></div>

<p>这会破坏 FlashAttention 不物化 attention matrix 的优势。</p>

<hr />

<h2 id="165-盲目增大-block_kv">16.5 盲目增大 block_kv</h2>

<p><code class="language-plaintext highlighter-rouge">block_kv=256</code> 不一定比 <code class="language-plaintext highlighter-rouge">128</code> 快。</p>

<p>需要看：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. MXU utilization
2. VMEM pressure
3. scores/p/dp/ds size
4. causal skip 粒度
5. latency
</code></pre></div></div>

<hr />

<h1 id="17-推荐优化顺序">17. 推荐优化顺序</h1>

<p>如果一个 Pallas attention backward kernel MXU utilization 很低，建议按下面顺序优化。</p>

<h2 id="step-1-分离-dq--dkv-profiling">Step 1: 分离 DQ / DKV profiling</h2>

<p>先确认：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DQ 慢，还是 DKV 慢？
</code></pre></div></div>

<p>记录：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DQ latency
DKV latency
DQ MXU utilization
DKV MXU utilization
</code></pre></div></div>

<hr />

<h2 id="step-2-加-block-level-causal-skip">Step 2: 加 block-level causal skip</h2>

<p>这是低风险高收益优化。</p>

<p>目标：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>完全 causal invalid 的 block 不做 QK^T。
</code></pre></div></div>

<hr />

<h2 id="step-3-消除-full-seq-blockspec">Step 3: 消除 full-seq BlockSpec</h2>

<p>将：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">q_seq_len</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">)</span>
</code></pre></div></div>

<p>改成：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="p">(</span><span class="n">block_b</span><span class="p">,</span> <span class="n">group_size</span><span class="p">,</span> <span class="n">block_q_major</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">)</span>
</code></pre></div></div>

<hr />

<h2 id="step-4-reduction-维放进-grid">Step 4: reduction 维放进 grid</h2>

<p>DQ：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>把 kv_block 放进 grid。
</code></pre></div></div>

<p>DKV：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>把 q_block 放进 grid。
</code></pre></div></div>

<hr />

<h2 id="step-5-用-scratch-做-partial-accumulation">Step 5: 用 scratch 做 partial accumulation</h2>

<p>DQ：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dq_scratch 跨 kv_block 累加。
</code></pre></div></div>

<p>DKV：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dk_scratch / dv_scratch 跨 q_block 累加。
</code></pre></div></div>

<hr />

<h2 id="step-6-引入-majorminor-tile">Step 6: 引入 major/minor tile</h2>

<p>把：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q
block_kv
</code></pre></div></div>

<p>拆成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q_major
block_kv_major
block_q_minor
block_kv_minor
</code></pre></div></div>

<hr />

<h2 id="step-7-系统-sweep-block-size">Step 7: 系统 sweep block size</h2>

<p>不要凭感觉判断最优 block。</p>

<p>至少 sweep：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_q_major = 64 / 128 / 256
block_kv_major = 128 / 256
block_q_minor = 64 / 128
block_kv_minor = 128
</code></pre></div></div>

<hr />

<h2 id="step-8-视情况引入-block_g">Step 8: 视情况引入 block_g</h2>

<p>只有当：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>group_size 很大
scores tile 太大
VMEM 压力明显
</code></pre></div></div>

<p>才考虑：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">block_g</span> <span class="o">&lt;</span> <span class="n">group_size</span>
</code></pre></div></div>

<hr />

<h1 id="18-代码评审-checklist">18. 代码评审 Checklist</h1>

<p>提交生产级 Pallas kernel 前，建议检查：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[ ] 数学公式是否写清楚？
[ ] 输入输出 shape 是否校验？
[ ] GQA head mapping 是否正确？
[ ] grid 是否暴露足够并行度？
[ ] 是否避免 full-seq BlockSpec？
[ ] 是否避免单 program 内部长 reduction loop？
[ ] 是否正确使用 scratch？
[ ] scratch init/store 条件是否正确？
[ ] reduction grid 维是否有正确 dimension semantics？
[ ] causal block-level skip 是否正确？
[ ] diagonal block 是否仍然做 element-level mask？
[ ] dtype cast 是否合理？
[ ] accumulation 是否使用 fp32？
[ ] 是否避免写出不必要的大 tensor？
[ ] padding 后是否 slice 回原 shape？
[ ] 是否有 reference correctness test？
[ ] 是否有 DQ/DKV 单独 benchmark？
[ ] 是否记录 latency / MXU utilization / effective TFLOP/s？
[ ] 是否明确 FLOP 统计口径？
[ ] 是否 sweep 过核心 block size？
[ ] 是否有 named_scope 方便 profile？
</code></pre></div></div>

<hr />

<h1 id="19-一句话总结">19. 一句话总结</h1>

<p>生产级 Pallas kernel 的核心不是“把公式翻译成 kernel”，而是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>把数学计算拆成硬件友好的 tile，
把长 reduction 暴露成 grid，
用 scratch 做局部累加，
避免 full-sequence input window，
尽量让 MXU 连续吃到规整 matmul，
同时控制 VMEM/register/HBM 压力。
</code></pre></div></div>

<p>对于 FlashAttention / GQA backward，最常用、最值得掌握的模式是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DQ:
  grid over q_block × kv_block
  dq_scratch 跨 kv_block 累加

DKV:
  grid over kv_block × q_block
  dk_scratch / dv_scratch 跨 q_block 累加

Causal:
  block-level skip + diagonal block element-level mask

Tiling:
  major block 控制 grid 粒度
  minor block 控制 program 内部 dot 粒度

Benchmark:
  latency 为主
  FLOP 口径固定
  DQ/DKV 分开看
</code></pre></div></div>]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[Pallas TPU Kernel 写法与优化文档]]></summary></entry><entry><title type="html">Pallas TPU Distributed Collectives 复习文档</title><link href="https://wqh011128.github.io/blog/2026/06/04/Pallas_TPU_Distributed_Collectives.html" rel="alternate" type="text/html" title="Pallas TPU Distributed Collectives 复习文档" /><published>2026-06-04T00:00:00+00:00</published><updated>2026-06-04T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/06/04/Pallas_TPU_Distributed_Collectives</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/06/04/Pallas_TPU_Distributed_Collectives.html"><![CDATA[<h1 id="pallas-tpu-distributed-collectives-复习文档">Pallas TPU Distributed Collectives 复习文档</h1>

<p>本文根据 JAX 官方文档 <a href="https://docs.jax.dev/en/latest/pallas/tpu/distributed.html">Distributed Computing in Pallas for TPUs</a> 整理，重点复习四个 collective 示例：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">lax.ppermute</code></li>
  <li><code class="language-plaintext highlighter-rouge">lax.all_gather</code></li>
  <li><code class="language-plaintext highlighter-rouge">lax.psum</code></li>
  <li><code class="language-plaintext highlighter-rouge">lax.psum_scatter</code></li>
</ul>

<p>说明：代码是“教学整理版”，保留官方示例的核心结构、API 和数据流，但注释与组织方式做了重写，方便复习。真实运行仍需要 TPU 环境。</p>

<hr />

<h2 id="0-共同背景">0. 共同背景</h2>

<h3 id="01-ring-设备模型">0.1 Ring 设备模型</h3>

<p>假设有 4 个 device：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>D0 -&gt; D1 -&gt; D2 -&gt; D3 -&gt; D0
</code></pre></div></div>

<p>右邻居：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">right</span> <span class="o">=</span> <span class="p">(</span><span class="n">my_id</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)</span> <span class="o">%</span> <span class="n">num_devices</span>
</code></pre></div></div>

<p>左邻居：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">left</span> <span class="o">=</span> <span class="p">(</span><span class="n">my_id</span> <span class="o">-</span> <span class="mi">1</span> <span class="o">+</span> <span class="n">num_devices</span><span class="p">)</span> <span class="o">%</span> <span class="n">num_devices</span>
</code></pre></div></div>

<p>在本文模拟中：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>D0 的本地数据记作 x0
D1 的本地数据记作 x1
D2 的本地数据记作 x2
D3 的本地数据记作 x3
</code></pre></div></div>

<h3 id="02-meshpartitionspecnamedsharding">0.2 <code class="language-plaintext highlighter-rouge">mesh</code>、<code class="language-plaintext highlighter-rouge">PartitionSpec</code>、<code class="language-plaintext highlighter-rouge">NamedSharding</code></h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">P</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">sharding</span><span class="p">.</span><span class="n">PartitionSpec</span>

<span class="n">mesh</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">make_mesh</span><span class="p">((</span><span class="n">num_devices</span><span class="p">,),</span> <span class="p">(</span><span class="s">"x"</span><span class="p">,))</span>
<span class="n">partition</span> <span class="o">=</span> <span class="n">P</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="s">"x"</span><span class="p">)</span>
<span class="n">sharding</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">sharding</span><span class="p">.</span><span class="n">NamedSharding</span><span class="p">(</span><span class="n">mesh</span><span class="p">,</span> <span class="n">partition</span><span class="p">)</span>
</code></pre></div></div>

<p>三者区别：</p>

<table>
  <thead>
    <tr>
      <th>名称</th>
      <th>含义</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">mesh</code></td>
      <td>设备组成的逻辑网格，例如一维 <code class="language-plaintext highlighter-rouge">x</code> 轴有 4 个 TPU device</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">PartitionSpec</code></td>
      <td>说明数组维度如何映射到 mesh 轴</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">NamedSharding</code></td>
      <td>把设备网格和切分规则合成一个真正可用于 <code class="language-plaintext highlighter-rouge">device_put</code> 的 sharding</td>
    </tr>
  </tbody>
</table>

<p>例子：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">partition</span> <span class="o">=</span> <span class="n">P</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="s">"x"</span><span class="p">)</span>
<span class="n">input_arr</span><span class="p">.</span><span class="n">shape</span> <span class="o">==</span> <span class="p">(</span><span class="mi">8</span><span class="p">,</span> <span class="mi">128</span> <span class="o">*</span> <span class="n">num_devices</span><span class="p">)</span>
</code></pre></div></div>

<p>表示：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>第 0 维不切分
第 1 维沿 x 轴切分
</code></pre></div></div>

<p>如果 <code class="language-plaintext highlighter-rouge">num_devices=4</code>，全局 shape 是 <code class="language-plaintext highlighter-rouge">(8, 512)</code>，每个 device 本地拿到 <code class="language-plaintext highlighter-rouge">(8, 128)</code>。</p>

<h3 id="03-tpu-pallas-通信rdma-是-push-only">0.3 TPU Pallas 通信：RDMA 是 push-only</h3>

<p>TPU Pallas 的 remote DMA 是 <strong>push-only</strong>：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>当前 device 可以把本地 src_ref 推送到远端 device 的 dst_ref。
当前 device 不能直接读取远端 device 的内存。
</code></pre></div></div>

<p>核心 API：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">copy</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_remote_copy</span><span class="p">(</span>
    <span class="n">src_ref</span><span class="o">=</span><span class="n">local_ref</span><span class="p">,</span>
    <span class="n">dst_ref</span><span class="o">=</span><span class="n">remote_ref</span><span class="p">,</span>
    <span class="n">send_sem</span><span class="o">=</span><span class="n">send_sem</span><span class="p">,</span>
    <span class="n">recv_sem</span><span class="o">=</span><span class="n">recv_sem</span><span class="p">,</span>
    <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">target_device</span><span class="p">,),</span>
    <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
<span class="p">)</span>
<span class="n">copy</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
<span class="n">copy</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>
</code></pre></div></div>

<p>关键点：</p>

<table>
  <thead>
    <tr>
      <th>参数</th>
      <th>含义</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">src_ref</code></td>
      <td>本 device 上要发送的数据</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">dst_ref</code></td>
      <td>目标 device 上接收的位置</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">send_sem</code></td>
      <td>发送端 DMA semaphore</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">recv_sem</code></td>
      <td>接收端 DMA semaphore</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">.start()</code></td>
      <td>发起异步 DMA</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">.wait()</code></td>
      <td>等发送和接收都完成</td>
    </tr>
  </tbody>
</table>

<h3 id="04-hbmvmemvreg-的直觉">0.4 HBM、VMEM、VREG 的直觉</h3>

<p>在这些示例里，常见内存角色是：</p>

<table>
  <thead>
    <tr>
      <th>名称</th>
      <th>直觉</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>HBM</td>
      <td>每个 device 自己的大内存，适合放通信 buffer</td>
    </tr>
    <tr>
      <td>VMEM</td>
      <td>TPU 上更靠近计算的内存，适合做局部计算</td>
    </tr>
    <tr>
      <td>VREG</td>
      <td>vector register，真正执行向量操作的位置</td>
    </tr>
  </tbody>
</table>

<p>一个常见流程：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>远端 DMA 写入本 device 的 HBM scratch
        ↓
local DMA 把 HBM scratch 拷到 VMEM scratch
        ↓
在 VMEM/VREG 中计算
        ↓
写回 HBM scratch 或最终 o_ref
</code></pre></div></div>

<hr />

<h2 id="1-laxppermute右移一格">1. <code class="language-plaintext highlighter-rouge">lax.ppermute</code>：右移一格</h2>

<h3 id="11-直观解释">1.1 直观解释</h3>

<p><code class="language-plaintext highlighter-rouge">ppermute</code> 是最简单的跨设备通信：每个 device 把自己的 shard 发给右邻居。</p>

<p>对 4 个 device：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>输入:
  D0: x0
  D1: x1
  D2: x2
  D3: x3

每个 device 发送到右邻居:
  D0 -&gt; D1
  D1 -&gt; D2
  D2 -&gt; D3
  D3 -&gt; D0

输出:
  D0: x3
  D1: x0
  D2: x1
  D3: x2
</code></pre></div></div>

<h3 id="12-配图">1.2 配图</h3>

<pre><code class="language-mermaid">flowchart LR
  D0["D0: x0"] --&gt; D1["D1 receives x0"]
  D10["D1: x1"] --&gt; D2["D2 receives x1"]
  D20["D2: x2"] --&gt; D3["D3 receives x2"]
  D30["D3: x3"] --&gt; D0R["D0 receives x3"]
</code></pre>

<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">jax</span>
<span class="kn">from</span> <span class="nn">jax</span> <span class="kn">import</span> <span class="n">lax</span>
<span class="kn">from</span> <span class="nn">jax</span> <span class="kn">import</span> <span class="n">numpy</span> <span class="k">as</span> <span class="n">jnp</span>
<span class="kn">from</span> <span class="nn">jax.experimental</span> <span class="kn">import</span> <span class="n">pallas</span> <span class="k">as</span> <span class="n">pl</span>
<span class="kn">from</span> <span class="nn">jax.experimental.pallas</span> <span class="kn">import</span> <span class="n">tpu</span> <span class="k">as</span> <span class="n">pltpu</span>

<span class="n">P</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">sharding</span><span class="p">.</span><span class="n">PartitionSpec</span>
<span class="n">num_devices</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">local_device_count</span><span class="p">()</span>

<span class="n">partition</span> <span class="o">=</span> <span class="n">P</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="s">"x"</span><span class="p">)</span>
<span class="n">mesh</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">make_mesh</span><span class="p">((</span><span class="n">num_devices</span><span class="p">,),</span> <span class="p">(</span><span class="s">"x"</span><span class="p">,))</span>
<span class="n">sharding</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">sharding</span><span class="p">.</span><span class="n">NamedSharding</span><span class="p">(</span><span class="n">mesh</span><span class="p">,</span> <span class="n">partition</span><span class="p">)</span>

<span class="n">input_arr</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">uniform</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">key</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span>
    <span class="n">shape</span><span class="o">=</span><span class="p">(</span><span class="mi">8</span><span class="p">,</span> <span class="mi">128</span> <span class="o">*</span> <span class="n">num_devices</span><span class="p">),</span>
<span class="p">)</span>
<span class="n">input_arr</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">device_put</span><span class="p">(</span><span class="n">input_arr</span><span class="p">,</span> <span class="n">sharding</span><span class="p">)</span>


<span class="k">def</span> <span class="nf">right_shift_kernel</span><span class="p">(</span><span class="n">x_ref</span><span class="p">,</span> <span class="n">y_ref</span><span class="p">,</span> <span class="n">send_sem</span><span class="p">,</span> <span class="n">recv_sem</span><span class="p">):</span>
    <span class="n">my_id</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">axis_index</span><span class="p">(</span><span class="s">"x"</span><span class="p">)</span>
    <span class="n">target</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">rem</span><span class="p">(</span><span class="n">my_id</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>

    <span class="n">dma</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_remote_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">x_ref</span><span class="p">,</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">y_ref</span><span class="p">,</span>
        <span class="n">send_sem</span><span class="o">=</span><span class="n">send_sem</span><span class="p">,</span>
        <span class="n">recv_sem</span><span class="o">=</span><span class="n">recv_sem</span><span class="p">,</span>
        <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">target</span><span class="p">,),</span>
        <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
    <span class="p">)</span>
    <span class="n">dma</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
    <span class="n">dma</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>


<span class="n">out_shape</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">ShapeDtypeStruct</span><span class="p">((</span><span class="mi">8</span><span class="p">,</span> <span class="mi">128</span><span class="p">),</span> <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>

<span class="n">grid_spec</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">PrefetchScalarGridSpec</span><span class="p">(</span>
    <span class="n">num_scalar_prefetch</span><span class="o">=</span><span class="mi">0</span><span class="p">,</span>
    <span class="n">in_specs</span><span class="o">=</span><span class="p">[</span><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">ANY</span><span class="p">)],</span>
    <span class="n">out_specs</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">ANY</span><span class="p">),</span>
    <span class="n">scratch_shapes</span><span class="o">=</span><span class="p">([</span><span class="n">pltpu</span><span class="p">.</span><span class="n">SemaphoreType</span><span class="p">.</span><span class="n">DMA</span><span class="p">]</span> <span class="o">*</span> <span class="mi">2</span><span class="p">),</span>
<span class="p">)</span>

<span class="n">right_shift</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">pallas_call</span><span class="p">(</span>
    <span class="n">right_shift_kernel</span><span class="p">,</span>
    <span class="n">out_shape</span><span class="o">=</span><span class="n">out_shape</span><span class="p">,</span>
    <span class="n">grid_spec</span><span class="o">=</span><span class="n">grid_spec</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">pallas_result</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">jit</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">shard_map</span><span class="p">(</span>
        <span class="n">right_shift</span><span class="p">,</span>
        <span class="n">mesh</span><span class="o">=</span><span class="n">mesh</span><span class="p">,</span>
        <span class="n">in_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
        <span class="n">out_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
        <span class="n">check_vma</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
    <span class="p">)</span>
<span class="p">)(</span><span class="n">input_arr</span><span class="p">)</span>

<span class="n">perm</span> <span class="o">=</span> <span class="nb">tuple</span><span class="p">((</span><span class="n">src</span><span class="p">,</span> <span class="p">(</span><span class="n">src</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)</span> <span class="o">%</span> <span class="n">num_devices</span><span class="p">)</span> <span class="k">for</span> <span class="n">src</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_devices</span><span class="p">))</span>
<span class="n">xla_result</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">jit</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">shard_map</span><span class="p">(</span>
        <span class="k">lambda</span> <span class="n">x</span><span class="p">:</span> <span class="n">lax</span><span class="p">.</span><span class="n">ppermute</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="s">"x"</span><span class="p">,</span> <span class="n">perm</span><span class="p">),</span>
        <span class="n">mesh</span><span class="o">=</span><span class="n">mesh</span><span class="p">,</span>
        <span class="n">in_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
        <span class="n">out_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
    <span class="p">)</span>
<span class="p">)(</span><span class="n">input_arr</span><span class="p">)</span>
</code></pre></div></div>

<h3 id="14-逐步模拟">1.4 逐步模拟</h3>

<table>
  <thead>
    <tr>
      <th>Device</th>
      <th>初始本地输入</th>
      <th>发送到</th>
      <th>最终本地输出</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">x0</code></td>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">x3</code></td>
    </tr>
    <tr>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">x1</code></td>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">x0</code></td>
    </tr>
    <tr>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">x2</code></td>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">x1</code></td>
    </tr>
    <tr>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">x3</code></td>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">x2</code></td>
    </tr>
  </tbody>
</table>

<p>完成这个操作的代码是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">target</span> <span class="o">=</span> <span class="p">(</span><span class="n">my_id</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)</span> <span class="o">%</span> <span class="n">num_devices</span>
<span class="n">make_async_remote_copy</span><span class="p">(</span><span class="n">src_ref</span><span class="o">=</span><span class="n">x_ref</span><span class="p">,</span> <span class="n">dst_ref</span><span class="o">=</span><span class="n">y_ref</span><span class="p">,</span> <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">target</span><span class="p">,))</span>
</code></pre></div></div>

<p>这里 <code class="language-plaintext highlighter-rouge">dst_ref=y_ref</code> 表示目标 device 上同一个 Pallas output buffer 的本地位置。</p>

<hr />

<h2 id="2-laxall_gather每个-device-收集所有-shard">2. <code class="language-plaintext highlighter-rouge">lax.all_gather</code>：每个 device 收集所有 shard</h2>

<h3 id="21-直观解释">2.1 直观解释</h3>

<p><code class="language-plaintext highlighter-rouge">all_gather</code> 的目标是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>每个 device 一开始只有自己的 shard。
最后每个 device 都有所有 shard。
</code></pre></div></div>

<p>4 个 device：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>初始:
  D0: x0
  D1: x1
  D2: x2
  D3: x3

最终每个 device:
  [x0, x1, x2, x3]
</code></pre></div></div>

<p>官方图：</p>

<p><img src="https://docs.jax.dev/en/latest/_images/all_gather.svg" alt="all_gather" /></p>

<h3 id="22-配图">2.2 配图</h3>

<pre><code class="language-mermaid">flowchart LR
  S0["step 0: 每个 device 发自己的 shard"] --&gt; S1["step 1: 发刚从左边收到的 shard"]
  S1 --&gt; S2["step 2: 继续转发"]
  S2 --&gt; OUT["每个 device 都拥有 x0,x1,x2,x3"]
</code></pre>

<h3 id="23-教学版完整代码">2.3 教学版完整代码</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">partition</span> <span class="o">=</span> <span class="n">P</span><span class="p">(</span><span class="s">"x"</span><span class="p">,</span> <span class="bp">None</span><span class="p">)</span>
<span class="n">mesh</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">make_mesh</span><span class="p">((</span><span class="n">num_devices</span><span class="p">,),</span> <span class="p">(</span><span class="s">"x"</span><span class="p">,))</span>
<span class="n">sharding</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">sharding</span><span class="p">.</span><span class="n">NamedSharding</span><span class="p">(</span><span class="n">mesh</span><span class="p">,</span> <span class="n">partition</span><span class="p">)</span>

<span class="n">input_arr</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">uniform</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">key</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span>
    <span class="n">shape</span><span class="o">=</span><span class="p">(</span><span class="mi">8</span> <span class="o">*</span> <span class="n">num_devices</span><span class="p">,</span> <span class="mi">128</span><span class="p">),</span>
<span class="p">)</span>
<span class="n">input_arr</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">device_put</span><span class="p">(</span><span class="n">input_arr</span><span class="p">,</span> <span class="n">sharding</span><span class="p">)</span>


<span class="k">def</span> <span class="nf">all_gather_kernel</span><span class="p">(</span><span class="n">x_ref</span><span class="p">,</span> <span class="n">out_ref</span><span class="p">,</span> <span class="n">local_sem</span><span class="p">,</span> <span class="n">send_sem</span><span class="p">,</span> <span class="n">recv_sems</span><span class="p">):</span>
    <span class="n">step</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>
    <span class="n">my_id</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">axis_index</span><span class="p">(</span><span class="s">"x"</span><span class="p">)</span>
    <span class="n">right</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">rem</span><span class="p">(</span><span class="n">my_id</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>

    <span class="n">slot</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">rem</span><span class="p">(</span><span class="n">my_id</span> <span class="o">-</span> <span class="n">step</span> <span class="o">+</span> <span class="n">num_devices</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">step</span> <span class="o">==</span> <span class="mi">0</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
        <span class="n">local</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_copy</span><span class="p">(</span>
            <span class="n">src_ref</span><span class="o">=</span><span class="n">x_ref</span><span class="p">,</span>
            <span class="n">dst_ref</span><span class="o">=</span><span class="n">out_ref</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">my_id</span><span class="p">],</span>
            <span class="n">sem</span><span class="o">=</span><span class="n">local_sem</span><span class="p">,</span>
        <span class="p">)</span>
        <span class="n">local</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
        <span class="n">local</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>

    <span class="n">remote</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_remote_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">out_ref</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">slot</span><span class="p">],</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">out_ref</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">slot</span><span class="p">],</span>
        <span class="n">send_sem</span><span class="o">=</span><span class="n">send_sem</span><span class="p">,</span>
        <span class="n">recv_sem</span><span class="o">=</span><span class="n">recv_sems</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">step</span><span class="p">],</span>
        <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">right</span><span class="p">,),</span>
        <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
    <span class="p">)</span>
    <span class="n">remote</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
    <span class="n">remote</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>


<span class="n">out_shape</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">ShapeDtypeStruct</span><span class="p">((</span><span class="n">num_devices</span><span class="p">,</span> <span class="mi">8</span><span class="p">,</span> <span class="mi">128</span><span class="p">),</span> <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>

<span class="n">grid_spec</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">PrefetchScalarGridSpec</span><span class="p">(</span>
    <span class="n">num_scalar_prefetch</span><span class="o">=</span><span class="mi">0</span><span class="p">,</span>
    <span class="n">in_specs</span><span class="o">=</span><span class="p">[</span><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">ANY</span><span class="p">)],</span>
    <span class="n">out_specs</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">ANY</span><span class="p">),</span>
    <span class="n">scratch_shapes</span><span class="o">=</span><span class="p">(</span>
        <span class="p">[</span><span class="n">pltpu</span><span class="p">.</span><span class="n">SemaphoreType</span><span class="p">.</span><span class="n">DMA</span><span class="p">]</span> <span class="o">*</span> <span class="mi">2</span>
        <span class="o">+</span> <span class="p">[</span><span class="n">pltpu</span><span class="p">.</span><span class="n">SemaphoreType</span><span class="p">.</span><span class="n">DMA</span><span class="p">((</span><span class="n">num_devices</span> <span class="o">-</span> <span class="mi">1</span><span class="p">,))]</span>
    <span class="p">),</span>
    <span class="n">grid</span><span class="o">=</span><span class="p">(</span><span class="n">num_devices</span> <span class="o">-</span> <span class="mi">1</span><span class="p">,),</span>
<span class="p">)</span>

<span class="n">all_gather</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">pallas_call</span><span class="p">(</span>
    <span class="n">all_gather_kernel</span><span class="p">,</span>
    <span class="n">out_shape</span><span class="o">=</span><span class="n">out_shape</span><span class="p">,</span>
    <span class="n">grid_spec</span><span class="o">=</span><span class="n">grid_spec</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">pallas_result</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">jit</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">shard_map</span><span class="p">(</span>
        <span class="n">all_gather</span><span class="p">,</span>
        <span class="n">mesh</span><span class="o">=</span><span class="n">mesh</span><span class="p">,</span>
        <span class="n">in_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
        <span class="n">out_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
        <span class="n">check_vma</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
    <span class="p">)</span>
<span class="p">)(</span><span class="n">input_arr</span><span class="p">)</span>

<span class="n">xla_result</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">jit</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">shard_map</span><span class="p">(</span>
        <span class="k">lambda</span> <span class="n">x</span><span class="p">:</span> <span class="n">lax</span><span class="p">.</span><span class="n">all_gather</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="s">"x"</span><span class="p">),</span>
        <span class="n">mesh</span><span class="o">=</span><span class="n">mesh</span><span class="p">,</span>
        <span class="n">in_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
        <span class="n">out_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
    <span class="p">)</span>
<span class="p">)(</span><span class="n">input_arr</span><span class="p">)</span>
</code></pre></div></div>

<h3 id="24-关键代码解释">2.4 关键代码解释</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">grid</span><span class="o">=</span><span class="p">(</span><span class="n">num_devices</span> <span class="o">-</span> <span class="mi">1</span><span class="p">,)</span>
</code></pre></div></div>

<p>每个 device 自己已经有自己的 shard，只需要再收到其他 <code class="language-plaintext highlighter-rouge">num_devices - 1</code> 个 shard。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">slot</span> <span class="o">=</span> <span class="p">(</span><span class="n">my_id</span> <span class="o">-</span> <span class="n">step</span><span class="p">)</span> <span class="o">%</span> <span class="n">num_devices</span>
</code></pre></div></div>

<p>当前这一轮要把哪个 slot 转发给右邻居。</p>

<p>例如 D0：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>step 0: 发送 slot 0
step 1: 发送 slot 3
step 2: 发送 slot 2
</code></pre></div></div>

<p>同时 D0 会从 D3 收到：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>step 0: slot 3
step 1: slot 2
step 2: slot 1
</code></pre></div></div>

<h3 id="25-遍历模拟">2.5 遍历模拟</h3>

<p>表格展示每轮结束后，每个 device 的 output slots。</p>

<table>
  <thead>
    <tr>
      <th>loop 后</th>
      <th>D0 slots</th>
      <th>D1 slots</th>
      <th>D2 slots</th>
      <th>D3 slots</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>初始</td>
      <td><code class="language-plaintext highlighter-rouge">[-,-,-,-]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[-,-,-,-]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[-,-,-,-]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[-,-,-,-]</code></td>
    </tr>
    <tr>
      <td>step 0</td>
      <td><code class="language-plaintext highlighter-rouge">[x0,-,-,x3]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[x0,x1,-,-]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[-,x1,x2,-]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[-,-,x2,x3]</code></td>
    </tr>
    <tr>
      <td>step 1</td>
      <td><code class="language-plaintext highlighter-rouge">[x0,-,x2,x3]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[x0,x1,-,x3]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[x0,x1,x2,-]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[-,x1,x2,x3]</code></td>
    </tr>
    <tr>
      <td>step 2</td>
      <td><code class="language-plaintext highlighter-rouge">[x0,x1,x2,x3]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[x0,x1,x2,x3]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[x0,x1,x2,x3]</code></td>
      <td><code class="language-plaintext highlighter-rouge">[x0,x1,x2,x3]</code></td>
    </tr>
  </tbody>
</table>

<p>操作来源：</p>

<table>
  <thead>
    <tr>
      <th>动作</th>
      <th>代码</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>step 0 本地写自己的 slot</td>
      <td><code class="language-plaintext highlighter-rouge">out_ref.at[my_id] = x_ref</code> via <code class="language-plaintext highlighter-rouge">make_async_copy</code></td>
    </tr>
    <tr>
      <td>每轮转发已有 slot</td>
      <td><code class="language-plaintext highlighter-rouge">src_ref=out_ref.at[slot]</code></td>
    </tr>
    <tr>
      <td>写到右邻居同编号 slot</td>
      <td><code class="language-plaintext highlighter-rouge">dst_ref=out_ref.at[slot]</code> with <code class="language-plaintext highlighter-rouge">device_id=(right,)</code></td>
    </tr>
  </tbody>
</table>

<p>为什么 <code class="language-plaintext highlighter-rouge">recv_sems</code> 有多个？</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">recv_sems</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">step</span><span class="p">]</span>
</code></pre></div></div>

<p>每一轮一个 receive semaphore，避免多个 in-flight DMA 复用同一个接收计数器。官方文档也提醒，大 kernel 中 device 更容易不同步，复用 receive semaphore 可能导致静默错误。</p>

<hr />

<h2 id="3-laxpsum单向-ring-all-reduce-sum">3. <code class="language-plaintext highlighter-rouge">lax.psum</code>：单向 ring all-reduce sum</h2>

<h3 id="31-直观解释">3.1 直观解释</h3>

<p><code class="language-plaintext highlighter-rouge">psum</code> 是 all-reduce sum：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>每个 device 一开始有自己的 x。
最后每个 device 都得到 x0 + x1 + x2 + x3。
</code></pre></div></div>

<p>这个 Pallas 示例不是先 all-gather 再求和，而是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>数据沿 ring 传递。
每个 device 收到一份数据就加进自己的 accumulator。
</code></pre></div></div>

<p>官方图：</p>

<p><img src="https://docs.jax.dev/en/latest/_images/reduce_sum_2.svg" alt="reduce_sum_2" /></p>

<h3 id="32-配图">3.2 配图</h3>

<pre><code class="language-mermaid">flowchart LR
  D0["D0: x0"] --&gt; D1["D1 accum"]
  D1 --&gt; D2["D2 accum"]
  D2 --&gt; D3["D3 accum"]
  D3 --&gt; D0R["D0 accum"]
</code></pre>

<p>每个 device 都在做同样的事情，所以一圈后每个 device 都累加了所有输入。</p>

<h3 id="33-教学版完整代码">3.3 教学版完整代码</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">partition</span> <span class="o">=</span> <span class="n">P</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="s">"x"</span><span class="p">)</span>
<span class="n">mesh</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">make_mesh</span><span class="p">((</span><span class="n">num_devices</span><span class="p">,),</span> <span class="p">(</span><span class="s">"x"</span><span class="p">,))</span>
<span class="n">sharding</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">sharding</span><span class="p">.</span><span class="n">NamedSharding</span><span class="p">(</span><span class="n">mesh</span><span class="p">,</span> <span class="n">partition</span><span class="p">)</span>

<span class="n">input_arr</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">uniform</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">key</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span>
    <span class="n">shape</span><span class="o">=</span><span class="p">(</span><span class="mi">8</span><span class="p">,</span> <span class="mi">128</span> <span class="o">*</span> <span class="n">num_devices</span><span class="p">),</span>
<span class="p">)</span>
<span class="n">input_arr</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">device_put</span><span class="p">(</span><span class="n">input_arr</span><span class="p">,</span> <span class="n">sharding</span><span class="p">)</span>


<span class="k">def</span> <span class="nf">neighbor_barrier</span><span class="p">(</span><span class="n">left</span><span class="p">,</span> <span class="n">right</span><span class="p">,</span> <span class="n">double_barrier</span><span class="o">=</span><span class="bp">True</span><span class="p">):</span>
    <span class="n">barrier</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">get_barrier_semaphore</span><span class="p">()</span>
    <span class="k">for</span> <span class="n">nbr</span> <span class="ow">in</span> <span class="p">[</span><span class="n">left</span><span class="p">,</span> <span class="n">right</span><span class="p">]:</span>
        <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_signal</span><span class="p">(</span>
            <span class="n">barrier</span><span class="p">,</span>
            <span class="n">inc</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span>
            <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">nbr</span><span class="p">,),</span>
            <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
        <span class="p">)</span>
    <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_wait</span><span class="p">(</span><span class="n">barrier</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>

    <span class="k">if</span> <span class="n">double_barrier</span><span class="p">:</span>
        <span class="o">@</span><span class="n">functools</span><span class="p">.</span><span class="n">partial</span><span class="p">(</span><span class="n">pl</span><span class="p">.</span><span class="n">run_scoped</span><span class="p">,</span> <span class="n">second</span><span class="o">=</span><span class="n">pltpu</span><span class="p">.</span><span class="n">SemaphoreType</span><span class="p">.</span><span class="n">REGULAR</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_</span><span class="p">(</span><span class="n">second</span><span class="p">):</span>
            <span class="k">for</span> <span class="n">nbr</span> <span class="ow">in</span> <span class="p">[</span><span class="n">left</span><span class="p">,</span> <span class="n">right</span><span class="p">]:</span>
                <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_signal</span><span class="p">(</span>
                    <span class="n">second</span><span class="p">,</span>
                    <span class="n">inc</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span>
                    <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">nbr</span><span class="p">,),</span>
                    <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
                <span class="p">)</span>
            <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_wait</span><span class="p">(</span><span class="n">second</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>


<span class="k">def</span> <span class="nf">psum_kernel</span><span class="p">(</span>
    <span class="n">x_ref</span><span class="p">,</span>
    <span class="n">o_ref</span><span class="p">,</span>
    <span class="n">hbm_buf</span><span class="p">,</span>
    <span class="n">local_copy_sem</span><span class="p">,</span>
    <span class="n">remote_recv_sem</span><span class="p">,</span>
    <span class="n">remote_send_sem</span><span class="p">,</span>
    <span class="n">capacity_sem</span><span class="p">,</span>
    <span class="n">vmem_recv</span><span class="p">,</span>
<span class="p">):</span>
    <span class="n">step</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>
    <span class="n">working</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">rem</span><span class="p">(</span><span class="n">step</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
    <span class="n">receiving</span> <span class="o">=</span> <span class="mi">1</span> <span class="o">-</span> <span class="n">working</span>

    <span class="n">my_id</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">axis_index</span><span class="p">(</span><span class="s">"x"</span><span class="p">)</span>
    <span class="n">right</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">rem</span><span class="p">(</span><span class="n">my_id</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>
    <span class="n">left</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">rem</span><span class="p">(</span><span class="n">my_id</span> <span class="o">-</span> <span class="mi">1</span> <span class="o">+</span> <span class="n">num_devices</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">step</span> <span class="o">==</span> <span class="mi">0</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
        <span class="n">neighbor_barrier</span><span class="p">(</span><span class="n">left</span><span class="p">,</span> <span class="n">right</span><span class="p">)</span>
        <span class="n">o_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">zeros_like</span><span class="p">(</span><span class="n">o_ref</span><span class="p">)</span>
        <span class="n">vmem_recv</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">zeros_like</span><span class="p">(</span><span class="n">vmem_recv</span><span class="p">)</span>

        <span class="n">first</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_remote_copy</span><span class="p">(</span>
            <span class="n">src_ref</span><span class="o">=</span><span class="n">x_ref</span><span class="p">,</span>
            <span class="n">dst_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">working</span><span class="p">],</span>
            <span class="n">send_sem</span><span class="o">=</span><span class="n">remote_send_sem</span><span class="p">,</span>
            <span class="n">recv_sem</span><span class="o">=</span><span class="n">remote_recv_sem</span><span class="p">,</span>
            <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">right</span><span class="p">,),</span>
            <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
        <span class="p">)</span>
        <span class="n">first</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
        <span class="n">first</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>

    <span class="c1"># 告诉左邻居：我已经准备好接收它下一次写入。
</span>    <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_signal</span><span class="p">(</span>
        <span class="n">capacity_sem</span><span class="p">,</span>
        <span class="n">inc</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span>
        <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">left</span><span class="p">,),</span>
        <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
    <span class="p">)</span>

    <span class="c1"># 本地 HBM -&gt; VMEM，用于累加。
</span>    <span class="n">local</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">working</span><span class="p">],</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">vmem_recv</span><span class="p">,</span>
        <span class="n">sem</span><span class="o">=</span><span class="n">local_copy_sem</span><span class="p">,</span>
    <span class="p">)</span>
    <span class="n">local</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>

    <span class="c1"># 写右邻居前，先等右邻居说它准备好了。
</span>    <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_wait</span><span class="p">(</span><span class="n">capacity_sem</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span>

    <span class="n">remote</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_remote_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">working</span><span class="p">],</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">receiving</span><span class="p">],</span>
        <span class="n">send_sem</span><span class="o">=</span><span class="n">remote_send_sem</span><span class="p">,</span>
        <span class="n">recv_sem</span><span class="o">=</span><span class="n">remote_recv_sem</span><span class="p">,</span>
        <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">right</span><span class="p">,),</span>
        <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
    <span class="p">)</span>
    <span class="n">remote</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>

    <span class="n">local</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>
    <span class="n">o_ref</span><span class="p">[...]</span> <span class="o">+=</span> <span class="n">vmem_recv</span><span class="p">[...]</span>
    <span class="n">remote</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>


<span class="n">out_shape</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">ShapeDtypeStruct</span><span class="p">((</span><span class="mi">8</span><span class="p">,</span> <span class="mi">128</span><span class="p">),</span> <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">ShapeDtypeStruct</span><span class="p">((</span><span class="mi">2</span><span class="p">,</span> <span class="mi">8</span><span class="p">,</span> <span class="mi">128</span><span class="p">),</span> <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
<span class="p">)</span>

<span class="n">grid_spec</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">PrefetchScalarGridSpec</span><span class="p">(</span>
    <span class="n">num_scalar_prefetch</span><span class="o">=</span><span class="mi">0</span><span class="p">,</span>
    <span class="n">in_specs</span><span class="o">=</span><span class="p">[</span><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pltpu</span><span class="p">.</span><span class="n">VMEM</span><span class="p">)],</span>
    <span class="n">out_specs</span><span class="o">=</span><span class="p">[</span>
        <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pltpu</span><span class="p">.</span><span class="n">VMEM</span><span class="p">),</span>
        <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">ANY</span><span class="p">),</span>
    <span class="p">],</span>
    <span class="n">grid</span><span class="o">=</span><span class="p">(</span><span class="n">num_devices</span><span class="p">,),</span>
    <span class="n">scratch_shapes</span><span class="o">=</span><span class="p">(</span>
        <span class="p">[</span><span class="n">pltpu</span><span class="p">.</span><span class="n">SemaphoreType</span><span class="p">.</span><span class="n">DMA</span><span class="p">]</span> <span class="o">*</span> <span class="mi">3</span>
        <span class="o">+</span> <span class="p">[</span><span class="n">pltpu</span><span class="p">.</span><span class="n">SemaphoreType</span><span class="p">.</span><span class="n">REGULAR</span><span class="p">]</span>
        <span class="o">+</span> <span class="p">[</span><span class="n">pltpu</span><span class="p">.</span><span class="n">VMEM</span><span class="p">((</span><span class="mi">8</span><span class="p">,</span> <span class="mi">128</span><span class="p">),</span> <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">)]</span>
    <span class="p">),</span>
<span class="p">)</span>

<span class="n">psum_pallas</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">pallas_call</span><span class="p">(</span>
    <span class="n">psum_kernel</span><span class="p">,</span>
    <span class="n">out_shape</span><span class="o">=</span><span class="n">out_shape</span><span class="p">,</span>
    <span class="n">grid_spec</span><span class="o">=</span><span class="n">grid_spec</span><span class="p">,</span>
    <span class="n">compiler_params</span><span class="o">=</span><span class="n">pltpu</span><span class="p">.</span><span class="n">CompilerParams</span><span class="p">(</span><span class="n">collective_id</span><span class="o">=</span><span class="mi">0</span><span class="p">),</span>
<span class="p">)</span>

<span class="n">pallas_result</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">jit</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">shard_map</span><span class="p">(</span>
        <span class="n">psum_pallas</span><span class="p">,</span>
        <span class="n">mesh</span><span class="o">=</span><span class="n">mesh</span><span class="p">,</span>
        <span class="n">in_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
        <span class="n">out_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
        <span class="n">check_vma</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
    <span class="p">)</span>
<span class="p">)(</span><span class="n">input_arr</span><span class="p">)[</span><span class="mi">0</span><span class="p">]</span>

<span class="n">xla_result</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">jit</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">shard_map</span><span class="p">(</span>
        <span class="k">lambda</span> <span class="n">x</span><span class="p">:</span> <span class="n">lax</span><span class="p">.</span><span class="n">psum</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="s">"x"</span><span class="p">),</span>
        <span class="n">mesh</span><span class="o">=</span><span class="n">mesh</span><span class="p">,</span>
        <span class="n">in_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
        <span class="n">out_specs</span><span class="o">=</span><span class="n">partition</span><span class="p">,</span>
    <span class="p">)</span>
<span class="p">)(</span><span class="n">input_arr</span><span class="p">)</span>
</code></pre></div></div>

<h3 id="34-关键区域">3.4 关键区域</h3>

<table>
  <thead>
    <tr>
      <th>区域</th>
      <th>memory space</th>
      <th>用途</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">x_ref</code></td>
      <td>VMEM</td>
      <td>本 device 的原始输入</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">o_ref</code></td>
      <td>VMEM</td>
      <td>本 device 的累加结果</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">hbm_buf[0]</code></td>
      <td>HBM</td>
      <td>通信双缓冲 slot 0</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">hbm_buf[1]</code></td>
      <td>HBM</td>
      <td>通信双缓冲 slot 1</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">vmem_recv</code></td>
      <td>VMEM</td>
      <td>从 HBM 拷出来后用于累加的临时 buffer</td>
    </tr>
  </tbody>
</table>

<h3 id="35-遍历模拟">3.5 遍历模拟</h3>

<p><code class="language-plaintext highlighter-rouge">H0 = hbm_buf[0]</code>，<code class="language-plaintext highlighter-rouge">H1 = hbm_buf[1]</code>。</p>

<table>
  <thead>
    <tr>
      <th>loop 后</th>
      <th>D0 <code class="language-plaintext highlighter-rouge">(H0,H1,o)</code></th>
      <th>D1 <code class="language-plaintext highlighter-rouge">(H0,H1,o)</code></th>
      <th>D2 <code class="language-plaintext highlighter-rouge">(H0,H1,o)</code></th>
      <th>D3 <code class="language-plaintext highlighter-rouge">(H0,H1,o)</code></th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>prologue 后</td>
      <td><code class="language-plaintext highlighter-rouge">(x3,-,0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x0,-,0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x1,-,0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x2,-,0)</code></td>
    </tr>
    <tr>
      <td>step 0</td>
      <td><code class="language-plaintext highlighter-rouge">(x3,x2,x3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x0,x3,x0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x1,x0,x1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x2,x1,x2)</code></td>
    </tr>
    <tr>
      <td>step 1</td>
      <td><code class="language-plaintext highlighter-rouge">(x1,x2,x3+x2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x2,x3,x0+x3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x3,x0,x1+x0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x0,x1,x2+x1)</code></td>
    </tr>
    <tr>
      <td>step 2</td>
      <td><code class="language-plaintext highlighter-rouge">(x1,x0,x3+x2+x1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x2,x1,x0+x3+x2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x3,x2,x1+x0+x3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x0,x3,x2+x1+x0)</code></td>
    </tr>
    <tr>
      <td>step 3</td>
      <td><code class="language-plaintext highlighter-rouge">(x3,x0,all)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x0,x1,all)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x1,x2,all)</code></td>
      <td><code class="language-plaintext highlighter-rouge">(x2,x3,all)</code></td>
    </tr>
  </tbody>
</table>

<p>其中：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>all = x0 + x1 + x2 + x3
</code></pre></div></div>

<p>每轮核心操作来源：</p>

<table>
  <thead>
    <tr>
      <th>动作</th>
      <th>代码</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>第 0 轮先发自己的输入给右邻居</td>
      <td><code class="language-plaintext highlighter-rouge">first = make_async_remote_copy(src_ref=x_ref, dst_ref=hbm_buf.at[working])</code></td>
    </tr>
    <tr>
      <td>读当前 working slot 到 VMEM</td>
      <td><code class="language-plaintext highlighter-rouge">make_async_copy(src_ref=hbm_buf.at[working], dst_ref=vmem_recv)</code></td>
    </tr>
    <tr>
      <td>把当前 working slot 继续发给右邻居</td>
      <td><code class="language-plaintext highlighter-rouge">make_async_remote_copy(src_ref=hbm_buf.at[working], dst_ref=hbm_buf.at[receiving])</code></td>
    </tr>
    <tr>
      <td>累加</td>
      <td><code class="language-plaintext highlighter-rouge">o_ref[...] += vmem_recv[...]</code></td>
    </tr>
    <tr>
      <td>防止跑快一轮覆盖邻居 working slot</td>
      <td><code class="language-plaintext highlighter-rouge">capacity_sem</code> signal/wait</td>
    </tr>
  </tbody>
</table>

<p><code class="language-plaintext highlighter-rouge">capacity_sem</code> 的作用：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>如果所有 device 严格同步，双缓冲不会冲突。
但不同 device 可以跑快或跑慢。
快的 device 下一轮 receiving_slot 可能正好是慢邻居还在读的 working_slot。
capacity_sem 让发送者必须等接收者确认“我已经进入对应轮次，可以写我的 receiving slot”。
</code></pre></div></div>

<hr />

<h2 id="4-laxpsum_scatter双向-reduce-scatter">4. <code class="language-plaintext highlighter-rouge">lax.psum_scatter</code>：双向 reduce-scatter</h2>

<h3 id="41-直观解释">4.1 直观解释</h3>

<p><code class="language-plaintext highlighter-rouge">psum_scatter</code> 的语义可以看成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>先对所有 device 的输入按 block 求和；
然后每个 device 只保留属于自己的那个 block。
</code></pre></div></div>

<p>但高效实现不是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>all-reduce 完整结果 -&gt; 再 scatter
</code></pre></div></div>

<p>而是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>每个 output block 的 partial sum 在 ring 上移动；
经过一个 device，就加上该 device 对这个 block 的贡献；
最后回到目标 device 时，这个 block 已经完整。
</code></pre></div></div>

<p>官方语义图：</p>

<p><img src="https://docs.jax.dev/en/latest/_images/reduce_scatter_1.svg" alt="reduce_scatter_1" /></p>

<p>官方通信图：</p>

<p><img src="https://docs.jax.dev/en/latest/_images/reduce_scatter_2.svg" alt="reduce_scatter_2" /></p>

<h3 id="42-配图">4.2 配图</h3>

<p>一个 block 被切成上下两半：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>T = top half    -&gt; 向左传
B = bottom half -&gt; 向右传
</code></pre></div></div>

<pre><code class="language-mermaid">flowchart LR
  subgraph Top["top half: 向左"]
    T0["D0 starts T0"] --&gt; T3["D3 adds"] --&gt; T2["D2 adds"] --&gt; T1["D1 adds"] --&gt; T0R["D0 receives full T0"]
  end
  subgraph Bottom["bottom half: 向右"]
    B0["D0 starts B0"] --&gt; B1["D1 adds"] --&gt; B2["D2 adds"] --&gt; B3["D3 adds"] --&gt; B0R["D0 receives full B0"]
  end
</code></pre>

<p>注意：这段教学代码的数据流是双向的，但 phase 调度是交替的：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>phase LEFT:
  发送上一阶段的 right-half
  计算当前 left-half

phase RIGHT:
  发送刚算好的 left-half
  计算当前 right-half
</code></pre></div></div>

<h3 id="43-教学版完整代码">4.3 教学版完整代码</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">partition</span> <span class="o">=</span> <span class="n">P</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="s">"x"</span><span class="p">)</span>
<span class="n">mesh</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">make_mesh</span><span class="p">((</span><span class="n">num_devices</span><span class="p">,),</span> <span class="p">(</span><span class="s">"x"</span><span class="p">,))</span>
<span class="n">sharding</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">sharding</span><span class="p">.</span><span class="n">NamedSharding</span><span class="p">(</span><span class="n">mesh</span><span class="p">,</span> <span class="n">partition</span><span class="p">)</span>

<span class="n">block_size</span> <span class="o">=</span> <span class="p">(</span><span class="mi">16</span><span class="p">,</span> <span class="mi">128</span><span class="p">)</span>
<span class="n">input_arr</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">uniform</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">key</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span>
    <span class="n">shape</span><span class="o">=</span><span class="p">(</span><span class="n">block_size</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="o">*</span> <span class="n">num_devices</span><span class="p">,</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">*</span> <span class="n">num_devices</span><span class="p">),</span>
<span class="p">)</span>
<span class="n">input_arr</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">device_put</span><span class="p">(</span><span class="n">input_arr</span><span class="p">,</span> <span class="n">sharding</span><span class="p">)</span>

<span class="n">LEFT</span> <span class="o">=</span> <span class="mi">0</span>
<span class="n">RIGHT</span> <span class="o">=</span> <span class="mi">1</span>


<span class="k">def</span> <span class="nf">mod</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">n</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">lax</span><span class="p">.</span><span class="n">rem</span><span class="p">(</span><span class="n">x</span> <span class="o">+</span> <span class="n">n</span><span class="p">,</span> <span class="n">n</span><span class="p">)</span>


<span class="k">def</span> <span class="nf">signal</span><span class="p">(</span><span class="n">direction</span><span class="p">,</span> <span class="n">sem</span><span class="p">):</span>
    <span class="n">my_id</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">axis_index</span><span class="p">(</span><span class="s">"x"</span><span class="p">)</span>
    <span class="k">if</span> <span class="n">direction</span> <span class="o">==</span> <span class="n">LEFT</span><span class="p">:</span>
        <span class="n">target</span> <span class="o">=</span> <span class="n">mod</span><span class="p">(</span><span class="n">my_id</span> <span class="o">-</span> <span class="mi">1</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="n">target</span> <span class="o">=</span> <span class="n">mod</span><span class="p">(</span><span class="n">my_id</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>
    <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_signal</span><span class="p">(</span>
        <span class="n">sem</span><span class="p">,</span>
        <span class="n">inc</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span>
        <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">target</span><span class="p">,),</span>
        <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
    <span class="p">)</span>


<span class="k">def</span> <span class="nf">reduce_scatter_kernel</span><span class="p">(</span>
    <span class="n">x_ref</span><span class="p">,</span>
    <span class="n">o_ref</span><span class="p">,</span>
    <span class="n">hbm_buf</span><span class="p">,</span>
    <span class="n">local_copy_sem</span><span class="p">,</span>
    <span class="n">left_recv_sem</span><span class="p">,</span>
    <span class="n">left_send_sem</span><span class="p">,</span>
    <span class="n">right_recv_sem</span><span class="p">,</span>
    <span class="n">right_send_sem</span><span class="p">,</span>
    <span class="n">left_capacity_sem</span><span class="p">,</span>
    <span class="n">right_capacity_sem</span><span class="p">,</span>
    <span class="n">accum</span><span class="p">,</span>
<span class="p">):</span>
    <span class="n">step</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>
    <span class="n">phase</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">program_id</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>

    <span class="n">is_first</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">logical_and</span><span class="p">(</span><span class="n">step</span> <span class="o">==</span> <span class="mi">0</span><span class="p">,</span> <span class="n">phase</span> <span class="o">==</span> <span class="n">LEFT</span><span class="p">)</span>
    <span class="n">is_last_step</span> <span class="o">=</span> <span class="n">step</span> <span class="o">==</span> <span class="n">pl</span><span class="p">.</span><span class="n">num_programs</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span> <span class="o">-</span> <span class="mi">1</span>

    <span class="n">working</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">rem</span><span class="p">(</span><span class="n">step</span><span class="p">,</span> <span class="mi">2</span><span class="p">)</span>
    <span class="n">receiving</span> <span class="o">=</span> <span class="mi">1</span> <span class="o">-</span> <span class="n">working</span>

    <span class="n">my_id</span> <span class="o">=</span> <span class="n">lax</span><span class="p">.</span><span class="n">axis_index</span><span class="p">(</span><span class="s">"x"</span><span class="p">)</span>
    <span class="n">right</span> <span class="o">=</span> <span class="n">mod</span><span class="p">(</span><span class="n">my_id</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>
    <span class="n">left</span> <span class="o">=</span> <span class="n">mod</span><span class="p">(</span><span class="n">my_id</span> <span class="o">-</span> <span class="mi">1</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>

    <span class="n">left_block</span> <span class="o">=</span> <span class="n">mod</span><span class="p">(</span><span class="n">my_id</span> <span class="o">+</span> <span class="n">step</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>
    <span class="n">right_block</span> <span class="o">=</span> <span class="n">mod</span><span class="p">(</span><span class="n">my_id</span> <span class="o">-</span> <span class="n">step</span> <span class="o">-</span> <span class="mi">1</span><span class="p">,</span> <span class="n">num_devices</span><span class="p">)</span>

    <span class="n">half_rows</span> <span class="o">=</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="o">//</span> <span class="mi">2</span>
    <span class="n">top</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">ds</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">half_rows</span><span class="p">)</span>
    <span class="n">bottom</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">ds</span><span class="p">(</span><span class="n">half_rows</span><span class="p">,</span> <span class="n">half_rows</span><span class="p">)</span>
    <span class="n">current_half</span> <span class="o">=</span> <span class="n">pl</span><span class="p">.</span><span class="n">ds</span><span class="p">(</span><span class="n">phase</span> <span class="o">*</span> <span class="n">half_rows</span><span class="p">,</span> <span class="n">half_rows</span><span class="p">)</span>

    <span class="n">init_left</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_remote_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">x_ref</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">my_id</span><span class="p">,</span> <span class="n">top</span><span class="p">],</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">working</span><span class="p">,</span> <span class="n">top</span><span class="p">],</span>
        <span class="n">send_sem</span><span class="o">=</span><span class="n">left_send_sem</span><span class="p">,</span>
        <span class="n">recv_sem</span><span class="o">=</span><span class="n">left_recv_sem</span><span class="p">,</span>
        <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">left</span><span class="p">,),</span>
        <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
    <span class="p">)</span>

    <span class="n">init_right</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_remote_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">x_ref</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">my_id</span><span class="p">,</span> <span class="n">bottom</span><span class="p">],</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">working</span><span class="p">,</span> <span class="n">bottom</span><span class="p">],</span>
        <span class="n">send_sem</span><span class="o">=</span><span class="n">right_send_sem</span><span class="p">,</span>
        <span class="n">recv_sem</span><span class="o">=</span><span class="n">right_recv_sem</span><span class="p">,</span>
        <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">right</span><span class="p">,),</span>
        <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
    <span class="p">)</span>

    <span class="n">send_left</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_remote_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">working</span><span class="p">,</span> <span class="n">top</span><span class="p">],</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">receiving</span><span class="p">,</span> <span class="n">top</span><span class="p">],</span>
        <span class="n">send_sem</span><span class="o">=</span><span class="n">left_send_sem</span><span class="p">,</span>
        <span class="n">recv_sem</span><span class="o">=</span><span class="n">left_recv_sem</span><span class="p">,</span>
        <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">left</span><span class="p">,),</span>
        <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
    <span class="p">)</span>

    <span class="n">send_right</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_remote_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">receiving</span><span class="p">,</span> <span class="n">bottom</span><span class="p">],</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">working</span><span class="p">,</span> <span class="n">bottom</span><span class="p">],</span>
        <span class="n">send_sem</span><span class="o">=</span><span class="n">right_send_sem</span><span class="p">,</span>
        <span class="n">recv_sem</span><span class="o">=</span><span class="n">right_recv_sem</span><span class="p">,</span>
        <span class="n">device_id</span><span class="o">=</span><span class="p">(</span><span class="n">right</span><span class="p">,),</span>
        <span class="n">device_id_type</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">DeviceIdType</span><span class="p">.</span><span class="n">MESH</span><span class="p">,</span>
    <span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">is_first</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
        <span class="n">neighbor_barrier</span><span class="p">(</span><span class="n">left</span><span class="p">,</span> <span class="n">right</span><span class="p">)</span>
        <span class="n">o_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">zeros_like</span><span class="p">(</span><span class="n">o_ref</span><span class="p">)</span>
        <span class="n">accum</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">zeros_like</span><span class="p">(</span><span class="n">accum</span><span class="p">)</span>

        <span class="n">init_left</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
        <span class="n">init_left</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>
        <span class="n">init_right</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>

        <span class="n">signal</span><span class="p">(</span><span class="n">LEFT</span><span class="p">,</span> <span class="n">right_capacity_sem</span><span class="p">)</span>
        <span class="n">signal</span><span class="p">(</span><span class="n">RIGHT</span><span class="p">,</span> <span class="n">left_capacity_sem</span><span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="o">~</span><span class="n">is_first</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
        <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">phase</span> <span class="o">==</span> <span class="n">LEFT</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
            <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_wait</span><span class="p">(</span><span class="n">right_capacity_sem</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span>
            <span class="n">send_right</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>

        <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">phase</span> <span class="o">==</span> <span class="n">RIGHT</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
            <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_wait</span><span class="p">(</span><span class="n">left_capacity_sem</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span>
            <span class="n">send_left</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>

    <span class="n">local</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">working</span><span class="p">,</span> <span class="n">current_half</span><span class="p">],</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">accum</span><span class="p">,</span>
        <span class="n">sem</span><span class="o">=</span><span class="n">local_copy_sem</span><span class="p">,</span>
    <span class="p">)</span>
    <span class="n">local</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
    <span class="n">local</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="o">~</span><span class="n">is_last_step</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
        <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">phase</span> <span class="o">==</span> <span class="n">LEFT</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
            <span class="n">accum</span><span class="p">[...]</span> <span class="o">+=</span> <span class="n">x_ref</span><span class="p">[</span><span class="n">left_block</span><span class="p">,</span> <span class="n">top</span><span class="p">]</span>

        <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">phase</span> <span class="o">==</span> <span class="n">RIGHT</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
            <span class="n">accum</span><span class="p">[...]</span> <span class="o">+=</span> <span class="n">x_ref</span><span class="p">[</span><span class="n">right_block</span><span class="p">,</span> <span class="n">bottom</span><span class="p">]</span>

    <span class="n">local</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">make_async_copy</span><span class="p">(</span>
        <span class="n">src_ref</span><span class="o">=</span><span class="n">accum</span><span class="p">,</span>
        <span class="n">dst_ref</span><span class="o">=</span><span class="n">hbm_buf</span><span class="p">.</span><span class="n">at</span><span class="p">[</span><span class="n">working</span><span class="p">,</span> <span class="n">current_half</span><span class="p">],</span>
        <span class="n">sem</span><span class="o">=</span><span class="n">local_copy_sem</span><span class="p">,</span>
    <span class="p">)</span>
    <span class="n">local</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
    <span class="n">local</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">is_first</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
        <span class="n">init_right</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="o">~</span><span class="n">is_first</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
        <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">phase</span> <span class="o">==</span> <span class="n">LEFT</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
            <span class="n">send_right</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>
            <span class="n">signal</span><span class="p">(</span><span class="n">LEFT</span><span class="p">,</span> <span class="n">right_capacity_sem</span><span class="p">)</span>

        <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">phase</span> <span class="o">==</span> <span class="n">RIGHT</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
            <span class="n">send_left</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>
            <span class="n">signal</span><span class="p">(</span><span class="n">RIGHT</span><span class="p">,</span> <span class="n">left_capacity_sem</span><span class="p">)</span>

    <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">is_last_step</span><span class="p">)</span>
    <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
        <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">phase</span> <span class="o">==</span> <span class="n">LEFT</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
            <span class="n">o_ref</span><span class="p">[</span><span class="n">top</span><span class="p">,</span> <span class="p">...]</span> <span class="o">=</span> <span class="n">accum</span><span class="p">[...]</span>
            <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_wait</span><span class="p">(</span><span class="n">right_capacity_sem</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span>

        <span class="o">@</span><span class="n">pl</span><span class="p">.</span><span class="n">when</span><span class="p">(</span><span class="n">phase</span> <span class="o">==</span> <span class="n">RIGHT</span><span class="p">)</span>
        <span class="k">def</span> <span class="nf">_</span><span class="p">():</span>
            <span class="n">o_ref</span><span class="p">[</span><span class="n">bottom</span><span class="p">,</span> <span class="p">...]</span> <span class="o">=</span> <span class="n">accum</span><span class="p">[...]</span>
            <span class="n">pl</span><span class="p">.</span><span class="n">semaphore_wait</span><span class="p">(</span><span class="n">left_capacity_sem</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span>


<span class="n">out_shape</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">ShapeDtypeStruct</span><span class="p">((</span><span class="n">block_size</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">1</span><span class="p">]),</span> <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">ShapeDtypeStruct</span><span class="p">((</span><span class="mi">2</span><span class="p">,</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">1</span><span class="p">]),</span> <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
<span class="p">)</span>

<span class="n">grid_spec</span> <span class="o">=</span> <span class="n">pltpu</span><span class="p">.</span><span class="n">PrefetchScalarGridSpec</span><span class="p">(</span>
    <span class="n">num_scalar_prefetch</span><span class="o">=</span><span class="mi">0</span><span class="p">,</span>
    <span class="n">in_specs</span><span class="o">=</span><span class="p">[</span><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pltpu</span><span class="p">.</span><span class="n">VMEM</span><span class="p">)],</span>
    <span class="n">out_specs</span><span class="o">=</span><span class="p">[</span>
        <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pltpu</span><span class="p">.</span><span class="n">VMEM</span><span class="p">),</span>
        <span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span><span class="n">memory_space</span><span class="o">=</span><span class="n">pl</span><span class="p">.</span><span class="n">ANY</span><span class="p">),</span>
    <span class="p">],</span>
    <span class="n">grid</span><span class="o">=</span><span class="p">(</span><span class="n">num_devices</span><span class="p">,</span> <span class="mi">2</span><span class="p">),</span>
    <span class="n">scratch_shapes</span><span class="o">=</span><span class="p">(</span>
        <span class="p">[</span><span class="n">pltpu</span><span class="p">.</span><span class="n">SemaphoreType</span><span class="p">.</span><span class="n">DMA</span><span class="p">]</span> <span class="o">*</span> <span class="mi">5</span>
        <span class="o">+</span> <span class="p">[</span><span class="n">pltpu</span><span class="p">.</span><span class="n">SemaphoreType</span><span class="p">.</span><span class="n">REGULAR</span><span class="p">]</span> <span class="o">*</span> <span class="mi">2</span>
        <span class="o">+</span> <span class="p">[</span><span class="n">pltpu</span><span class="p">.</span><span class="n">VMEM</span><span class="p">((</span><span class="n">block_size</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="o">//</span> <span class="mi">2</span><span class="p">,</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">1</span><span class="p">]),</span> <span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">)]</span>
    <span class="p">),</span>
<span class="p">)</span>


<span class="k">def</span> <span class="nf">pallas_reduce_scatter</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="p">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">num_devices</span><span class="p">,</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">1</span><span class="p">])</span>
    <span class="k">return</span> <span class="n">pl</span><span class="p">.</span><span class="n">pallas_call</span><span class="p">(</span>
        <span class="n">reduce_scatter_kernel</span><span class="p">,</span>
        <span class="n">out_shape</span><span class="o">=</span><span class="n">out_shape</span><span class="p">,</span>
        <span class="n">grid_spec</span><span class="o">=</span><span class="n">grid_spec</span><span class="p">,</span>
        <span class="n">compiler_params</span><span class="o">=</span><span class="n">pltpu</span><span class="p">.</span><span class="n">CompilerParams</span><span class="p">(</span><span class="n">collective_id</span><span class="o">=</span><span class="mi">0</span><span class="p">),</span>
    <span class="p">)(</span><span class="n">x</span><span class="p">)[</span><span class="mi">0</span><span class="p">]</span>


<span class="n">pallas_result</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">jit</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">shard_map</span><span class="p">(</span>
        <span class="n">pallas_reduce_scatter</span><span class="p">,</span>
        <span class="n">mesh</span><span class="o">=</span><span class="n">mesh</span><span class="p">,</span>
        <span class="n">in_specs</span><span class="o">=</span><span class="n">P</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="s">"x"</span><span class="p">),</span>
        <span class="n">out_specs</span><span class="o">=</span><span class="n">P</span><span class="p">(</span><span class="s">"x"</span><span class="p">,</span> <span class="bp">None</span><span class="p">),</span>
        <span class="n">check_vma</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
    <span class="p">)</span>
<span class="p">)(</span><span class="n">input_arr</span><span class="p">)</span>


<span class="k">def</span> <span class="nf">xla_reduce_scatter</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="p">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">num_devices</span><span class="p">,</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">1</span><span class="p">])</span>
    <span class="k">return</span> <span class="n">lax</span><span class="p">.</span><span class="n">psum_scatter</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="s">"x"</span><span class="p">)</span>


<span class="n">xla_result</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">jit</span><span class="p">(</span>
    <span class="n">jax</span><span class="p">.</span><span class="n">shard_map</span><span class="p">(</span>
        <span class="n">xla_reduce_scatter</span><span class="p">,</span>
        <span class="n">mesh</span><span class="o">=</span><span class="n">mesh</span><span class="p">,</span>
        <span class="n">in_specs</span><span class="o">=</span><span class="n">P</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="s">"x"</span><span class="p">),</span>
        <span class="n">out_specs</span><span class="o">=</span><span class="n">P</span><span class="p">(</span><span class="s">"x"</span><span class="p">,</span> <span class="bp">None</span><span class="p">),</span>
    <span class="p">)</span>
<span class="p">)(</span><span class="n">input_arr</span><span class="p">)</span>
</code></pre></div></div>

<h3 id="44-输入为什么是设备倍数">4.4 输入为什么是设备倍数</h3>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">input_arr</span><span class="p">.</span><span class="n">shape</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">block_size</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="o">*</span> <span class="n">num_devices</span><span class="p">,</span>
    <span class="n">block_size</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">*</span> <span class="n">num_devices</span><span class="p">,</span>
<span class="p">)</span>
<span class="n">partition</span> <span class="o">=</span> <span class="n">P</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="s">"x"</span><span class="p">)</span>
</code></pre></div></div>

<p>第 1 维被切到所有 device，所以每个 device 本地宽度是 <code class="language-plaintext highlighter-rouge">block_size[1]</code>。</p>

<p>每个 device 本地 shape 是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>(block_size[0] * num_devices, block_size[1])
</code></pre></div></div>

<p>然后 reshape：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">x</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">num_devices</span><span class="p">,</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">block_size</span><span class="p">[</span><span class="mi">1</span><span class="p">])</span>
</code></pre></div></div>

<p>本地变成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x_d[0], x_d[1], x_d[2], x_d[3]
</code></pre></div></div>

<p>其中：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x_d[b] = device d 对最终 block b 的贡献
</code></pre></div></div>

<p>最终：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>D0 得到 sum_d x_d[0]
D1 得到 sum_d x_d[1]
D2 得到 sum_d x_d[2]
D3 得到 sum_d x_d[3]
</code></pre></div></div>

<h3 id="45-关键区域">4.5 关键区域</h3>

<table>
  <thead>
    <tr>
      <th>区域</th>
      <th>memory space</th>
      <th>用途</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">x_ref</code></td>
      <td>VMEM</td>
      <td>当前 device 的所有 block 贡献</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">o_ref</code></td>
      <td>VMEM</td>
      <td>当前 device 最终保留的 reduced block</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">hbm_buf[0]</code></td>
      <td>HBM</td>
      <td>通信 slot 0</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">hbm_buf[1]</code></td>
      <td>HBM</td>
      <td>通信 slot 1</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">accum</code></td>
      <td>VMEM</td>
      <td>当前 half-block 的临时累加器</td>
    </tr>
  </tbody>
</table>

<h3 id="46-遍历模拟">4.6 遍历模拟</h3>

<p>假设：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>num_devices = 4
D0, D1, D2, D3
T = top half
B = bottom half
</code></pre></div></div>

<p>记号：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>T2(2+1+0) = block 2 的 top half，已经累加 device 2、1、0 的贡献
B0(0+1+2+3) = block 0 的 bottom half 完整结果
</code></pre></div></div>

<h4 id="loop-0-left">Loop <code class="language-plaintext highlighter-rouge">(0, LEFT)</code></h4>

<p>执行：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">init_left</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
<span class="n">init_left</span><span class="p">.</span><span class="n">wait</span><span class="p">()</span>
<span class="n">init_right</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>

<span class="n">local_copy</span><span class="p">:</span> <span class="n">H0</span><span class="p">.</span><span class="n">T</span> <span class="o">-&gt;</span> <span class="n">accum</span>
<span class="n">accum</span> <span class="o">+=</span> <span class="n">x_ref</span><span class="p">[</span><span class="n">left_block</span><span class="p">,</span> <span class="n">T</span><span class="p">]</span>
<span class="n">accum</span> <span class="o">-&gt;</span> <span class="n">H0</span><span class="p">.</span><span class="n">T</span>
</code></pre></div></div>

<table>
  <thead>
    <tr>
      <th>Device</th>
      <th>H0.T</th>
      <th>H0.B</th>
      <th>H1.T</th>
      <th>H1.B</th>
      <th>accum</th>
      <th>o_ref</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
  </tbody>
</table>

<h4 id="loop-0-right">Loop <code class="language-plaintext highlighter-rouge">(0, RIGHT)</code></h4>

<p>执行：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">send_left</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
<span class="n">local_copy</span><span class="p">:</span> <span class="n">H0</span><span class="p">.</span><span class="n">B</span> <span class="o">-&gt;</span> <span class="n">accum</span>
<span class="n">accum</span> <span class="o">+=</span> <span class="n">x_ref</span><span class="p">[</span><span class="n">right_block</span><span class="p">,</span> <span class="n">B</span><span class="p">]</span>
<span class="n">accum</span> <span class="o">-&gt;</span> <span class="n">H0</span><span class="p">.</span><span class="n">B</span>
</code></pre></div></div>

<table>
  <thead>
    <tr>
      <th>Device</th>
      <th>H0.T</th>
      <th>H0.B</th>
      <th>H1.T</th>
      <th>H1.B</th>
      <th>accum</th>
      <th>o_ref</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
  </tbody>
</table>

<h4 id="loop-1-left">Loop <code class="language-plaintext highlighter-rouge">(1, LEFT)</code></h4>

<p>执行：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">send_right</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
<span class="n">local_copy</span><span class="p">:</span> <span class="n">H1</span><span class="p">.</span><span class="n">T</span> <span class="o">-&gt;</span> <span class="n">accum</span>
<span class="n">accum</span> <span class="o">+=</span> <span class="n">x_ref</span><span class="p">[</span><span class="n">left_block</span><span class="p">,</span> <span class="n">T</span><span class="p">]</span>
<span class="n">accum</span> <span class="o">-&gt;</span> <span class="n">H1</span><span class="p">.</span><span class="n">T</span>
</code></pre></div></div>

<table>
  <thead>
    <tr>
      <th>Device</th>
      <th>H0.T</th>
      <th>H0.B</th>
      <th>H1.T</th>
      <th>H1.B</th>
      <th>accum</th>
      <th>o_ref</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
  </tbody>
</table>

<h4 id="loop-1-right">Loop <code class="language-plaintext highlighter-rouge">(1, RIGHT)</code></h4>

<p>执行：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">send_left</span><span class="p">.</span><span class="n">start</span><span class="p">()</span>
<span class="n">local_copy</span><span class="p">:</span> <span class="n">H1</span><span class="p">.</span><span class="n">B</span> <span class="o">-&gt;</span> <span class="n">accum</span>
<span class="n">accum</span> <span class="o">+=</span> <span class="n">x_ref</span><span class="p">[</span><span class="n">right_block</span><span class="p">,</span> <span class="n">B</span><span class="p">]</span>
<span class="n">accum</span> <span class="o">-&gt;</span> <span class="n">H1</span><span class="p">.</span><span class="n">B</span>
</code></pre></div></div>

<table>
  <thead>
    <tr>
      <th>Device</th>
      <th>H0.T</th>
      <th>H0.B</th>
      <th>H1.T</th>
      <th>H1.B</th>
      <th>accum</th>
      <th>o_ref</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
  </tbody>
</table>

<h4 id="loop-2-left">Loop <code class="language-plaintext highlighter-rouge">(2, LEFT)</code></h4>

<table>
  <thead>
    <tr>
      <th>Device</th>
      <th>H0.T</th>
      <th>H0.B</th>
      <th>H1.T</th>
      <th>H1.B</th>
      <th>accum</th>
      <th>o_ref</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
  </tbody>
</table>

<h4 id="loop-2-right">Loop <code class="language-plaintext highlighter-rouge">(2, RIGHT)</code></h4>

<table>
  <thead>
    <tr>
      <th>Device</th>
      <th>H0.T</th>
      <th>H0.B</th>
      <th>H1.T</th>
      <th>H1.B</th>
      <th>accum</th>
      <th>o_ref</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
    <tr>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">-</code></td>
    </tr>
  </tbody>
</table>

<h4 id="loop-3-left">Loop <code class="language-plaintext highlighter-rouge">(3, LEFT)</code></h4>

<p>最后一轮不再加本地贡献，只把完整 top half 写入输出。</p>

<table>
  <thead>
    <tr>
      <th>Device</th>
      <th>H0.T</th>
      <th>H0.B</th>
      <th>H1.T</th>
      <th>H1.B</th>
      <th>accum</th>
      <th>o_ref</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(all)</code></td>
    </tr>
    <tr>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(all)</code></td>
    </tr>
    <tr>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(all)</code></td>
    </tr>
    <tr>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(all)</code></td>
    </tr>
  </tbody>
</table>

<h4 id="loop-3-right">Loop <code class="language-plaintext highlighter-rouge">(3, RIGHT)</code></h4>

<p>最后一轮不再加本地贡献，只把完整 bottom half 写入输出。</p>

<table>
  <thead>
    <tr>
      <th>Device</th>
      <th>H0.T</th>
      <th>H0.B</th>
      <th>H1.T</th>
      <th>H1.B</th>
      <th>accum</th>
      <th>o_ref</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>D0</td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T0(all)+B0(all)</code></td>
    </tr>
    <tr>
      <td>D1</td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(1+0+3+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B1(1+2+3+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T1(all)+B1(all)</code></td>
    </tr>
    <tr>
      <td>D2</td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(2+1+0+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B2(2+3+0+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T2(all)+B2(all)</code></td>
    </tr>
    <tr>
      <td>D3</td>
      <td><code class="language-plaintext highlighter-rouge">T0(0+3+2+1)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B0(0+1+2+3)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(3+2+1+0)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">B3(3+0+1+2)</code></td>
      <td><code class="language-plaintext highlighter-rouge">T3(all)+B3(all)</code></td>
    </tr>
  </tbody>
</table>

<h3 id="47-最后复习口诀">4.7 最后复习口诀</h3>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>ppermute:
  一跳。每个 device 把自己的 shard 发给右邻居。

all_gather:
  多跳收集。每轮转发一个已有 slot，最终每个 device 收齐所有 slots。

psum:
  单向 all-reduce。数据沿 ring 转，每个 device 收到就加到自己的 o_ref。

psum_scatter:
  直接 reduce-scatter，不是真的先 all-reduce 再 scatter。
  partial sum 在路上传，经过 device 就加本地贡献。
  top half 向左，bottom half 向右。
</code></pre></div></div>

<hr />

<h2 id="5-复习时最容易错的点">5. 复习时最容易错的点</h2>

<h3 id="51-x_refi-不是读-device-i-的内存">5.1 <code class="language-plaintext highlighter-rouge">x_ref[i]</code> 不是读 device i 的内存</h3>

<p>在 <code class="language-plaintext highlighter-rouge">shard_map</code> 内，<code class="language-plaintext highlighter-rouge">x_ref</code> 永远是当前 device 的本地 shard。</p>

<p>所以：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">x_ref</span><span class="p">[</span><span class="n">left_block</span><span class="p">,</span> <span class="n">top</span><span class="p">]</span>
</code></pre></div></div>

<p>意思是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>当前 device 本地保存的 block left_block 的 top half
</code></pre></div></div>

<p>不是远端 device <code class="language-plaintext highlighter-rouge">left_block</code> 的内存。</p>

<h3 id="52-grid-循环不是天然全局同步">5.2 grid 循环不是天然全局同步</h3>

<p>Pallas TPU grid 可以理解成本 device 上按顺序执行：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">for</span> <span class="n">step</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(...):</span>
    <span class="n">kernel</span><span class="p">(...)</span>
</code></pre></div></div>

<p>但不同 device 之间不保证每一轮严格同步。同步需要：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>DMA semaphores
regular semaphores
barrier semaphores
</code></pre></div></div>

<h3 id="53-双缓冲只解决同轮读写冲突">5.3 双缓冲只解决同轮读写冲突</h3>

<p>如果所有设备同轮执行：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>working_slot 和 receiving_slot 不冲突
</code></pre></div></div>

<p>但如果一个 device 跑快一轮：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>快 device 以为它在写邻居 receiving_slot
慢邻居可能还在读同一个 slot 作为 working_slot
</code></pre></div></div>

<p>所以 <code class="language-plaintext highlighter-rouge">capacity_sem</code> 用来做接收方确认。</p>

<h3 id="54-psum_scatter-传的是-accumulator不是输入本身">5.4 <code class="language-plaintext highlighter-rouge">psum_scatter</code> 传的是 accumulator，不是输入本身</h3>

<p>在 <code class="language-plaintext highlighter-rouge">psum</code> 中：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>传输入 shard，accumulator 留在本地。
</code></pre></div></div>

<p>在 <code class="language-plaintext highlighter-rouge">psum_scatter</code> 中：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>传 partial sum，输入贡献留在本地。
</code></pre></div></div>

<p>这是理解 reduce-scatter 的关键。</p>

<hr />

<h2 id="6-官方参考">6. 官方参考</h2>

<ul>
  <li>主页面：<a href="https://docs.jax.dev/en/latest/pallas/tpu/distributed.html">Distributed Computing in Pallas for TPUs</a></li>
  <li><code class="language-plaintext highlighter-rouge">make_async_remote_copy</code> API：<a href="https://docs.jax.dev/en/latest/_autosummary/jax.experimental.pallas.tpu.make_async_remote_copy.html">jax.experimental.pallas.tpu.make_async_remote_copy</a></li>
  <li><code class="language-plaintext highlighter-rouge">ppermute</code> 小节：<a href="https://docs.jax.dev/en/latest/pallas/tpu/distributed.html#example-right-permute-lax-ppermute">Example: Right Permute</a></li>
  <li><code class="language-plaintext highlighter-rouge">all_gather</code> 小节：<a href="https://docs.jax.dev/en/latest/pallas/tpu/distributed.html#example-all-gather-lax-all-gather">Example: All-gather</a></li>
  <li><code class="language-plaintext highlighter-rouge">psum</code> 小节：<a href="https://docs.jax.dev/en/latest/pallas/tpu/distributed.html#example-all-reduce-sum-lax-psum">Example: All-Reduce Sum</a></li>
  <li><code class="language-plaintext highlighter-rouge">psum_scatter</code> 小节：<a href="https://docs.jax.dev/en/latest/pallas/tpu/distributed.html#example-bi-directional-reduce-scatter-lax-psum-scatter">Example: Bi-directional Reduce-Scatter</a></li>
</ul>]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[Pallas TPU Distributed Collectives 复习文档]]></summary></entry><entry><title type="html">JAX 到 VLIW，以及 Pallas / Splash Attention 复习文档</title><link href="https://wqh011128.github.io/blog/2026/06/04/JAX_to_VLIW_Pallas_Splash_Attention_%E5%A4%8D%E4%B9%A0%E6%96%87%E6%A1%A3.html" rel="alternate" type="text/html" title="JAX 到 VLIW，以及 Pallas / Splash Attention 复习文档" /><published>2026-06-04T00:00:00+00:00</published><updated>2026-06-04T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/06/04/JAX_to_VLIW_Pallas_Splash_Attention_%E5%A4%8D%E4%B9%A0%E6%96%87%E6%A1%A3</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/06/04/JAX_to_VLIW_Pallas_Splash_Attention_%E5%A4%8D%E4%B9%A0%E6%96%87%E6%A1%A3.html"><![CDATA[<h1 id="jax-到-vliw以及-pallas--splash-attention-复习笔记">JAX 到 VLIW，以及 Pallas / Splash Attention 复习笔记</h1>

<p>本文整理两篇 Patrick Toulme 的文章：</p>

<ul>
  <li><a href="https://patricktoulme.substack.com/p/from-jax-to-vliw-tracing-a-computation">From JAX to VLIW: Tracing a Computation Through the TPU Compiler</a></li>
  <li><a href="https://patricktoulme.substack.com/p/when-xla-isnt-enough-from-pallas">When XLA Isn’t Enough: From Pallas to VLIW</a></li>
</ul>

<p>目标不是逐字翻译，而是按模型框架开发者的视角复习：重点理解 JAX/HLO/XLA/Pallas 这些上层如何影响底层执行；LLO/VLIW 只学会基础读法，不陷入指令细节。</p>

<hr />

<h2 id="0-两条总线路jax-to-vliw-与-pallas-to-vliw">0. 两条总线路：JAX to VLIW 与 Pallas to VLIW</h2>

<p>先把两篇文章压缩成两张图。</p>

<h3 id="01-普通-jax--xla-路线">0.1 普通 JAX / XLA 路线</h3>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Python JAX function
  -&gt; JAX tracing / Jaxpr
  -&gt; StableHLO / HLO
  -&gt; HLO optimization passes
       algebraic simplification
       layout assignment
       tiling
       fusion
       memory assignment
       copy scheduling
  -&gt; TPU backend LLO
  -&gt; VLIW bundles
  -&gt; TPU hardware
       HBM / VMEM / MXU / VPU / XLU / DMA
</code></pre></div></div>

<p>直觉：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>JAX 代码描述“我要算什么”
HLO 描述“张量图是什么”
XLA 优化“怎样重排、融合、分配内存”
LLO 描述“TPU 上具体用哪些硬件动作”
VLIW bundle 描述“同一个时刻哪些硬件指令一起发射”
</code></pre></div></div>

<h3 id="02-pallas--mosaic-路线">0.2 Pallas / Mosaic 路线</h3>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Python Pallas kernel
  -&gt; pallas_call
  -&gt; HLO custom-call
       XLA 只看到一个 opaque call
  -&gt; Mosaic MLIR
       Mosaic 能看到 Pallas kernel body
  -&gt; TPU backend LLO
  -&gt; VLIW bundles
  -&gt; TPU hardware
</code></pre></div></div>

<p>直觉：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>普通 JAX:
  你写完整 tensor program，让 XLA 自动找 fusion / tiling。

Pallas:
  你直接写 tiled kernel program，让 Mosaic 和 TPU backend 继续降到底层。
</code></pre></div></div>

<h3 id="03-两篇文章合起来的核心">0.3 两篇文章合起来的核心</h3>

<p>第一篇的主题：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>XLA 很强。
普通 JAX 程序可以被自动优化成 TPU 上的高质量 VLIW 程序。
</code></pre></div></div>

<p>第二篇的主题：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>XLA 有边界。
XLA 会优化你写出来的算法，但通常不会自动发明另一个算法。
Splash / FlashAttention 这类 online softmax 是算法级改写，因此需要 Pallas 表达。
</code></pre></div></div>

<p>最重要的对比：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>XLA fusion:
  减少已经存在的中间 tensor 的读写。

Splash / FlashAttention:
  改写算法，让完整 attention matrix 根本不成为程序语义里的 tensor。
</code></pre></div></div>

<hr />

<h1 id="第一篇from-jax-to-vliw">第一篇：From JAX to VLIW</h1>

<h2 id="1-例子在算什么">1. 例子在算什么</h2>

<p>文章使用一个很小的 attention-like block，便于观察整个 TPU 编译链路：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">h</span> <span class="o">=</span> <span class="n">x</span> <span class="o">@</span> <span class="n">w1</span>
<span class="n">h</span> <span class="o">=</span> <span class="n">h</span> <span class="o">/</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">mean</span><span class="p">(</span><span class="n">h</span> <span class="o">**</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">eps</span><span class="p">)</span>
<span class="n">h</span> <span class="o">=</span> <span class="n">softmax</span><span class="p">(</span><span class="n">h</span><span class="p">)</span>
<span class="n">out</span> <span class="o">=</span> <span class="n">h</span> <span class="o">@</span> <span class="n">w2</span>
</code></pre></div></div>

<p>形状大致是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x:   [16, 64]
w1:  [64, 64]
w2:  [64, 32]
out: [16, 32]
</code></pre></div></div>

<p>它包含：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>matmul_1
RMSNorm-like normalization
softmax
matmul_2
</code></pre></div></div>

<p>作者使用 <code class="language-plaintext highlighter-rouge">jax.jit</code> 触发编译，用 dump flags 输出 HLO 和 TPU backend 的 LLO。<code class="language-plaintext highlighter-rouge">jax.named_call</code> 的作用是让 IR metadata 中保留人类可读名字，例如 <code class="language-plaintext highlighter-rouge">matmul_1</code>、<code class="language-plaintext highlighter-rouge">rms_norm</code>、<code class="language-plaintext highlighter-rouge">softmax</code>。</p>

<p>对框架开发者来说，这个例子的价值在于：它虽然小，但覆盖了模型里常见的关键模式：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>matmul producer
elementwise chain
reduce
broadcast
multi-output use
softmax
second matmul
</code></pre></div></div>

<hr />

<h2 id="2-hlo-是什么">2. HLO 是什么</h2>

<p>HLO 是 High Level Optimizer IR。可以把它理解成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>带 shape / dtype / layout / metadata 的 SSA 张量计算图
</code></pre></div></div>

<p>SSA 表示每个中间值只定义一次：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>%dot = dot(%x, %w1)
%square = multiply(%dot, %dot)
%sum = reduce(%square)
%sqrt = sqrt(...)
</code></pre></div></div>

<p>HLO 仍然是比较高层的 IR。它关心：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>op 类型：dot / reduce / broadcast / exp / sqrt
tensor shape：f32[16,64]
layout：{1,0}
reduce 维度
metadata：来自哪一行 Python，来自哪个 named_call
</code></pre></div></div>

<p>它暂时不关心：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>哪个 MXU 执行 matmul
哪个 cycle 发射指令
何时 vmatpush / vpop
哪个 VLIW bundle 同时发射 DMA 和 vector op
</code></pre></div></div>

<p>因此，HLO 是模型框架开发者很值得读的层级。很多性能问题在 HLO 层已经能判断：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>是否 materialize 了巨大中间矩阵
是否有不必要 broadcast
是否存在 producer 被多个 consumer 使用
是否可以 fusion
layout 是否合理
copy-start / copy-done 是否能 overlap
</code></pre></div></div>

<hr />

<h2 id="3-初始-hlo从-python-语义到显式张量图">3. 初始 HLO：从 Python 语义到显式张量图</h2>

<h3 id="31-matmul">3.1 Matmul</h3>

<p>Python：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">h</span> <span class="o">=</span> <span class="n">x</span> <span class="o">@</span> <span class="n">w1</span>
</code></pre></div></div>

<p>HLO 中一般表现为：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dot(x, w1)
</code></pre></div></div>

<p>HLO 会带上 contracting dimensions，例如：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>lhs_contracting_dims={1}
rhs_contracting_dims={0}
</code></pre></div></div>

<p>也就是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x:  [16,64]
w1: [64,64]
沿 x 的第 1 维和 w1 的第 0 维相乘求和
输出 [16,64]
</code></pre></div></div>

<h3 id="32-rmsnorm-like-部分">3.2 RMSNorm-like 部分</h3>

<p>Python 直觉：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">square</span> <span class="o">=</span> <span class="n">h</span> <span class="o">**</span> <span class="mi">2</span>
<span class="n">mean</span> <span class="o">=</span> <span class="nb">sum</span><span class="p">(</span><span class="n">square</span><span class="p">,</span> <span class="n">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdims</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span> <span class="o">/</span> <span class="mi">64</span>
<span class="n">rms</span> <span class="o">=</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">mean</span> <span class="o">+</span> <span class="n">eps</span><span class="p">)</span>
<span class="n">h_norm</span> <span class="o">=</span> <span class="n">h</span> <span class="o">/</span> <span class="n">rms</span>
</code></pre></div></div>

<p>HLO 里会显式拆成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>multiply(h, h)
reduce_sum(axis=1)
multiply by 1/64 或 divide by 64
add eps
sqrt
broadcast
divide
</code></pre></div></div>

<p>注意 <code class="language-plaintext highlighter-rouge">broadcast</code>：Python 写 <code class="language-plaintext highlighter-rouge">keepdims=True</code> 或隐式 broadcasting 时看起来很自然，但 HLO 必须显式说明哪个 shape 扩展到哪个 shape。</p>

<h3 id="33-softmax">3.3 Softmax</h3>

<p>Python：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">m</span> <span class="o">=</span> <span class="nb">max</span><span class="p">(</span><span class="n">h</span><span class="p">,</span> <span class="n">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdims</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">e</span> <span class="o">=</span> <span class="n">exp</span><span class="p">(</span><span class="n">h</span> <span class="o">-</span> <span class="n">m</span><span class="p">)</span>
<span class="n">p</span> <span class="o">=</span> <span class="n">e</span> <span class="o">/</span> <span class="nb">sum</span><span class="p">(</span><span class="n">e</span><span class="p">,</span> <span class="n">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdims</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
</code></pre></div></div>

<p>HLO：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>reduce_max
broadcast max
subtract
exp
reduce_sum
broadcast sum
divide
</code></pre></div></div>

<p>HLO 中 <code class="language-plaintext highlighter-rouge">reduce</code> 通常会带一个小 computation region，表示规约函数是 <code class="language-plaintext highlighter-rouge">add</code> 还是 <code class="language-plaintext highlighter-rouge">maximum</code>。</p>

<hr />

<h2 id="4-hlo-optimization-passes-做了什么">4. HLO optimization passes 做了什么</h2>

<p>文章展示了几类重要优化。</p>

<h3 id="41-algebraic-simplification">4.1 Algebraic simplification</h3>

<p>例如：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x / 64
</code></pre></div></div>

<p>会变成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x * 0.015625
</code></pre></div></div>

<p>因为乘法通常比除法更便宜。这类优化在 HLO 层完成，不是 TPU 特有。</p>

<h3 id="42-layout-assignment">4.2 Layout assignment</h3>

<p>HLO shape 可能出现：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>f32[16,64]{1,0}
</code></pre></div></div>

<p>解释：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>f32[16,64]  dtype + shape
{1,0}       physical layout order
</code></pre></div></div>

<p><code class="language-plaintext highlighter-rouge">{1,0}</code> 可以粗略理解为最后一维更连续，接近 row-major 的直觉。</p>

<h3 id="43-tpu-tiling">4.3 TPU tiling</h3>

<p>优化后可能出现：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>f32[16,64]{1,0:T(8,128)}
</code></pre></div></div>

<p><code class="language-plaintext highlighter-rouge">T(8,128)</code> 是 TPU tile annotation。它和 TPU VPU 的向量组织有关：常见向量寄存器结构可以理解为 8 个 sublanes，每个 sublane 128 lanes。</p>

<p>这层含义是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>HLO tensor 不是只作为抽象 ndarray 存放，
backend 已经开始决定怎样把它切成更贴近 TPU 硬件的 tile。
</code></pre></div></div>

<h3 id="44-fusion">4.4 Fusion</h3>

<p>Fusion 是第一篇的主角之一。</p>

<p>没有 fusion：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>op A 产生中间 tensor
写出去
op B 读回来
产生另一个中间 tensor
写出去
op C 再读回来
</code></pre></div></div>

<p>有 fusion：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>op A / B / C 进入同一个 fusion kernel
中间值尽量留在寄存器、VMEM 或临时片上状态
只写 fusion 的对外输出
</code></pre></div></div>

<p>文章里的最终 HLO 大致被压成几个 fusion：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>multiply_reduce_fusion:
  matmul_1 + square + reduce_sum
  -&gt; (sum_of_squares, matmul_result)

add_sqrt_fusion:
  sum_of_squares * 1/64 + eps -&gt; sqrt

fusion.5:
  normalized h -&gt; reduce_max

fusion.2:
  exp(h - max) -&gt; reduce_sum

fusion:
  normalize softmax + matmul_2 -&gt; output
</code></pre></div></div>

<hr />

<h2 id="5-multi-output-fusion为什么重要">5. Multi-output fusion：为什么重要</h2>

<p>第一次 matmul 的结果：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">h</span> <span class="o">=</span> <span class="n">x</span> <span class="o">@</span> <span class="n">w1</span>
</code></pre></div></div>

<p>有两个用途：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>用途 1：h -&gt; h ** 2 -&gt; reduce_sum，用于 RMSNorm
用途 2：h -&gt; 后续 normalize / softmax
</code></pre></div></div>

<p>如果处理不好，编译器可能遇到两难：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>要么重复计算 matmul
要么把 matmul result 写到大 buffer，后续再读
</code></pre></div></div>

<p>multi-output fusion 的做法是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>multiply_reduce_fusion -&gt; (reduce_sum(h*h), h)
</code></pre></div></div>

<p>它让 matmul 只算一次，同时产出：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 每行平方和
2. matmul result 本身
</code></pre></div></div>

<p>对模型框架开发者来说，这对应一个常见图优化问题：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>一个 producer 有多个 consumer。
</code></pre></div></div>

<p>好的优化器需要在以下选择间做代价判断：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>duplicate producer
materialize producer
multi-output fusion
recompute
</code></pre></div></div>

<hr />

<h2 id="6-memory-assignmenthbmvmemcopy-startcopy-done">6. Memory assignment：HBM、VMEM、copy-start、copy-done</h2>

<p>文章中 HLO 会出现类似 memory space 的 annotation。粗略理解：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>S(0): HBM，片外大内存，通常是默认
S(1): VMEM，片上 SRAM
S(2): sync token 或 backend-specific space
</code></pre></div></div>

<p>HBM 和 VMEM 的关系：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>HBM:
  大，慢，跨 kernel 稳定可见

VMEM:
  小，快，片上，适合 tile 和短生命周期中间值
</code></pre></div></div>

<p>HLO 中还会出现：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>copy-start(w1)
copy-done(w1)
</code></pre></div></div>

<p>含义：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>copy-start:
  发起从 HBM 到 VMEM 的异步拷贝

copy-done:
  等待这个拷贝完成，之后才能使用
</code></pre></div></div>

<p>这允许 backend 做 overlap：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>一边计算当前阶段
一边 DMA 搬下一阶段需要的数据
</code></pre></div></div>

<p>例如：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>先 copy-start(w1)
copy-done(w1) 后调用第一个 matmul fusion

尽早 copy-start(w2)
在 RMSNorm / softmax 期间后台搬 w2
最后 matmul 前 copy-done(w2)
</code></pre></div></div>

<p>这就是编译器自动做的 prefetch / overlap。</p>

<hr />

<h2 id="7-fusion-结束是否一定写回-hbm">7. Fusion 结束是否一定写回 HBM</h2>

<p>这是非常容易误解的点。</p>

<p>更准确的说法：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>每个 HLO op 结束不一定写 HBM。
fusion 内部中间值通常不写 HBM。
fusion 的对外输出必须 materialize 到某个 buffer。
这个 buffer 可能是 HBM，也可能是 VMEM，取决于 memory assignment、大小、生命周期和后端能力。
</code></pre></div></div>

<p>因此：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>不是“每个 op 写 HBM”
也不是“每个 fusion 一定写 HBM”
而是“fusion/custom-call 边界通常是 materialization 边界”
</code></pre></div></div>

<p>在第一篇 toy program 里，一些中间 buffer 可以放到 VMEM。TLP 可以把 VMEM buffer 地址传给后续 fusion，所以并不是所有 fusion output 都落 HBM。</p>

<p>但在第二篇 naive attention 中，<code class="language-plaintext highlighter-rouge">scores</code> 形状是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[heads, seq_len, seq_len]
</code></pre></div></div>

<p>具体例子：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[8, 2048, 2048] f32
= 8 * 2048 * 2048 * 4 bytes
= 128MB
</code></pre></div></div>

<p>这个中间矩阵太大，而且跨多个 fusion 被使用：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>fusion.5 产生 scores
fusion.2 需要 scores 来做 exp 和 reduce_sum
final fusion 需要 softmax 后的信息继续乘 V
</code></pre></div></div>

<p>所以它通常会 materialize 到 HBM。第二篇的性能问题本质就是这里。</p>

<p>判断规则：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>小的、短生命周期的、编译器能 schedule 在片上的值：
  可能留在 VMEM / register / scratch。

大的、跨 fusion 使用的、生命周期较长的 tensor：
  往往落 HBM。

Flash/Splash Attention 的胜利：
  不是让 128MB scores 写得更快，
  而是让完整 scores tensor 不存在。
</code></pre></div></div>

<hr />

<h2 id="8-从-hlo-到-llo层级突然变低">8. 从 HLO 到 LLO：层级突然变低</h2>

<p>HLO 还像数学图：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dot
reduce
sqrt
exp
fusion
</code></pre></div></div>

<p>LLO 开始像硬件动作：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>vld
vst
vmatpush
vmatmul
vpop
vxpose
vrot.slane
dma.hbm_to_vmem
</code></pre></div></div>

<p>LLO 是 TPU-specific Low Level Operators。你不需要把每条指令都背下来，但要能识别几个模式。</p>

<h3 id="81-tpu-硬件单元粗略图">8.1 TPU 硬件单元粗略图</h3>

<table>
  <thead>
    <tr>
      <th>单元</th>
      <th>作用</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>MXU</td>
      <td>Matrix Multiply Unit，负责矩阵乘，TPU 主要算力来源</td>
    </tr>
    <tr>
      <td>VPU</td>
      <td>Vector Processing Unit，负责 elementwise、vector add/mul/select 等</td>
    </tr>
    <tr>
      <td>XLU</td>
      <td>负责 transpose、shuffle、cross-lane movement</td>
    </tr>
    <tr>
      <td>Scalar Unit</td>
      <td>标量计算、地址、控制流</td>
    </tr>
    <tr>
      <td>DMA</td>
      <td>HBM 与 VMEM 之间异步搬运</td>
    </tr>
  </tbody>
</table>

<h3 id="82-llo-中识别-matmul">8.2 LLO 中识别 matmul</h3>

<p>常见模式：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>vld        从 VMEM load tile
vmatpush   把 tile push 到 MXU
vmatmul    触发矩阵乘
vpop       从 MXU 取回 accumulator / result
</code></pre></div></div>

<p>如果看到：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>vmatpush
vmatmul
vpop.mrf / vpop.f32
</code></pre></div></div>

<p>基本可以判断这里在做 MXU matmul。</p>

<h3 id="83-llo-中识别-vector-compute">8.3 LLO 中识别 vector compute</h3>

<p>常见：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>vadd
vmul
vsub
vsel
vcmp
vpow2
vrsqrt
</code></pre></div></div>

<p>这些通常是 VPU 上的向量操作。比如 RMSNorm 的 <code class="language-plaintext highlighter-rouge">sqrt</code>，后端可能用 <code class="language-plaintext highlighter-rouge">vrsqrt</code> 先算 reciprocal sqrt，再组合得到需要的结果。</p>

<h3 id="84-llo-中识别-reduction--transpose">8.4 LLO 中识别 reduction / transpose</h3>

<p>RMSNorm 要做：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>sum(h ** 2, axis=-1)
</code></pre></div></div>

<p>TPU 的 lane/sublane 数据布局和 Python 行列不完全一致，所以 reduction 之前常会出现：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>vxpose
vpop.trf
</code></pre></div></div>

<p>然后是 tree reduction：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>vadd
vrot.slane 4
vadd
vrot.slane 2
vadd
vrot.slane 1
vadd
</code></pre></div></div>

<p>这是并行规约：8 个 sublane 求和，不是串行加 7 次，而是 3 轮合并。</p>

<hr />

<h2 id="9-vliw-bundle-的特点">9. VLIW bundle 的特点</h2>

<p>VLIW 是 Very Long Instruction Word。</p>

<p>一个 bundle 可以包含多条指令：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>bundle:
  dma.hbm_to_vmem
  vld
  vld
  vmov
  scalar address op
  maybe vmatmul
</code></pre></div></div>

<p>它的核心特点：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. bundle 内部的指令可以并行发射。
2. 编译器静态决定哪些指令能放进同一个 bundle。
3. 硬件不依赖复杂乱序执行，而是按编译器安排执行。
4. 一个 bundle 可以同时使用多个硬件单元，例如 DMA、MXU、VPU、Scalar。
5. 好的 schedule 会把数据搬运、地址计算、当前 tile 计算、下一 tile 预取交织起来。
</code></pre></div></div>

<p>所以看到 bundle 时，不要把它理解成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>bundle 中的指令一条条串行执行
</code></pre></div></div>

<p>更应该理解成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>这是编译器打包好的并行发射包。
</code></pre></div></div>

<p>文章第一篇说明，toy program 最终被拆成多个 fusion 和一个 TLP；<code class="language-plaintext highlighter-rouge">multiply_reduce_fusion</code> 这样的 kernel 会生成几十个 bundle，TLP 负责把多个 fusion 和 DMA 调度起来。</p>

<p>bundle count 有参考意义，但不是唯一性能指标：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>bundle 少不一定快
bundle 多不一定慢
HBM traffic、MXU 利用率、DMA overlap、stall、register pressure 都会影响性能
</code></pre></div></div>

<hr />

<h2 id="10-tlptop-level-program">10. TLP：Top Level Program</h2>

<p>TLP 是整个 compiled computation 的总控程序。它不是某个 fusion 的内部实现，而是调度多个 fusion/custom-call 和 DMA。</p>

<p>第一篇中可以粗略理解为：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>copy w1 HBM -&gt; VMEM
call multiply_reduce_fusion

start copy w2 HBM -&gt; VMEM

call add_sqrt_fusion
call reduce_max fusion
call exp/reduce_sum fusion

wait w2 copy done
call final matmul fusion

final sync
</code></pre></div></div>

<p>对框架开发者来说：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>HLO graph:
  决定有哪些 computation 节点和依赖。

HLO fusion:
  决定哪些 op 合成 kernel。

TLP:
  决定 kernel 和 DMA 的 top-level schedule。

LLO/VLIW:
  决定 kernel 内部如何使用 TPU 硬件单元。
</code></pre></div></div>

<hr />

<h1 id="第二篇when-xla-isnt-enough">第二篇：When XLA Isn’t Enough</h1>

<h2 id="11-第二篇的问题意识">11. 第二篇的问题意识</h2>

<p>第一篇告诉我们：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>XLA/TPU compiler 能把普通 JAX 程序优化得很深。
</code></pre></div></div>

<p>第二篇问：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>既然 XLA 这么强，为什么还需要 Pallas？
</code></pre></div></div>

<p>答案：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>XLA 能优化你写出来的计算图，
但通常不会自动把 naive attention 改写成 FlashAttention / Splash Attention。
</code></pre></div></div>

<p>这是“图优化”和“算法改写”的边界。</p>

<hr />

<h2 id="12-naive-attention-的问题">12. Naive attention 的问题</h2>

<p>标准 attention：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><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">T</span>
<span class="n">scores</span> <span class="o">=</span> <span class="n">scores</span> <span class="o">+</span> <span class="n">mask</span>
<span class="n">m</span> <span class="o">=</span> <span class="nb">max</span><span class="p">(</span><span class="n">scores</span><span class="p">,</span> <span class="n">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdims</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">e</span> <span class="o">=</span> <span class="n">exp</span><span class="p">(</span><span class="n">scores</span> <span class="o">-</span> <span class="n">m</span><span class="p">)</span>
<span class="n">l</span> <span class="o">=</span> <span class="nb">sum</span><span class="p">(</span><span class="n">e</span><span class="p">,</span> <span class="n">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">,</span> <span class="n">keepdims</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">p</span> <span class="o">=</span> <span class="n">e</span> <span class="o">/</span> <span class="n">l</span>
<span class="n">out</span> <span class="o">=</span> <span class="n">p</span> <span class="o">@</span> <span class="n">V</span>
</code></pre></div></div>

<p>核心中间矩阵：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>scores: [heads, q_len, kv_len]
</code></pre></div></div>

<p>文章使用的例子中：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>heads = 8
seq_len = 2048
head_dim = 128
</code></pre></div></div>

<p>所以：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>scores = [8, 2048, 2048] f32
       = 128MB
</code></pre></div></div>

<p>XLA 会做真正的优化，例如：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>fusion.5:
  Q @ K^T + mask + reduce_max
  -&gt; (max, scores)

fusion.2:
  exp(scores - max) + reduce_sum
  -&gt; sum

fusion:
  normalize + S @ V
  -&gt; output
</code></pre></div></div>

<p>这里的 XLA 已经很努力：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>matmul 和 mask / max 融合
exp 和 sum 融合
normalize 和 final matmul 融合
</code></pre></div></div>

<p>但问题仍然存在：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>完整 scores matrix 仍然被创建出来，并跨 fusion 使用。
</code></pre></div></div>

<p>所以 attention 的瓶颈不只是“op 没 fusion”，而是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>算法本身 materialize 了巨大的 [H,S,S] 矩阵。
</code></pre></div></div>

<hr />

<h2 id="13-为什么不能简单依赖-xla">13. 为什么不能简单依赖 XLA</h2>

<p>XLA 的优化通常是在等价计算图范围内做：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>融合 producer/consumer
消除冗余 broadcast
代数化简
layout assignment
tile scheduling
memory placement
</code></pre></div></div>

<p>但 FlashAttention / Splash Attention 本质上做了更深的事情：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>把 attention 的执行方式改成 streaming over KV blocks。
</code></pre></div></div>

<p>换句话说：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>naive:
  先完整生成 scores
  再 softmax
  再乘 V

Splash:
  每次只生成一个 scores tile
  边扫 KV tile 边维护 online softmax 状态
  从不生成完整 scores
</code></pre></div></div>

<p>这属于算法级重写。普通 XLA pass 很难从任意 attention graph 自动证明并改写成这种形式，尤其还要考虑 mask、sparsity、GQA/MQA、numerical stability、tile size、memory constraints。</p>

<hr />

<h2 id="14-先解决一个直觉误区kv-tile-结果为什么是累加不是-cat">14. 先解决一个直觉误区：KV tile 结果为什么是累加，不是 cat</h2>

<p>对一个 query 向量 <code class="language-plaintext highlighter-rouge">q</code>，attention 是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>s_j = q @ k_j
p_j = softmax(s)_j
out = sum_j p_j * v_j
</code></pre></div></div>

<p>如果 KV 被切成两个 tile：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>tile 0: k0, k1, k2
tile 1: k3, k4, k5
</code></pre></div></div>

<p>那么 scores 可以概念上拼接：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>scores_tile0 = [s0, s1, s2]
scores_tile1 = [s3, s4, s5]
scores = cat([scores_tile0, scores_tile1])
</code></pre></div></div>

<p>但是最终输出不是 scores。最终输出是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>out =
  p0 * v0 + p1 * v1 + p2 * v2
  + p3 * v3 + p4 * v4 + p5 * v5
</code></pre></div></div>

<p>按 tile 写：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>out =
  sum_{j in tile0} p_j * v_j
  + sum_{j in tile1} p_j * v_j
</code></pre></div></div>

<p>所以：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>scores 可以在 KV 维度 cat。
output 是沿 KV 维度 weighted sum，因此不同 KV tile 对 output 的贡献要相加。
</code></pre></div></div>

<p>如果把每个 tile 的 output cat 起来，shape 都错了。对每个 query，attention 最终只输出一个 <code class="language-plaintext highlighter-rouge">head_dim</code> 向量。</p>

<hr />

<h2 id="15-online-softmax为什么需要-mlo">15. Online softmax：为什么需要 m、l、o</h2>

<p>先只看一个 query。它要 attend 到很多 key/value：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>key/value 位置：1, 2, 3, ..., N
score：a_1, a_2, a_3, ..., a_N
value：v_1, v_2, v_3, ..., v_N
</code></pre></div></div>

<p>标准 attention 输出是：</p>

\[\mathrm{Out}
=
\sum_{t=1}^{N}
\frac{e^{a_t}}{\sum_{r=1}^{N} e^{a_r}}
v_t\]

<p>为了数值稳定，通常减去最大值：</p>

\[M = \max_{1 \le t \le N} a_t\]

<p>于是：</p>

\[\mathrm{Out}
=
\frac{
\sum_{t=1}^{N} e^{a_t - M} v_t
}{
\sum_{t=1}^{N} e^{a_t - M}
}\]

<p>这里可以拆出两个量：</p>

\[O = \sum_{t=1}^{N} e^{a_t - M} v_t\]

\[L = \sum_{t=1}^{N} e^{a_t - M}\]

<p>最后：</p>

\[\mathrm{Out} = \frac{O}{L}\]

<p>所以：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>L = softmax denominator，分母，也可以理解成总权重
O = softmax @ V 的 numerator，分子，也就是未归一化的加权 V 总和
</code></pre></div></div>

<p>注意 shape：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>L:
  对每个 query 是一个标量。
  对一个 Q block 是 [bq] 或 [bq, 1]。

O:
  对每个 query 是一个 head_dim 向量。
  对一个 Q block 是 [bq, head_dim]。
</code></pre></div></div>

<h3 id="151-分块之后local_sum-和-local_out-是什么">15.1 分块之后，local_sum 和 local_out 是什么</h3>

<p>现在把 KV 分成两个 block：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block A: 位置 1,2,3
block B: 位置 4,5,6
</code></pre></div></div>

<p>如果已经处理完 block A，我们维护：</p>

\[m_A = \max(a_1,a_2,a_3)\]

\[L_A =
e^{a_1-m_A} +
e^{a_2-m_A} +
e^{a_3-m_A}\]

\[O_A =
e^{a_1-m_A}v_1 +
e^{a_2-m_A}v_2 +
e^{a_3-m_A}v_3\]

<p>这里：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>m_A = 已处理 blocks 的最大 score
L_A = 已处理 blocks 的分母累计
O_A = 已处理 blocks 的加权 V 分子累计
</code></pre></div></div>

<p>现在来了 block B：</p>

\[m_B = \max(a_4,a_5,a_6)\]

<p>新的全局最大值是：</p>

\[m_{AB} = \max(m_A, m_B)\]

<p>block B 自己的贡献必须用新的最大值 (m_{AB}) 来算：</p>

\[L_B =
e^{a_4-m_{AB}} +
e^{a_5-m_{AB}} +
e^{a_6-m_{AB}}\]

\[O_B =
e^{a_4-m_{AB}}v_4 +
e^{a_5-m_{AB}}v_5 +
e^{a_6-m_{AB}}v_6\]

<p>这两个就是代码里的：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>local_sum = L_B
local_out = O_B
</code></pre></div></div>

<p>也就是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>local_sum:
  当前 KV block 对 softmax 分母的贡献。

local_out:
  当前 KV block 对 attention 输出分子的贡献。
</code></pre></div></div>

<p>对应到 tile 写法：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>scores_block = Q_block @ K_block^T
p_block = exp(scores_block - m_new)

local_sum = reduce_sum(p_block, axis=KV_block_dim)
local_out = p_block @ V_block
</code></pre></div></div>

<h3 id="152-为什么旧的-ol-要乘-correction">15.2 为什么旧的 O/L 要乘 correction</h3>

<p>问题是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>旧的 L_A / O_A 是用旧最大值 m_A 算的。
现在新的最大值变成了 m_AB。
</code></pre></div></div>

<p>为了把旧贡献换到新的 max 坐标系，需要缩放：</p>

\[\alpha = e^{m_A - m_{AB}}\]

<p>于是：</p>

\[L_A' = \alpha L_A\]

\[O_A' = \alpha O_A\]

<p>这是因为：</p>

\[e^{a_t - m_{AB}}
=
e^{a_t - m_A} \cdot e^{m_A - m_{AB}}\]

<p>所以旧累计值都要乘同一个缩放因子。</p>

<h3 id="153-l_new-和-o_new-是什么">15.3 l_new 和 o_new 是什么</h3>

<p>合并旧 blocks 与当前 block：</p>

\[L_{AB}
=
\alpha L_A + L_B\]

\[O_{AB}
=
\alpha O_A + O_B\]

<p>这就是代码：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">l_new</span> <span class="o">=</span> <span class="n">correction</span> <span class="o">*</span> <span class="n">l_prev</span> <span class="o">+</span> <span class="n">local_sum</span>
<span class="n">o_new</span> <span class="o">=</span> <span class="n">correction</span> <span class="o">*</span> <span class="n">o_prev</span> <span class="o">+</span> <span class="n">local_out</span>
</code></pre></div></div>

<p>变量对照：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>correction = alpha
l_prev     = L_A
o_prev     = O_A
local_sum  = L_B
local_out  = O_B
l_new      = L_AB
o_new      = O_AB
</code></pre></div></div>

<p>最后扫完所有 KV blocks：</p>

\[\mathrm{Out}
=
\frac{O_{\mathrm{final}}}{L_{\mathrm{final}}}\]

<p>也就是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">output</span> <span class="o">=</span> <span class="n">o_final</span> <span class="o">/</span> <span class="n">l_final</span>
</code></pre></div></div>

<p>一句话：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Attention = 加权 value 之和 / 权重之和

O = 加权 value 之和，也就是分子
L = 权重之和，也就是分母

online softmax 只是分块处理时，边走边维护 O 和 L。
</code></pre></div></div>

<hr />

<h2 id="16-pallas-是什么">16. Pallas 是什么</h2>

<p>Pallas 是 JAX 的 custom kernel language。官方文档将它描述为 JAX extension，用于为 GPU/TPU 写 custom kernels，同时保留一部分 JAX tracing 和 <code class="language-plaintext highlighter-rouge">jax.numpy</code> 风格。Pallas API 仍然是 experimental。</p>

<p>简化理解：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>JAX:
  写 tensor program。

Pallas:
  写 tile program。
</code></pre></div></div>

<p>Pallas 不是手写汇编。你不直接写 VLIW 指令，而是显式表达：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>grid 怎么划分
每个 program 处理哪个 tile
输入输出如何 block
scratch 如何跨 iterations 保存
哪些 metadata 放 SMEM
</code></pre></div></div>

<p>后端仍然负责：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>MXU/VPU 映射
DMA scheduling
VMEM layout
VLIW bundle packing
</code></pre></div></div>

<hr />

<h2 id="17-pallas-kernel-语法核心">17. Pallas kernel 语法核心</h2>

<p>一个典型 Pallas kernel 长这样：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">kernel</span><span class="p">(</span><span class="n">q_ref</span><span class="p">,</span> <span class="n">k_ref</span><span class="p">,</span> <span class="n">v_ref</span><span class="p">,</span> <span class="n">o_ref</span><span class="p">,</span> <span class="n">scratch_ref</span><span class="p">):</span>
    <span class="n">q</span> <span class="o">=</span> <span class="n">q_ref</span><span class="p">[...]</span>
    <span class="n">k</span> <span class="o">=</span> <span class="n">k_ref</span><span class="p">[...]</span>
    <span class="n">v</span> <span class="o">=</span> <span class="n">v_ref</span><span class="p">[...]</span>

    <span class="n">result</span> <span class="o">=</span> <span class="p">...</span>

    <span class="n">o_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">result</span>
</code></pre></div></div>

<p>这里的参数不是普通 JAX array，而是 <code class="language-plaintext highlighter-rouge">Ref</code>。</p>

<p><code class="language-plaintext highlighter-rouge">Ref</code> 可以理解成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>kernel 内部看到的一块可读写内存视图
</code></pre></div></div>

<p>读取：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">x</span> <span class="o">=</span> <span class="n">x_ref</span><span class="p">[...]</span>
</code></pre></div></div>

<p>写入：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">o_ref</span><span class="p">[...]</span> <span class="o">=</span> <span class="n">y</span>
</code></pre></div></div>

<p>调用时用：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">pl</span><span class="p">.</span><span class="n">pallas_call</span><span class="p">(</span>
    <span class="n">kernel</span><span class="p">,</span>
    <span class="n">out_shape</span><span class="o">=</span><span class="p">...,</span>
    <span class="n">grid</span><span class="o">=</span><span class="p">...,</span>
    <span class="n">in_specs</span><span class="o">=</span><span class="p">...,</span>
    <span class="n">out_specs</span><span class="o">=</span><span class="p">...,</span>
    <span class="n">scratch_shapes</span><span class="o">=</span><span class="p">...,</span>
    <span class="n">compiler_params</span><span class="o">=</span><span class="p">...,</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="p">...)</span>
</code></pre></div></div>

<p>核心概念表：</p>

<table>
  <thead>
    <tr>
      <th>概念</th>
      <th>作用</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">Ref</code></td>
      <td>kernel 内部的可读写内存视图</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">grid</code></td>
      <td>kernel program 的多维迭代空间</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">BlockSpec</code></td>
      <td>每个 grid point 如何映射到 input/output 的 tile</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">index_map</code></td>
      <td>从 grid indices 返回 array block indices</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">scratch_shapes</code></td>
      <td>跨 grid iteration 持久存在的临时 buffer</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">memory_space</code></td>
      <td>指定 VMEM / SMEM 等</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">compiler_params</code></td>
      <td>给 TPU backend 的编译 hint</td>
    </tr>
  </tbody>
</table>

<hr />

<h2 id="18-grid-怎么理解">18. grid 怎么理解</h2>

<p>文章中的 Splash grid 可以抽象成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>grid = (num_heads, num_q_blocks, num_kv_blocks)
</code></pre></div></div>

<p>一个 grid point：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>(h, i, j)
</code></pre></div></div>

<p>表示：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>h: 第几个 attention head
i: 第几个 Q block
j: 第几个 KV block
</code></pre></div></div>

<p>概念上像三层循环：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">for</span> <span class="n">h</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_heads</span><span class="p">):</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="n">num_q_blocks</span><span class="p">):</span>
        <span class="n">initialize</span> <span class="n">m</span><span class="p">,</span> <span class="n">l</span><span class="p">,</span> <span class="n">o</span> <span class="k">for</span> <span class="n">this</span> <span class="p">(</span><span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="p">)</span>

        <span class="k">for</span> <span class="n">j</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_kv_blocks</span><span class="p">):</span>
            <span class="n">scores</span> <span class="o">=</span> <span class="n">Q</span><span class="p">[</span><span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="p">]</span> <span class="o">@</span> <span class="n">K</span><span class="p">[</span><span class="n">h</span><span class="p">,</span> <span class="n">j</span><span class="p">].</span><span class="n">T</span>
            <span class="n">update</span> <span class="n">m</span><span class="p">,</span> <span class="n">l</span><span class="p">,</span> <span class="n">o</span>

        <span class="n">O</span><span class="p">[</span><span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="p">]</span> <span class="o">=</span> <span class="n">o</span> <span class="o">/</span> <span class="n">l</span>
</code></pre></div></div>

<p>实际 Pallas/TPU backend 可以根据 dimension semantics 和 pipeline 策略重排或并行化某些维度。</p>

<hr />

<h2 id="19-dimension_semantics-怎么理解">19. dimension_semantics 怎么理解</h2>

<p>文章提到类似：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dimension 0: heads       -&gt; parallel
dimension 1: q_blocks    -&gt; arbitrary
dimension 2: kv_blocks   -&gt; arbitrary
</code></pre></div></div>

<p>先不要把它理解成 tensor shape 的维度。它是 <code class="language-plaintext highlighter-rouge">grid</code> 的维度。</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>grid = (h, i, j)
</code></pre></div></div>

<p>因此：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dimension 0 = h = heads
dimension 1 = i = q_blocks
dimension 2 = j = kv_blocks
</code></pre></div></div>

<h3 id="191-heads-维为什么-parallel">19.1 heads 维为什么 parallel</h3>

<p>不同 attention head 之间没有依赖：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>head 0 不需要 head 1 的 m/l/o
head 1 不需要 head 0 的 m/l/o
</code></pre></div></div>

<p>所以可以告诉编译器：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>heads 维 iteration 可以并行或自由调度。
</code></pre></div></div>

<h3 id="192-q_blocks-维为什么数学上独立">19.2 q_blocks 维为什么数学上独立</h3>

<p>不同 Q block 也有自己的输出和 scratch：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Q block 0 有自己的 m/l/o
Q block 1 有自己的 m/l/o
</code></pre></div></div>

<p>数学上它们是独立的。不过文章示例中可能仍把它标成 <code class="language-plaintext highlighter-rouge">arbitrary</code>，这是保守选择，表示不要让编译器对这一维做过强假设。</p>

<h3 id="193-kv_blocks-维为什么有依赖">19.3 kv_blocks 维为什么有依赖</h3>

<p>对固定 <code class="language-plaintext highlighter-rouge">(h, i)</code>，要沿 <code class="language-plaintext highlighter-rouge">j</code> 扫过 KV blocks：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>j = 0 -&gt; 得到 m0, l0, o0
j = 1 -&gt; 需要读 m0, l0, o0，更新成 m1, l1, o1
j = 2 -&gt; 需要读 m1, l1, o1，更新成 m2, l2, o2
</code></pre></div></div>

<p>所以 KV block 维不能随便并行或乱序，因为 online softmax 的状态沿这一维传递。</p>

<p>一句话：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>parallel across heads
independent across Q blocks
accumulate across KV blocks
</code></pre></div></div>

<hr />

<h2 id="20-blockspec-怎么理解">20. BlockSpec 怎么理解</h2>

<p><code class="language-plaintext highlighter-rouge">BlockSpec</code> 定义：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>某个 grid point 应该看到数组的哪一个 tile。
</code></pre></div></div>

<p>例如 Q：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="n">block_shape</span><span class="o">=</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="n">bq</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">index_map</span><span class="o">=</span><span class="k">lambda</span> <span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">,</span> <span class="o">*</span><span class="n">_</span><span class="p">:</span> <span class="p">(</span><span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="p">,</span> <span class="mi">0</span><span class="p">),</span>
<span class="p">)</span>
</code></pre></div></div>

<p>假设 Q shape 是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Q: [num_heads, q_len, head_dim]
</code></pre></div></div>

<p><code class="language-plaintext highlighter-rouge">index_map</code> 返回：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>(h, i, 0)
</code></pre></div></div>

<p>分别对应：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>axis 0: head 轴       -&gt; 第 h 个 block
axis 1: sequence 轴   -&gt; 第 i 个 Q block
axis 2: head_dim 轴   -&gt; 第 0 个 feature block
</code></pre></div></div>

<p>结合：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block_shape=(None, bq, head_dim)
</code></pre></div></div>

<p>可以近似理解为：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">Q</span><span class="p">[</span><span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="o">*</span><span class="n">bq</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="o">*</span><span class="n">bq</span><span class="p">,</span> <span class="mi">0</span><span class="p">:</span><span class="n">head_dim</span><span class="p">]</span>
</code></pre></div></div>

<p>注意：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>None 不是“读完整维度”。
None 表示该维度取 size-1 的 slice，并在 kernel 内部 squeeze 掉。
</code></pre></div></div>

<p>所以 kernel 内部看到的 Q tile 通常是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[bq, head_dim]
</code></pre></div></div>

<p>而不是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[1, bq, head_dim]
</code></pre></div></div>

<h3 id="201-为什么-lambda-h-i-j-h-j-0-中的-0-表示读完整-head_dim">20.1 为什么 <code class="language-plaintext highlighter-rouge">lambda h, i, j: (h, j, 0)</code> 中的 0 表示读完整 head_dim</h3>

<p>它不是因为 <code class="language-plaintext highlighter-rouge">0</code> 有“完整”的特殊含义。</p>

<p>对于 K：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>K: [num_kv_heads, kv_len, head_dim]
</code></pre></div></div>

<p>BlockSpec：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="n">block_shape</span><span class="o">=</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="n">bkv</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span>
    <span class="n">index_map</span><span class="o">=</span><span class="k">lambda</span> <span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">:</span> <span class="p">(</span><span class="n">h</span><span class="p">,</span> <span class="n">j</span><span class="p">,</span> <span class="mi">0</span><span class="p">),</span>
<span class="p">)</span>
</code></pre></div></div>

<p>真实 slice 近似是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">K</span><span class="p">[</span><span class="n">h</span><span class="p">,</span> <span class="n">j</span><span class="o">*</span><span class="n">bkv</span><span class="p">:(</span><span class="n">j</span><span class="o">+</span><span class="mi">1</span><span class="p">)</span><span class="o">*</span><span class="n">bkv</span><span class="p">,</span> <span class="mi">0</span><span class="o">*</span><span class="n">head_dim</span><span class="p">:</span><span class="mi">1</span><span class="o">*</span><span class="n">head_dim</span><span class="p">]</span>
</code></pre></div></div>

<p>也就是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">K</span><span class="p">[</span><span class="n">h</span><span class="p">,</span> <span class="n">j</span><span class="o">*</span><span class="n">bkv</span><span class="p">:(</span><span class="n">j</span><span class="o">+</span><span class="mi">1</span><span class="p">)</span><span class="o">*</span><span class="n">bkv</span><span class="p">,</span> <span class="mi">0</span><span class="p">:</span><span class="n">head_dim</span><span class="p">]</span>
</code></pre></div></div>

<p>所以：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>0 只是 head_dim 这一轴的 block index。
因为 block_shape 的最后一维恰好等于完整 head_dim，
所以从第 0 个 block 开始读，就读到了完整 feature 维。
</code></pre></div></div>

<p>如果 block_shape 写成：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">block_shape</span><span class="o">=</span><span class="p">(</span><span class="bp">None</span><span class="p">,</span> <span class="n">bkv</span><span class="p">,</span> <span class="mi">32</span><span class="p">)</span>
</code></pre></div></div>

<p>而 <code class="language-plaintext highlighter-rouge">head_dim=128</code>，那么：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>0 -&gt; 读 0:32
1 -&gt; 读 32:64
2 -&gt; 读 64:96
3 -&gt; 读 96:128
</code></pre></div></div>

<p>此时 <code class="language-plaintext highlighter-rouge">0</code> 就不代表完整 head_dim 了。</p>

<hr />

<h2 id="21-kv-的-sparse-indirection">21. K/V 的 sparse indirection</h2>

<p>Dense attention 里，K/V 的 index map 可以很简单：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">lambda</span> <span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">:</span> <span class="p">(</span><span class="n">h</span><span class="p">,</span> <span class="n">j</span><span class="p">,</span> <span class="mi">0</span><span class="p">)</span>
</code></pre></div></div>

<p>意思是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>当前 program 是 (h, i, j)
就读取第 h 个 head、第 j 个 KV block
</code></pre></div></div>

<p>Sparse attention 里，某些 KV block 对当前 Q block 完全被 mask，没必要加载和计算。Splash 使用 <code class="language-plaintext highlighter-rouge">data_next_ref</code> 做间接寻址：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">k_index_map</span><span class="p">(</span><span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">,</span> <span class="n">data_next_ref</span><span class="p">,</span> <span class="n">block_mask_ref</span><span class="p">,</span> <span class="n">mask_next_ref</span><span class="p">):</span>
    <span class="n">next_j</span><span class="p">,</span> <span class="o">*</span><span class="n">_</span> <span class="o">=</span> <span class="n">_next_nonzero</span><span class="p">(</span>
        <span class="n">h</span><span class="p">,</span> <span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">,</span>
        <span class="n">data_next_ref</span><span class="p">,</span>
        <span class="n">block_mask_ref</span><span class="p">,</span>
        <span class="n">mask_next_ref</span><span class="p">,</span>
    <span class="p">)</span>
    <span class="k">return</span> <span class="p">(</span><span class="n">h</span> <span class="o">//</span> <span class="n">q_heads_per_kv_head</span><span class="p">,</span> <span class="n">next_j</span><span class="p">,</span> <span class="mi">0</span><span class="p">)</span>
</code></pre></div></div>

<p>直觉：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>本来要读 KV block j
但如果 j 是 fully masked block
就通过 data_next_ref 跳到下一个有效 block next_j
</code></pre></div></div>

<p>例如：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>j:       0  1  2  3  4  5
valid:   1  0  0  1  0  1
</code></pre></div></div>

<p>那么：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>data_next[h,i,1] 可能指向 3
data_next[h,i,2] 也可能指向 3
</code></pre></div></div>

<p>这样 kernel 在寻址阶段就跳过无效块：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>不是 load 之后发现不用，
而是根本不 load 那些 fully masked KV blocks。
</code></pre></div></div>

<p><code class="language-plaintext highlighter-rouge">h // q_heads_per_kv_head</code> 用于 GQA/MQA：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Q heads 可能比 KV heads 多。
多个 Q heads 共享同一个 KV head。
</code></pre></div></div>

<p>例如：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Q heads = 8
KV heads = 2
q_heads_per_kv_head = 4
</code></pre></div></div>

<p>则：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Q head 0,1,2,3 -&gt; KV head 0
Q head 4,5,6,7 -&gt; KV head 1
</code></pre></div></div>

<hr />

<h2 id="22-pallas-暴露的-tpu-memory-hierarchy">22. Pallas 暴露的 TPU memory hierarchy</h2>

<p>文章中的核心内存概念：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>HBM:
  片外大内存。
  Q/K/V/O 原始大数组通常在这里。

VMEM:
  片上 SRAM。
  默认 BlockSpec refs 指向 VMEM tile。

SMEM:
  scalar memory。
  适合小索引、mask metadata、next pointer、control-flow decision。

Scratch:
  kernel iterations 之间持久存在的临时 buffer。
</code></pre></div></div>

<h3 id="221-vmemblockspec-refs-默认在这里">22.1 VMEM：BlockSpec refs 默认在这里</h3>

<p>例如：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">((</span><span class="n">bq</span><span class="p">,</span> <span class="n">head_dim</span><span class="p">),</span> <span class="n">index_map</span><span class="p">)</span>
</code></pre></div></div>

<p>外面的 Q/K/V 大数组在 HBM，进入 kernel 时当前 tile 被搬到 VMEM。kernel 内部看到的是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Q tile: [bq, head_dim]
K tile: [bkv, head_dim]
V tile: [bkv, head_dim]
</code></pre></div></div>

<h3 id="222-smem放小的控制信息">22.2 SMEM：放小的控制信息</h3>

<p>例如：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">pl</span><span class="p">.</span><span class="n">BlockSpec</span><span class="p">(</span>
    <span class="p">(</span><span class="n">num_heads</span><span class="p">,),</span>
    <span class="k">lambda</span> <span class="o">*</span><span class="n">_</span><span class="p">:</span> <span class="p">(</span><span class="mi">0</span><span class="p">,),</span>
    <span class="n">memory_space</span><span class="o">=</span><span class="n">pltpu</span><span class="p">.</span><span class="n">SMEM</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>适合：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>block mask metadata
data_next pointer
mask_next pointer
小整数索引
控制流判断
</code></pre></div></div>

<p>这些不是大矩阵数据，而是告诉 kernel：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>当前 block 是否有效
下一个有效 block 是谁
怎样跳过 masked region
</code></pre></div></div>

<h3 id="223-scratchonline-softmax-的记忆">22.3 Scratch：online softmax 的记忆</h3>

<p>文章中的：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">m_scratch_ref</span>
<span class="n">l_scratch_ref</span>
<span class="n">o_scratch_ref</span>
</code></pre></div></div>

<p>分别保存：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>m: running max
l: running denominator
o: running unnormalized output numerator
</code></pre></div></div>

<p>对固定 <code class="language-plaintext highlighter-rouge">(h, i)</code>，沿 <code class="language-plaintext highlighter-rouge">j</code> 扫 KV blocks：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>j=0:
  初始化并写 m/l/o scratch

j=1:
  读旧 m/l/o scratch
  更新成新的 m/l/o

j=2:
  继续

最后:
  output = o / l
  写最终 O
</code></pre></div></div>

<p>关键：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>这些 scratch 在 VMEM 里跨 grid iterations 保留，
不需要每个 KV tile 都把 m/l/o 写回 HBM 再读回来。
</code></pre></div></div>

<p>这就是 online softmax 能高效实现的关键之一。</p>

<hr />

<h2 id="23-splash-attention-的执行逻辑">23. Splash Attention 的执行逻辑</h2>

<p>对一个 Q block 和一个 KV block：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Q_tile: [bq, head_dim]
K_tile: [bkv, head_dim]
V_tile: [bkv, head_dim]
</code></pre></div></div>

<p>当前 tile scores：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>scores_tile = Q_tile @ K_tile.T
</code></pre></div></div>

<p>shape：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>[bq, bkv]
</code></pre></div></div>

<p>这个 scores tile 只在 VMEM 中短暂存在。然后：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>更新 running max m
更新 running denominator l
更新 running output numerator o
</code></pre></div></div>

<p>再处理下一个 KV tile。</p>

<p>最终：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>O_tile = o / l
</code></pre></div></div>

<p>写回 HBM。</p>

<p>对比 naive：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>naive:
  scores_full = Q @ K.T
  scores_full shape = [H, S, S]
  scores_full 会跨 fusion materialize

Splash:
  scores_tile = Q_tile @ K_tile.T
  scores_tile shape = [bq, bkv]
  scores_tile 不跨 kernel 边界
</code></pre></div></div>

<p>这就是 HBM traffic 大幅下降的原因。</p>

<hr />

<h2 id="24-pallas-kernel-到-hlo为什么是-custom-call">24. Pallas kernel 到 HLO：为什么是 custom-call</h2>

<p>Pallas kernel 在 HLO 中通常表现为：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>custom-call(...)
custom_call_target="tpu_custom_call"
</code></pre></div></div>

<p>这意味着：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>XLA HLO graph 不再展开 Pallas kernel 内部细节。
</code></pre></div></div>

<p>XLA 看到的是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>这里有一个 custom-call
输入是 Q/K/V/mask metadata
输出是 O
</code></pre></div></div>

<p>Pallas kernel body 作为 MLIR payload 交给 Mosaic。之后 Mosaic 和 TPU backend 继续降到：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>LLO -&gt; VLIW bundles
</code></pre></div></div>

<p>所以 Pallas 的定位不是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>绕开整个 TPU compiler
</code></pre></div></div>

<p>而是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>绕开 XLA 高层图优化对 kernel 内部算法的表达限制，
但仍然使用 Mosaic / TPU backend 做底层 lowering 和 scheduling。
</code></pre></div></div>

<hr />

<h1 id="25-两篇文章的最终对比">25. 两篇文章的最终对比</h1>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>普通 JAX / XLA</th>
      <th>Pallas / Splash</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>你写的东西</td>
      <td>高层 tensor program</td>
      <td>tiled kernel program</td>
    </tr>
    <tr>
      <td>XLA 是否看见内部</td>
      <td>看见完整 HLO graph</td>
      <td>HLO 只看见 custom-call</td>
    </tr>
    <tr>
      <td>优化主要来自</td>
      <td>fusion、layout、tiling、memory assignment、VLIW scheduling</td>
      <td>你显式表达算法级 tiling / streaming，后端继续调度</td>
    </tr>
    <tr>
      <td>attention scores</td>
      <td>完整 <code class="language-plaintext highlighter-rouge">[H,S,S]</code> tensor 存在</td>
      <td>只有 <code class="language-plaintext highlighter-rouge">[bq,bkv]</code> tile 短暂存在</td>
    </tr>
    <tr>
      <td>softmax</td>
      <td>对完整 scores 分阶段处理</td>
      <td>online softmax，维护 <code class="language-plaintext highlighter-rouge">m/l/o</code></td>
    </tr>
    <tr>
      <td>中间矩阵 HBM traffic</td>
      <td>大</td>
      <td>小</td>
    </tr>
    <tr>
      <td>程序员负担</td>
      <td>低，主要写 JAX</td>
      <td>高，需要理解 grid、BlockSpec、scratch、memory space</td>
    </tr>
    <tr>
      <td>适用场景</td>
      <td>常规模型算子、可由 XLA 表达的图优化</td>
      <td>编译器不会自动发现的算法级优化</td>
    </tr>
  </tbody>
</table>

<hr />

<h1 id="26-给模型框架开发者的复习清单">26. 给模型框架开发者的复习清单</h1>

<h2 id="261-hlo-层重点看什么">26.1 HLO 层重点看什么</h2>

<p>读 HLO 时优先看：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 是否出现巨大中间 tensor
   例如 attention scores [H,S,S]

2. producer/consumer 是否 fusion
   elementwise + reduction + broadcast 是否合并

3. common producer 如何处理
   multi-output fusion / duplicate / materialize / recompute

4. layout 是否合理
   是否出现 TPU tile annotation

5. memory assignment
   大 tensor 是否 HBM
   小中间结果是否 VMEM

6. copy-start/copy-done
   是否有异步拷贝和 compute overlap
</code></pre></div></div>

<h2 id="262-llo-层只需先会识别模式">26.2 LLO 层只需先会识别模式</h2>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>vmatpush / vmatmul / vpop
  -&gt; MXU matmul

vadd / vmul / vsel / vcmp / vpow2 / vrsqrt
  -&gt; VPU vector compute

vxpose / vpop.trf / vrot.slane
  -&gt; transpose / shuffle / reduction

dma.hbm_to_vmem / dma.done.wait
  -&gt; HBM 与 VMEM 之间数据搬运

bundle { ... }
  -&gt; VLIW 静态并行发射包
</code></pre></div></div>

<h2 id="263-pallas-层重点看什么">26.3 Pallas 层重点看什么</h2>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. grid
   每个 program id 代表什么？
   哪些维度 independent，哪些维度 carry state？

2. BlockSpec
   每个 grid point 看到数组的哪个 tile？

3. index_map
   是否存在 sparse indirection？
   是否有 GQA/MQA head mapping？

4. scratch
   哪些状态跨 grid iteration 保留？
   是否避免 HBM round-trip？

5. memory_space
   哪些 metadata 放 SMEM？
   哪些 tile 放 VMEM？

6. custom-call
   HLO 是否只看见 opaque call？
   kernel body 是否进入 Mosaic？
</code></pre></div></div>

<hr />

<h1 id="27-最后用一句话复习">27. 最后用一句话复习</h1>

<p>第一篇：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>JAX 写自然的 tensor program，XLA 可以把它优化成带 fusion、VMEM placement、DMA overlap 和 VLIW bundle 的 TPU 程序。
</code></pre></div></div>

<p>第二篇：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>当性能瓶颈来自算法级 materialization，例如完整 attention matrix，XLA 的 fusion 仍然不够；Pallas 让我们直接表达 streaming tiled algorithm，用 online softmax 避免完整 scores 落 HBM。
</code></pre></div></div>

<p>最核心判断：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>如果问题是“这些 op 能不能合并得更好”，先看 XLA/HLO fusion。
如果问题是“这个巨大中间 tensor 是否本来就不该存在”，考虑算法改写和 Pallas。
</code></pre></div></div>

<hr />

<h2 id="参考资料">参考资料</h2>

<ul>
  <li>Patrick Toulme, <a href="https://patricktoulme.substack.com/p/from-jax-to-vliw-tracing-a-computation">From JAX to VLIW: Tracing a Computation Through the TPU Compiler</a></li>
  <li>Patrick Toulme, <a href="https://patricktoulme.substack.com/p/when-xla-isnt-enough-from-pallas">When XLA Isn’t Enough: From Pallas to VLIW</a></li>
  <li>JAX documentation, <a href="https://docs.jax.dev/en/latest/pallas/index.html">Pallas: a JAX kernel language</a></li>
  <li>JAX documentation, <a href="https://docs.jax.dev/en/latest/pallas/quickstart.html">Pallas Quickstart</a></li>
  <li>JAX documentation, <a href="https://docs.jax.dev/en/latest/pallas/grid_blockspec.html">Grids and BlockSpecs</a></li>
  <li>JAX documentation, <a href="https://docs.jax.dev/en/latest/_autosummary/jax.experimental.pallas.BlockSpec.html">jax.experimental.pallas.BlockSpec</a></li>
  <li>JAX documentation, <a href="https://docs.jax.dev/en/latest/pallas/tpu/details.html">Pallas TPU details</a></li>
  <li>JAX documentation, <a href="https://docs.jax.dev/en/latest/pallas/tpu/pipelining.html">Pallas TPU pipelining</a></li>
</ul>]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[JAX 到 VLIW，以及 Pallas / Splash Attention 复习笔记]]></summary></entry><entry><title type="html">MoE 优化的探索</title><link href="https://wqh011128.github.io/blog/2026/05/10/MoE-Optimization.html" rel="alternate" type="text/html" title="MoE 优化的探索" /><published>2026-05-10T00:00:00+00:00</published><updated>2026-05-10T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/05/10/MoE-Optimization</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/05/10/MoE-Optimization.html"><![CDATA[<h1 id="moe-优化的探索从-minimax-01cometflashmoe-到-deepseek-v4"><strong>MoE 优化的探索：从 MiniMax-01、Comet、FlashMoE 到 DeepSeek-V4</strong></h1>

<p><strong>Link:</strong></p>

<ul>
  <li>MiniMax-01: <a href="https://arxiv.org/pdf/2501.08313">arXiv PDF</a></li>
  <li>Comet: <a href="https://arxiv.org/abs/2502.19811">arXiv</a></li>
  <li>FlashMoE: <a href="https://flash-moe.github.io/">Project Page</a> / <a href="https://arxiv.org/abs/2506.04667">arXiv</a></li>
  <li>DeepSeek-V4: <a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/resolve/main/DeepSeek_V4.pdf">Technical Report</a></li>
  <li>DeepSeek MegaMoE2: <a href="https://github.com/deepseek-ai/DeepGEMM/pull/304">DeepGEMM PR #304</a> / <a href="https://github.com/deepseek-ai/DeepGEMM/pull/316">Benchmark PR #316</a></li>
</ul>

<p>[toc]</p>

<hr />

<h2 id="main-idea"><strong>Main Idea</strong></h2>

<p>MoE 的核心收益来自 sparse computation：每个 token 只激活少数 experts，因此模型总参数可以很大，但单 token 计算量仍然可控。</p>

<p>问题是，MoE 在分布式训练和推理里会引入很重的通信。一个 token 被 router 分到某个 expert 后，那个 expert 可能在另一张 GPU 上，于是系统必须先把 token 发过去，再把 expert 的输出发回来。这个过程通常叫：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Dispatch -&gt; Expert Compute -&gt; Combine
</code></pre></div></div>

<p>当模型规模变大时，瓶颈不再只是 expert 的矩阵乘，而是 <strong>通信和计算之间的空转</strong>。如果所有 GPU 都先等 <code class="language-plaintext highlighter-rouge">Dispatch All-to-All</code> 完成，再开始 expert GEMM，算完之后又一起等 <code class="language-plaintext highlighter-rouge">Combine All-to-All</code>，那么 GPU 会在通信阶段闲着，网络也会在计算阶段闲着。</p>

<p>MoE 优化的主线就是：</p>

<blockquote>
  <p>不要让通信和计算串行排队，而是把 token、expert、tile 或 kernel 切成更细粒度，让“正在算的一部分”和“正在传的一部分”重叠起来。</p>
</blockquote>

<hr />

<h2 id="1-moe-基本流程"><strong>1. MoE 基本流程</strong></h2>

<p>以常见的 Transformer MoE FFN 为例，一个 MoE 层可以分成四类操作。</p>

<table>
  <thead>
    <tr>
      <th>阶段</th>
      <th>做什么</th>
      <th>主要瓶颈</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>Router</td>
      <td>为每个 token 选择 top-k experts</td>
      <td>负载均衡、路由开销</td>
    </tr>
    <tr>
      <td>Dispatch</td>
      <td>把 token 发到 expert 所在设备</td>
      <td>all-to-all 通信</td>
    </tr>
    <tr>
      <td>Expert FFN</td>
      <td>每个 expert 对收到的 token 做 MLP</td>
      <td>GEMM 计算</td>
    </tr>
    <tr>
      <td>Combine</td>
      <td>把 expert 输出发回原位置并加权合并</td>
      <td>all-to-all 通信</td>
    </tr>
  </tbody>
</table>

<p>Expert FFN 本身通常又包含两个大 Linear：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Linear-1: X_e -&gt; hidden expansion
Activation: SiLU / SwiGLU
Linear-2: hidden expansion -&gt; output
</code></pre></div></div>

<p>如果写成 SwiGLU 风格：</p>

\[\begin{aligned}
g_e &amp;= X_e W_{e,\mathrm{gate}} \\
u_e &amp;= X_e W_{e,\mathrm{up}} \\
z_e &amp;= \mathrm{SiLU}(g_e) \odot u_e \\
Y_e &amp;= z_e W_{e,\mathrm{down}}
\end{aligned}\]

<p>其中 $X_e$ 是分配给 expert $e$ 的 token 表示。<code class="language-plaintext highlighter-rouge">Linear-1</code> 可以理解为 $W_{e,\mathrm{gate}}$ 和 $W_{e,\mathrm{up}}$ 这一侧的 GEMM，<code class="language-plaintext highlighter-rouge">Linear-2</code> 是 $W_{e,\mathrm{down}}$ 的 GEMM。</p>

<p>这里的 <strong>GEMM</strong> 是 General Matrix Multiplication，也就是通用矩阵乘。GPU 对 GEMM 极度优化，所以 expert compute 的核心就是如何让这些 GEMM 一直吃满算力。</p>

<hr />

<h2 id="2-minimax-01token-group--process-group-级别的-overlap"><strong>2. MiniMax-01：token group / process group 级别的 overlap</strong></h2>

<p>在MiniMax-01的总结中解读了EP，EP overlap，以及MiniMax的MoE创新。其优化重点不是把 expert 内部拆成 <code class="language-plaintext highlighter-rouge">Linear-1 / Linear-2</code> 的流水线，而是围绕 <strong>Expert Parallelism (EP)</strong>、<strong>Expert Tensor Parallelism (ETP)</strong> 和 <strong>Expert Data Parallelism (EDP)</strong> 做更合理的通信计算重叠。</p>

<h3 id="21-token-grouping-based-overlap">2.1 Token-grouping-based overlap</h3>

<p>MiniMax-01 先把 tokens 切成多个 group。每个 group 都要经历：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>a2a-dispatch -&gt; expert compute -&gt; a2a-combine
</code></pre></div></div>

<p>如果完全串行，时间线是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>group 0: dispatch -&gt; compute -&gt; combine
group 1: dispatch -&gt; compute -&gt; combine
group 2: dispatch -&gt; compute -&gt; combine
</code></pre></div></div>

<p>MiniMax-01 希望改成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>group 0: dispatch -&gt; compute -&gt; combine
group 1:          dispatch -&gt; compute -&gt; combine
group 2:                   dispatch -&gt; compute -&gt; combine
</code></pre></div></div>

<p>这样，某个 group 在做 expert compute 时，另一个 group 可以做 dispatch 或 combine。</p>

<p>这是一种 <strong>token group 粒度</strong> 的 overlap。它没有改变 expert FFN 的数学结构，也没有把单个 expert 的 GEMM 拆成 wave（像DeepSeek v4那样）。它只是把 token batch 切小，让通信和计算不要在整层级别完全串行。</p>

<h3 id="22-为什么还要-etp--edp">2.2 为什么还要 ETP / EDP</h3>

<p>MiniMax-01 进一步指出，仅靠 EP 不一定够。当 expert 参数太大时，可以用 <strong>Expert Tensor Parallelism (ETP)</strong> 把单个 expert 的参数也切到多个设备上。</p>

<p>这时 MoE 层的流程会变成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>a2a-dispatch -&gt; allgather -&gt; expert compute -&gt; reduce-scatter -&gt; a2a-combine
</code></pre></div></div>

<p>这里的 <code class="language-plaintext highlighter-rouge">allgather</code> 和 <code class="language-plaintext highlighter-rouge">reduce-scatter</code> 来自 tensor parallelism：为了让多个设备共同计算一个 expert，需要先收集输入或中间张量，再把结果规约分散回去。</p>

<p>MiniMax-01 还引入 <strong>Expert Data Parallelism (EDP)</strong>，本质上是在 expert 维度做复制，缓解某些 expert 负载过高的问题。</p>

<h3 id="23-minimax-01-的定位">2.3 MiniMax-01 的定位</h3>

<p>MiniMax-01 更像是训练系统级别的 MoE 并行优化：</p>

<ul>
  <li>重点是 EP / ETP / EDP 的组合。</li>
  <li>overlap 粒度主要是 token group 和 process group。</li>
  <li>它处理的是大规模训练时不同并行策略之间的通信压力。</li>
  <li>它不是 Comet / FlashMoE / DeepSeek 那种更细的 expert GEMM pipeline 或 kernel-level 调度。</li>
</ul>

<hr />

<h2 id="3-cometfine-grained-computation-communication-overlapping"><strong>3. Comet：fine-grained computation-communication overlapping</strong></h2>

<p>这一节开始单独精读 <strong>Comet: Fine-grained Computation-communication Overlapping for Mixture-of-Experts</strong>。这篇文章可以看成是在回答一个很具体的问题：</p>

<blockquote>
  <p>传统 EP overlap 已经把 token batch 切成 chunk 了，为什么 MoE 层里还是有明显 GPU idle time？如果继续往下切，应该切什么，怎么切，谁先算，谁先传？</p>
</blockquote>

<p>Comet 的答案不是简单地“再把 chunk 切小一点”。它真正做的是：找到 MoE layer 中通信算子和计算算子之间共享的 buffer，也就是 <strong>shared tensor</strong>，然后根据 consumer 的依赖关系决定沿哪个维度切分，并重新安排 GroupGEMM 的 tile 执行顺序。</p>

<h3 id="31-传统-ep-overlap-为什么还是-coarse-grained">3.1 传统 EP overlap 为什么还是 coarse-grained</h3>

<p>Comet 在 Introduction 里先分析了传统做法。一个分布式 MoE 层通常可以抽象成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Receive / Dispatch
-&gt; Expert computation
-&gt; Send / Combine
</code></pre></div></div>

<p>如果不 overlap，就是先收完所有 token，再做 expert GEMM，最后再发回结果。传统 EP overlap 会把 expert computation kernel 切成几个 chunk，让某个 chunk 的计算和另一个 chunk 的通信同时发生。</p>

<p><img src="https://arxiv.org/html/2502.19811/x1.png" alt="Comet Figure 1: MoE execution analysis" /></p>

<p>Figure 1(b) 想表达的是：把输入拆成 chunk 后，确实可以让一部分通信和一部分计算重叠，但这种 overlap 仍然是 <strong>coarse-grained</strong>。原因有三个。</p>

<p>第一，chunk 仍然必须作为一个整体 ready。即使 chunk 已经比完整 batch 小，一个 chunk 内部仍然可能有很多 token；只要这个 chunk 需要的 token 还没全部到齐，expert computation 就不能启动。</p>

<p>第二，chunk 变小会伤害 GEMM 效率。论文里提到，原本完整 expert computation 的时间是 $t$，切成两个 chunk 后可能变成 $t_1+t_2&gt;t$。这不是数学计算量变多了，而是小 GEMM 更难吃满 tensor core，还会带来更多调度和访存开销。</p>

<p>第三，MoE 是动态的。这里的动态不是模型结构变了，而是 router 每次会把不同 token 分给不同 experts。某个 step 里 Expert0 可能收到很多 token，Expert1 很少；下个 step 又可能反过来。于是每个 expert 的输入形状、通信量、计算量都在运行时变化，导致“通信 chunk”和“计算 chunk”的时间很难稳定对齐。</p>

<p>所以传统 EP overlap 的核心问题是：它把通信和计算装进不同 kernel / stream 里，让它们粗粒度并行，但对 GPU thread blocks、GEMM tile 顺序、remote I/O 这些底层资源缺少精细控制。</p>

<h3 id="32-moe-structure论文里的-forward-过程">3.2 MoE Structure：论文里的 forward 过程</h3>

<p>Comet 的 Figure 2 给了一个普通 MoE layer 的执行例子：两个 GPU，总共四个 experts，GPU0 放 Expert0/Expert1，GPU1 放 Expert2/Expert3。每个 token 被 router 分到 $k$ 个 experts。图里 Token A 被路由到 Expert0、Expert1、Expert3。</p>

<p><img src="https://arxiv.org/html/2502.19811/x2.png" alt="Comet Figure 2: MoE layer across two GPUs" /></p>

<p>论文中几个重要符号可以这样读：</p>

<table>
  <thead>
    <tr>
      <th>符号</th>
      <th>含义</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>$E$</td>
      <td>expert 总数</td>
    </tr>
    <tr>
      <td>$k$</td>
      <td>每个 token 被路由到的 expert 数量，也就是 top-k</td>
    </tr>
    <tr>
      <td>$TP$</td>
      <td>tensor parallel size</td>
    </tr>
    <tr>
      <td>$EP$</td>
      <td>expert parallel size</td>
    </tr>
    <tr>
      <td>$TP \times EP$</td>
      <td>总并行 world size</td>
    </tr>
    <tr>
      <td>$M$</td>
      <td>GEMM 中的 row 维度，在 MoE 里通常对应 token / token-expert rows</td>
    </tr>
    <tr>
      <td>$K$</td>
      <td>GEMM reduction 维度，例如输入 hidden 或 FFN intermediate 维</td>
    </tr>
    <tr>
      <td>$N$</td>
      <td>GEMM output column 维度，例如输出 hidden 或 FFN intermediate 维</td>
    </tr>
    <tr>
      <td>$T_M$</td>
      <td>GEMM tile 在 $M$ 维的大小</td>
    </tr>
    <tr>
      <td>$T_N$</td>
      <td>GEMM tile 在 $N$ 维的大小</td>
    </tr>
  </tbody>
</table>

<p>MoE 的每个 expert FFN 有两层 GEMM：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>layer0: first expert GEMM, usually up/gate projection
activation: SiLU / SwiGLU
layer1: second expert GEMM, usually down projection
</code></pre></div></div>

<p>论文把 MoE 的执行过程分成两类 pipeline。</p>

<p><strong>communication-computation pipeline:</strong> 对应 MoE layer0。</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Dispatch communication -&gt; layer0 GroupGEMM
</code></pre></div></div>

<p>这里通信是 producer，layer0 GroupGEMM 是 consumer。shared tensor 是 dispatch 之后、即将被 layer0 GEMM 消费的 expert input buffer。</p>

<p><strong>computation-communication pipeline:</strong> 对应 MoE layer1。</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>layer1 GroupGEMM -&gt; top-k routed reduction / combine communication
</code></pre></div></div>

<p>这里 layer1 GroupGEMM 是 producer，reduction/combine 是 consumer。shared tensor 是 layer1 GEMM 产生、即将被合并和通信的 expert output buffer。</p>

<h3 id="33-tp-是什么为什么-moe-会用">3.3 TP 是什么，为什么 MoE 会用</h3>

<p><strong>Tensor Parallelism (TP)</strong> 是把同一个线性层的权重切到多个 GPU 上。它和 EP 的区别是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>EP: 不同 experts 放在不同 GPU 上，每个 expert 权重通常是完整的。
TP: 同一个 expert / linear 的权重沿 hidden dimension 切开，多张 GPU 一起算一个矩阵乘。
</code></pre></div></div>

<p>比如一个线性层：</p>

\[Y = XW,\quad W\in\mathbb{R}^{K\times N}\]

<p>如果按 $N$ 维做 column parallel，GPU0 负责 $W[:,0:N/2]$，GPU1 负责 $W[:,N/2:N]$。如果按 $K$ 维做 row parallel，多个 GPU 分别计算部分乘积，再做 reduce。</p>

<p>MoE 会用 TP，是因为单个 expert 的 FFN 也可能很大。只用 EP 时，每个 expert 权重完整放在某张 GPU 上；如果 expert 太大，或者希望提升单 expert GEMM 的吞吐，就需要 TP 继续切 expert 内部权重。实际大模型里常常是 <strong>EP + TP 混合并行</strong>。</p>

<h3 id="34-granularity-mismatchtoken-level-communication-vs-tile-level-computation">3.4 Granularity mismatch：token-level communication vs tile-level computation</h3>

<p>Comet 的第一个关键观察是 <strong>granularity mismatch between computation and communication</strong>。</p>

<p>在 MoE 里，通信的基本单位通常是 token。Router 决定某个 token 要去哪个 expert，于是系统把这个 token 发到 expert 所在 GPU。</p>

<p>但高性能 GEMM 的基本单位不是单 token，而是 tile。论文里 Figure 2 的紫色块就是一个 computation tile，例如 $128\times128$。这意味着一个 expert 的某个 GEMM tile 可能需要 128 个 token rows，而这些 token 由 router 决定，可能随机分布在多个 GPU 上。</p>

<p>这就产生了依赖：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>一个 GEMM tile 需要的 token rows 没有全部 ready
-&gt; 这个 tile 不能开始计算
</code></pre></div></div>

<p>所以问题不是“tile 大小不一样导致完成时间不一样”，而是：</p>

<blockquote>
  <p>每个 tile 依赖的 token 来源不同，ready 时间不同；coarse-grained dispatch 会让 tile 等整个 chunk 或整个 expert input buffer。</p>
</blockquote>

<p>Comet 因此提出 fine-grained communication：每个 computation tile 通过 <strong>Unified Virtual Address (UVA)</strong> 直接读/写它需要的数据。</p>

<p>UVA 的实际作用是提供统一虚拟地址空间。在支持 GPU peer access 的情况下，GPU kernel 可以拿到远端 GPU buffer 的地址，并发起细粒度 remote load/store。它不是让 tile “自己有智能”，而是让 kernel 里的 communication thread blocks 可以根据路由 metadata，把某个 tile 需要的 remote token rows 拉到本地，或者把某个输出 tile 写回目标位置。</p>

<p>但 fine-grained remote I/O 很慢。如果把远程读写塞进 GEMM compute thread block，会破坏 tensor core pipeline。Comet 后面才需要 thread block specialization：通信 block 专门做 remote I/O，计算 block 保持高效 GEMM。</p>

<h3 id="35-design-overviewshared-tensor-是桥">3.5 Design overview：shared tensor 是桥</h3>

<p>Comet 的 Figure 3 是整篇文章的设计总览。</p>

<p><img src="https://arxiv.org/html/2502.19811/x3.png" alt="Comet Figure 3: Design overview" /></p>

<p>Comet 有两个核心设计：</p>

<table>
  <thead>
    <tr>
      <th>机制</th>
      <th>解决什么</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>Shared tensor based dependency resolving</td>
      <td>分析 producer/consumer 之间的真实数据依赖，决定 shared tensor 沿哪个维度切，并重排 tile 顺序</td>
    </tr>
    <tr>
      <td>Adaptive workload assignment</td>
      <td>在 fused kernel 内动态分配 thread blocks 给通信和计算，减少 pipeline bubble</td>
    </tr>
  </tbody>
</table>

<p>这里的 <strong>shared tensor</strong> 可以简单理解成：producer 和 consumer 共用的那块中间 buffer。</p>

<p>对 layer0：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>producer: dispatch communication
shared tensor: expert input X_e
consumer: layer0 GroupGEMM
</code></pre></div></div>

<p>对 layer1：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>producer: layer1 GroupGEMM
shared tensor: expert output Y_e
consumer: top-k routed reduction + combine communication
</code></pre></div></div>

<p>shared tensor 重要，是因为 overlap 只有在 producer 和 consumer 能处理 shared tensor 的不同独立部分时才成立。如果 consumer 必须等完整 tensor，overlap 就退化成普通串行。</p>

<h3 id="36-311how-to-decompose-the-shared-tensor">3.6 3.1.1：How to decompose the shared tensor</h3>

<p>Figure 4 把 layer0 和 layer1 都建模成 producer-consumer 关系。</p>

<p><img src="https://arxiv.org/html/2502.19811/x4.png" alt="Comet Figure 4: Producer-consumer modeling" /></p>

<p>Comet 的原则是：</p>

<blockquote>
  <p>沿 consumer 视角下相互独立的维度切 shared tensor。</p>
</blockquote>

<h4 id="361-layer0-为什么沿-m-切">3.6.1 Layer0 为什么沿 $M$ 切</h4>

<p>Layer0 的 shared tensor 是 layer0 GEMM 的输入矩阵：</p>

\[X_e \in \mathbb{R}^{M_e\times K}\]

<p>其中 $M_e$ 是 expert $e$ 收到的 token rows 数量，$K$ 是 token embedding / hidden dimension。</p>

<p>Layer0 的 consumer 是 GEMM：</p>

\[H_e = X_e W_{e,0}\]

<p>对 GEMM 来说，不同 token rows 之间相互独立。也就是说，先算 $X_e[M_0,:]$ 和后算 $X_e[M_1,:]$ 不会改变结果。因此 layer0 可以沿 $M$ 维切：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>X_e[M0, :] -&gt; layer0 GroupGEMM
X_e[M1, :] -&gt; layer0 GroupGEMM
X_e[M2, :] -&gt; layer0 GroupGEMM
</code></pre></div></div>

<p>但不能沿 $K$ 维随便切，因为 GEMM 对 $K$ 维做 reduction。算一个输出元素需要完整的 $K$ 维乘加：</p>

\[H_{e,i,n}=\sum_{k}X_{e,i,k}W_{e,0,k,n}\]

<p>如果切 $K$，不同分块之间还要额外做 partial sum reduction，consumer 不能直接独立消费。</p>

<h4 id="362-layer1-为什么沿-n-切">3.6.2 Layer1 为什么沿 $N$ 切</h4>

<p>Layer1 的 shared tensor 是 layer1 GEMM 的输出：</p>

\[Y_e \in \mathbb{R}^{M_e\times N}\]

<p>Layer1 后面的 consumer 不是普通逐元素操作，而是 <strong>top-k routed reduction + combine</strong>。注意，这里的 top-k reduction 不是重新选择 top-k；top-k 在 router 阶段已经完成了。这里的意思是：对同一个原始 token 的多个 expert 输出按 router weight 做加权合并。</p>

<p>例如 top-2：</p>

\[O_t = w_{t,e_1}Y_{t,e_1}+w_{t,e_2}Y_{t,e_2}\]

<p>如果沿 $M$ 切，可能把同一个 token 的两个 expert 输出拆到不同块：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>M tile 0: token A from Expert0
M tile 1: token A from Expert3
</code></pre></div></div>

<p>这时 consumer 处理 <code class="language-plaintext highlighter-rouge">M tile 0</code> 时拿不到 token A 的完整 top-k routed outputs，因此 $M$ 维存在 interdependency。</p>

<p>但 $N$ 维是 output feature column。不同 feature columns 的 weighted reduction 相互独立：</p>

\[O_{t,n}=w_{t,e_1}Y_{t,e_1,n}+w_{t,e_2}Y_{t,e_2,n}\]

<p>所以 layer1 可以沿 $N$ 切：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Y[:, N0] -&gt; reduction + combine
Y[:, N1] -&gt; reduction + combine
Y[:, N2] -&gt; reduction + combine
</code></pre></div></div>

<p>这就是论文中“layer0 沿 $M$ 分解，layer1 沿 $N$ 分解”的根本原因。</p>

<h3 id="37-312how-to-reschedule-the-decomposed-shared-tensor">3.7 3.1.2：How to reschedule the decomposed shared tensor</h3>

<p>只知道沿哪个维度切还不够。Comet 还要决定切完之后怎么排执行顺序。论文给了两个原则：</p>

<ol>
  <li>sub-tensors 要尽量对齐原始 GEMM tile granularity，否则 GEMM 效率会下降。</li>
  <li>优先执行 producer 已经产出、consumer 可以立即使用的部分，让 consumer 尽早启动。</li>
</ol>

<h4 id="371-layer0按-m-切后先算-local-token-tiles">3.7.1 Layer0：按 $M$ 切后，先算 local-token tiles</h4>

<p>Layer0 是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Dispatch -&gt; layer0 GroupGEMM
</code></pre></div></div>

<p>Figure 5 画的是 Rank0 上有三个 experts，每个 expert 都需要 local data 和 remote data。</p>

<p><img src="https://arxiv.org/html/2502.19811/x5.png" alt="Comet Figure 5: Decompose and reschedule layer0 shared tensor" /></p>

<p>Comet 会先按 source rank 对 token 排序。直觉上：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>local tokens | remote rank 1 tokens | remote rank 2 tokens | ...
</code></pre></div></div>

<p>然后 GroupGEMM 的 tile compute sequence 会优先从 local tokens 所在 tile 开始。这样本地 tile 可以马上计算，同时远程 token 还在通过 communication blocks 传输。</p>

<p>时间线可以理解成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>t0: compute local-token tiles
    communicate remote-rank-1 token rows

t1: compute remote-rank-1 tiles
    communicate remote-rank-2 token rows

t2: compute remote-rank-2 tiles
</code></pre></div></div>

<p>这里的 tile 通常类似 $T_M\times T_N$，例如论文前文举的 $128\times128$。但要注意，这个 tile 是 <strong>某个 expert GEMM 内部的 tile</strong>，不是把不同 experts 的 token 混在一起乘同一个权重。</p>

<h4 id="372-关键疑问不同-experts-权重不同128times128-tile-怎么算">3.7.2 关键疑问：不同 experts 权重不同，$128\times128$ tile 怎么算</h4>

<p>这是理解 GroupGEMM 的关键。</p>

<p>GroupGEMM 不是把所有 experts 的 token 拼成一个大矩阵，然后乘同一个权重。它是一组独立 GEMM 的调度：</p>

\[H_e = X_e W_{e,0},\quad e\in\mathcal{E}_{\mathrm{local}}\]

<p>也就是说，Rank0 上如果有 Expert0、Expert1、Expert2，那么 GroupGEMM 实际上在执行：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Expert0: X_0 @ W_0
Expert1: X_1 @ W_1
Expert2: X_2 @ W_2
</code></pre></div></div>

<p>每个 computation tile 都带着 expert id。属于 Expert0 的 tile 只会用 $W_0$，属于 Expert1 的 tile 只会用 $W_1$。高性能 grouped GEMM kernel 会通过 metadata / pointer arrays / offsets 找到每个 expert 对应的 input pointer、weight pointer 和 output pointer。</p>

<p>所以 Figure 5 中的 $M$ 切分不是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>把不同 expert 的 128 行混成一个 tile，用同一个 W 去乘
</code></pre></div></div>

<p>而是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>在每个 expert 自己的 X_e[M_e, K] 里按 M tile 切；
GroupGEMM 只是把多个 expert 的 tile 放到同一个 kernel 里统一调度。
</code></pre></div></div>

<p>如果某个 expert 收到的 token 不足一个完整 $T_M$，实现上可以用 partial tile、padding 或 grouped GEMM 的 ragged shape metadata 处理。数学上仍然是每个 expert 使用自己的权重。</p>

<p>这也解释了为什么 layer0 按 $M$ 切完不需要“还原成原始 token 顺序”再进入 layer1。Layer0 输出仍然按 expert 分组保存：</p>

\[H_e[M_e, K']\]

<p>中间 activation 是逐元素的，不需要跨 expert 或跨 token 重新排列。Layer1 继续对同一个 expert 的 $H_e$ 做：</p>

\[Y_e = H_e W_{e,1}\]

<p>真正需要恢复到原始 token 顺序，是 layer1 结束后的 routed reduction / combine 阶段。</p>

<h4 id="373-layer1按-n-切后column-wise-执行-groupgemm">3.7.3 Layer1：按 $N$ 切后，column-wise 执行 GroupGEMM</h4>

<p>Layer1 是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>layer1 GroupGEMM -&gt; reduction + combine communication
</code></pre></div></div>

<p>Figure 6 说明了 Comet 如何重排 layer1 的 GroupGEMM。</p>

<p><img src="https://arxiv.org/html/2502.19811/x6.png" alt="Comet Figure 6: Rescheduled compute sequence for layer1" /></p>

<p>如果不重排，GroupGEMM 可能按 expert 顺序执行：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Expert0: N0 -&gt; N1 -&gt; N2 -&gt; N3
Expert1: N0 -&gt; N1 -&gt; N2 -&gt; N3
Expert2: N0 -&gt; N1 -&gt; N2 -&gt; N3
</code></pre></div></div>

<p>这样 consumer 很难提前开始，因为它想处理某个 column block 时，需要相关 experts 的同一段 columns 都已经产生。</p>

<p>Comet 改成 column-wise：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>N0 group:
  Expert0:N0 -&gt; Expert1:N0 -&gt; Expert2:N0

N1 group:
  Expert0:N1 -&gt; Expert1:N1 -&gt; Expert2:N1

N2 group:
  Expert0:N2 -&gt; Expert1:N2 -&gt; Expert2:N2
</code></pre></div></div>

<p>注意，这里的 <code class="language-plaintext highlighter-rouge">N0</code> 不是“对 128 列做 top-k selection”。Top-k selection 已经在 router 完成。这里做的是：对已经确定的 top-k experts，在 <code class="language-plaintext highlighter-rouge">N0</code> 这一段 output features 上做 weighted reduction 和 combine。</p>

<p>如果 $N_0$ 表示 columns $0:T_N$，那么 consumer 可以先做：</p>

\[O_{t,N_0}=\sum_{e\in\mathrm{TopK}(t)}w_{t,e}Y_{t,e,N_0}\]

<p>同时 layer1 GroupGEMM 继续计算 $Y[:,N_1]$。这样就形成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>compute Y[:, N0]
-&gt; reduce/combine Y[:, N0]

while

compute Y[:, N1]
</code></pre></div></div>

<p>所以 Figure 6 的核心不是“列维度上重新选 top-k”，而是：按 column block 提前产出可被 consumer 完整处理的一段 output features。</p>

<h3 id="38-32adaptive-workload-assignment">3.8 3.2：Adaptive Workload Assignment</h3>

<p>经过 3.1 的 dependency resolving 后，Comet 已经知道哪些数据可以先算、哪些数据可以先传。但还有一个问题：fine-grained remote I/O 很慢，GEMM 又很吃 tensor core。谁来做通信，谁来做计算，分配多少 GPU 资源，不能拍脑袋。</p>

<p>Figure 7 展示的是 Comet 在 Hopper 上的 fused kernel 设计。</p>

<p><img src="https://arxiv.org/html/2502.19811/x7.png" alt="Comet Figure 7: Thread block specialized kernel" /></p>

<h4 id="381-thread-block-specialization">3.8.1 Thread block specialization</h4>

<p>最直接的融合方式叫 vertical fusion：每个 thread block 既做 GEMM，也在 prologue / epilogue 里做通信 I/O。</p>

<p>问题是 remote I/O 延迟远高于本地显存访问。如果把 remote read/write 插进 GEMM thread block，可能会阻塞后续 tensor core 计算，尤其 Hopper 上 GEMM 通常利用 TMA 建立异步 compute pipeline，长延迟 remote I/O 会破坏这个 pipeline。</p>

<p>Comet 因此把 thread blocks 隔离成两类：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>compute thread blocks: 负责 GEMM，尽量复用默认 CUTLASS GEMM 实现
communication thread blocks: 负责 remote I/O、top-k routed reduction、local/remote writeback
</code></pre></div></div>

<p>这样做的代价是会多一些 global memory 读写，但收益是通信不会污染 GEMM 的关键路径，而且系统可以精确控制多少 blocks 做通信、多少 blocks 做计算。</p>

<h4 id="382-adaptive-thread-block-assignment">3.8.2 Adaptive thread block assignment</h4>

<p>论文 3.2.2 解决的是：通信 block 和计算 block 到底分多少？</p>

<p>假设一个 fused kernel 总共有 $N_{\mathrm{TB}}$ 个 thread blocks，其中：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>N_{\mathrm{comp}}: compute blocks
N_{\mathrm{comm}}: communication blocks
N_{\mathrm{TB}} = N_{\mathrm{comp}} + N_{\mathrm{comm}}
</code></pre></div></div>

<p>如果 $N_{\mathrm{comm}}$ 太少，远程 I/O 跟不上，GEMM 算完后会等通信。
如果 $N_{\mathrm{comm}}$ 太多，GEMM blocks 变少，计算吞吐下降。
最佳分界点和输入 token length、TP/EP 配置、expert shape、硬件带宽都有关。</p>

<p>Figure 8 说明不同配置下最优 $N_{\mathrm{comm}}$ 不同。</p>

<p><img src="https://arxiv.org/html/2502.19811/x8.png" alt="Comet Figure 8: Adaptive thread block assignment" /></p>

<p>论文给的例子是：当输入 token length 从 4096 变到 16384 时，最优通信 block 数会变化；当 TP 从 8 调到 4 时，最优分配点也会明显变化。</p>

<p>所以 Comet 的做法不是运行时在线搜索，而是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 预编译多个 kernel，每个 kernel 使用不同的 compute/communication block division point。
2. 部署前 profile 不同模型配置和输入形状，记录最优配置 metadata。
3. 运行时根据 metadata 选择合适 kernel。
</code></pre></div></div>

<p>这就是 adaptive workload assignment。它的目标不是改变数学计算，而是让 fine-grained pipeline 的通信段和计算段时间尽量对齐，减少 pipeline bubbles。</p>

<h3 id="39-把我们讨论过的几个易错点放在一起">3.9 把我们讨论过的几个易错点放在一起</h3>

<p><strong>Tile 不是 token chunk。</strong> token chunk 是 EP overlap 的 coarse 粒度；tile 是 GEMM kernel 的计算分块，例如 $128\times128$。一个 tile 通常包含多个 token rows 和一段 output columns。</p>

<p><strong>Tile 不是 expert 的最大容量。</strong> 一个 expert 实际收到多少 token 由 router 决定，记作 $M_e$。GEMM tile 是在 $M_e\times K$ 或 $M_e\times N$ 这个矩阵内部继续切出来的计算块。</p>

<p><strong>Layer0 的 $M$ 切分发生在 layer0 计算前。</strong> 它切的是 dispatch 后的 expert input $X_e$，目的是 token rows 到一块、算一块。</p>

<p><strong>Layer1 的 $N$ 切分发生在 layer1 计算过程中。</strong> 它不是先完整算出 $Y_e$ 再切，而是调整 GroupGEMM 顺序，直接先产出 $Y[:,N_0]$，让 reduction/combine 提前消费。</p>

<p><strong>Layer0 到 layer1 中间不需要还原成原始 token 顺序。</strong> layer0 输出仍然按 expert 分组，activation 和 layer1 都可以在 expert-local layout 里继续做。只有最终 combine 时才需要根据 routing metadata 回到原 token 位置。</p>

<p><strong>UVA 不等于免费远程访问。</strong> UVA 只是让 kernel 可以用统一地址访问远端 GPU buffer；真正的通信仍然有高延迟，所以 Comet 才要用 communication thread blocks 隔离远程 I/O。</p>

<h3 id="310-comet-的定位">3.10 Comet 的定位</h3>

<p>Comet 的核心不是“把 MoE 拆成几个阶段”，而是：</p>

<blockquote>
  <p>找到 MoE 中通信和 GEMM 之间共享的 tensor，分析 consumer 在哪个维度上可以独立消费，再按这个维度切分并重排 GroupGEMM tile 顺序，最后用 adaptive thread block assignment 让通信和计算更稳定地重叠。</p>
</blockquote>

<p>它比 MiniMax-01 更细，因为它已经进入 GEMM tile / thread block 级别；但它还没有像 FlashMoE 那样把整个 MoE operator 改造成一个 persistent GPU runtime。Comet 更像是在现有 MoE/GEMM 执行栈上，通过 dependency resolving 和 fused kernel scheduling，把 coarse EP overlap 推到 tile-level overlap。</p>

<hr />

<h2 id="4-flashmoepersistent-kernel--gpu-resident-scheduling"><strong>4. FlashMoE：persistent kernel + GPU-resident scheduling</strong></h2>

<p>如果说 Comet 是“把 MoE 中通信和 GEMM 之间的 shared tensor 拆到 tile 级，并重排 producer-consumer 顺序”，那 FlashMoE 再往前走了一步：</p>

<blockquote>
  <p>只优化某两段 pipeline 还不够。只要 MoE 仍然依赖 host-managed scheduling、bulk-synchronous collectives 和大量短 kernel launch，就仍然会有系统级 idle gap。FlashMoE 想把整个 distributed MoE operator 放进一个 GPU-resident persistent kernel 里。</p>
</blockquote>

<p>这篇文章适合在读完 Comet 后继续看，因为它们都在讲 fine-grained overlap，但抽象层级不同：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Comet:
  找 shared tensor -&gt; 按 consumer 依赖切分 -&gt; 重排 GEMM tile -&gt; 专门 thread blocks 做通信/计算

FlashMoE:
  把 MoE operator 变成 GPU 常驻 runtime
  用 actor、task、symmetric tensor layout 管理 dispatch / expert compute / combine
</code></pre></div></div>

<h3 id="41-读前概念gpu-kernelmegakernelkernel-launch-和-sm">4.1 读前概念：GPU kernel、megakernel、kernel launch 和 SM</h3>

<p>FlashMoE 比 Comet 难读，很大一部分原因是它不只是讲 MoE，而是在讲 GPU runtime。先把几个底层词拆开。</p>

<p><strong>GPU kernel</strong> 是运行在 GPU 上的函数。CPU 端调用 CUDA runtime，把一个 kernel 发射到 GPU 上执行，这个动作叫 <strong>kernel launch</strong>。一次 launch 会指定 grid / block / thread 结构，例如：</p>

<div class="language-cpp highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">moe_kernel</span><span class="o">&lt;&lt;&lt;</span><span class="n">num_blocks</span><span class="p">,</span> <span class="n">threads_per_block</span><span class="o">&gt;&gt;&gt;</span><span class="p">(</span><span class="n">args</span><span class="p">);</span>
</code></pre></div></div>

<p>这个调用发生在 CPU host 侧。GPU 收到任务后，把很多 thread blocks 分配到各个 <strong>SM (Streaming Multiprocessor)</strong> 上执行。SM 可以理解成 GPU 里的计算工厂，一个 GPU 有很多个 SM；每个 SM 同时驻留若干 thread blocks，内部以 warp 为调度单位执行指令。</p>

<p><strong>kernel launch overhead</strong> 指的是：CPU 发起 kernel、CUDA runtime 排队、GPU 接收任务、kernel 之间同步和数据交接带来的额外时间。单个 launch 的开销看起来可能不大，但 MoE 层如果由很多短 kernel 和 collective 串起来，launch gap 会变得明显。</p>

<p><strong>megakernel</strong> 通常指把原本多个 GPU kernels / operators 合并成一个更大的 kernel，让这个 kernel 内部自己完成多个阶段的工作。FlashMoE 的 persistent kernel 可以看成一种 MoE megakernel：它不是只算一个 GEMM，而是在一个常驻 kernel 里调度 dispatch、expert FFN、combine 等任务。</p>

<p>这里的“减少 CPU launch”不是说 kernel 可以像 Python 函数那样反复复用，而是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>传统做法:
  CPU launch dispatch kernel
  CPU launch GEMM kernel
  CPU launch activation kernel
  CPU launch GEMM kernel
  CPU launch combine kernel

FlashMoE:
  CPU launch one persistent kernel
  persistent kernel 内部自己调度 MoE tasks
</code></pre></div></div>

<p>论文 Table 1 里说 FlashMoE 的 <code class="language-plaintext highlighter-rouge">#GPU Ops</code> 是 1，而 Comet / Megatron / DeepEP 是几十到几百个。这里的 GPU op 可以理解成需要 host/runtime 调度的一次 GPU operator 或 kernel-level operation。其他框架不是“每个 expert 必然启动一个 kernel”这么简单，实际可能是 group GEMM、通信 collective、elementwise、permute、combine 等很多 operators；但它们仍然需要多次 kernel launch 或 runtime scheduling。FlashMoE 的关键就是把这些阶段压进一个 persistent execution context。</p>

<p>这也是为什么如果不理解 kernel launch，很难感受到 FlashMoE 的收益。它省的不只是某个 GEMM 的计算时间，而是整条 MoE operator 里反复出现的：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>host launch -&gt; GPU 执行 -&gt; 等待/同步 -&gt; host launch 下一个 op
</code></pre></div></div>

<h3 id="42-几个底层库和机制nvshmemdmain-device-blascutlass">4.2 几个底层库和机制：NVSHMEM、DMA、in-device BLAS、CUTLASS</h3>

<p>FlashMoE 还频繁提到几个系统名词。</p>

<table>
  <thead>
    <tr>
      <th>名词</th>
      <th>怎么理解</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">NVSHMEM</code></td>
      <td>NVIDIA 的 GPU-side PGAS 通信库。它让多个 GPU 拥有一个对称的 shared address space，并支持 GPU kernel 内部发起 put/get/signal 等 one-sided communication。</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">DMA</code></td>
      <td>Direct Memory Access，直接内存访问。意思是数据搬运可以由硬件/设备引擎完成，不必由 CPU 一字节一字节参与。FlashMoE 语境下重点是 device-initiated remote memory transfer。</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">in-device BLAS</code></td>
      <td>在 GPU device 侧调用/实现的线性代数操作。BLAS 是矩阵乘、向量运算等基础线性代数接口，FlashMoE 需要在 persistent kernel 内部做类似 GEMM/FFN 的计算。</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">CUTLASS</code></td>
      <td>NVIDIA 开源的 CUDA C++ 模板库，用来写高性能 GEMM / convolution 等 kernels。可以把它看成构建自定义矩阵乘 kernel 的积木。</td>
    </tr>
  </tbody>
</table>

<p>如果用一句话连接这些概念：</p>

<blockquote>
  <p>FlashMoE 用 NVSHMEM / UVA / DMA 处理跨 GPU tile 搬运，用 CUTLASS 风格的 device-side GEMM 处理 expert FFN，并把它们组织进一个 persistent MoE megakernel。</p>
</blockquote>

<h3 id="43-从-comet-视角看-flashmoe-的动机">4.3 从 Comet 视角看 FlashMoE 的动机</h3>

<p>FlashMoE 论文的 Introduction 和 Motivation 里强调三个传统 MoE 系统问题。</p>

<p><img src="https://flash-moe.github.io/static/images/intro-fig.png" alt="FlashMoE overview" /></p>

<p>第一是 <strong>synchronous communication</strong>。传统 MoE 常用 <code class="language-plaintext highlighter-rouge">AllToAll</code> / <code class="language-plaintext highlighter-rouge">AllGather</code>。这些 collective 是 bulk-synchronous 的：参与通信的 GPU 都要进入同一个集体操作，慢的 GPU 会拖住快的 GPU。MoE router 又是动态的，每个 step 分到各 expert 的 token 数不同，所以 straggler 很常见。</p>

<p>第二是 <strong>kernel launch overhead</strong>。一个 MoE forward 可能包含：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>router kernel
dispatch communication kernel
expert GEMM kernel
activation kernel
expert GEMM kernel
combine communication kernel
</code></pre></div></div>

<p>如果每个阶段都要 host 端发起 kernel，并在 kernel 之间交接，短 kernel 和同步点会堆出很明显的 CUDA API / launch gap。FlashMoE 论文里强调，它的目标是一个 persistent kernel，而不是一串短命 kernel。</p>

<p>第三是 <strong>task locality 没有被充分利用</strong>。Comet 已经告诉我们：数据到达和 GEMM tile ready 是更细的粒度。但如果调度仍然由 CPU 或粗粒度 collective 管，GPU 很难做到“这个 tile ready 了就马上安排一个 block 去算/传”。</p>

<p>所以 FlashMoE 的目标不是再提出一种新的 router，也不是改变 MoE 数学，而是把 MoE 的执行形态从：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>host launches many kernels + collective barriers
</code></pre></div></div>

<p>变成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>one persistent GPU kernel + device-side tasks + one-sided communication
</code></pre></div></div>

<h3 id="44-moe-数学没有变变的是执行系统">4.4 MoE 数学没有变：变的是执行系统</h3>

<p>FlashMoE 仍然处理普通的 MoE FFN。对每个 token $x_i$，gate 选择 top-k experts，并给出 combine weights。一个简化的 MoE 输出可以写成：</p>

\[y_i = \sum_{j=1}^{k} w_{i,e_j} \cdot E_{e_j}(x_i)\]

<p>其中 $E_{e_j}$ 是第 $j$ 个被选中的 expert，$w_{i,e_j}$ 是 gate 给出的权重。</p>

<p>每个 expert 本身还是普通 FFN：</p>

\[E_e(x)=W_{e,2}\,\sigma(W_{e,1}x+b_{e,1})+b_{e,2}\]

<p>所以要记住：</p>

<blockquote>
  <p>FlashMoE 改的是 MoE operator 的 runtime，不是 MoE 的函数形式。</p>
</blockquote>

<p>这和 Comet 一样：Comet 的 shared tensor decomposition 也不改变数学结果，只改变数据到达、GEMM tile 和 consumer operator 的执行顺序。</p>

<h3 id="45-flashmoe-的核心架构figure-5-和-single-persistent-kernel">4.5 FlashMoE 的核心架构：Figure 5 和 single persistent kernel</h3>

<p>项目页里的架构图和论文 Figure 5 很适合作为第一张地图。</p>

<p><img src="https://flash-moe.github.io/static/images/architecture.jpg" alt="FlashMoE architecture" /></p>

<p>FlashMoE 的 persistent kernel 在每张 GPU 上长期运行。它内部不是所有 thread blocks 都做同一件事，而是用 actor model 把 blocks / warps 分成几类角色：</p>

<table>
  <thead>
    <tr>
      <th>角色</th>
      <th>作用</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>Processor</td>
      <td>执行真正的计算和通信任务，例如 GEMM、element-wise、combine、tile transfer</td>
    </tr>
    <tr>
      <td>Subscriber</td>
      <td>接收 peer GPU 发来的 tile packets，并解码成 task descriptors</td>
    </tr>
    <tr>
      <td>Scheduler</td>
      <td>维护 ready queue，把已经 ready 的 task 分配给 Processor</td>
    </tr>
  </tbody>
</table>

<p>论文实现里，大部分 thread blocks 是 Processor。最后一个 block 被当作类似 “OS block” 的管理块，其中三条 warps 做 Subscriber，一条 warp 做 Scheduler。这个细节不太显眼，但很重要：FlashMoE 并不是让所有 blocks 都参与调度，而是用很少的 GPU 资源做管理，把大部分 SM 留给计算任务。</p>

<p>这句话可以拆成一个更具体的执行画面：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>GPU 上有很多 SM。
persistent kernel 启动后，许多 Processor blocks 驻留在这些 SM 上。
Processor blocks 不固定属于某个 expert，也不固定只做 dispatch 或 GEMM。

OS block:
  Subscriber warps: 监听其他 GPU 发来的 messages / packets。
  Scheduler warp:   把已经 ready 的 task 放进队列，并分配给空闲 Processor。

Processor block:
  从队列拿 task。
  task ready 就执行，不 ready 就换下一个。
</code></pre></div></div>

<p>所以原文说 “ensuring that no GPU SM remains idle throughout the lifetime of the MoE operator” 的意思不是物理上 100% 没有任何空泡，而是：FlashMoE 不希望某些 SM 因为等某个大 collective 或某个固定 expert 而长期闲着。只要队列里还有 ready tasks，Scheduler 就可以把它们分给空闲 Processor blocks。</p>

<p>传统做法更像：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>所有 GPU 等 dispatch 完成
然后所有 GPU 做 GEMM
然后所有 GPU 等 combine
</code></pre></div></div>

<p>FlashMoE 更像：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>哪个 tile 到了，就生成哪个 task。
哪个 Processor 空了，就拿一个 ready task。
</code></pre></div></div>

<p>这就是 Figure 5 的核心：它把 MoE operator 从“固定阶段流水线”变成了“GPU 内部任务系统”。</p>

<p>这和 Comet 的 thread block specialization 有继承关系：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Comet:
  compute blocks + communication blocks

FlashMoE:
  Processor blocks + Subscriber/Scheduler management warps
</code></pre></div></div>

<p>Comet 的重点是隔离 remote I/O 和 GEMM；FlashMoE 的重点是进一步把 task 的产生、通知、调度、执行都放进 GPU kernel 内部。</p>

<h3 id="46-actor-model-of-concurrent-computation-是什么">4.6 Actor model of concurrent computation 是什么</h3>

<p>Actor model 是一种并发计算范式，不是机器学习模型。它最早来自并发系统设计，核心思想是：系统由很多 actor 组成，每个 actor 有自己的状态，通过 message 通信，收到 message 后执行动作、修改状态、继续发送 message。</p>

<p>放到 FlashMoE 里，可以这样对应：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Processor 是执行 actor:
  收到 task descriptor 后执行 tile compute / tile transfer / combine。

Subscriber 是消息入口 actor:
  接收来自其他 GPU 的 messages，并把它们解码成本地可调度的 tasks。

Scheduler 是调度 actor:
  维护 ready tasks，把任务分派给 Processor。
</code></pre></div></div>

<p>这和普通 CUDA kernel 的区别是：普通 kernel 通常是 “每个 thread/block 按固定索引算自己的那一块”；FlashMoE 的 actor-style kernel 则是 “block 从任务队列拿活干”。它更接近 GPU 上的小型 runtime。</p>

<p>论文 Figure 6 里的 actor dependencies 可以按这条链理解：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>remote GPU sends message
-&gt; Subscriber decodes message into task
-&gt; Scheduler observes task readiness
-&gt; Processor executes task
-&gt; Processor may produce another message/task
</code></pre></div></div>

<p>这个图的重点不是 actor 名字本身，而是 <strong>任务依赖不再由 CPU host 串行推进，而是在 GPU 内部通过 message 和 ready queue 推进</strong>。</p>

<h3 id="47-unified-task-abstractiontile-是-runtime-的基本工作单元">4.7 Unified task abstraction：tile 是 runtime 的基本工作单元</h3>

<p>Comet 里我们已经反复说过，tile 不是 token chunk，而是 GEMM / tensor 的静态分块。FlashMoE 沿用这个视角，但把 tile 进一步封装成 runtime task。</p>

<p>论文把 task descriptor 理解成一组 metadata + operator。一个 task 至少要告诉 Processor：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>我要处理哪块 tile
这个 tile 属于哪个 expert / 哪张 GPU / 哪个通信阶段
要执行的是 FFN、combine，还是 tile transfer
输入地址和输出地址在哪里
依赖是否已经 ready
</code></pre></div></div>

<p>用一个很简化的伪代码表达：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">while</span> <span class="n">persistent_kernel_is_running</span><span class="p">:</span>
    <span class="n">task</span> <span class="o">=</span> <span class="n">scheduler</span><span class="p">.</span><span class="n">pop_ready_task</span><span class="p">()</span>

    <span class="k">if</span> <span class="n">task</span><span class="p">.</span><span class="nb">type</span> <span class="o">==</span> <span class="s">"dispatch_tile"</span><span class="p">:</span>
        <span class="n">processor</span><span class="p">.</span><span class="n">transfer_tile</span><span class="p">(</span><span class="n">task</span><span class="p">)</span>

    <span class="k">elif</span> <span class="n">task</span><span class="p">.</span><span class="nb">type</span> <span class="o">==</span> <span class="s">"ffn_tile"</span><span class="p">:</span>
        <span class="n">processor</span><span class="p">.</span><span class="n">gemm_and_activation</span><span class="p">(</span><span class="n">task</span><span class="p">)</span>

    <span class="k">elif</span> <span class="n">task</span><span class="p">.</span><span class="nb">type</span> <span class="o">==</span> <span class="s">"combine_tile"</span><span class="p">:</span>
        <span class="n">processor</span><span class="p">.</span><span class="n">weighted_accumulate_and_writeback</span><span class="p">(</span><span class="n">task</span><span class="p">)</span>
</code></pre></div></div>

<p>这里的关键不是伪代码本身，而是 <strong>ready task</strong> 这个概念。传统 MoE 是阶段式：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>所有 dispatch 完成 -&gt; 所有 expert GEMM 开始 -&gt; 所有 combine 开始
</code></pre></div></div>

<p>FlashMoE 是任务式：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>某个 dispatch tile 到了 -&gt; 生成 FFN task
某个 FFN output tile 完成 -&gt; 生成 combine task
某个 remote packet 到了 -&gt; Subscriber 解码成 task
</code></pre></div></div>

<p>这就把 Comet 的“tile ready 就尽早消费”推广成一个 GPU 内部的任务系统。</p>

<p>所谓 <strong>Unified task abstraction</strong>，就是把 MoE 中看起来不同的事情统一成同一种 task 格式。</p>

<p>比如传统看法里，下面几件事是不同 operators：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>dispatch token tile
run expert FFN tile
apply activation
combine output tile
write back remote output
</code></pre></div></div>

<p>FlashMoE 会尽量把它们都描述成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>task = {
  operator_type,
  input_pointer,
  output_pointer,
  expert_id,
  source_rank,
  target_rank,
  tile_coordinate,
  dependency_state
}
</code></pre></div></div>

<p>这样 Scheduler 不需要理解“这是 MoE 第几阶段的大 op”，它只需要判断：这个 task 的输入 ready 了吗？ready 就交给 Processor。</p>

<p>这个抽象的好处是动态性。MoE router 让每个 expert 的 token 数动态变化，传统固定 pipeline 很容易某些阶段等人；task abstraction 让运行时只关心 ready tasks，而不是强行按固定阶段推进。</p>

<h3 id="48-tile-dimensions为什么不是越大越好">4.8 Tile dimensions：为什么不是越大越好</h3>

<p>FlashMoE 论文里有一个很值得记的工程细节：它选择的 tile size 是 $(128,64)$。</p>

<p>这和我们讨论 Comet 时常说的 $128\times128$ 不矛盾。不同 kernel、不同 operator、不同寄存器压力下，最优 tile shape 会变。FlashMoE 给出的直觉是：</p>

<ul>
  <li>tile 太小：每个 task 的计算量太少，GPU 利用率低，调度开销相对变大。</li>
  <li>tile width 太大：每个线程需要保存更多中间值，寄存器使用量上升，可能触发 register spill。</li>
  <li>tile height 太大：如果 thread 数不变，每个 thread 负责的元素更多，单线程工作量增加，整体并行度下降。</li>
  <li>thread block 太大：同一个 SM 上能同时驻留的 blocks 变少，SM occupancy 降低。</li>
  <li>block 内线程越多，同步点 <code class="language-plaintext highlighter-rouge">__syncthreads()</code> 的成本也更高，因为更多线程要一起等到 barrier。</li>
</ul>

<p>这里的 register spill 是理解这段话的关键。GPU 上每个 thread 有很快的寄存器，但寄存器数量有限。如果一个 tile width 太大，线程需要同时保存更多 accumulator / pointer / metadata，中间变量装不下，就会 spill 到 local memory。local memory 名字听起来像本地，其实通常落到显存路径，慢很多。</p>

<p>SM occupancy 指一个 SM 上同时驻留多少 warps / thread blocks。occupancy 高不一定总是更快，但太低会让 GPU 难以隐藏访存延迟。假设 thread block 太大，一个 SM 只能驻留很少 blocks；某个 block 等内存或 barrier 时，SM 没有足够其他 warps 可切换，就容易空转。</p>

<p>所以 FlashMoE 选择 $(128,64)$ 不是因为它在数学上最自然，而是一个平衡：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>tile 足够大:
  每个 task 有足够计算量，可以摊薄调度成本。

tile 不太大:
  register pressure 不爆，occupancy 不掉太多，同步成本可控。
</code></pre></div></div>

<p>所以 tile size 不是“expert 能处理的最大 token 数”，而是 kernel/hardware co-design 的结果。这个点和 Comet 的 granularity mismatch 可以接上：tile 是 GPU 计算和调度的单位，不是 MoE 逻辑层的 token chunk。</p>

<h3 id="49-one-sided-communicationuvanvshmem-和-device-initiated-dma">4.9 One-sided communication：UVA、NVSHMEM 和 device-initiated DMA</h3>

<p>FlashMoE 和 Comet 都想摆脱纯 collective 的粗粒度等待，但 FlashMoE 更强调 <strong>one-sided, device-initiated communication</strong>。</p>

<p>传统 <code class="language-plaintext highlighter-rouge">AllToAll</code> 是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>所有 GPU 进入 collective
runtime 统一搬数据
所有参与方等待完成
</code></pre></div></div>

<p>FlashMoE 使用 NVSHMEM 建立跨 GPU 的 global address space，并在可用时利用 UVA 做 DMA / RDMA 风格的数据搬运。直觉上，GPU kernel 里的 Processor 或 communication task 可以直接把某个 tile 写到目标 GPU 的某个地址，而不是等 host 发起一个大 collective。</p>

<p>但这里有两个容易误解的点。</p>

<p>第一，one-sided 不等于没有通信成本。远端读写仍然有延迟和带宽限制，只是它不再要求所有 GPU 同步进入同一个 collective。</p>

<p>第二，one-sided 需要非常小心的内存布局。如果两个 GPU 同时往目标 GPU 的同一块 buffer 写，就会发生 write-write conflict。FlashMoE 的 symmetric tensor layout 就是为了解决这个问题。</p>

<h3 id="410-symmetric-tensor-layoutfigure-7-为什么符号很多">4.10 Symmetric Tensor Layout：Figure 7 为什么符号很多</h3>

<p>FlashMoE 的 Figure 7 讲的是 symmetric tensor layout。它是全篇一个不太容易一眼看懂、但非常关键的设计。</p>

<p><img src="https://arxiv.org/html/2506.04667/x11.png" alt="FlashMoE Figure 7(a): Symmetric tensor layout" /></p>

<p><img src="https://arxiv.org/html/2506.04667/x12.png" alt="FlashMoE Figure 7(b): DMA/RDMA state machine" /></p>

<p>论文里可以把这个 layout 粗略读成：</p>

\[L \sim [P, r, 2, s, n, C_{\mathrm{up}}, h]\]

<p>这里不是逐字符复刻论文公式，而是帮助理解每个维度的含义：</p>

<table>
  <thead>
    <tr>
      <th>符号</th>
      <th>含义</th>
      <th>为什么需要</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>$P$</td>
      <td>expert-parallel world size</td>
      <td>要区分来自哪个 peer GPU / 写向哪个 peer GPU</td>
    </tr>
    <tr>
      <td>$r$</td>
      <td>communication rounds</td>
      <td>MoE forward 里可能有 dispatch round、combine round，不同轮不能覆盖</td>
    </tr>
    <tr>
      <td>$2$</td>
      <td>incoming / outgoing 两个方向</td>
      <td>发送出去的数据和接收进来的数据要分开</td>
    </tr>
    <tr>
      <td>$s$</td>
      <td>staging buffers</td>
      <td>temporal buffering，让不同时间段的 writes 落到不同 slot</td>
    </tr>
    <tr>
      <td>$n$</td>
      <td>每张 GPU 上 local experts 数</td>
      <td>一个 GPU 上可能有多个 experts</td>
    </tr>
    <tr>
      <td>$C_{\mathrm{up}}$</td>
      <td>upscaled expert capacity</td>
      <td>给每个 expert 的 token capacity 预留空间</td>
    </tr>
    <tr>
      <td>$h$</td>
      <td>token hidden dimension</td>
      <td>每个 token vector 的 hidden size</td>
    </tr>
  </tbody>
</table>

<p>为什么要这么复杂？因为 FlashMoE 要支持 fully non-blocking one-sided writes。不同 GPU 的 task 会同时向 symmetric layout 中写 tile。如果只用一个普通 buffer，就很难避免冲突；要么加锁，要么同步，要么冒险覆盖。</p>

<p>FlashMoE 的思路是加 temporal dimensions：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>不同通信轮次用不同 slots
outgoing 和 incoming 分开
不同 staging buffer 分开
不同 peer/local expert/capacity slot 分开
</code></pre></div></div>

<p>这样一个 remote write 的目标地址由 source GPU、target GPU、round、direction、expert slot、capacity slot 等 metadata 唯一决定。论文还给了 theorem：这个布局是 write-write conflict-free。</p>

<p>如果觉得符号多，可以只抓住 Figure 7 的三个问题：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. 谁写？source peer rank / direction。
2. 写到哪个阶段？communication round / staging buffer。
3. 写给谁算？local expert id / capacity slot / hidden vector。
</code></pre></div></div>

<p>可以把 Figure 7(a) 想成一个多维停车场：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>第 1 层: 这是 dispatch 轮还是 combine 轮？
第 2 层: 这是 outgoing slot 还是 incoming slot？
第 3 层: 这是第几个 staging buffer？
第 4 层: 这是写给哪个 local expert？
第 5 层: 这是这个 expert capacity 里的第几个 token slot？
第 6 层: 这个 token 的 hidden vector。
</code></pre></div></div>

<p>远程 GPU 写入时，不是说“随便找一个地方写”，而是根据 routing metadata 算出唯一坐标：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>target_address =
  L[peer_rank, round, direction, stage, local_expert, capacity_slot, hidden_slice]
</code></pre></div></div>

<p>Figure 7(b) 的状态机则是在说明 DMA / RDMA 两种路径都遵循“写数据 + 发 signal + 接收端消费”的顺序。Subscriber 看到 signal 后，才会把对应 packet 解码成 task descriptor，再通知 Scheduler。这样可以避免 Processor 在数据还没写完整时就去读。</p>

<p>Figure 7 复杂，是因为 FlashMoE 要让很多 GPU 同时做 one-sided writes。如果布局少一个维度，就可能出现两个 writer 写到同一个地址，或者 dispatch 和 combine 的数据互相覆盖。它不是为了让 tensor 看起来复杂，而是为了让地址计算天然避免冲突。</p>

<p>这个设计和 Comet 的 shared tensor 思路有微妙区别：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Comet shared tensor:
  关注 producer 和 consumer 怎么共享一块中间 tensor，并按依赖切分。

FlashMoE symmetric tensor:
  关注跨 GPU one-sided writes 怎样在不加同步的情况下安全落位。
</code></pre></div></div>

<p>换句话说，Comet 主要解决“什么时候可以算”；FlashMoE 还要解决“数据可以安全写到哪里”。</p>

<h3 id="411-in-place-padding为什么-padding-也会浪费通信">4.11 In-place padding：为什么 padding 也会浪费通信</h3>

<p>MoE 里常见一个 capacity 概念：每个 expert 最多接收多少 token。为了让 buffer shape 固定，很多实现会把 expert input padding 到 capacity。</p>

<p>例如某个 expert capacity 是 128，但这次只收到 37 个真实 token。传统实现可能会构造一个 128 行的 buffer：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>37 real tokens + 91 null tokens
</code></pre></div></div>

<p>如果这些 null tokens 也被跨 GPU 传输，就浪费网络带宽。FlashMoE 的 payload efficiency 指的就是避免发送这种无意义 payload。</p>

<p>FlashMoE 的 in-place padding 思路是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>只把真实 token 通过网络发送过去
padding 在本地 symmetric tensor buffer 里完成
</code></pre></div></div>

<p>这样既满足后续 Processor 按固定 tile shape / aligned capacity 读取的需求，又避免把 null token 当作真实数据走网络。</p>

<p>这个点有点“不显眼但很实用”：MoE 的动态路由会导致每个 expert 的真实 token 数远小于或不等于 capacity，如果系统总是按 capacity 传输，那么稀疏计算节省下来的部分收益会被无效通信吃掉。</p>

<h3 id="412-flashmoe-的执行流程从一个-token-tile-的视角看">4.12 FlashMoE 的执行流程：从一个 token tile 的视角看</h3>

<p>可以用一个 tile 的生命周期来理解 FlashMoE。</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>1. Gate / routing 产生 token -&gt; expert 的映射。

2. 对某个目标 expert，Processor 生成 dispatch tile task。

3. 如果 expert 在远端 GPU：
   Processor 通过 NVSHMEM / UVA 把 tile 写入目标 GPU 的 symmetric tensor slot。
   同时发送 signal，通知目标 GPU 有 tile 到达。

4. 目标 GPU 的 Subscriber 收到 packet / signal，
   解码成 task descriptor。

5. Scheduler 把 ready task 分给空闲 Processor。

6. Processor 执行 expert FFN tile：
   Linear-1 -&gt; activation -&gt; Linear-2。

7. 输出 tile ready 后，生成 combine task。

8. combine task 根据 routing weight 做 weighted accumulation，
   如果原 token 在远端，再通过 one-sided write 写回。
</code></pre></div></div>

<p>这条链路的核心是：每个 tile 的 arrival、compute、combine 都可以独立推进，而不是被整个 MoE layer 的阶段边界卡住。</p>

<h3 id="413-和-comet-的关系不是替代而是更底层的-runtime-化">4.13 和 Comet 的关系：不是替代，而是更底层的 runtime 化</h3>

<p>如果只读过 Comet，很容易把 FlashMoE 理解成“另一个更激进的 Comet”。这个说法有一半对，一半不够准确。</p>

<p>相同点：</p>

<ul>
  <li>都认为 MoE 的瓶颈来自通信和计算的粒度不匹配。</li>
  <li>都把粒度降到 tile/task 级别。</li>
  <li>都不满足于传统 EP overlap 的 token chunk 级 pipeline。</li>
  <li>都需要把通信和 GEMM 资源隔离开，避免 remote I/O 破坏计算。</li>
</ul>

<p>不同点：</p>

<table>
  <thead>
    <tr>
      <th>维度</th>
      <th>Comet</th>
      <th>FlashMoE</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>核心抽象</td>
      <td>shared tensor dependency resolving</td>
      <td>GPU-resident actor/task runtime</td>
    </tr>
    <tr>
      <td>主要对象</td>
      <td>两条 producer-consumer pipeline</td>
      <td>整个 distributed MoE operator</td>
    </tr>
    <tr>
      <td>通信方式</td>
      <td>更细粒度 P2P / kernel 协调</td>
      <td>one-sided device-initiated DMA/RDMA</td>
    </tr>
    <tr>
      <td>调度位置</td>
      <td>fused kernel 内的 block assignment</td>
      <td>persistent kernel 内 Scheduler 分配 tasks</td>
    </tr>
    <tr>
      <td>内存布局重点</td>
      <td>shared tensor 的切分方向和执行顺序</td>
      <td>symmetric tensor layout 保证 non-blocking writes</td>
    </tr>
    <tr>
      <td>解决的额外问题</td>
      <td>granularity mismatch</td>
      <td>kernel launch、collective barrier、payload inefficiency</td>
    </tr>
  </tbody>
</table>

<p>所以我会这样理解：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Comet:
  把 MoE 的通信-计算 overlap 做到 tile-aware。

FlashMoE:
  把 MoE 的执行环境改成 tile-task runtime。
</code></pre></div></div>

<h3 id="414-读-flashmoe-时值得记录的不太显眼的细节">4.14 读 FlashMoE 时值得记录的“不太显眼”的细节</h3>

<p><strong>第一，single kernel 不等于把所有代码机械粘在一个 kernel 里。</strong> 真正关键是 persistent kernel 内部有 task abstraction、actor role 和 ready queue。否则只是把阶段写进一个大 kernel，仍然可能内部空转。</p>

<p><strong>第二，Subscriber / Scheduler 是成本，不是免费午餐。</strong> FlashMoE 特意只用最后一个 block 做 OS block，就是为了把管理成本压低。调度能力太弱会供不上 tasks；调度资源太多又会抢走计算 blocks。</p>

<p><strong>第三，tile size 是硬件约束下的折中。</strong> 论文选择 $(128,64)$，背后是 register pressure、shared memory、SM occupancy、block synchronization 的平衡，不是理论上越大越好。</p>

<p><strong>第四，symmetric tensor layout 是正确性设计，不只是性能优化。</strong> 如果没有这个布局，one-sided writes 可能需要同步或锁；一旦同步变多，FlashMoE 的核心优势就会被抵消。</p>

<p><strong>第五，payload efficiency 是 MoE 特有问题。</strong> 因为 router 动态分配 token，capacity padding 很常见。FlashMoE 的 in-place padding 让网络只传真实 token，不传 null token。</p>

<p><strong>第六，实验结果要看清边界。</strong> 论文的 evaluation 主要测单个 MoE layer forward，硬件是 8 张 H100；FlashMoE 用 FP32，而很多 baseline 用 FP16，论文认为这反而让 FlashMoE 更吃亏。但这也意味着，读结果时要把精度、forward-only、单层 MoE operator 和完整训练系统区分开。</p>

<h3 id="415-flashmoe-的定位">4.15 FlashMoE 的定位</h3>

<p>FlashMoE 可以理解成：</p>

<blockquote>
  <p>把 distributed MoE layer 从“host 发起的一串 kernels + collectives”改造成“GPU 内部常驻的 tile-task runtime”。</p>
</blockquote>

<p>它比 Comet 更系统级、更激进。Comet 解决的是 shared tensor 的依赖切分和 tile 重排；FlashMoE 解决的是整个 MoE operator 如何在 GPU 内部自行调度、通信、计算和写回。</p>

<p>这也是为什么 DeepSeek-V4 那种 expert-wave pipeline 读起来会更像工程折中：它吸收了 fine-grained overlap 的思想，但没有把论文重点放在一个完整 persistent MoE runtime 上。而 FlashMoE 的主张更明确：<strong>要突破 MoE 系统瓶颈，不能只看 GEMM，也不能只看通信，必须把 kernel launch、调度、远程写、内存布局和 padding 一起设计。</strong></p>

<hr />

<h2 id="5-deepseek-v4expert-wave-级别的-moe-pipeline"><strong>5. DeepSeek-V4：expert-wave 级别的 MoE pipeline</strong></h2>

<p>DeepSeek-V4 的 MoE 优化小节叫 <strong>Fine-Grained Communication-Computation Overlap in Expert Parallelism</strong>。它借鉴了前面这些工作，但落点更工程化：围绕 DeepSeek-V4 自己的 expert parallelism，把 MoE 层拆成可以流水的 expert waves。</p>

<h3 id="51-deepseek-的五段流程不是-minimax-那个五段">5.1 DeepSeek 的五段流程：不是 MiniMax 那个“五段”</h3>

<p>DeepSeek-V4 把 MoE 层执行写成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Dispatch All-to-All
-&gt; Linear-1 GEMM
-&gt; SwiGLU / FP8 Cast
-&gt; Linear-2 GEMM
-&gt; Combine All-to-All
</code></pre></div></div>

<p>这里的 <code class="language-plaintext highlighter-rouge">Linear-1</code> 和 <code class="language-plaintext highlighter-rouge">Linear-2</code> 是 expert FFN 的两个自然 GEMM。</p>

<p>如果 expert 是 SwiGLU 结构，那么 <code class="language-plaintext highlighter-rouge">Linear-1</code> 可以理解成同时做 gate/up 两个 projection：</p>

\[\begin{aligned}
g_e &amp;= X_e W_{e,\mathrm{gate}} \\
u_e &amp;= X_e W_{e,\mathrm{up}} \\
z_e &amp;= \mathrm{SiLU}(g_e)\odot u_e \\
Y_e &amp;= z_e W_{e,\mathrm{down}}
\end{aligned}\]

<p>所以 DeepSeek-V4 Figure 5 里的 <code class="language-plaintext highlighter-rouge">Linear-1</code> 对应前两行，<code class="language-plaintext highlighter-rouge">SwiGLU / FP8 Cast</code> 对应第三行和量化转换，<code class="language-plaintext highlighter-rouge">Linear-2</code> 对应最后一行。它拆的是 <strong>expert FFN 执行路径</strong>，不是 MiniMax-01 里 <code class="language-plaintext highlighter-rouge">a2a -&gt; allgather -&gt; compute -&gt; reduce-scatter -&gt; a2a</code> 那种并行通信路径。</p>

<h3 id="52-什么是-wave">5.2 什么是 wave</h3>

<p>一个 wave 可以理解成一小批 experts。例如某个 MoE 层有很多 experts，不必等所有 experts 的 dispatch 都完成后再一起计算，而是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>wave 0 的 token 到了 -&gt; 先算 wave 0 的 experts
wave 1 的 token 正在传 -&gt; 传完后接着算 wave 1
wave 0 算完后 -&gt; 立刻 combine wave 0 的输出
</code></pre></div></div>

<p>理想化时间线：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>t0: dispatch wave 0
t1: compute  wave 0 | dispatch wave 1
t2: combine  wave 0 | compute wave 1 | dispatch wave 2
t3: combine  wave 1 | compute wave 2 | dispatch wave 3
</code></pre></div></div>

<p>也就是说，DeepSeek-V4 的切分粒度是 <strong>expert wave</strong>。它既不是 MiniMax-01 的 token group，也不是 Comet 的 shared tensor tile，也不是 FlashMoE 的完整 persistent task runtime。</p>

<p>这里有一个很容易混淆的点：wave 的切分对象更接近 <strong>local experts</strong>，而不是 token 序列本身。router 会让不同 experts 收到不同数量的 tokens；某个 wave 只要它包含的 local experts 已经拿到足够的输入，就可以开始执行对应 expert GEMM。没有必要等本 rank 上所有 experts 都完成 dispatch。</p>

<p>因此它想解决的是 MoE 里典型的 long-tail 问题：某些 experts 的 token 较多或通信到达更慢，如果整层同步等待，其他已经 ready 的 experts 会空转；如果按 wave 推进，ready 的部分先算，没 ready 的部分继续通信。</p>

<h3 id="53-为什么它能隐藏通信">5.3 为什么它能隐藏通信</h3>

<p>如果计算足够重，通信就可以被计算盖住。DeepSeek-V4 给出的判断可以理解成：</p>

\[\frac{C}{B} \leq \frac{V_{\mathrm{comp}}}{V_{\mathrm{comm}}}\]

<p>其中：</p>

<ul>
  <li>$C$：设备峰值计算能力。</li>
  <li>$B$：通信带宽。</li>
  <li>$V_{\mathrm{comp}}$：MoE 层计算量。</li>
  <li>$V_{\mathrm{comm}}$：MoE 层通信量。</li>
</ul>

<p>如果右边足够大，说明每传一点数据都能对应很多计算，那么通信更容易被计算隐藏。DeepSeek-V4-Pro 里每个 token-expert pair 大约需要 $6hd$ FLOPs，通信量大约是 $3h$ bytes，所以比例大约是：</p>

\[\frac{V_{\mathrm{comp}}}{V_{\mathrm{comm}}} \approx 2d\]

<p>当 $d=3072$ 时，就是约 $6144$ FLOPs/Byte。这说明在合适硬件条件下，MoE expert 计算可以覆盖很大一部分通信。</p>

<p>更具体地说，DeepSeek-V4-Pro 的一个 token-expert pair 需要大约 $6hd$ FLOPs：gate projection、up projection、down projection 各贡献一部分。通信量大约是 $3h$ bytes：dispatch 输入和 combine 输出不完全同精度，论文按 FP8 dispatch 与 BF16 combine 做估算。这个比例说明：只要系统能把通信切细并和 GEMM 对齐，通信就有机会被 Tensor Core 计算遮住。</p>

<h3 id="54-公开代码megamoe2-对应什么">5.4 公开代码：MegaMoE2 对应什么</h3>

<p>论文说已经公开 CUDA-based mega-kernel <strong>MegaMoE2</strong>，对应 DeepGEMM PR #304。PR 描述里明确写到：它把 <code class="language-plaintext highlighter-rouge">Dispatch -&gt; Linear-1 -&gt; SwiGLU -&gt; Linear-2 -&gt; Combine</code> 融合进一个 mega-kernel，并重叠 NVLink communication 与 Tensor Core computation。</p>

<p>这个公开代码给我们几个很重要的判断依据。</p>

<p>第一，DeepSeek 这里不是只写了一个普通 PyTorch MoE。Hugging Face model repo 里的 <code class="language-plaintext highlighter-rouge">inference/model.py</code> 更像结构可读版本；真正的 MoE runtime 优化在 DeepGEMM 的 MegaMoE2 kernel 里。</p>

<p>第二，MegaMoE2 的 Python 测试把 fused path 和 legacy baseline 放在一起。legacy baseline 的逻辑可以概括成：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">recv_x</span> <span class="o">=</span> <span class="n">dispatch</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">topk_idx</span><span class="p">,</span> <span class="n">topk_weights</span><span class="p">)</span>
<span class="n">l1_y</span> <span class="o">=</span> <span class="n">grouped_gemm</span><span class="p">(</span><span class="n">recv_x</span><span class="p">,</span> <span class="n">l1_weights</span><span class="p">)</span>
<span class="n">l1_y</span> <span class="o">=</span> <span class="n">swiglu_and_cast</span><span class="p">(</span><span class="n">l1_y</span><span class="p">,</span> <span class="n">topk_weights</span><span class="p">)</span>
<span class="n">l2_y</span> <span class="o">=</span> <span class="n">grouped_gemm</span><span class="p">(</span><span class="n">l1_y</span><span class="p">,</span> <span class="n">l2_weights</span><span class="p">)</span>
<span class="n">y</span> <span class="o">=</span> <span class="n">combine</span><span class="p">(</span><span class="n">l2_y</span><span class="p">)</span>
</code></pre></div></div>

<p>而 fused path 变成：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nb">buffer</span> <span class="o">=</span> <span class="n">get_symmetric_buffer</span><span class="p">(...)</span>
<span class="n">w1</span><span class="p">,</span> <span class="n">w2</span> <span class="o">=</span> <span class="n">transform_weights_for_mega_moe</span><span class="p">(...)</span>
<span class="n">fp8_fp4_mega_moe</span><span class="p">(</span>
    <span class="n">output</span><span class="p">,</span>
    <span class="n">w1</span><span class="p">,</span>
    <span class="n">w2</span><span class="p">,</span>
    <span class="nb">buffer</span><span class="p">,</span>
    <span class="n">cumulative_local_expert_recv_stats</span><span class="o">=</span><span class="p">...,</span>
<span class="p">)</span>
</code></pre></div></div>

<p>这不是逐行引用源码，而是把公开测试里的结构抽象出来。重点是：MegaMoE2 不再让 host/runtime 一段一段驱动 dispatch、GEMM、activation、GEMM、combine，而是把这些阶段放进一个 fused kernel/runtime 里调度。</p>

<p>第三，Python API 层没有暴露 <code class="language-plaintext highlighter-rouge">num_waves</code> 这种模型超参数。wave 更像 kernel 内部的调度粒度：kernel 根据 <code class="language-plaintext highlighter-rouge">topk_idx</code>、<code class="language-plaintext highlighter-rouge">topk_weights</code>、local expert receive statistics、symmetric buffer 和变换后的 expert weights，决定哪些 experts 已经 ready，哪些结果可以 combine。</p>

<h3 id="55-结合-pr-316-看配置和收益">5.5 结合 PR #316 看配置和收益</h3>

<p>DeepGEMM Benchmark PR #316 给了和 DeepSeek-V4 对齐的测试配置：</p>

<table>
  <thead>
    <tr>
      <th>模型</th>
      <th style="text-align: right">Experts</th>
      <th style="text-align: right">Top-k</th>
      <th style="text-align: right">Hidden size $h$</th>
      <th style="text-align: right">Intermediate size $d$</th>
      <th style="text-align: right">EP</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>DeepSeek-V4-Flash</td>
      <td style="text-align: right">256</td>
      <td style="text-align: right">6</td>
      <td style="text-align: right">4096</td>
      <td style="text-align: right">2048</td>
      <td style="text-align: right">8</td>
    </tr>
    <tr>
      <td>DeepSeek-V4-Pro</td>
      <td style="text-align: right">384</td>
      <td style="text-align: right">6</td>
      <td style="text-align: right">7168</td>
      <td style="text-align: right">3072</td>
      <td style="text-align: right">8</td>
    </tr>
  </tbody>
</table>

<p>这里的 <code class="language-plaintext highlighter-rouge">top-k=6</code> 表示每个 token 会路由到 6 个 experts；EP8 表示 experts 分布在 8 个 expert-parallel ranks 上。benchmark 里 MegaMoE2 相对 legacy baseline 的加速大致在 1.50x 到 1.96x 之间，小 batch / latency-sensitive 场景收益尤其明显。</p>

<p>这也解释了为什么 DeepSeek-V4 论文特别提到 RL rollout 和 high-speed agent serving：这些场景 batch 未必很大，但 latency 很敏感；如果整层 MoE 等待所有 experts 同步，通信 bubble 会更明显。按 wave 推进之后，局部 ready 的 experts 可以先算，上一批结果可以先回传，pipeline 更容易保持饱和。</p>

<h3 id="56-论文里几个容易漏掉的系统细节">5.6 论文里几个容易漏掉的系统细节</h3>

<p><strong>第一，dispatch 是 pull-based。</strong> DeepSeek-V4 论文说 dispatch 阶段由每个 GPU 主动从远端 GPU 读取 activations，而不是让远端 GPU 细粒度 push 过来。原因是 fine-grained push 需要大量低延迟通知，当前硬件上 signaling overhead 很难忽略；pull-based 让当前 GPU 根据本地需要发起读取，更容易和本地 wave 调度配合。</p>

<p><strong>第二，带宽不是唯一瓶颈。</strong> 论文用 $\frac{C}{B}\leq \frac{V_{\mathrm{comp}}}{V_{\mathrm{comm}}}$ 说明只要计算通信比足够高，通信可以被计算遮住。此时继续只提高互联带宽不一定最优，反而要关注 compute、HBM、NVLink 和 power budget 是否能同时撑住。</p>

<p><strong>第三，SwiGLU / cast 也在流水线上。</strong> Figure 5 不是只画了两次 GEMM，中间还有 <code class="language-plaintext highlighter-rouge">SwiGLU / FP8 Cast</code>。如果这个 element-wise 阶段太慢，Linear-1 和 Linear-2 之间也会产生 bubble。论文因此提到，未来更低成本的 activation 可能有利于这种 fine-grained overlap。</p>

<h3 id="57-deepseek-的定位">5.7 DeepSeek 的定位</h3>

<p>DeepSeek-V4 不是在论文里提出一个通用 MoE runtime 系统，而是在自己的大模型系统里落地了一个适合 large-scale EP 的 expert-wave pipeline。</p>

<p>它的重点是：</p>

<ul>
  <li>把 expert computation 切成 waves。</li>
  <li>让 dispatch、expert GEMM、combine 同时推进。</li>
  <li>尽量让 all-to-all 通信隐藏在 GEMM 计算后面。</li>
  <li>结合 FP8/FP4、DeepGEMM 等底层优化服务高吞吐训练和推理。</li>
</ul>

<hr />

<h2 id="6-四篇论文的粒度对比"><strong>6. 四篇论文的粒度对比</strong></h2>

<p>最重要的是不要把所有“overlap”都理解成同一种切法。</p>

<table>
  <thead>
    <tr>
      <th>工作</th>
      <th>切分粒度</th>
      <th>主要对象</th>
      <th>解决的问题</th>
      <th>一句话</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>MiniMax-01</td>
      <td>token group / process group</td>
      <td>EP、ETP、EDP 通信</td>
      <td>大规模训练中的 a2a、allgather、reduce-scatter 开销</td>
      <td>group-level overlap</td>
    </tr>
    <tr>
      <td>Comet</td>
      <td>shared tensor / GEMM tile / thread block</td>
      <td>Dispatch-L1、L2-Combine 的 producer-consumer 链路</td>
      <td>通信粒度和 GEMM tile 粒度不匹配</td>
      <td>tile-level dependency-aware overlap</td>
    </tr>
    <tr>
      <td>FlashMoE</td>
      <td>tile task / actor / persistent kernel</td>
      <td>整个 distributed MoE operator</td>
      <td>collective barrier 和 kernel launch 开销</td>
      <td>GPU-resident task scheduling</td>
    </tr>
    <tr>
      <td>DeepSeek-V4</td>
      <td>expert wave</td>
      <td>Dispatch、Linear-1、Activation、Linear-2、Combine</td>
      <td>expert parallelism 下通信阻塞计算</td>
      <td>expert-wave pipeline</td>
    </tr>
  </tbody>
</table>

<p>再把它们放到一条“越来越细”的轴上：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>MiniMax-01
  token group / process group
    -&gt;
Comet
  shared tensor / GEMM tile / thread block
    -&gt;
DeepSeek-V4
  expert wave pipeline for production EP
    -&gt;
FlashMoE
  persistent kernel + GPU-resident task runtime
</code></pre></div></div>

<p>这条线不是严格的优劣排序，而是抽象层级不同。</p>

<hr />

<h2 id="7-一个统一的理解框架"><strong>7. 一个统一的理解框架</strong></h2>

<p>MoE 优化可以按三层看。</p>

<h3 id="71-并行策略层">7.1 并行策略层</h3>

<p>这一层决定 expert 放在哪些 GPU 上，token 怎么发过去。</p>

<p>典型问题：</p>

<ul>
  <li>expert 数量如何映射到 GPU？</li>
  <li>top-k 后 token 分布不均怎么办？</li>
  <li>EP、ETP、EDP 怎么组合？</li>
  <li>all-to-all 和 allgather/reduce-scatter 怎么重叠？</li>
</ul>

<p>MiniMax-01 主要在这一层。</p>

<h3 id="72-kernel--gemm-调度层">7.2 Kernel / GEMM 调度层</h3>

<p>这一层决定矩阵乘和通信怎么交错。</p>

<p>典型问题：</p>

<ul>
  <li>token 到了一部分，GEMM 能不能先算？</li>
  <li><code class="language-plaintext highlighter-rouge">Linear-2</code> 算出一部分，combine 能不能先发？</li>
  <li>哪些 thread block 做通信，哪些做计算？</li>
  <li>group GEMM 的 tile 顺序怎么排？</li>
</ul>

<p>Comet 和 DeepSeek-V4 主要在这一层，只是切分对象不同。</p>

<h3 id="73-runtime-层">7.3 Runtime 层</h3>

<p>这一层决定整个 MoE operator 是否还依赖外部 collective 和一串 kernel launch。</p>

<p>典型问题：</p>

<ul>
  <li>能不能一个 persistent kernel 管完整个 MoE？</li>
  <li>GPU 能不能自己发起远端读写？</li>
  <li>task ready 之后能不能马上调度？</li>
  <li>如何避免 barrier 和 straggler？</li>
</ul>

<p>FlashMoE 主要在这一层。</p>

<hr />]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[MoE 优化的探索：从 MiniMax-01、Comet、FlashMoE 到 DeepSeek-V4]]></summary></entry><entry><title type="html">DeepSeek V4笔记</title><link href="https://wqh011128.github.io/blog/2026/05/08/DeepSeekV4.html" rel="alternate" type="text/html" title="DeepSeek V4笔记" /><published>2026-05-08T00:00:00+00:00</published><updated>2026-05-08T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/05/08/DeepSeekV4</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/05/08/DeepSeekV4.html"><![CDATA[<h1 id="deepseek-v4-towards-highly-efficient-million-token-context-intelligence">DeepSeek V4: Towards Highly Efficient Million-Token Context Intelligence</h1>

<p><strong>Link:</strong></p>

<ul>
  <li>Technical Report: <a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/resolve/main/DeepSeek_V4.pdf">DeepSeek_V4.pdf</a></li>
  <li>Hugging Face Model Card: <a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro">deepseek-ai/DeepSeek-V4-Pro</a></li>
  <li>Official Inference Code: <a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/blob/main/inference/model.py">inference/model.py</a></li>
  <li>Transformers Code: <a href="https://github.com/huggingface/transformers/blob/main/src/transformers/models/deepseek_v4/modeling_deepseek_v4.py">modeling_deepseek_v4.py</a></li>
  <li>Transformers Docs: <a href="https://huggingface.co/docs/transformers/main/model_doc/deepseek_v4">DeepSeek-V4</a></li>
</ul>

<p>[toc]</p>

<hr />

<h2 id="main-contributions"><strong>Main Contributions</strong></h2>

<ul>
  <li><strong>1M Context:</strong> DeepSeek-V4 系列把上下文长度推到 1M tokens。重点不是单纯把 RoPE 拉长，而是围绕长上下文重新设计 attention、KV cache 和 sparse selection。</li>
  <li><strong>Hybrid Attention Architecture:</strong> 使用 <strong>Compressed Sparse Attention (CSA)</strong> 和 <strong>Heavily Compressed Attention (HCA)</strong> 混合结构，把近邻 token 精确保留，把远程 token 压缩后选择性读取。</li>
  <li><strong>Manifold-Constrained Hyper-Connections (mHC):</strong> 把传统 residual stream 扩成多条 hidden streams，并用受约束的 mixing 矩阵进行合并与传播。</li>
  <li><strong>MoE Routing Update:</strong> MoE router 支持 <code class="language-plaintext highlighter-rouge">sqrt softplus</code> 分数函数和前若干层的 hash routing。</li>
  <li><strong>MTP:</strong> 继续使用 Multi-Token Prediction 思路，为训练提供 next-n token 信号，也为后续推理辅助留下结构入口。</li>
  <li><strong>Training / Post-training:</strong> 官方模型卡中提到 32T+ tokens 预训练、Muon optimizer，以及 SFT + GRPO + on-policy distillation 的后训练路线。</li>
</ul>

<p>官方给出的模型规模如下：</p>

<table>
  <thead>
    <tr>
      <th>Model</th>
      <th style="text-align: right">Total Params</th>
      <th style="text-align: right">Activated Params</th>
      <th style="text-align: right">Context Length</th>
      <th>Precision</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>DeepSeek-V4-Flash</td>
      <td style="text-align: right">284B</td>
      <td style="text-align: right">13B</td>
      <td style="text-align: right">1M</td>
      <td>FP4 + FP8 mixed</td>
    </tr>
    <tr>
      <td>DeepSeek-V4-Pro</td>
      <td style="text-align: right">1.6T</td>
      <td style="text-align: right">49B</td>
      <td style="text-align: right">1M</td>
      <td>FP4 + FP8 mixed</td>
    </tr>
    <tr>
      <td>DeepSeek-V4-Flash-Base</td>
      <td style="text-align: right">284B</td>
      <td style="text-align: right">13B</td>
      <td style="text-align: right">1M</td>
      <td>FP8 mixed</td>
    </tr>
    <tr>
      <td>DeepSeek-V4-Pro-Base</td>
      <td style="text-align: right">1.6T</td>
      <td style="text-align: right">49B</td>
      <td style="text-align: right">1M</td>
      <td>FP8 mixed</td>
    </tr>
  </tbody>
</table>

<p><img src="/img/deepseek-v4-performance.png" alt="DeepSeek V4 performance" /></p>

<hr />

<h2 id="1-整体架构"><strong>1. 整体架构</strong></h2>

<p>我先把 DeepSeek-V4 看成三条线：</p>

<ol>
  <li><strong>Attention 线：</strong> 用 sliding window 处理最近 token，用 CSA / HCA 处理远程 token。</li>
  <li><strong>Residual 线：</strong> 用 mHC 替代简单 residual，让 hidden state 在多条 stream 之间混合。</li>
  <li><strong>MoE 线：</strong> 继续走 sparse MoE，但 router 分数和部分层路由方式有变化。</li>
</ol>

<p>一个比较接近源码的 block 级流程如下：</p>

<pre><code class="language-mermaid">flowchart TD
    A[input_ids] --&gt; B[Token Embedding]
    B --&gt; C[Expand to n_hc streams]
    C --&gt; D[mHC pre mix for attention]
    D --&gt; E[RMSNorm]
    E --&gt; F[Hybrid Attention: SWA + CSA/HCA]
    F --&gt; G[mHC post mix]
    G --&gt; H[mHC pre mix for MoE]
    H --&gt; I[RMSNorm]
    I --&gt; J[MoE FFN: hash/top-k router + experts]
    J --&gt; K[mHC post mix]
    K --&gt; L[Hyper Head + RMSNorm + LM Head]
</code></pre>

<p>在官方 inference 代码里，这个结构大概对应：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">Transformer</code>: embedding、主干 blocks、lm head、MTP blocks。</li>
  <li><code class="language-plaintext highlighter-rouge">Block</code>: attention + MoE，两处都包了一层 hyper-connection。</li>
  <li><code class="language-plaintext highlighter-rouge">Attention</code>: sliding window cache + optional compressor + optional indexer。</li>
  <li><code class="language-plaintext highlighter-rouge">MoE</code>: router、routed experts、shared expert。</li>
</ul>

<p>需要注意：Transformers 版本为了兼容 <code class="language-plaintext highlighter-rouge">generate</code>、cache 和 HF API，类名更细；官方 inference 版本更像“论文结构的最小可读实现”。这篇笔记会同时参考两者，但讲概念时优先贴近官方 inference 写法。</p>

<hr />

<h2 id="2-和-deepseek-v3--v32-的关系"><strong>2. 和 DeepSeek-V3 / V3.2 的关系</strong></h2>

<p>DeepSeek-V3 的两个核心词是 <strong>MLA</strong> 和 <strong>DeepSeekMoE</strong>。V4 没有丢掉 MoE 路线，但长上下文这件事变成了中心问题。</p>

<p>可以先这样理解：</p>

<table>
  <thead>
    <tr>
      <th>模型</th>
      <th>主要 attention / cache 思路</th>
      <th>MoE</th>
      <th>长上下文处理重点</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>DeepSeek-V2</td>
      <td>MLA 降低 KV cache</td>
      <td>DeepSeekMoE</td>
      <td>经济型推理</td>
    </tr>
    <tr>
      <td>DeepSeek-V3</td>
      <td>MLA + MTP + MoE 训练优化</td>
      <td>DeepSeekMoE</td>
      <td>训练效率和推理能力</td>
    </tr>
    <tr>
      <td>DeepSeek-V4</td>
      <td>CSA + HCA hybrid attention</td>
      <td>MoE + hash/top-k router</td>
      <td>1M context 下压缩和稀疏选择</td>
    </tr>
  </tbody>
</table>

<p>V4 的关键不是“把所有历史 token 都完整记下来”。如果 1M tokens 全量 KV cache 直接进 attention，显存和算力都不现实。V4 的路线是：</p>

<ul>
  <li>最近窗口：保留完整 KV，保证局部上下文精度。</li>
  <li>更远上下文：压缩成少量 compressed entries。</li>
  <li>需要读远程信息时：用 indexer 从 compressed entries 里选 top-k。</li>
</ul>

<p>所以 V4 的 attention 更像一个分层记忆系统：近处看原文，远处看摘要，必要时再从摘要索引里挑重点。</p>

<hr />

<h2 id="3-mtp-multi-token-prediction"><strong>3. MTP: Multi-Token Prediction</strong></h2>

<h3 id="31-mtp-是什么">3.1 MTP 是什么</h3>

<p>MTP = <strong>Multi-Token Prediction</strong>。</p>

<p>普通 causal LM 的训练目标是：</p>

\[p(x_{t+1}\mid x_{\le t})\]

<p>也就是每个位置只预测下一个 token。MTP 的思路是：在主 next-token loss 之外，再让模型预测更远的 token，例如：</p>

\[p(x_{t+2}\mid x_{\le t}),\quad p(x_{t+3}\mid x_{\le t}),\quad \ldots\]

<p>可以写成一个简化目标：</p>

\[\mathcal{L}_{\mathrm{MTP}}
= \frac{1}{D}\sum_{k=1}^{D}
\mathrm{CE}\big(p_k(x_{t+k}\mid x_{\le t}), x_{t+k}\big)\]

<p>其中 $D$ 是额外预测的步数。</p>

<p>直觉上，next-token prediction 只要求模型把眼前一步做好；MTP 会逼模型在 hidden state 里保留更多“接下来几步可能怎么走”的信息。DeepSeek-V3 技术报告里已经使用过 MTP，V4 继续保留了这个方向。</p>

<h3 id="32-deepseek-v4-里的-mtpblock">3.2 DeepSeek-V4 里的 MTPBlock</h3>

<p>官方 inference 代码里有 <code class="language-plaintext highlighter-rouge">n_mtp_layers</code> 和 <code class="language-plaintext highlighter-rouge">MTPBlock</code>。从结构上看，MTPBlock 会拿两类信息：</p>

<ul>
  <li>当前主干模型的 hidden state。</li>
  <li>shifted input token 的 embedding。</li>
</ul>

<p>先对齐源码接口。官方 inference 里的 <code class="language-plaintext highlighter-rouge">MTPBlock.forward</code> 是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><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">start_pos</span><span class="p">,</span> <span class="n">input_ids</span><span class="p">):</span>
    <span class="n">e</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">embed</span><span class="p">(</span><span class="n">input_ids</span><span class="p">)</span>
    <span class="n">e</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">enorm</span><span class="p">(</span><span class="n">e</span><span class="p">)</span>
    <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">hnorm</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="bp">self</span><span class="p">.</span><span class="n">e_proj</span><span class="p">(</span><span class="n">e</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="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">h_proj</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="nb">super</span><span class="p">().</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">start_pos</span><span class="p">,</span> <span class="n">input_ids</span><span class="p">)</span>
    <span class="n">logits</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">head</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">hc_head_fn</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">hc_head_scale</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">hc_head_base</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">norm</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">logits</span>
</code></pre></div></div>

<p>也就是说，源码里的 <code class="language-plaintext highlighter-rouge">MTPBlock</code> <strong>不会自己构造 shifted tokens</strong>。它只要求调用方传入一段 <code class="language-plaintext highlighter-rouge">input_ids</code>，然后把这段 token 的 embedding 和上一层 hidden streams 融合。真正的 shift 逻辑属于训练数据/调用逻辑：第 1 个 MTPBlock 传入 <code class="language-plaintext highlighter-rouge">x[t+1]</code>，目标是 <code class="language-plaintext highlighter-rouge">x[t+2]</code>；第 2 个 MTPBlock 再传入 <code class="language-plaintext highlighter-rouge">x[t+2]</code>，目标是 <code class="language-plaintext highlighter-rouge">x[t+3]</code>。</p>

<p>这里的 <strong>shifted input token</strong> 指的是“向后错位后的真实 token”。假设原始序列是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>x0, x1, x2, x3, x4, x5
</code></pre></div></div>

<p>普通 next-token prediction 在位置 <code class="language-plaintext highlighter-rouge">t=2</code> 使用 <code class="language-plaintext highlighter-rouge">x0, x1, x2</code> 的上下文预测 <code class="language-plaintext highlighter-rouge">x3</code>。</p>

<p>MTP 想继续预测更远的位置。为了预测 <code class="language-plaintext highlighter-rouge">x4</code>，MTPBlock 会额外看到一个 shifted token，也就是 <code class="language-plaintext highlighter-rouge">x3</code> 的 embedding。这里不是把 <code class="language-plaintext highlighter-rouge">x3</code> 当成模型生成结果，而是在训练时把真实序列右移后喂给 MTP 模块，帮助它构造“预测下一步的下一步”的中间状态。</p>

<p>可以把它想成：</p>

<table>
  <thead>
    <tr>
      <th>位置</th>
      <th>主模型看到的上下文</th>
      <th>主 LM head 目标</th>
      <th>第 1 个 MTPBlock 额外输入</th>
      <th>第 1 个 MTPBlock 目标</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">t=0</code></td>
      <td><code class="language-plaintext highlighter-rouge">x0</code></td>
      <td><code class="language-plaintext highlighter-rouge">x1</code></td>
      <td><code class="language-plaintext highlighter-rouge">x1</code></td>
      <td><code class="language-plaintext highlighter-rouge">x2</code></td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">t=1</code></td>
      <td><code class="language-plaintext highlighter-rouge">x0, x1</code></td>
      <td><code class="language-plaintext highlighter-rouge">x2</code></td>
      <td><code class="language-plaintext highlighter-rouge">x2</code></td>
      <td><code class="language-plaintext highlighter-rouge">x3</code></td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">t=2</code></td>
      <td><code class="language-plaintext highlighter-rouge">x0, x1, x2</code></td>
      <td><code class="language-plaintext highlighter-rouge">x3</code></td>
      <td><code class="language-plaintext highlighter-rouge">x3</code></td>
      <td><code class="language-plaintext highlighter-rouge">x4</code></td>
    </tr>
  </tbody>
</table>

<p>所以 <code class="language-plaintext highlighter-rouge">shifted_input_ids</code> 不是一个新 token，也不是预测 token，而是训练数据中已经存在的 token，只是相对当前位置向后平移了一格。如果有多个 MTP 层，第 2 个 MTPBlock 会继续用更后面的 shifted token，去预测更远的目标。</p>

<p>然后 MTPBlock 会把 shifted token embedding 和主 hidden state 分别投影、相加，再过一个 block，最后接同一个 head 输出 logits。</p>

<p>用等价伪代码表示：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">tokens</span> <span class="o">=</span> <span class="p">[</span><span class="n">x0</span><span class="p">,</span> <span class="n">x1</span><span class="p">,</span> <span class="n">x2</span><span class="p">,</span> <span class="n">x3</span><span class="p">,</span> <span class="n">x4</span><span class="p">,</span> <span class="n">x5</span><span class="p">]</span>
<span class="n">T</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="n">tokens</span><span class="p">)</span>

<span class="c1"># 主模型：位置 t 的 hidden 用来预测 x[t+1]
# 有效位置：t = 0, 1, 2, 3, 4
</span><span class="n">main_hidden</span> <span class="o">=</span> <span class="n">backbone</span><span class="p">(</span><span class="n">tokens</span><span class="p">[:</span><span class="o">-</span><span class="mi">1</span><span class="p">])</span>
<span class="n">main_logits</span> <span class="o">=</span> <span class="n">lm_head</span><span class="p">(</span><span class="n">main_hidden</span><span class="p">)</span>
<span class="n">main_loss</span> <span class="o">=</span> <span class="mi">0</span>
<span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">T</span> <span class="o">-</span> <span class="mi">1</span><span class="p">):</span>
    <span class="n">main_loss</span> <span class="o">+=</span> <span class="n">CE</span><span class="p">(</span><span class="n">main_logits</span><span class="p">[</span><span class="n">t</span><span class="p">],</span> <span class="n">tokens</span><span class="p">[</span><span class="n">t</span> <span class="o">+</span> <span class="mi">1</span><span class="p">])</span>

<span class="c1"># 第 1 个 MTPBlock：额外输入 tokens[t+1]，目标变成 tokens[t+2]
# 有效位置：t = 0, 1, 2, 3
</span><span class="n">shifted_1</span> <span class="o">=</span> <span class="n">tokens</span><span class="p">[</span><span class="mi">1</span><span class="p">:]</span>               <span class="c1"># x1, x2, x3, x4, x5
</span><span class="n">target_1</span> <span class="o">=</span> <span class="n">tokens</span><span class="p">[</span><span class="mi">2</span><span class="p">:]</span>                <span class="c1"># x2, x3, x4, x5
</span><span class="n">mtp_hidden_1</span> <span class="o">=</span> <span class="n">mtp_block_1</span><span class="p">(</span><span class="n">main_hidden</span><span class="p">[:</span><span class="o">-</span><span class="mi">1</span><span class="p">],</span> <span class="n">shifted_1</span><span class="p">[:</span><span class="o">-</span><span class="mi">1</span><span class="p">])</span>
<span class="n">mtp_logits_1</span> <span class="o">=</span> <span class="n">lm_head</span><span class="p">(</span><span class="n">mtp_hidden_1</span><span class="p">)</span>
<span class="n">mtp_loss_1</span> <span class="o">=</span> <span class="mi">0</span>
<span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="nb">range</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="n">mtp_loss_1</span> <span class="o">+=</span> <span class="n">CE</span><span class="p">(</span><span class="n">mtp_logits_1</span><span class="p">[</span><span class="n">t</span><span class="p">],</span> <span class="n">target_1</span><span class="p">[</span><span class="n">t</span><span class="p">])</span>

<span class="c1"># 如果有第 2 个 MTPBlock：继续往后错一格，目标变成 tokens[t+3]
# 有效位置：t = 0, 1, 2
</span><span class="n">shifted_2</span> <span class="o">=</span> <span class="n">tokens</span><span class="p">[</span><span class="mi">2</span><span class="p">:]</span>               <span class="c1"># x2, x3, x4, x5
</span><span class="n">target_2</span> <span class="o">=</span> <span class="n">tokens</span><span class="p">[</span><span class="mi">3</span><span class="p">:]</span>                <span class="c1"># x3, x4, x5
</span><span class="n">mtp_hidden_2</span> <span class="o">=</span> <span class="n">mtp_block_2</span><span class="p">(</span><span class="n">mtp_hidden_1</span><span class="p">[:</span><span class="o">-</span><span class="mi">1</span><span class="p">],</span> <span class="n">shifted_2</span><span class="p">[:</span><span class="o">-</span><span class="mi">1</span><span class="p">])</span>
<span class="n">mtp_logits_2</span> <span class="o">=</span> <span class="n">lm_head</span><span class="p">(</span><span class="n">mtp_hidden_2</span><span class="p">)</span>
<span class="n">mtp_loss_2</span> <span class="o">=</span> <span class="mi">0</span>
<span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">T</span> <span class="o">-</span> <span class="mi">3</span><span class="p">):</span>
    <span class="n">mtp_loss_2</span> <span class="o">+=</span> <span class="n">CE</span><span class="p">(</span><span class="n">mtp_logits_2</span><span class="p">[</span><span class="n">t</span><span class="p">],</span> <span class="n">target_2</span><span class="p">[</span><span class="n">t</span><span class="p">])</span>
</code></pre></div></div>

<p>再把一个 <code class="language-plaintext highlighter-rouge">MTPBlock</code> 本身展开看，它内部大概是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">MTPBlock</span><span class="p">:</span>
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="n">prev_hidden_streams</span><span class="p">,</span> <span class="n">shifted_input_ids</span><span class="p">):</span>
        <span class="n">shifted_emb</span> <span class="o">=</span> <span class="n">embed</span><span class="p">(</span><span class="n">shifted_input_ids</span><span class="p">)</span>
        <span class="n">token_part</span> <span class="o">=</span> <span class="n">e_proj</span><span class="p">(</span><span class="n">norm_embed</span><span class="p">(</span><span class="n">shifted_emb</span><span class="p">))</span>
        <span class="n">hidden_part</span> <span class="o">=</span> <span class="n">h_proj</span><span class="p">(</span><span class="n">norm_hidden</span><span class="p">(</span><span class="n">prev_hidden_streams</span><span class="p">))</span>

        <span class="n">mtp_hidden</span> <span class="o">=</span> <span class="n">token_part</span><span class="p">[:,</span> <span class="p">:,</span> <span class="bp">None</span><span class="p">,</span> <span class="p">:]</span> <span class="o">+</span> <span class="n">hidden_part</span>
        <span class="n">mtp_hidden</span> <span class="o">=</span> <span class="n">transformer_block</span><span class="p">(</span><span class="n">mtp_hidden</span><span class="p">,</span> <span class="n">shifted_input_ids</span><span class="p">)</span>

        <span class="n">logits</span> <span class="o">=</span> <span class="n">shared_lm_head</span><span class="p">(</span><span class="n">hyper_head</span><span class="p">(</span><span class="n">mtp_hidden</span><span class="p">))</span>
        <span class="k">return</span> <span class="n">logits</span>
</code></pre></div></div>

<p>这里有一个容易混淆的点：MTP 不是把模型一次 forward 变成“直接吐多个 token”。在公开 inference 代码里，主 <code class="language-plaintext highlighter-rouge">Transformer.forward</code> 默认仍然返回 next-token logits；<code class="language-plaintext highlighter-rouge">MTPBlock</code> 更像是额外结构，让 hidden state 多承担几个未来 token 的预测任务。</p>

<hr />

<h2 id="4-moe-routersoftplus-和-hash-route"><strong>4. MoE Router、Softplus 和 Hash Route</strong></h2>

<h3 id="41-softplus-是什么">4.1 Softplus 是什么</h3>

<p>Softplus 的定义是：</p>

\[\mathrm{softplus}(x)=\log(1+e^x)\]

<p>它可以看成 ReLU 的平滑版本：</p>

\[\mathrm{ReLU}(x)=\max(0,x)\]

<p>区别是：</p>

<ul>
  <li>ReLU 在 $x&lt;0$ 时直接变成 0。</li>
  <li>Softplus 始终大于 0，并且在 0 附近是平滑的。</li>
</ul>

<p>这对 router 分数有用，因为 MoE router 最后要得到非负权重。如果使用 softplus，负 logits 不会被硬切成 0，而是保留一个很小的正值。</p>

<h3 id="42-hash-route-是什么">4.2 Hash Route 是什么</h3>

<p>普通 top-k router 是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">indices</span> <span class="o">=</span> <span class="n">topk</span><span class="p">(</span><span class="n">scores</span><span class="p">,</span> <span class="n">k</span><span class="p">)</span>
</code></pre></div></div>

<p>Hash route 则是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">indices</span> <span class="o">=</span> <span class="n">tid2eid</span><span class="p">[</span><span class="n">input_ids</span><span class="p">]</span>
</code></pre></div></div>

<p>也就是说，某些层的 expert 选择不再由当前 hidden state 动态决定，而是由 token id 查表得到。</p>

<p>这像是把“这个 token 去哪些专家”提前写进一个表里。对应到源码，hash layer 会有一个 <code class="language-plaintext highlighter-rouge">tid2eid</code> 参数，它的形状大致是：</p>

\[[\mathrm{vocab\_size},\ \mathrm{num\_experts\_per\_tok}]\]

<h3 id="43-为什么要-hash-route">4.3 为什么要 Hash Route</h3>

<p>这部分论文和源码能确认的是：DeepSeek-V4 确实支持前若干层使用 hash routing。至于动机，我的理解要稍微保守一点：</p>

<ul>
  <li>早期层的 token 表示还比较接近 token identity，用 token id 路由有一定合理性。</li>
  <li>hash routing 可以减少动态 top-k 路由带来的计算和通信波动。</li>
  <li>固定路由可能让早期 expert 分配更稳定，降低某些 batch 下的 routing collapse 风险。</li>
  <li>但它也会牺牲一部分上下文自适应能力，所以不适合所有层都这么做。</li>
</ul>

<p>因此我会把 Hash Route 理解成：<strong>早期层用更便宜、更稳定的 token-id 路由，后续层再交给 hidden-state-based top-k router。</strong></p>

<h3 id="44-hugging-face-router-代码">4.4 Hugging Face Router 代码</h3>

<p>Hugging Face / 官方 inference 的 router 可以按这两步读：先算 routing scores，再决定 selected experts。Hash route 改的是 selected experts 的来源，不是把 expert weights 全部变成常数。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">logits</span> <span class="o">=</span> <span class="n">linear</span><span class="p">(</span><span class="n">hidden_states</span><span class="p">,</span> <span class="n">router_weight</span><span class="p">)</span>
<span class="n">scores</span> <span class="o">=</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">softplus</span><span class="p">(</span><span class="n">logits</span><span class="p">))</span>

<span class="k">if</span> <span class="n">hash_routing</span><span class="p">:</span>
    <span class="n">selected_experts</span> <span class="o">=</span> <span class="n">tid2eid</span><span class="p">[</span><span class="n">input_ids</span><span class="p">]</span>
<span class="k">else</span><span class="p">:</span>
    <span class="n">selected_experts</span> <span class="o">=</span> <span class="n">topk</span><span class="p">(</span><span class="n">scores</span> <span class="o">+</span> <span class="n">correction_bias</span><span class="p">,</span> <span class="n">k</span><span class="p">)</span>

<span class="n">router_weights</span> <span class="o">=</span> <span class="n">gather</span><span class="p">(</span><span class="n">scores</span><span class="p">,</span> <span class="n">selected_experts</span><span class="p">)</span>
<span class="n">router_weights</span> <span class="o">=</span> <span class="n">router_weights</span> <span class="o">/</span> <span class="n">router_weights</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">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
</code></pre></div></div>

<p>这里 <code class="language-plaintext highlighter-rouge">correction_bias</code> 可以影响 top-k 选择，但真正的 <code class="language-plaintext highlighter-rouge">router_weights</code> 仍然来自未加 bias 的 <code class="language-plaintext highlighter-rouge">scores</code>。也就是说，selection 和 weighting 可以使用略不同的信号。</p>

<hr />

<h2 id="5-mhc-manifold-constrained-hyper-connections"><strong>5. mHC: Manifold-Constrained Hyper-Connections</strong></h2>

<p>论文第 2.2 节的写法是：先介绍 <strong>Standard Hyper-Connections (HC)</strong>，再说明 <strong>mHC</strong> 相比 HC 多了哪些稳定性约束。这里我也按论文顺序记。</p>

<h3 id="51-standard-hyper-connections">5.1 Standard Hyper-Connections</h3>

<p>标准 HC 的第一步是把 residual stream 的“宽度”扩成 $n_{\mathrm{hc}}$ 条。论文里说的是：</p>

\[\mathbb{R}^{d}\rightarrow\mathbb{R}^{n_{\mathrm{hc}}\times d}\]

<p>其中 $d$ 是真实喂给 Transformer layer 的 hidden size，$n_{\mathrm{hc}}$ 是 Hyper-Connection 的 stream 数量。Hugging Face 配置里对应字段叫 <code class="language-plaintext highlighter-rouge">hc_mult</code>，DeepSeek-V4 常见设置是：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>hc_mult = 4
</code></pre></div></div>

<p>对第 $l$ 层之前的 residual state，论文定义：</p>

\[\begin{aligned}
X_l
&amp;=
\left[
x_{l,1};
\ldots;
x_{l,n_{\mathrm{hc}}}
\right]^T
&amp;\in
\mathbb{R}^{n_{\mathrm{hc}}\times d}
\end{aligned}\]

<p>这里每个 $x_{l,i}\in\mathbb{R}^{d}$。所以它不是把 $d$ 切成 $d/n_{\mathrm{hc}}$，而是维护 $n_{\mathrm{hc}}$ 条完整的 $d$ 维 residual streams。</p>

<p>HC 引入三个线性映射：</p>

\[A_l\in\mathbb{R}^{1\times n_{\mathrm{hc}}}\]

\[B_l\in\mathbb{R}^{n_{\mathrm{hc}}\times n_{\mathrm{hc}}}\]

\[C_l\in\mathbb{R}^{n_{\mathrm{hc}}\times 1}\]

<p>更新公式是论文的式 (1)：</p>

\[\begin{aligned}
X_{l+1}
&amp;=
B_lX_l+C_lF_l(A_lX_l)
\end{aligned}\]

<p>逐项看维度：</p>

<ul>
  <li>$A_lX_l\in\mathbb{R}^{1\times d}$：把 $n_{\mathrm{hc}}$ 条 stream collapse 成一条，作为 layer input。</li>
  <li>$F_l(A_lX_l)\in\mathbb{R}^{1\times d}$：第 $l$ 个实际 layer 的输出，仍然是 $d$ 维。</li>
  <li>$C_lF_l(A_lX_l)\in\mathbb{R}^{n_{\mathrm{hc}}\times d}$：把 layer 输出写回 $n_{\mathrm{hc}}$ 条 stream。</li>
  <li>$B_lX_l\in\mathbb{R}^{n_{\mathrm{hc}}\times d}$：旧 residual streams 之间的线性混合。</li>
</ul>

<p>所以 HC 的关键点是：<strong>inner layer 仍然只处理 $d$ 维输入输出，额外扩展的是 residual stream 的宽度。</strong> 这给模型增加了一条不同于 hidden size 的扩展轴，但因为 $n_{\mathrm{hc}}$ 通常很小，所以开销相对可控。</p>

<h3 id="52-mhc-相比-hc-改了什么">5.2 mHC 相比 HC 改了什么</h3>

<p>论文说标准 HC 在多层堆叠时可能出现 numerical instability。mHC 的核心改动是：<strong>不再让 $B_l$ 成为任意矩阵，而是把 $B_l$ 约束到 doubly stochastic matrices 的流形上。</strong></p>

<p>论文定义的流形是 Birkhoff polytope：</p>

\[\begin{aligned}
\mathcal{M}
&amp;=
\Bigl\{
M\in\mathbb{R}^{n_{\mathrm{hc}}\times n_{\mathrm{hc}}}
\ \Big|
M\mathbf{1}_{n_{\mathrm{hc}}}=\mathbf{1}_{n_{\mathrm{hc}}},
\ \mathbf{1}_{n_{\mathrm{hc}}}^{\top}M=\mathbf{1}_{n_{\mathrm{hc}}}^{\top},
\ M\ge 0
\Bigr\}
\end{aligned}\]

<p>然后约束：</p>

\[B_l\in\mathcal{M}\]

<p>这里的含义很具体：</p>

<ul>
  <li>$M\mathbf{1}=\mathbf{1}$：每一行和为 1。</li>
  <li>$\mathbf{1}^{T}M=\mathbf{1}^{T}$：每一列和为 1。</li>
  <li>$M\ge 0$：所有元素非负。</li>
</ul>

<p>论文给出的理由是：这个约束会让 $\lVert B_l\rVert_2$ 被 1 bound 住，因此 residual transformation 是 non-expansive 的。换句话说，$B_l$ 混合多条 stream 时不会无限放大信号。并且 $\mathcal{M}$ 对矩阵乘法封闭，所以多层堆叠时仍然能保持稳定性。</p>

<p>此外，论文还约束 $A_l$ 和 $C_l$ 非负且有界，避免不同 stream 之间因为符号相反而发生 signal cancellation。</p>

<h3 id="53-dynamic-parameterization">5.3 Dynamic Parameterization</h3>

<p>mHC 里的 $A_l,B_l,C_l$ 不是固定参数，而是动态生成的。论文把它们分成：</p>

<ul>
  <li>dynamic component：由当前输入 $X_l$ 生成。</li>
  <li>static component：输入无关的可学习 bias。</li>
</ul>

<p>给定：</p>

\[X_l\in\mathbb{R}^{n_{\mathrm{hc}}\times d}\]

<p>先 flatten 并 RMSNorm：</p>

\[\begin{aligned}
\hat{X}_l
&amp;=
\mathrm{RMSNorm}(\mathrm{vec}(X_l)) \\
&amp;\in
\mathbb{R}^{1\times n_{\mathrm{hc}}d}
\end{aligned}\]

<p>然后生成 unconstrained raw parameters：</p>

\[\begin{aligned}
\tilde{A}_l
&amp;=
\alpha_l^{\mathrm{pre}}
\cdot
(\hat{X}_lW_l^{\mathrm{pre}})
+
S_l^{\mathrm{pre}}
\end{aligned}\]

\[\begin{aligned}
\tilde{B}_l
&amp;=
\alpha_l^{\mathrm{res}}
\cdot
\mathrm{Mat}(\hat{X}_lW_l^{\mathrm{res}})
+
S_l^{\mathrm{res}}
\end{aligned}\]

\[\begin{aligned}
\tilde{C}_l
&amp;=
\alpha_l^{\mathrm{post}}
\cdot
(\hat{X}_lW_l^{\mathrm{post}})^T
+
S_l^{\mathrm{post}}
\end{aligned}\]

<p>对应维度是：</p>

\[\begin{aligned}
W_l^{\mathrm{pre}},W_l^{\mathrm{post}}
&amp;\in
\mathbb{R}^{n_{\mathrm{hc}}d\times n_{\mathrm{hc}}}
\end{aligned}\]

\[\begin{aligned}
W_l^{\mathrm{res}}
&amp;\in
\mathbb{R}^{n_{\mathrm{hc}}d\times n_{\mathrm{hc}}^2}
\end{aligned}\]

\[S_l^{\mathrm{pre}}\in\mathbb{R}^{1\times n_{\mathrm{hc}}},
\quad
S_l^{\mathrm{res}}\in\mathbb{R}^{n_{\mathrm{hc}}\times n_{\mathrm{hc}}},
\quad
S_l^{\mathrm{post}}\in\mathbb{R}^{n_{\mathrm{hc}}\times 1}\]

<p>其中 $\mathrm{Mat}(\cdot)$ 把 $1\times n_{\mathrm{hc}}^2$ 的向量 reshape 成 $n_{\mathrm{hc}}\times n_{\mathrm{hc}}$ 矩阵。$\alpha_l^{\mathrm{pre}},\alpha_l^{\mathrm{res}},\alpha_l^{\mathrm{post}}\in\mathbb{R}$ 是可学习 gating factors，论文说它们初始化为较小值。这样做的直觉是：训练初期更接近 static component，等训练稳定后再逐步放大 input-dependent dynamic component。</p>

<h3 id="54-applying-parameter-constraints">5.4 Applying Parameter Constraints</h3>

<p>得到 raw parameters 之后，mHC 才施加约束。</p>

<p>对 input mapping 和 output mapping：</p>

\[A_l=\sigma(\tilde{A}_l)\]

\[C_l=2\sigma(\tilde{C}_l)\]

<p>这里 $A_l$ 的范围是 $(0,1)$，$C_l$ 的范围是 $(0,2)$。它们都是非负且有界的。</p>

<p>对 residual mapping，先保证正数：</p>

\[M^{(0)}=\exp(\tilde{B}_l)\]

<p>然后用 Sinkhorn-Knopp 反复做列归一化和行归一化：</p>

\[\begin{aligned}
M^{(t)}
&amp;=
T_r\left(T_c\left(M^{(t-1)}\right)\right)
\end{aligned}\]

<p>其中 $T_c$ 表示 column normalization，$T_r$ 表示 row normalization。最终：</p>

\[B_l=M^{(t_{\max})}\]

<p>论文中取：</p>

\[t_{\max}=20\]

<p>这回答了之前那个疑问：20 次 Sinkhorn 发生在 $n_{\mathrm{hc}}\times n_{\mathrm{hc}}$ 的小矩阵上。DeepSeek-V4 常见 $n_{\mathrm{hc}}=4$，所以每个 token 的 $B_l$ 是 $4\times4$，不是 $d\times d$，也不是 $(n_{\mathrm{hc}}d)\times(n_{\mathrm{hc}}d)$。</p>

<h3 id="55-和-hugging-face-变量的对应">5.5 和 Hugging Face 变量的对应</h3>

<p>HF Transformers 里的 <code class="language-plaintext highlighter-rouge">DeepseekV4DecoderLayer</code> 有两个 hyper-connection 模块：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">attn_hc</code>：包住 self-attention。</li>
  <li><code class="language-plaintext highlighter-rouge">ffn_hc</code>：包住 MoE / FFN。</li>
</ul>

<p>调用形式可以概括为：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">post</span><span class="p">,</span> <span class="n">comb</span><span class="p">,</span> <span class="n">collapsed</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">attn_hc</span><span class="p">(</span><span class="n">hidden_states</span><span class="p">)</span>
</code></pre></div></div>

<p>对应论文符号：</p>

<table>
  <thead>
    <tr>
      <th>HF 变量</th>
      <th>论文符号</th>
      <th>维度</th>
      <th>作用</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">hidden_states</code></td>
      <td>$X_l$</td>
      <td>$[B_{\mathrm{sz}},S,n_{\mathrm{hc}},d]$</td>
      <td>多条 residual streams</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">pre</code></td>
      <td>$A_l$</td>
      <td>$[B_{\mathrm{sz}},S,1,n_{\mathrm{hc}}]$ 或 broadcast 等价形态</td>
      <td>input mapping</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">collapsed</code></td>
      <td>$A_lX_l$</td>
      <td>$[B_{\mathrm{sz}},S,d]$</td>
      <td>真正喂给 attention / FFN 的输入</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">comb</code></td>
      <td>$B_l$</td>
      <td>$[B_{\mathrm{sz}},S,n_{\mathrm{hc}},n_{\mathrm{hc}}]$</td>
      <td>constrained residual mapping</td>
    </tr>
    <tr>
      <td><code class="language-plaintext highlighter-rouge">post</code></td>
      <td>$C_l$</td>
      <td>$[B_{\mathrm{sz}},S,n_{\mathrm{hc}},1]$ 或 broadcast 等价形态</td>
      <td>output mapping</td>
    </tr>
  </tbody>
</table>

<p><code class="language-plaintext highlighter-rouge">pre</code> 没有返回，是因为它已经在 <code class="language-plaintext highlighter-rouge">collapsed = A_lX_l</code> 这一步用掉了。decoder layer 拿到 <code class="language-plaintext highlighter-rouge">collapsed</code> 后调用 attention / FFN，再把输出写回：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">post</span><span class="p">,</span> <span class="n">comb</span><span class="p">,</span> <span class="n">collapsed</span> <span class="o">=</span> <span class="n">attn_hc</span><span class="p">(</span><span class="n">hidden_states</span><span class="p">)</span>
<span class="n">attn_output</span> <span class="o">=</span> <span class="n">self_attn</span><span class="p">(</span><span class="n">norm</span><span class="p">(</span><span class="n">collapsed</span><span class="p">))</span>
<span class="n">hidden_states</span> <span class="o">=</span> <span class="n">post</span><span class="p">[...,</span> <span class="bp">None</span><span class="p">]</span> <span class="o">*</span> <span class="n">attn_output</span><span class="p">[...,</span> <span class="bp">None</span><span class="p">,</span> <span class="p">:]</span> <span class="o">+</span> <span class="n">comb</span> <span class="o">@</span> <span class="n">hidden_states</span>
</code></pre></div></div>

<p>这就是论文公式：</p>

\[\begin{aligned}
X_{l+1}
&amp;=
B_lX_l+C_lF_l(A_lX_l)
\end{aligned}\]

<p>这里为了代码 broadcast，<code class="language-plaintext highlighter-rouge">post[..., None] * attn_output[..., None, :]</code> 写法看起来和矩阵公式不同，但含义就是 $C_lF_l(A_lX_l)$。</p>

<hr />

<h2 id="6-csa-compressed-sparse-attention"><strong>6. CSA: Compressed Sparse Attention</strong></h2>

<p>论文第 2.3 节先给 Figure 3，再按 <strong>Compressed Key-Value Entries → Lightning Indexer → Shared Key-Value Multi-Query Attention → Grouped Output Projection → Attention Sink</strong> 的顺序拆 CSA。这一节也按这个顺序记。</p>

<p><img src="/img/deepseek-v4-csa.png" alt="DeepSeek-V4 CSA architecture" /></p>

<p>Figure 3 里有三条流：</p>

<ol>
  <li><strong>Sliding Window KV Entries:</strong> 最近一小段 token 不压缩，直接参加 attention，用来保留局部细粒度信息。</li>
  <li><strong>Compressed KV Entries:</strong> 历史 KV 先由 token-level compressor 压缩，再由 top-k selector 选一部分进入 attention。</li>
  <li><strong>Lightning Indexer:</strong> 不直接输出 hidden state，只给 compressed KV entries 打分并选择 top-k。</li>
</ol>

<p>一句话概括：CSA 不是把所有历史 token 都扔进 attention，而是先把每 $m$ 个 KV 压成一个 entry，再让每个 query 只 attend 到 top-k 个 compressed entries，同时额外保留最近 sliding window 的原始 KV。</p>

<h3 id="61-compressed-key-value-entries">6.1 Compressed Key-Value Entries</h3>

<p>论文从输入 hidden states 开始：</p>

\[H\in\mathbb{R}^{n\times d}\]

<p>其中 $n$ 是 sequence length，$d$ 是 hidden size。CSA 先计算两组 KV entries：</p>

\[C^a,C^b\in\mathbb{R}^{n\times c}\]

<p>以及两组对应的 compression weights：</p>

\[Z^a,Z^b\in\mathbb{R}^{n\times c}\]

<p>这里 $c$ 是 head dimension。更细地说：</p>

<ul>
  <li>$C^a,C^b$ 是两套 original KV entries，后面会被加权压缩成 compressed KV entries。</li>
  <li>$Z^a,Z^b$ 不是最终 KV，而是和 $C^a,C^b$ 同形状的 compression logits / weights，用来决定窗口内不同 token 的压缩权重。</li>
  <li>上标 $a,b$ 表示两组 interleaved / overlapping series，不是 attention head 编号。</li>
</ul>

<p>对第 $i$ 个 compressed entry，论文把两个长度为 $m$ 的连续片段拼起来看：</p>

\[\begin{aligned}
C^b_{m(i-1):mi-1}
&amp;\in
\mathbb{R}^{m\times c} \\
C^a_{mi:m(i+1)-1}
&amp;\in
\mathbb{R}^{m\times c}
\end{aligned}\]

<p>也就是说，$b$ 组的 $m$ 个 token 在前，$a$ 组的 $m$ 个 token 紧接在后。对 $C_i^{\mathrm{Comp}}$ 来说，它实际聚合了一个长度为 $2m$ 的邻域信息：</p>

\[\left[
m(i-1),\;
m(i+1)-1
\right]\]

<p>但输出仍然只有一个 compressed entry。</p>

<p>论文公式可以理解成两步。第一步，先在这 $2m$ 个位置上做 compression softmax：</p>

\[\begin{aligned}
\left[
S^a_{mi:m(i+1)-1};
S^b_{m(i-1):mi-1}
\right]
&amp;=
\mathrm{Softmax}_{\mathrm{row}}
\left(
\left[
Z^a_{mi:m(i+1)-1};
Z^b_{m(i-1):mi-1}
\right]
+
\left[
b^a;
b^b
\right]
\right)
\end{aligned}\]

<p>其中：</p>

<ul>
  <li>$S^a,S^b$ 是真正用于加权求和的 compression scores。</li>
  <li>$b^a,b^b\in\mathbb{R}^{m\times c}$ 是 learnable compression biases。</li>
  <li>$\mathrm{Softmax}_{\mathrm{row}}$ 表示沿 token/window 这一维归一化；对每个 head dimension，都在这 $2m$ 个候选 token 上分配权重。</li>
</ul>

<p>第二步，用这些 scores 加权两组 KV entries：</p>

\[\begin{aligned}
C_i^{\mathrm{Comp}}
&amp;=
\sum_{j=mi}^{m(i+1)-1}
S^a_j\odot C^a_j
+
\sum_{j=m(i-1)}^{mi-1}
S^b_j\odot C^b_j
\end{aligned}\]

<p>其中 $\odot$ 是逐维乘法。因为每个 $C^a_j,C^b_j\in\mathbb{R}^{c}$，所以：</p>

\[C_i^{\mathrm{Comp}}\in\mathbb{R}^{c}\]

<p>现在再回答“为什么有 overlap 还能压缩到 $1/m$”：因为输出 index 还是按 $i$ 走，每 $m$ 个 token 产生一个 $C_i^{\mathrm{Comp}}$。overlap 只改变一个 compressed entry 聚合时能看的范围，让它看见 $2m$ 个 token；它不改变输出 stride。于是 entries 数量仍然约为：</p>

\[\frac{n}{m}\]

<p>所以压缩率仍然约为：</p>

\[\frac{1}{m}\]

<p>一个更具体的例子：如果 $m=4$，那么 $C_i^{\mathrm{Comp}}$ 聚合的是前一个 block 的 4 个 token 和当前 block 的 4 个 token，一共 8 个 token 的信息；但它仍然只输出 1 个 $c$ 维 compressed KV entry。</p>

<p>边界条件也要注意：当 $i=0$ 时，$m(i-1):mi-1$ 这一段不存在。论文说这个 undefined 的 segment 会被忽略。实现上通常可以理解为 mask 掉非法位置，让 softmax 不给这些位置分配权重。</p>

<h3 id="62-lightning-indexer-for-sparse-selection">6.2 Lightning Indexer for Sparse Selection</h3>

<p>压缩后的 $C^{\mathrm{Comp}}$ 仍然很多。例如 1M context、$m=4$ 时，远程 compressed entries 仍然是十万级。CSA 不会把所有 compressed entries 都送进 core attention，而是借用 DSA strategy，先用一个轻量 indexer 做 top-k selection。</p>

<p>这一节最容易混淆的点是：<strong>Indexer 不是最终 attention。它只负责从 $C^{\mathrm{Comp}}$ 中选出 $C_t^{\mathrm{SprsComp}}$，然后 core attention 再读取这些 selected entries。</strong></p>

<h4 id="621-indexer-keys">6.2.1 Indexer keys</h4>

<p>论文先对历史 hidden states 做一次和 $C^{\mathrm{Comp}}$ 类似的压缩，但这里压缩出来的是 indexer keys，而不是最终参与 attention 的 KV entries：</p>

\[K^{\mathrm{IComp}}\in \mathbb{R}^{\frac{n}{m}\times c^I}\]

<p>其中 $c^I$ 是 indexer head dimension。第 $s$ 个 compressed block 的 indexer key 记作 $K_s^{\mathrm{IComp}}\in \mathbb{R}^{c^I}$。可以把它理解成每个 compressed block 的“检索向量”。真正被选中的内容仍然来自 $C_s^{\mathrm{Comp}}$。</p>

<h4 id="622-indexer-queries">6.2.2 Indexer queries</h4>

<p>对 query token $t$，论文从它的 hidden state $h_t$ 生成 indexer queries。第一步是 down-projection：</p>

\[c_t^Q = h_t \cdot W^{DQ}\]

<p>然后再 up-projection 成 $n^{I}_{h}$ 个 indexer query heads：</p>

\[\begin{aligned}
\left[q^{I}_{t,1};q^{I}_{t,2};\ldots;q^{I}_{t,n^{I}_{h}}\right]
&amp;= q_t^I
= c_t^Q \cdot W^{IUQ}
\end{aligned}\]

<p>维度对应关系如下：</p>

<ul>
  <li>$h_t\in \mathbb{R}^{d}$：query token $t$ 的输入 hidden state。</li>
  <li>$c_t^Q\in \mathbb{R}^{d_c}$：query compressed latent vector。</li>
  <li>$W^{DQ}\in \mathbb{R}^{d\times d_c}$：indexer query 的 down-projection matrix。</li>
  <li>$W^{IUQ}\in \mathbb{R}^{d_c\times c^I n^{I}_{h}}$：indexer query 的 up-projection matrix。</li>
  <li>$q^{I}_{t,h}\in \mathbb{R}^{c^I}$：第 $h$ 个 indexer query head。</li>
</ul>

<p>这里的“先 down 再 up”不是为了还原原始 hidden state，而是为了用较低秩的 $c_t^Q$ 生成多头 indexer queries。后面的 core attention query 也会复用这个 $c_t^Q$。</p>

<h4 id="623-dsa-index-score">6.2.3 DSA index score</h4>

<p>接下来，论文给每个 indexer query head 一个 query-dependent weight：</p>

\[\begin{aligned}
\left[w^{I}_{t,1};w^{I}_{t,2};\ldots;w^{I}_{t,n^{I}_{h}}\right]
&amp;= w_t^I
= h_t \cdot W^w
\end{aligned}\]

<p>其中 $W^w\in \mathbb{R}^{d\times n^{I}<em>{h}}$，$w^{I}</em>{t,h}\in \mathbb{R}$。然后计算 query token $t$ 对 preceding compressed block $s$ 的 index score：</p>

\[\begin{aligned}
I_{t,s}
&amp;=
\sum_{h=1}^{n^{I}_{h}}
w^{I}_{t,h}\cdot
\mathrm{ReLU}\left(q^{I}_{t,h}\cdot K_s^{\mathrm{IComp}}\right)
\end{aligned}\]

<p>这里 $s&lt;\left\lfloor t/m\right\rfloor$，表示只从 query token 之前的 compressed blocks 中选。这个公式可以拆成三层看：</p>

<ol>
  <li>$q^{I}_{t,h}\cdot K_s^{\mathrm{IComp}}$：第 $h$ 个 indexer head 对 block $s$ 的相似度。</li>
  <li>$\mathrm{ReLU}(\cdot)$：只保留正相关的 block。</li>
  <li>$w^{I}_{t,h}$：让不同 query token 动态决定哪些 indexer heads 更重要。</li>
</ol>

<p>得到所有 $I_{t,:}$ 之后，用 top-k selector 保留对应的 compressed KV entries：</p>

\[\begin{aligned}
C_t^{\mathrm{SprsComp}}
&amp;=
\left\{
C_s^{\mathrm{Comp}}
\mid
I_{t,s}\in \mathrm{Top}\text{-}k\left(I_{t,:}\right)
\right\}
\end{aligned}\]

<p>注意 top-k 选的是 $C_s^{\mathrm{Comp}}$，不是 $K_s^{\mathrm{IComp}}$。$K^{\mathrm{IComp}}$ 只负责打分，$C^{\mathrm{Comp}}$ 才是后面 core attention 的 key/value 来源。</p>

<h4 id="624-伪代码">6.2.4 伪代码</h4>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">lightning_indexer</span><span class="p">(</span><span class="n">h_t</span><span class="p">,</span> <span class="n">K_IComp</span><span class="p">,</span> <span class="n">C_Comp</span><span class="p">,</span> <span class="n">top_k</span><span class="p">):</span>
    <span class="c1"># Eq. (13): low-rank latent for the query token
</span>    <span class="n">c_Q_t</span> <span class="o">=</span> <span class="n">h_t</span> <span class="o">@</span> <span class="n">W_DQ</span>

    <span class="c1"># Eq. (14): n_h^I indexer query heads, each with dimension c^I
</span>    <span class="n">q_I_t</span> <span class="o">=</span> <span class="p">(</span><span class="n">c_Q_t</span> <span class="o">@</span> <span class="n">W_IUQ</span><span class="p">).</span><span class="n">reshape</span><span class="p">(</span><span class="n">n_h_I</span><span class="p">,</span> <span class="n">c_I</span><span class="p">)</span>

    <span class="c1"># Eq. (15): query-dependent weights for indexer heads
</span>    <span class="n">w_I_t</span> <span class="o">=</span> <span class="n">h_t</span> <span class="o">@</span> <span class="n">W_w</span>

    <span class="n">scores</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="k">for</span> <span class="n">s</span><span class="p">,</span> <span class="n">K_s_IComp</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">K_IComp</span><span class="p">):</span>
        <span class="c1"># Eq. (16): DSA-style sparse selection score
</span>        <span class="n">score</span> <span class="o">=</span> <span class="mf">0.0</span>
        <span class="k">for</span> <span class="n">h</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">n_h_I</span><span class="p">):</span>
            <span class="n">score</span> <span class="o">+=</span> <span class="n">w_I_t</span><span class="p">[</span><span class="n">h</span><span class="p">]</span> <span class="o">*</span> <span class="n">relu</span><span class="p">(</span><span class="n">dot</span><span class="p">(</span><span class="n">q_I_t</span><span class="p">[</span><span class="n">h</span><span class="p">],</span> <span class="n">K_s_IComp</span><span class="p">))</span>
        <span class="n">scores</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">score</span><span class="p">)</span>

    <span class="n">selected_blocks</span> <span class="o">=</span> <span class="n">topk</span><span class="p">(</span><span class="n">scores</span><span class="p">,</span> <span class="n">k</span><span class="o">=</span><span class="n">top_k</span><span class="p">)</span>

    <span class="c1"># Eq. (17): selected compressed KV entries for core attention
</span>    <span class="k">return</span> <span class="p">[</span><span class="n">C_Comp</span><span class="p">[</span><span class="n">s</span><span class="p">]</span> <span class="k">for</span> <span class="n">s</span> <span class="ow">in</span> <span class="n">selected_blocks</span><span class="p">]</span>
</code></pre></div></div>

<h3 id="63-shared-key-value-multi-query-attention">6.3 Shared Key-Value Multi-Query Attention</h3>

<p>选出 $C_t^{\mathrm{SprsComp}}$ 之后，CSA 才进入真正的 core attention。论文这里采用 Multi-Query Attention 的形式：每个 compressed KV entry 同时作为 key 和 value，而 query 有 $n_h$ 个 heads。</p>

<p>先从前面已经得到的 $c_t^Q$ 生成 core attention queries：</p>

\[\begin{aligned}
\left[q_{t,1};q_{t,2};\ldots;q_{t,n_h}\right]
&amp;= q_t
= c_t^Q \cdot W^{UQ}
\end{aligned}\]

<p>其中 $n_h$ 是 query heads 数量，$W^{UQ}\in \mathbb{R}^{d_c\times c n_h}$，每个 $q_{t,i}\in \mathbb{R}^{c}$。论文特别说明：这里的 latent query vector $c_t^Q$ 和 indexer queries 使用的是同一个 $c_t^Q$。</p>

<p>然后对每个 query head 做 core attention：</p>

\[\begin{aligned}
o_{t,i}
&amp;=
\mathrm{CoreAttn}
\left(
\mathrm{query}=q_{t,i},
\mathrm{key}=C_t^{\mathrm{SprsComp}},
\mathrm{value}=C_t^{\mathrm{SprsComp}}
\right)
\end{aligned}\]

<p>其中 $o_{t,i}\in \mathbb{R}^{c}$。Figure 3 里还会把 selected compressed KV entries 与 sliding window KV entries concat 后送入 shared key-value MQA；公式 (19) 写的是 sparse compressed branch 的核心形式。保留 sliding window KV 的目的，是让最近 token 的 fine-grained dependencies 不完全依赖压缩表示。</p>

<h3 id="64-grouped-output-projection">6.4 Grouped Output Projection</h3>

<p>Core attention 之后会得到 $n_h$ 个 head outputs：</p>

\[\begin{aligned}
\left[o_{t,1};o_{t,2};\ldots;o_{t,n_h}\right]
&amp;= o_t
\in \mathbb{R}^{c n_h}
\end{aligned}\]

<p>如果直接把 $o_t\in \mathbb{R}^{c n_h}$ 投影回 $d$ 维 hidden state，计算量大约是：</p>

\[\mathcal{O}\left(c n_h d\right)\]

<p>论文说 DeepSeek-V4 中 $c n_h$ 很大，所以直接做 output projection 会很贵。Grouped Output Projection 的做法是先把 $n_h$ 个 outputs 分成 $g$ 组。第 $i$ 组输出满足：</p>

\[o^G_{t,i}\in \mathbb{R}^{c\frac{n_h}{g}}\]

<p>每一组先投影到一个更小的 intermediate output：</p>

\[o^{G'}_{t,i}\in \mathbb{R}^{d_g},
\qquad
d_g &lt; c\frac{n_h}{g}\]

<p>最后再把所有组的 intermediate outputs 拼接：</p>

\[\begin{aligned}
\left[o^{G'}_{t,1};o^{G'}_{t,2};\ldots;o^{G'}_{t,g}\right]
&amp;\in \mathbb{R}^{d_g g}
\end{aligned}\]

<p>并投影成最终 attention output：</p>

\[\hat{o}_t\in \mathbb{R}^{d}\]

<p>直观计算量可以写成两段：</p>

\[\mathcal{O}\left(c n_h d_g\right)
+
\mathcal{O}\left(g d_g d\right)\]

<p>因为 $d_g &lt; c\frac{n_h}{g}$，第一段相当于先在每组内部做 bottleneck compression；第二段再做跨组混合。它比直接 $c n_h\rightarrow d$ 更省的核心原因是：不再让一个巨大矩阵同时处理所有 heads 到 $d$ 维，而是先把每组 head outputs 压到较小的 $d_g$。</p>

<h3 id="65-attention-sink">6.5 Attention Sink</h3>

<p>论文在 CSA / HCA 的 core attention 里还用了 <strong>Attention Sink</strong>。这一节的 $h$ 是 attention head 的编号，不是 mHC 里的 $n_{\mathrm{hc}}$。</p>

<p>普通 attention 会把一个 head 的全部权重都分给真实 token / compressed block：</p>

\[\sum_j s_{h,i,j}=1\]

<p>但在长上下文压缩场景里，有些 compressed block 对当前 query 可能没用。如果没有 sink，这个 head 还是必须把 100% 概率质量硬分给它们，可能引入噪声。</p>

<p>论文的做法是给每个 attention head 加一个可学习的 sink logit $z’_h$，并把 $\exp(z’_h)$ 加进 softmax 分母：</p>

\[\begin{aligned}
s_{h,i,j}
&amp;=
\frac{\exp(z_{h,i,j})}
{\sum_k\exp(z_{h,i,k})+\exp(z'_h)}
\end{aligned}\]

<p>其中：</p>

<ul>
  <li>$z_{h,i,j}$：第 $h$ 个 attention head 中，第 $i$ 个 query token 对第 $j$ 个 preceding token / compressed block 的 attention logit。</li>
  <li>$s_{h,i,j}$：对应的 attention score。</li>
  <li>$z’_h$：第 $h$ 个 attention head 的 learnable sink logit。</li>
</ul>

<p>这样真实 KV 上的总 attention mass 变成：</p>

\[\begin{aligned}
\sum_j s_{h,i,j}
&amp;=
\frac{\sum_j\exp(z_{h,i,j})}
{\sum_k\exp(z_{h,i,k})+\exp(z'_h)}
\le 1
\end{aligned}\]

<p>剩下的概率质量被 sink 吃掉，不参与 value 聚合。直觉上，这允许某个 head 表达“这次这条 attention branch 没什么可看的”，而不是强迫它在一堆低相关 compressed entries 里硬选。</p>

<h3 id="66-hugging-face-对应代码">6.6 Hugging Face 对应代码</h3>

<p>HF / 官方 inference 里的 CSA 关键路径可以分成两块：<code class="language-plaintext highlighter-rouge">Compressor</code> 负责生成 compressed entries，<code class="language-plaintext highlighter-rouge">Indexer</code> 负责从这些 entries 中选 top-k。下面是保留变量含义后的删减版。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># Compressor: C / Z / compression bias
</span><span class="n">coff</span> <span class="o">=</span> <span class="mi">1</span> <span class="o">+</span> <span class="n">overlap</span>
<span class="n">ape</span> <span class="o">=</span> <span class="n">Parameter</span><span class="p">([</span><span class="n">compress_ratio</span><span class="p">,</span> <span class="n">coff</span> <span class="o">*</span> <span class="n">head_dim</span><span class="p">])</span>
<span class="n">wkv</span> <span class="o">=</span> <span class="n">Linear</span><span class="p">(</span><span class="n">dim</span><span class="p">,</span> <span class="n">coff</span> <span class="o">*</span> <span class="n">head_dim</span><span class="p">)</span>
<span class="n">wgate</span> <span class="o">=</span> <span class="n">Linear</span><span class="p">(</span><span class="n">dim</span><span class="p">,</span> <span class="n">coff</span> <span class="o">*</span> <span class="n">head_dim</span><span class="p">)</span>

<span class="n">kv</span> <span class="o">=</span> <span class="n">wkv</span><span class="p">(</span><span class="n">hidden_states</span><span class="p">)</span>       <span class="c1"># corresponds to C entries
</span><span class="n">score</span> <span class="o">=</span> <span class="n">wgate</span><span class="p">(</span><span class="n">hidden_states</span><span class="p">)</span>  <span class="c1"># corresponds to Z weights
</span>
<span class="n">kv</span> <span class="o">=</span> <span class="n">kv</span><span class="p">.</span><span class="n">unflatten</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">compress_ratio</span><span class="p">))</span>
<span class="n">score</span> <span class="o">=</span> <span class="n">score</span><span class="p">.</span><span class="n">unflatten</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">compress_ratio</span><span class="p">))</span> <span class="o">+</span> <span class="n">ape</span>

<span class="k">if</span> <span class="n">overlap</span><span class="p">:</span>
    <span class="n">kv</span> <span class="o">=</span> <span class="n">overlap_transform</span><span class="p">(</span><span class="n">kv</span><span class="p">,</span> <span class="n">value</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
    <span class="n">score</span> <span class="o">=</span> <span class="n">overlap_transform</span><span class="p">(</span><span class="n">score</span><span class="p">,</span> <span class="n">value</span><span class="o">=-</span><span class="n">inf</span><span class="p">)</span>

<span class="n">kv</span> <span class="o">=</span> <span class="p">(</span><span class="n">kv</span> <span class="o">*</span> <span class="n">score</span><span class="p">.</span><span class="n">softmax</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="nb">sum</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="n">kv</span> <span class="o">=</span> <span class="n">norm</span><span class="p">(</span><span class="n">kv</span><span class="p">)</span>
</code></pre></div></div>

<p>当 <code class="language-plaintext highlighter-rouge">overlap=True</code> 时，<code class="language-plaintext highlighter-rouge">overlap_transform</code> 对应论文里的 $C^b_{m(i-1):mi-1}$ 和 $C^a_{mi:m(i+1)-1}$ 拼接；当 <code class="language-plaintext highlighter-rouge">compress_ratio=4</code> 时，这就是 CSA 的 compressed KV 生成方式。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># Indexer: Eq. (13)-(17)
</span><span class="n">c_Q_t</span> <span class="o">=</span> <span class="n">h_t</span> <span class="o">@</span> <span class="n">W_DQ</span>
<span class="n">q_I_t</span> <span class="o">=</span> <span class="p">(</span><span class="n">c_Q_t</span> <span class="o">@</span> <span class="n">W_IUQ</span><span class="p">).</span><span class="n">reshape</span><span class="p">(</span><span class="n">n_h_I</span><span class="p">,</span> <span class="n">c_I</span><span class="p">)</span>
<span class="n">w_I_t</span> <span class="o">=</span> <span class="n">h_t</span> <span class="o">@</span> <span class="n">W_w</span>

<span class="n">raw</span> <span class="o">=</span> <span class="n">q_I_t</span> <span class="o">@</span> <span class="n">K_IComp</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="o">-</span><span class="mi">2</span><span class="p">)</span>
<span class="n">scores</span> <span class="o">=</span> <span class="p">(</span><span class="n">relu</span><span class="p">(</span><span class="n">raw</span><span class="p">)</span> <span class="o">*</span> <span class="n">w_I_t</span><span class="p">[:,</span> <span class="bp">None</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">0</span><span class="p">)</span>

<span class="n">selected_indices</span> <span class="o">=</span> <span class="n">topk</span><span class="p">(</span><span class="n">scores</span><span class="p">,</span> <span class="n">k</span><span class="o">=</span><span class="n">index_topk</span><span class="p">)</span>
<span class="n">C_t_SprsComp</span> <span class="o">=</span> <span class="n">gather</span><span class="p">(</span><span class="n">C_Comp</span><span class="p">,</span> <span class="n">selected_indices</span><span class="p">)</span>
</code></pre></div></div>

<p>最后 Attention Sink 对应的是每个 head 一个 learnable sink logit。它不会产生新的 value，只改 softmax 分母：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">denom</span> <span class="o">=</span> <span class="n">exp</span><span class="p">(</span><span class="n">attn_logits</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">keepdim</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span> <span class="o">+</span> <span class="n">exp</span><span class="p">(</span><span class="n">sink_logits</span><span class="p">)</span>
<span class="n">attn_scores</span> <span class="o">=</span> <span class="n">exp</span><span class="p">(</span><span class="n">attn_logits</span><span class="p">)</span> <span class="o">/</span> <span class="n">denom</span>
</code></pre></div></div>

<hr />

<h2 id="7-hca-heavily-compressed-attention"><strong>7. HCA: Heavily Compressed Attention</strong></h2>

<p>HCA = <strong>Heavily Compressed Attention</strong>。</p>

<p>相比 CSA，HCA 更直接：</p>

<ul>
  <li>不做 overlap。</li>
  <li>不做 indexer。</li>
  <li>用更大的压缩窗口，例如 $m’=128$。</li>
  <li>每 $m’$ 个 token 压成一个 compressed entry。</li>
</ul>

<p>对应公式更接近：</p>

\[\begin{aligned}
C^{\mathrm{HCA}}_i
&amp;= \sum_{j=0}^{m'-1}
\mathrm{softmax}(Z_{i,j}+B_j)\, C_{i,j}
\end{aligned}\]

<p>HCA 的定位不是“精确找远程细节”，而是“用很小代价保存远程上下文的粗粒度记忆”。</p>

<table>
  <thead>
    <tr>
      <th>Layer type</th>
      <th style="text-align: right">Compress rate</th>
      <th>Overlap</th>
      <th>Indexer</th>
      <th>KV 形态</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>Sliding window attention</td>
      <td style="text-align: right">0</td>
      <td>No</td>
      <td>No</td>
      <td>最近窗口原始 KV</td>
    </tr>
    <tr>
      <td>CSA</td>
      <td style="text-align: right">4</td>
      <td>Yes</td>
      <td>Yes</td>
      <td>原始窗口 KV + selected compressed KV</td>
    </tr>
    <tr>
      <td>HCA</td>
      <td style="text-align: right">128</td>
      <td>No</td>
      <td>No</td>
      <td>原始窗口 KV + 全部 heavily compressed KV</td>
    </tr>
  </tbody>
</table>

<p>官方 inference 默认 <code class="language-plaintext highlighter-rouge">compress_ratios</code> 类似：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>(0, 0, 4, 128, 4, 128, 4, 0)
</code></pre></div></div>

<p>这说明不同层混合使用无压缩、CSA 和 HCA。这个设计挺合理：不是所有层都需要同一种长程记忆，一些层负责精确，一些层负责便宜地扫远处。</p>

<p>HCA 在代码上仍然复用 <code class="language-plaintext highlighter-rouge">Compressor</code>，关键差异是压缩率更大、不开 overlap、也不走 indexer：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">compress_ratio</span> <span class="o">=</span> <span class="mi">128</span>
<span class="n">overlap</span> <span class="o">=</span> <span class="bp">False</span>

<span class="n">kv</span> <span class="o">=</span> <span class="n">wkv</span><span class="p">(</span><span class="n">hidden_states</span><span class="p">)</span>
<span class="n">score</span> <span class="o">=</span> <span class="n">wgate</span><span class="p">(</span><span class="n">hidden_states</span><span class="p">)</span> <span class="o">+</span> <span class="n">ape</span>

<span class="n">kv</span> <span class="o">=</span> <span class="n">kv</span><span class="p">.</span><span class="n">unflatten</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">compress_ratio</span><span class="p">))</span>
<span class="n">score</span> <span class="o">=</span> <span class="n">score</span><span class="p">.</span><span class="n">unflatten</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">compress_ratio</span><span class="p">))</span>

<span class="n">kv</span> <span class="o">=</span> <span class="p">(</span><span class="n">kv</span> <span class="o">*</span> <span class="n">score</span><span class="p">.</span><span class="n">softmax</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="nb">sum</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="n">kv</span> <span class="o">=</span> <span class="n">norm</span><span class="p">(</span><span class="n">kv</span><span class="p">)</span>
</code></pre></div></div>

<p>所以 HCA 的代码路径比 CSA 少两件事：没有 <code class="language-plaintext highlighter-rouge">overlap_transform</code>，也没有 <code class="language-plaintext highlighter-rouge">Indexer -&gt; topk -&gt; gather(C_Comp)</code>。</p>

<hr />

<h2 id="8-muon-optimizer"><strong>8. Muon Optimizer</strong></h2>

<p>DeepSeek-V4 的更新不只在 attention 和 MoE 上，也包括训练优化器。模型卡和论文摘要里都把 <strong>Muon optimizer</strong> 列为训练侧改动：它服务于更快收敛和更稳定的大规模预训练，而不是推理时的网络结构。</p>

<p>Muon 通常展开为 <strong>Momentum Orthogonalized by Newton-Schulz</strong>。它的核心不是像 AdamW 那样给每个参数做逐元素自适应缩放，而是把二维权重矩阵的梯度更新看成一个矩阵，再对 momentum update 做近似正交化。</p>

<p>可以用下面的抽象公式理解：</p>

\[\begin{aligned}
G_t &amp;= \nabla_W \mathcal{L}(W_t) \\
M_t &amp;= \mu M_{t-1}+G_t \\
\hat{M}_t &amp;= \mathrm{NewtonSchulz}(M_t) \\
W_{t+1} &amp;= W_t-\eta \hat{M}_t
\end{aligned}\]

<p>其中：</p>

<ul>
  <li>$G_t$：当前 step 对权重矩阵 $W_t$ 的梯度。</li>
  <li>$M_t$：带 momentum 的更新方向。</li>
  <li>$\mathrm{NewtonSchulz}(\cdot)$：用 Newton-Schulz iteration 做近似正交化，让更新方向在矩阵谱上更均衡。</li>
  <li>$\hat{M}_t$：正交化后的更新方向。</li>
</ul>

<p>直觉上，Muon 希望避免某些奇异方向被更新得过强、另一些方向几乎不动。对大模型里的线性层矩阵来说，这种矩阵级更新比逐元素缩放更贴近“层”的结构。</p>

<p>一个很粗略的 Newton-Schulz 伪代码是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">G</span> <span class="o">=</span> <span class="n">grad</span><span class="p">(</span><span class="n">W</span><span class="p">)</span>
<span class="n">M</span> <span class="o">=</span> <span class="n">momentum</span> <span class="o">*</span> <span class="n">M</span> <span class="o">+</span> <span class="n">G</span>

<span class="n">X</span> <span class="o">=</span> <span class="n">normalize</span><span class="p">(</span><span class="n">M</span><span class="p">)</span>
<span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_ns_steps</span><span class="p">):</span>
    <span class="n">A</span> <span class="o">=</span> <span class="n">X</span> <span class="o">@</span> <span class="n">X</span><span class="p">.</span><span class="n">T</span>
    <span class="n">X</span> <span class="o">=</span> <span class="n">a</span> <span class="o">*</span> <span class="n">X</span> <span class="o">+</span> <span class="n">b</span> <span class="o">*</span> <span class="p">(</span><span class="n">A</span> <span class="o">@</span> <span class="n">X</span><span class="p">)</span> <span class="o">+</span> <span class="n">c</span> <span class="o">*</span> <span class="p">(</span><span class="n">A</span> <span class="o">@</span> <span class="n">A</span> <span class="o">@</span> <span class="n">X</span><span class="p">)</span>

<span class="n">W</span> <span class="o">=</span> <span class="n">W</span> <span class="o">-</span> <span class="n">lr</span> <span class="o">*</span> <span class="n">X</span>
</code></pre></div></div>

<p>这段伪代码只表达机制，不等价于 DeepSeek-V4 的完整训练实现。实际训练里还要处理参数分组、混合精度、分布式并行和不同参数类型的 optimizer 分配。对这篇笔记来说，最重要的结论是：<strong>Muon 是训练稳定性和收敛速度的优化，不改变推理时的 forward graph；CSA、HCA、mHC、MoE router 才是模型结构里的主要变化。</strong></p>

<hr />

<h2 id="9-moe-runtime-optimization"><strong>9. MoE Runtime Optimization</strong></h2>

<p>DeepSeek-V4 论文里还有一个容易被忽略但很重要的系统优化：<strong>Fine-Grained Communication-Computation Overlap in Expert Parallelism</strong>。它不是改 MoE 的数学定义，也不是改变 router 选专家的规则，而是在 expert parallelism 下重新安排 MoE 层的执行顺序。</p>

<p>论文把一个 MoE 层拆成通信和计算交替出现的五段：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Dispatch All-to-All
-&gt; Linear-1 GEMM
-&gt; SwiGLU + FP8 Cast
-&gt; Linear-2 GEMM
-&gt; Combine All-to-All
</code></pre></div></div>

<p>其中 <code class="language-plaintext highlighter-rouge">Linear-1</code> 和 <code class="language-plaintext highlighter-rouge">Linear-2</code> 是 expert FFN 内部天然存在的两个 GEMM，不是额外引入的新结构。如果 expert 使用 SwiGLU，那么 <code class="language-plaintext highlighter-rouge">Linear-1</code> 通常同时产生 gate/up 两路中间表示，<code class="language-plaintext highlighter-rouge">SwiGLU + FP8 Cast</code> 把中间激活重新量化后交给 <code class="language-plaintext highlighter-rouge">Linear-2</code>。</p>

<h3 id="91-wave-到底是什么">9.1 wave 到底是什么</h3>

<p>DeepSeek-V4 的关键做法是把 local experts 进一步切成多个 <strong>expert waves</strong>。一个 wave 不是 token chunk，也不是某个 expert 的一个 GEMM tile，而是一小批 experts 及其对应的 routed tokens。</p>

<p>论文里的调度逻辑可以理解成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>wave 0 的 experts 完成 dispatch -&gt; 立刻开始 Linear-1 / Act / Linear-2
wave 1 的 experts 仍在接收 token -&gt; 通信继续推进
wave -1 的 experts 已经算完 -&gt; output 立刻 combine 回原 rank
</code></pre></div></div>

<p>进入稳态后，同一时刻可以同时发生三件事：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>combine previous wave | compute current wave | dispatch next wave
</code></pre></div></div>

<p>这就是它和普通 EP overlap 的区别。普通做法经常是“整层 dispatch 完成后再算，整层算完后再 combine”；DeepSeek-V4 让每个 wave 准备好就进入计算，不等其他 experts。这样可以减少 long-tail expert 或小 batch 场景下的空转，论文特别提到 RL rollout 和 high-speed agent serving 这类 latency-sensitive 场景会受益。</p>

<h3 id="92-为什么通信能被盖住">9.2 为什么通信能被盖住</h3>

<p>论文给出的判断不是“带宽越高越好”，而是看计算通信比：</p>

\[\frac{C}{B}\leq \frac{V_{\mathrm{comp}}}{V_{\mathrm{comm}}}\]

<p>其中 $C$ 是设备峰值计算能力，$B$ 是互联带宽，$V_{\mathrm{comp}}$ 是 MoE 计算量，$V_{\mathrm{comm}}$ 是 MoE 通信量。</p>

<p>对 DeepSeek-V4-Pro，一个 token-expert pair 的 SwiGLU expert 大约需要：</p>

\[V_{\mathrm{comp}}\approx 6hd\]

<p>这里的 $h$ 是 hidden size，$d$ 是 intermediate hidden size。三项来自 gate projection、up projection、down projection。通信量约为：</p>

\[V_{\mathrm{comm}}\approx 3h\]

<p>这里对应 FP8 dispatch 和 BF16 combine。因此：</p>

\[\frac{V_{\mathrm{comp}}}{V_{\mathrm{comm}}}\approx 2d\]

<p>当 $d=3072$ 时，阈值约为：</p>

\[2d=6144\ \mathrm{FLOPs/Byte}\]

<p>含义是：如果硬件的 $C/B$ 不超过这个比例，那么通信理论上可以被 expert GEMM 覆盖；继续堆更多互联带宽的收益会下降，反而要保证计算、显存和网络能同时高负载运行。</p>

<h3 id="93-公开代码能看到什么">9.3 公开代码能看到什么</h3>

<p>论文说已经开源 CUDA-based mega-kernel <strong>MegaMoE2</strong>，位置在 DeepGEMM。DeepGEMM PR #304 对这个 kernel 的描述很直接：它把 <code class="language-plaintext highlighter-rouge">dispatch / linear 1 / SwiGLU / linear 2 / combine</code> 融合进一个 mega-kernel，并重叠 NVLink communication 和 tensor core computation。</p>

<p>公开测试文件 <code class="language-plaintext highlighter-rouge">tests/test_mega_moe.py</code> 能看到两个路径的对照：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># fused path
</span><span class="nb">buffer</span> <span class="o">=</span> <span class="n">get_symm_buffer_for_mega_moe</span><span class="p">(...)</span>
<span class="n">transformed_l1_weights</span><span class="p">,</span> <span class="n">transformed_l2_weights</span> <span class="o">=</span> <span class="n">transform_weights_for_mega_moe</span><span class="p">(...)</span>
<span class="n">fp8_fp4_mega_moe</span><span class="p">(</span>
    <span class="n">y</span><span class="p">,</span>
    <span class="n">transformed_l1_weights</span><span class="p">,</span>
    <span class="n">transformed_l2_weights</span><span class="p">,</span>
    <span class="nb">buffer</span><span class="p">,</span>
    <span class="n">cumulative_local_expert_recv_stats</span><span class="o">=</span><span class="p">...,</span>
<span class="p">)</span>

<span class="c1"># legacy baseline
</span><span class="n">recv_x</span><span class="p">,</span> <span class="n">handle</span> <span class="o">=</span> <span class="n">ep_buffer</span><span class="p">.</span><span class="n">dispatch</span><span class="p">(...)</span>
<span class="n">l1_y</span> <span class="o">=</span> <span class="n">grouped_fp8_fp4_gemm</span><span class="p">(</span><span class="n">recv_x</span><span class="p">,</span> <span class="n">l1_weights</span><span class="p">,</span> <span class="p">...)</span>
<span class="n">l1_y</span> <span class="o">=</span> <span class="n">swiglu_and_cast</span><span class="p">(</span><span class="n">l1_y</span><span class="p">,</span> <span class="n">topk_weights</span><span class="p">,</span> <span class="p">...)</span>
<span class="n">l2_y</span> <span class="o">=</span> <span class="n">grouped_fp8_fp4_gemm</span><span class="p">(</span><span class="n">l1_y</span><span class="p">,</span> <span class="n">l2_weights</span><span class="p">,</span> <span class="p">...)</span>
<span class="n">y</span> <span class="o">=</span> <span class="n">ep_buffer</span><span class="p">.</span><span class="n">combine</span><span class="p">(</span><span class="n">l2_y</span><span class="p">,</span> <span class="n">handle</span><span class="o">=</span><span class="n">handle</span><span class="p">)</span>
</code></pre></div></div>

<p>这段对照很有用：它说明 MegaMoE2 不是单独优化某一个 GEMM，而是把 EP dispatch、两次 expert GEMM、SwiGLU/cast、combine 都放进同一个 fused runtime 里处理。Python API 层没有暴露 <code class="language-plaintext highlighter-rouge">num_waves</code> 这种模型超参数；wave 更像 kernel 内部的调度粒度，由 routed token metadata、local expert receive stats、symmetric buffer 和 transformed expert weights 共同驱动。</p>

<p>PR #316 的 benchmark 也和论文说法一致：在 EP8 下，DeepSeek-V4-Flash 使用 256 experts、top-k=6、hidden size 4096、intermediate size 2048；DeepSeek-V4-Pro 使用 384 experts、top-k=6、hidden size 7168、intermediate size 3072。不同 batch size 下相对 legacy baseline 大约有 1.50x 到 1.96x 的加速，其中小 batch latency 场景收益更明显。</p>

<h3 id="94-论文给硬件和-kernel-的几个观察">9.4 论文给硬件和 kernel 的几个观察</h3>

<p>DeepSeek-V4 这节最后还有几个系统侧观察，值得单独记下来：</p>

<ul>
  <li><strong>不是无限堆带宽。</strong> 当 $C/B$ 已经低于 $2d$ 这个阈值后，继续增加互联带宽的边际收益下降；更重要的是让 compute、memory、network 同时保持高负载。</li>
  <li><strong>power budget 会变成瓶颈。</strong> 极端融合会让 Tensor Core、HBM、NVLink 同时繁忙，功耗墙可能比单看 FLOPs 或带宽更早出现。</li>
  <li><strong>dispatch 采用 pull-based 思路。</strong> 论文说 dispatch 阶段由每个 GPU 主动读取远端 activations，避免 fine-grained push 带来的高通知延迟；如果未来硬件有更低延迟的跨 GPU signaling，push 才会更自然。</li>
  <li><strong>SwiGLU 本身也可能阻塞 pipeline。</strong> 因为 SwiGLU 有指数和除法相关开销，论文建议未来可以考虑更低成本的 element-wise activation，减少 post-GEMM 处理对 GEMM pipeline 的打断。</li>
</ul>

<p>这类优化和 MiniMax-01、Comet、FlashMoE 都有关，但粒度不同：MiniMax-01 更偏 token group / process group 级 overlap；Comet 更偏 shared tensor / GEMM tile 级重排；FlashMoE 更激进，把 MoE 执行放进 persistent kernel 里调度；DeepSeek-V4 则采用 expert-wave pipeline 来服务自己的大规模 MoE 推理与训练系统。</p>

<p>更完整的比较我单独整理在这里：<a href="/blog/2026/05/10/MoE-Optimization.html">MoE 优化的探索</a>。</p>

<hr />

<h2 id="10-总结表"><strong>10. 总结表</strong></h2>

<table>
  <thead>
    <tr>
      <th>模块</th>
      <th>解决的问题</th>
      <th>核心机制</th>
      <th>代价</th>
      <th>源码入口</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>MTP</td>
      <td>next-token 信号太局部</td>
      <td>预测未来多个 token</td>
      <td>额外 MTP block/head 计算</td>
      <td><code class="language-plaintext highlighter-rouge">MTPBlock</code></td>
    </tr>
    <tr>
      <td>MoE Router</td>
      <td>MoE routing scores 需要平滑非负，并选择 experts</td>
      <td>softplus/sqrtsoftplus scores + top-k/hash route</td>
      <td>selection 和 weighting 逻辑更复杂</td>
      <td><code class="language-plaintext highlighter-rouge">Gate</code> / <code class="language-plaintext highlighter-rouge">DeepseekV4TopKRouter</code></td>
    </tr>
    <tr>
      <td>Hash Route</td>
      <td>早期层动态路由成本和波动</td>
      <td>token id 查表得到 experts</td>
      <td>上下文自适应弱一些</td>
      <td><code class="language-plaintext highlighter-rouge">tid2eid</code> / <code class="language-plaintext highlighter-rouge">DeepseekV4HashRouter</code></td>
    </tr>
    <tr>
      <td>MoE Runtime Optimization</td>
      <td>expert parallelism 下通信阻塞计算</td>
      <td>expert waves + MegaMoE2 fused runtime</td>
      <td>kernel 调度和实现复杂度更高</td>
      <td>DeepGEMM <code class="language-plaintext highlighter-rouge">fp8_fp4_mega_moe</code></td>
    </tr>
    <tr>
      <td>mHC</td>
      <td>深层 residual 信号传播</td>
      <td>多 stream + Sinkhorn mixing</td>
      <td>小矩阵归一化开销</td>
      <td><code class="language-plaintext highlighter-rouge">DeepseekV4HyperConnection</code></td>
    </tr>
    <tr>
      <td>CSA</td>
      <td>远程上下文太长</td>
      <td>近邻原始 KV + $C_t^{\mathrm{SprsComp}}$ + shared KV MQA</td>
      <td>需要 compressor 和 indexer</td>
      <td><code class="language-plaintext highlighter-rouge">DeepseekV4CSACompressor</code></td>
    </tr>
    <tr>
      <td>HCA</td>
      <td>更便宜的远程记忆</td>
      <td>大窗口压缩，无 indexer</td>
      <td>细节损失更大</td>
      <td><code class="language-plaintext highlighter-rouge">DeepseekV4HCACompressor</code></td>
    </tr>
    <tr>
      <td>Attention Sink</td>
      <td>compressed KV 无关时不应强制分配满权重</td>
      <td>分母加入 learnable sink logit</td>
      <td>每个 head 多一个 sink 参数</td>
      <td>attention score computation</td>
    </tr>
    <tr>
      <td>Lightning Indexer</td>
      <td>compressed entries 仍然太多</td>
      <td>$K^{\mathrm{IComp}}$ 与 $q^I_t$ 计算 $I_{t,s}$ 后 top-k</td>
      <td>可能漏选远程关键信息</td>
      <td><code class="language-plaintext highlighter-rouge">Indexer</code> / <code class="language-plaintext highlighter-rouge">DeepseekV4Indexer</code></td>
    </tr>
    <tr>
      <td>Grouped Output Projection</td>
      <td>$c n_h\rightarrow d$ 直接投影太贵</td>
      <td>先按 $g$ 组压到 $d_g$，再投影到 $d$</td>
      <td>增加 per-group bottleneck</td>
      <td>attention output projection</td>
    </tr>
    <tr>
      <td>Muon Optimizer</td>
      <td>大规模预训练收敛和稳定性</td>
      <td>momentum update + Newton-Schulz 近似正交化</td>
      <td>训练侧额外 optimizer 计算</td>
      <td>training optimizer</td>
    </tr>
  </tbody>
</table>

<hr />

<h2 id="11-源码与资料入口"><strong>11. 源码与资料入口</strong></h2>

<ul>
  <li>DeepSeek-V4 Technical Report: <a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/resolve/main/DeepSeek_V4.pdf">DeepSeek_V4.pdf</a></li>
  <li>DeepSeek-V4-Pro Model Card: <a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro">deepseek-ai/DeepSeek-V4-Pro</a></li>
  <li>Official inference code: <a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/blob/main/inference/model.py">inference/model.py</a></li>
  <li>Official encoding code: <a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/blob/main/encoding/encoding_dsv4.py">encoding/encoding_dsv4.py</a></li>
  <li>DeepGEMM MegaMoE2 PR: <a href="https://github.com/deepseek-ai/DeepGEMM/pull/304">deepseek-ai/DeepGEMM#304</a></li>
  <li>DeepGEMM MegaMoE2 Benchmark PR: <a href="https://github.com/deepseek-ai/DeepGEMM/pull/316">deepseek-ai/DeepGEMM#316</a></li>
  <li>Transformers implementation: <a href="https://github.com/huggingface/transformers/blob/main/src/transformers/models/deepseek_v4/modeling_deepseek_v4.py">modeling_deepseek_v4.py</a></li>
  <li>Transformers config: <a href="https://github.com/huggingface/transformers/blob/main/src/transformers/models/deepseek_v4/configuration_deepseek_v4.py">configuration_deepseek_v4.py</a></li>
  <li>DeepSeek-V3 Technical Report, for MTP background: <a href="https://hf.co/papers/2412.19437">hf.co/papers/2412.19437</a></li>
  <li>mHC-lite discussion: <a href="https://hf.co/papers/2601.05732">hf.co/papers/2601.05732</a></li>
  <li>Sparse indexer related work: <a href="https://hf.co/papers/2603.12201">IndexCache</a>, <a href="https://hf.co/papers/2603.28458">HISA</a></li>
</ul>

<hr />

<h2 id="12-最后的一句话理解"><strong>12. 最后的一句话理解</strong></h2>

<p>DeepSeek-V4 的主线不是“又把模型做大了”。它更像是在问一个工程问题：当上下文来到 1M tokens，模型到底应该怎样保存、压缩、索引和读取历史信息？</p>

<p>CSA、HCA、Lightning Indexer 解决的是 <strong>长上下文的读写成本</strong>；mHC 解决的是 <strong>深层网络里的信息传播稳定性</strong>；MTP、router 和 Muon 则分别服务于 <strong>训练信号密度、sparse computation 与预训练优化</strong>。把这些合起来，DeepSeek-V4 才有了 1M context 下还能工作的结构基础。</p>]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[DeepSeek V4: Towards Highly Efficient Million-Token Context Intelligence]]></summary></entry><entry><title type="html">ZERO训练技术</title><link href="https://wqh011128.github.io/blog/2026/05/07/ZERO.html" rel="alternate" type="text/html" title="ZERO训练技术" /><published>2026-05-07T00:00:00+00:00</published><updated>2026-05-07T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/05/07/ZERO</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/05/07/ZERO.html"><![CDATA[<h1 id="zero训练技术">ZERO训练技术</h1>

<p>这篇笔记里我统一使用 <code class="language-plaintext highlighter-rouge">ZeRO</code> 这个写法，它指的是 DeepSpeed 提出的 <code class="language-plaintext highlighter-rouge">Zero Redundancy Optimizer</code>。理解 ZeRO 最好的方式，不是背“Stage 1/2/3 分别切了什么”，而是顺着一次训练 step 去看：</p>

<ol>
  <li>前向时每张卡手里有什么。</li>
  <li>反向时梯度怎么产生、怎么通信、怎么留下来。</li>
  <li>参数更新时谁来更、更新完怎么让所有卡重新对齐。</li>
</ol>

<p>只要把这三件事看顺了，<code class="language-plaintext highlighter-rouge">ZeRO-1</code>、<code class="language-plaintext highlighter-rouge">ZeRO-2</code>、<code class="language-plaintext highlighter-rouge">ZeRO-3</code> 的差异就会非常清楚。</p>

<hr />

<h2 id="先立一个基线普通-data-parallel-怎么训练">先立一个基线：普通 Data Parallel 怎么训练</h2>

<p>在普通的数据并行 <code class="language-plaintext highlighter-rouge">Data Parallel, DP</code> 里，假设有 <code class="language-plaintext highlighter-rouge">N</code> 张卡，那么每一张卡都会保存完整的一份：</p>

<ul>
  <li>model parameters</li>
  <li>gradients</li>
  <li>optimizer states</li>
</ul>

<p>这里的 optimizer states 以 Adam 为例，通常包括：</p>

<ul>
  <li>参数本身对应的 FP16/BF16 权重</li>
  <li>FP32 master weights</li>
  <li>一阶矩 <code class="language-plaintext highlighter-rouge">m</code></li>
  <li>二阶矩 <code class="language-plaintext highlighter-rouge">v</code></li>
</ul>

<p>所以普通 DP 的问题不是“算不动”，而是“每张卡都重复保存一整套训练状态，显存浪费很大”。</p>

<h3 id="一次训练-step-的流程">一次训练 step 的流程</h3>

<h4 id="1-前向">1. 前向</h4>

<p>每张卡都拿着<strong>完整模型参数</strong>，各自处理不同的 micro-batch。<br />
例如 rank 0 处理 batch 的一部分，rank 1 处理另一部分，但它们前向时用的模型权重是一模一样的完整副本。</p>

<h4 id="2-反向">2. 反向</h4>

<p>每张卡先根据自己的 micro-batch 算出一套<strong>本地完整梯度</strong>。<br />
注意这时的梯度虽然只来自本卡数据，但张量形状是完整的，因为每张卡前向时持有的就是完整模型。</p>

<p>然后所有卡对这些完整梯度做一次 <code class="language-plaintext highlighter-rouge">all-reduce</code>：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">all-reduce</code> 的作用是：各卡把同一个梯度张量做求和或求平均；</li>
  <li>操作结束后，每张卡都会拿到<strong>一份完全相同的完整梯度</strong>。</li>
</ul>

<h4 id="3-参数更新">3. 参数更新</h4>

<p>由于每张卡现在都有：</p>

<ul>
  <li>完整参数</li>
  <li>完整梯度</li>
  <li>完整 optimizer states</li>
</ul>

<p>所以每张卡都会独立执行一遍完全一样的 Adam update。<br />
因为输入完全一样，所以更新后的参数也完全一样，不需要额外同步。</p>

<h3 id="普通-dp-的特点">普通 DP 的特点</h3>

<ul>
  <li>好处：逻辑最直接，实现简单。</li>
  <li>代价：参数、梯度、optimizer states 在每张卡上都完整复制，显存冗余最大。</li>
</ul>

<p>后面的 ZeRO，本质上就是一步一步把这三类状态的“完整复制”改成“按 rank 分片保存”。</p>

<hr />

<h2 id="一张总对比表先建立全局感觉">一张总对比表先建立全局感觉</h2>

<table>
  <thead>
    <tr>
      <th>方案</th>
      <th>参数是否分片</th>
      <th>梯度是否分片</th>
      <th>Optimizer state 是否分片</th>
      <th>前向是否需要 <code class="language-plaintext highlighter-rouge">all-gather</code></th>
      <th>反向主要通信</th>
      <th>更新后如何恢复一致参数副本</th>
      <th>通信开销趋势</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>普通 DP</td>
      <td>否</td>
      <td>否</td>
      <td>否</td>
      <td>否</td>
      <td><code class="language-plaintext highlighter-rouge">all-reduce</code> 完整梯度</td>
      <td>不需要额外恢复，所有卡本地更新结果天然一致</td>
      <td>基线</td>
    </tr>
    <tr>
      <td>ZeRO-1</td>
      <td>否</td>
      <td>否</td>
      <td>是</td>
      <td>否</td>
      <td><code class="language-plaintext highlighter-rouge">all-reduce</code> 完整梯度</td>
      <td>更新后的参数分片通过 <code class="language-plaintext highlighter-rouge">all-gather</code> 重建完整参数</td>
      <td>略高于 DP</td>
    </tr>
    <tr>
      <td>ZeRO-2</td>
      <td>否</td>
      <td>是</td>
      <td>是</td>
      <td>否</td>
      <td>layer-by-layer <code class="language-plaintext highlighter-rouge">reduce-scatter</code> 梯度</td>
      <td>通常仍需 <code class="language-plaintext highlighter-rouge">all-gather</code> 让各卡重新拿到完整参数副本</td>
      <td>高于 ZeRO-1</td>
    </tr>
    <tr>
      <td>ZeRO-3</td>
      <td>是</td>
      <td>是</td>
      <td>是</td>
      <td>是</td>
      <td>参数 <code class="language-plaintext highlighter-rouge">all-gather</code> + 梯度 <code class="language-plaintext highlighter-rouge">reduce-scatter</code></td>
      <td>参数平时就是分片存放，需要时再 <code class="language-plaintext highlighter-rouge">all-gather</code></td>
      <td>最高</td>
    </tr>
  </tbody>
</table>

<p>如果只记一句话：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-1</code> 切 optimizer states。</li>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-2</code> 在 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 基础上再切 gradients。</li>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-3</code> 在 <code class="language-plaintext highlighter-rouge">ZeRO-2</code> 基础上再切 parameters。</li>
</ul>

<p>下面按训练流程逐个展开。</p>

<hr />

<h2 id="zero-1只切-optimizer-states">ZeRO-1：只切 optimizer states</h2>

<p><code class="language-plaintext highlighter-rouge">ZeRO-1</code> 的核心是：<strong>模型参数和梯度仍然完整复制，但 optimizer states 不再在每张卡上都保存完整副本，而是按 rank 分片保存。</strong></p>

<h3 id="先回答最关键的一句">先回答最关键的一句</h3>

<p>在 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 里：</p>

<ul>
  <li>反向传播结束后，<strong>每张卡保留的梯度仍然是完整的</strong>；</li>
  <li>被切分的重点是 optimizer states，而不是梯度；</li>
  <li>每张卡只更新自己负责的参数分片，是因为 optimizer states 的 ownership 已经按参数分片分配给不同 rank 了。</li>
</ul>

<h3 id="optimizer-state-怎么切分">optimizer state 怎么切分</h3>

<p>假设参数向量被逻辑上切成 <code class="language-plaintext highlighter-rouge">N</code> 份，对应 <code class="language-plaintext highlighter-rouge">N</code> 个 rank：</p>

<ul>
  <li>rank 0 负责第 0 段参数对应的 optimizer states</li>
  <li>rank 1 负责第 1 段参数对应的 optimizer states</li>
  <li>…</li>
  <li>rank <code class="language-plaintext highlighter-rouge">N-1</code> 负责最后一段</li>
</ul>

<p>这里“负责”指的是：</p>

<ul>
  <li>这段参数的 <code class="language-plaintext highlighter-rouge">m</code></li>
  <li>这段参数的 <code class="language-plaintext highlighter-rouge">v</code></li>
  <li>这段参数的 FP32 master copy</li>
</ul>

<p>只保存在对应的 rank 上，而不是每张卡都保存一份。</p>

<p>所以 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 节省的显存主要来自 Adam 状态，因为 Adam 的状态量通常比只看参数本身更大。</p>

<h3 id="一次训练-step-的流程-1">一次训练 step 的流程</h3>

<h4 id="1-前向-1">1. 前向</h4>

<p>和普通 DP 完全一样：</p>

<ul>
  <li>每张卡都有完整参数副本；</li>
  <li>每张卡用自己的 micro-batch 做前向。</li>
</ul>

<p>因此前向阶段并不需要额外参数通信。</p>

<h4 id="2-反向-1">2. 反向</h4>

<p>和普通 DP 也基本一样：</p>

<ul>
  <li>每张卡先得到一份本地完整梯度；</li>
  <li>然后对完整梯度做 <code class="language-plaintext highlighter-rouge">all-reduce</code>，得到所有卡一致的完整梯度。</li>
</ul>

<p>所以到反向结束时，每张卡手里仍然有<strong>一整份完整梯度</strong>，而不是只保留自己的梯度分片。</p>

<h4 id="3-更新">3. 更新</h4>

<p>这一步才是 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 和普通 DP 真正拉开差异的地方。</p>

<p>虽然每张卡都有完整梯度，但并不是每张卡都去更新整套参数，而是：</p>

<ol>
  <li>每个 rank 只更新自己负责的参数分片。</li>
  <li>这次更新所需的 optimizer states 也只在这个 rank 上存在。</li>
  <li>其他不归自己负责的参数段，本 rank 不做 optimizer update。</li>
</ol>

<p>换句话说：</p>

<ul>
  <li>梯度是完整可见的；</li>
  <li>但 optimizer states 是分片持有的；</li>
  <li>所以每个 rank 只对自己那一片参数拥有真正的“更新权限”。</li>
</ul>

<h3 id="更新后怎么得到统一模型">更新后怎么得到统一模型</h3>

<p>每个 rank 只更新了自己的参数分片之后，所有卡手里的参数副本其实暂时是不完整的“新状态”。<br />
这时需要把每张卡更新好的参数分片重新汇总，恢复成一致的完整模型副本。</p>

<p>通常用的算子是 <code class="language-plaintext highlighter-rouge">all-gather</code>。</p>

<h4 id="all-gather-是什么"><code class="language-plaintext highlighter-rouge">all-gather</code> 是什么</h4>

<p>可以把 <code class="language-plaintext highlighter-rouge">all-gather</code> 理解成：</p>

<ul>
  <li>每张卡拿出自己持有的一段张量；</li>
  <li>通信结束后，所有卡都拿到这些分片拼起来的完整张量。</li>
</ul>

<p>所以在 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 里，<code class="language-plaintext highlighter-rouge">all-gather</code> 的作用就是：</p>

<ul>
  <li>rank 0 拿出自己更新后的参数分片；</li>
  <li>rank 1 拿出自己更新后的参数分片；</li>
  <li>…</li>
  <li>最后所有卡都重建出一份一致的完整参数副本。</li>
</ul>

<h4 id="all-gather-和-all-reduce-的区别"><code class="language-plaintext highlighter-rouge">all-gather</code> 和 <code class="language-plaintext highlighter-rouge">all-reduce</code> 的区别</h4>

<ul>
  <li><code class="language-plaintext highlighter-rouge">all-reduce</code> 是“同位置元素做规约，比如求和/求平均，然后每张卡都得到同一个规约结果”。</li>
  <li><code class="language-plaintext highlighter-rouge">all-gather</code> 是“每张卡拿出不同分片，最后所有卡都拿到拼接后的完整结果”。</li>
</ul>

<p>在 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 里：</p>

<ul>
  <li>梯度同步靠 <code class="language-plaintext highlighter-rouge">all-reduce</code></li>
  <li>更新后参数重建靠 <code class="language-plaintext highlighter-rouge">all-gather</code></li>
</ul>

<h3 id="相比普通-dpzero-1-到底改了什么">相比普通 DP，ZeRO-1 到底改了什么</h3>

<p>相比普通 DP，<code class="language-plaintext highlighter-rouge">ZeRO-1</code> 唯一改变的是：<br />
<strong>optimizer states 从“每卡完整保存”变成了“按 rank 分片保存”；参数和梯度依然是完整副本。</strong></p>

<hr />

<h2 id="zero-2在-zero-1-基础上再切-gradients">ZeRO-2：在 ZeRO-1 基础上再切 gradients</h2>

<p><code class="language-plaintext highlighter-rouge">ZeRO-2</code> 的核心是：<strong>不仅 optimizer states 分片，梯度也分片。</strong><br />
参数在大多数时刻仍然是完整副本，这一点和 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 一样。</p>

<p>最重要的区别是：<code class="language-plaintext highlighter-rouge">ZeRO-2</code> 不再让“整网完整梯度”在每张卡上一直保留到反向结束，而是尽量在反向过程中就把梯度规约并切走。</p>

<h3 id="一次训练-step-的流程-2">一次训练 step 的流程</h3>

<h4 id="1-前向-2">1. 前向</h4>

<p>前向和 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 一样：</p>

<ul>
  <li>每张卡都持有完整参数副本；</li>
  <li>每张卡对自己的 micro-batch 做前向。</li>
</ul>

<p>因此前向阶段仍然不需要额外参数 <code class="language-plaintext highlighter-rouge">all-gather</code>。</p>

<h4 id="2-反向重点是-layer-by-layer-的-reduce-scatter">2. 反向：重点是 layer-by-layer 的 <code class="language-plaintext highlighter-rouge">reduce-scatter</code></h4>

<p>这里是 <code class="language-plaintext highlighter-rouge">ZeRO-2</code> 的重点。</p>

<p>在普通 DP 或 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 里，一个直观想法是：</p>

<ul>
  <li>整个模型先反向完；</li>
  <li>每张卡手里留着整网完整梯度；</li>
  <li>最后统一做梯度同步。</li>
</ul>

<p>而在 <code class="language-plaintext highlighter-rouge">ZeRO-2</code> 里，更像是：</p>

<ol>
  <li>某一层一旦反向完成，这一层的梯度就已经产生了。</li>
  <li>这时不等全模型其他层全部反向结束，就尽快对这一层梯度做 <code class="language-plaintext highlighter-rouge">reduce-scatter</code>。</li>
  <li>规约并切分完成后，本卡只保留自己负责的那一段梯度分片。</li>
  <li>然后继续上一层的反向。</li>
</ol>

<p>所以它是一个非常典型的 <strong>layer-by-layer 梯度规约与切分</strong> 的过程。</p>

<h3 id="reduce-scatter-是什么"><code class="language-plaintext highlighter-rouge">reduce-scatter</code> 是什么</h3>

<p><code class="language-plaintext highlighter-rouge">reduce-scatter</code> 可以理解成两步的融合：</p>

<ol>
  <li>先做 reduce：把各卡对应梯度做求和或求平均。</li>
  <li>再做 scatter：把规约后的完整结果按分片发给不同 rank。</li>
</ol>

<p>因此它很像：</p>

<p><code class="language-plaintext highlighter-rouge">all-reduce + scatter</code></p>

<p>但作为一个融合算子，它避免了先在每张卡上留下完整规约结果、再手动切分出去的中间步骤。</p>

<h3 id="反向后每张卡最终保留什么">反向后每张卡最终保留什么</h3>

<p>这是 <code class="language-plaintext highlighter-rouge">ZeRO-2</code> 最重要的结论之一：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-1</code> 反向结束后，每张卡保留的是<strong>完整梯度</strong>；</li>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-2</code> 反向结束后，每张卡保留的是<strong>自己负责参数分片对应的那一部分已规约梯度</strong>。</li>
</ul>

<p>也就是说，在 <code class="language-plaintext highlighter-rouge">ZeRO-2</code> 里，完整梯度不会像 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 那样长期驻留在每张卡显存里。</p>

<p>这就是它比 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 更省显存的核心原因。</p>

<h4 id="3-更新-1">3. 更新</h4>

<p>到了更新阶段，<code class="language-plaintext highlighter-rouge">ZeRO-2</code> 的逻辑反而比 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 更“对齐”：</p>

<ul>
  <li>optimizer states 本来就是按分片存的；</li>
  <li>梯度现在也已经按同样的 ownership 分片了；</li>
  <li>所以每个 rank 直接拿自己那部分梯度，更新自己那部分参数和 optimizer states。</li>
</ul>

<p>这一步不再需要“虽然梯度完整，但我只更新其中一部分”的逻辑解释，因为梯度本身也已经被切到对应 rank 了。</p>

<h4 id="4-更新后参数如何一致">4. 更新后参数如何一致</h4>

<p>由于参数在训练大部分时候仍按完整副本形式参与前向，所以更新结束后，通常仍需要把各 rank 更新后的参数分片重新同步成一致的完整参数副本，以便下一轮前向继续直接使用完整参数。</p>

<p>这里可以继续理解为依赖参数分片的 <code class="language-plaintext highlighter-rouge">all-gather</code> 重建。</p>

<h3 id="为什么-reduce-scatter-比最后留完整梯度更省显存">为什么 <code class="language-plaintext highlighter-rouge">reduce-scatter</code> 比“最后留完整梯度”更省显存</h3>

<p>因为它改变了梯度在显存中的驻留方式：</p>

<ul>
  <li>不是整网所有层的完整梯度一直堆在每张卡上；</li>
  <li>而是哪一层梯度算出来，就尽快规约并切成分片；</li>
  <li>本卡只留下自己需要负责更新的那一部分。</li>
</ul>

<p>所以显存里长期保留的梯度体积明显下降。</p>

<h3 id="相比-zero-1zero-2-到底改了什么">相比 ZeRO-1，ZeRO-2 到底改了什么</h3>

<p>相比 <code class="language-plaintext highlighter-rouge">ZeRO-1</code>，<code class="language-plaintext highlighter-rouge">ZeRO-2</code> 新增的本质变化就是：<br />
<strong>梯度不再完整复制，而是在反向过程中按 layer-by-layer 的 <code class="language-plaintext highlighter-rouge">reduce-scatter</code> 变成分片梯度。</strong></p>

<hr />

<h2 id="zero-3在-zero-2-基础上再切-parameters">ZeRO-3：在 ZeRO-2 基础上再切 parameters</h2>

<p><code class="language-plaintext highlighter-rouge">ZeRO-3</code> 的核心是：<strong>参数、梯度、optimizer states 三者全部分片。</strong></p>

<p>这也是为什么它显存节省最激进，但通信也最重。</p>

<h3 id="参数怎么切分">参数怎么切分</h3>

<p>在 <code class="language-plaintext highlighter-rouge">ZeRO-3</code> 里，参数平时不是“每张卡都有完整副本”，而是：</p>

<ul>
  <li>rank 0 保存一部分参数</li>
  <li>rank 1 保存另一部分参数</li>
  <li>…</li>
  <li>每张卡只长期保存自己拥有的参数分片</li>
</ul>

<p>所以 <code class="language-plaintext highlighter-rouge">ZeRO-3</code> 和 <code class="language-plaintext highlighter-rouge">ZeRO-2</code> 最大的区别在于：<br />
<strong>连前向所需的参数也不再常驻完整副本。</strong></p>

<h3 id="一次训练-step-的流程-3">一次训练 step 的流程</h3>

<h4 id="1-前向按层-all-gather-参数">1. 前向：按层 <code class="language-plaintext highlighter-rouge">all-gather</code> 参数</h4>

<p>因为每张卡平时只有参数分片，而前向计算某一层时通常需要该层完整权重，所以流程会变成：</p>

<ol>
  <li>当前层要做前向之前，各 rank 先把这层参数分片通过 <code class="language-plaintext highlighter-rouge">all-gather</code> 拼成临时完整参数。</li>
  <li>每张卡拿着这层完整参数做前向计算。</li>
  <li>这一层算完后，非本 rank 所拥有的那部分参数可以释放掉，只保留本地参数分片。</li>
</ol>

<p>所以 <code class="language-plaintext highlighter-rouge">ZeRO-3</code> 的前向不是“开局就有一整套完整模型”，而是“算哪一层，就临时 gather 哪一层”。</p>

<h4 id="2-反向先需要完整参数再做梯度-reduce-scatter">2. 反向：先需要完整参数，再做梯度 <code class="language-plaintext highlighter-rouge">reduce-scatter</code></h4>

<p>反向时同样要按层理解。</p>

<p>某一层反向时，需要该层参数参与梯度计算，因此通常也要保证这层的完整参数在计算时可见。<br />
完成该层反向后，这层产生的梯度再像 <code class="language-plaintext highlighter-rouge">ZeRO-2</code> 一样走 <code class="language-plaintext highlighter-rouge">reduce-scatter</code>：</p>

<ol>
  <li>该层反向得到梯度。</li>
  <li>对这层梯度做 <code class="language-plaintext highlighter-rouge">reduce-scatter</code>。</li>
  <li>每张卡最终只保留自己负责参数分片对应的那部分梯度。</li>
</ol>

<p>所以 <code class="language-plaintext highlighter-rouge">ZeRO-3</code> 的反向同时包含两类事情：</p>

<ul>
  <li>为了算这一层，需要临时拿到这层完整参数；</li>
  <li>为了节省梯度显存，梯度出来后又尽快分片规约掉。</li>
</ul>

<h4 id="3-更新-2">3. 更新</h4>

<p>到了更新阶段，逻辑和 <code class="language-plaintext highlighter-rouge">ZeRO-2</code> 一脉相承：</p>

<ul>
  <li>每个 rank 只持有自己的参数分片；</li>
  <li>只持有自己那部分梯度；</li>
  <li>只持有自己那部分 optimizer states；</li>
  <li>因此只更新自己的参数分片。</li>
</ul>

<h3 id="zero-3-的通信量是不是变大">ZeRO-3 的通信量是不是变大</h3>

<p>答案是：<strong>是的，通常会明显变大，而且会更频繁。</strong></p>

<p>原因不是“梯度通信突然变得神秘地更多了”，而是参数本身也进入了按层通信流程：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-2</code> 主要增加的是梯度 <code class="language-plaintext highlighter-rouge">reduce-scatter</code>；</li>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-3</code> 则在前向和反向两边，都要频繁为当前层参数做 <code class="language-plaintext highlighter-rouge">all-gather</code>。</li>
</ul>

<p>这会带来两个后果：</p>

<ol>
  <li>通信从更粗粒度的 step 级，同步成了更细粒度的 layer 级。</li>
  <li>参数 gather 和梯度 scatter 都要和计算过程紧密交织。</li>
</ol>

<p>所以 <code class="language-plaintext highlighter-rouge">ZeRO-3</code> 通常对：</p>

<ul>
  <li>GPU 间互联带宽</li>
  <li>通信与计算重叠 <code class="language-plaintext highlighter-rouge">overlap</code></li>
  <li>bucket 化和 prefetch</li>
</ul>

<p>会更加敏感。</p>

<h3 id="为什么说-zero-3-最省显存">为什么说 ZeRO-3 最省显存</h3>

<p>因为三类主要训练状态都不再完整复制：</p>

<ul>
  <li>parameters 分片</li>
  <li>gradients 分片</li>
  <li>optimizer states 分片</li>
</ul>

<p>这使得每张卡长期持有的 model states 最少。<br />
代价就是：为了把“长期常驻显存”压到最低，必须接受更重的动态通信。</p>

<h3 id="相比-zero-2zero-3-到底改了什么">相比 ZeRO-2，ZeRO-3 到底改了什么</h3>

<p>相比 <code class="language-plaintext highlighter-rouge">ZeRO-2</code>，<code class="language-plaintext highlighter-rouge">ZeRO-3</code> 的本质新增变化就是：<br />
<strong>参数也从“完整副本常驻”变成了“平时按分片保存，需要某层计算时再临时 <code class="language-plaintext highlighter-rouge">all-gather</code>”。</strong></p>

<hr />

<h2 id="最后用一句话串起来">最后用一句话串起来</h2>

<p>如果把三种 ZeRO 放成一条连续演进路线，可以这样记：</p>

<ol>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-1</code>：参数和梯度还是完整的，先把 optimizer states 从复制改成分片。</li>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-2</code>：在 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 基础上，把梯度也从复制改成分片，关键算子是 layer-by-layer 的 <code class="language-plaintext highlighter-rouge">reduce-scatter</code>。</li>
  <li><code class="language-plaintext highlighter-rouge">ZeRO-3</code>：在 <code class="language-plaintext highlighter-rouge">ZeRO-2</code> 基础上，把参数也从复制改成分片，关键代价是前向和反向都要更频繁地做参数 <code class="language-plaintext highlighter-rouge">all-gather</code>。</li>
</ol>

<p>所以 ZeRO 的本质，不是某一个神奇算子，而是一套越来越激进的状态分片策略：</p>

<ul>
  <li>先切 optimizer states</li>
  <li>再切 gradients</li>
  <li>最后切 parameters</li>
</ul>

<p>显存越省，通信越重，这就是 <code class="language-plaintext highlighter-rouge">ZeRO-1</code> 到 <code class="language-plaintext highlighter-rouge">ZeRO-3</code> 最核心的主线。</p>

<hr />

<h2 id="补充">补充</h2>

<p>这篇笔记先聚焦 <code class="language-plaintext highlighter-rouge">ZeRO-1 / ZeRO-2 / ZeRO-3</code> 的训练流程本身。<br />
像 <code class="language-plaintext highlighter-rouge">ZeRO-Offload</code>、<code class="language-plaintext highlighter-rouge">ZeRO-Infinity</code>、和 tensor parallel / pipeline parallel 的组合方式，这里先不展开，后面如果需要可以单独再开一篇。</p>]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[ZERO训练技术]]></summary></entry><entry><title type="html">强化学习笔记</title><link href="https://wqh011128.github.io/blog/2026/04/23/RL.html" rel="alternate" type="text/html" title="强化学习笔记" /><published>2026-04-23T00:00:00+00:00</published><updated>2026-04-23T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/04/23/RL</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/04/23/RL.html"><![CDATA[<h1 id="强化学习初学者速通文档">强化学习初学者速通文档</h1>

<h2 id="执行摘要">执行摘要</h2>

<p>如果你把这份文档只读一遍，最应该建立的认知有四条。第一，<strong>RLHF 不是一个单一算法，而是一条训练管线</strong>：先做 SFT，让模型学会“基本回答”；再收集偏好比较数据训练奖励模型；最后用 PPO 之类的策略优化算法，让模型在不偏离参考模型太远的前提下提高“人类更喜欢”的回答概率。OpenAI 的 InstructGPT 就是这条路线的经典代表，而且论文报告称，<strong>1.3B 的 InstructGPT 在人工偏好评测中优于 175B 的原始 GPT-3</strong>。</p>

<p>第二，<strong>PPO 是“稳着改策略”</strong>。它本质上是策略梯度方法的工程化强化版：先根据旧策略采样，再用一个“剪切后的替代目标”更新新策略，防止一次更新迈得太大。PPO 很通用，也很成熟，但在 LLM 对齐里通常还要配合参考模型、奖励模型、价值模型、KL 约束和大量训练细节，所以它常被认为“强但重”。</p>

<p>第三，<strong>DPO 是“把 RLHF 里的一部分 RL 数学化简掉”</strong>。在 Bradley–Terry 偏好模型和 KL 约束设定下，DPO 论文把“奖励模型 + RL”改写成了一个直接作用于策略的分类损失，因此<strong>不需要显式奖励模型，也不需要在线采样回路</strong>；论文与官方实现都强调其训练更简单、更稳定。对初学者来说，DPO 往往是最容易真正跑起来、也最容易从 loss 曲线和 chosen/rejected 对数概率中建立直觉的方法。</p>

<p>第四，<strong>GRPO 是“去掉 critic 的 PPO 亲戚”</strong>。它在 DeepSeekMath 中被正式提出，并在 DeepSeek-R1 系列训练中被采用。核心做法是：对同一个 prompt 生成一组回答，用组内相对奖励构造 advantage，而不是单独训练一个价值模型；因此它尤其适合<strong>有可验证奖励</strong>的任务，如数学、代码、规则格式约束。但你也要知道，GRPO 目前<strong>并不是一个像 PPO 那样完全定型、单一版本的算法</strong>：原始 DeepSeekMath 公式、DeepSeek-R1 实践和 Hugging Face TRL 当前实现之间，已经出现了对 KL 是否启用、奖励是否按组标准差归一化、损失是否做长度修正等细节分化。初学者看到不同文章写法不一致，不一定是谁错了，更可能是“同一家族的不同工程变体”。</p>

<h2 id="阅读前的最小背景与统一符号">阅读前的最小背景与统一符号</h2>

<p>为了把 RLHF、PPO、DPO、GRPO 放在同一个坐标系里，我们先统一最常见的符号。对大语言模型而言，状态通常不再写成经典 RL 的 $s_t$，而更常写成“<strong>prompt + 已生成前缀</strong>”；策略 $\pi_\theta$ 表示模型在当前上下文下生成下一个 token 的条件概率；$\pi_{\text{ref}}$ 是参考策略，通常来自 SFT 模型；$r$ 是奖励，可能来自奖励模型，也可能来自规则验证器；$V$ 是价值函数，只在 PPO/Actor-Critic 一类方法里需要；$A$ 是 advantage，表示“这个动作比当前平均水平好多少”。这些写法在 PPO、DPO、GRPO 的原始论文与 OpenAI Spinning Up 的策略优化教程中是一致的，只是 LLM 场景把“动作”换成了 token 序列。</p>

<table>
  <thead>
    <tr>
      <th>符号</th>
      <th>含义</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>$x$</td>
      <td>输入 prompt</td>
    </tr>
    <tr>
      <td>$y$</td>
      <td>一条完整回答</td>
    </tr>
    <tr>
      <td>$y^+, y^-$</td>
      <td>同一 prompt 下的优选回答 / 弱选回答</td>
    </tr>
    <tr>
      <td>$\pi_\theta(y\mid x)$</td>
      <td>当前策略模型</td>
    </tr>
    <tr>
      <td>$\pi_{\text{ref}}(y\mid x)$</td>
      <td>参考模型，通常是 SFT 模型</td>
    </tr>
    <tr>
      <td>$r_\phi(x,y)$</td>
      <td>奖励模型给出的分数</td>
    </tr>
    <tr>
      <td>$V_\psi$</td>
      <td>价值模型</td>
    </tr>
    <tr>
      <td>$A$</td>
      <td>advantage，表示相对优劣</td>
    </tr>
    <tr>
      <td>$\beta$</td>
      <td>控制与参考模型偏离程度的系数</td>
    </tr>
    <tr>
      <td>$\epsilon$</td>
      <td>PPO/GRPO 中剪切范围</td>
    </tr>
    <tr>
      <td>$G$</td>
      <td>GRPO 中每个 prompt 采样的候选回答数</td>
    </tr>
  </tbody>
</table>

<p>把四个概念放进一张图里，初学者最需要记住的是：<strong>RLHF 是总流程，PPO 是流程中的一个优化器；DPO 是把这个流程中的一部分目标直接改写；GRPO 则是另一种在线策略优化路线，重点是用“组内相对优势”替代 critic。</strong></p>

<pre><code class="language-mermaid">flowchart LR
    A[预训练模型] --&gt; B[SFT 监督微调]
    B --&gt; C[采样多个回答]
    C --&gt; D[人类或规则给出偏好/分数]

    D --&gt; E[奖励模型 RM]
    E --&gt; F[PPO 式在线优化]
    B --&gt; F
    F --&gt; G[经典 RLHF 模型]

    D --&gt; H[DPO 直接偏好优化]
    B --&gt; H
    H --&gt; I[DPO 模型]

    D --&gt; J[GRPO 组内相对优势]
    B --&gt; J
    J --&gt; K[GRPO 模型]
</code></pre>

<p>上图对应的是行业里最常见的理解方式：经典 RLHF 走“<strong>SFT → 偏好比较 → RM → PPO</strong>”，DPO 走“<strong>SFT → 偏好比较 → 直接优化策略</strong>”，GRPO 更偏“<strong>SFT/基座模型 → 在线采样多个回答 → 组内相对奖励 → 更新策略</strong>”。其中 InstructGPT 和 TL;DR 总结工作展示了经典 RLHF 管线，DPO 论文给出了直接偏好优化的闭式改写，DeepSeekMath 与 DeepSeek-R1 则展示了 GRPO 在推理任务上的用法。</p>

<h2 id="rlhf">RLHF</h2>

<h3 id="通俗直观解释">通俗直观解释</h3>

<p>RLHF 可以理解成“<strong>先让模型学会说话，再让人类教它什么叫说得更好</strong>”。只做预训练时，模型学到的是“互联网上像人类一样继续写下去”；但这不等于“遵从指令、诚实、无害、对用户有帮助”。因此 OpenAI 在 InstructGPT 里采用了三步法：先做人类演示的 SFT，再做人类比较偏好的奖励模型，最后用 PPO 优化策略。OpenAI 在论文和博客中都把这条路线当作把“语言建模目标”改成“更符合用户意图目标”的关键办法。</p>

<p>这条路线为什么有效？因为在很多真实任务里，<strong>“哪个回答更好”比“正确答案唯一是什么”更容易标注</strong>。比如总结、对话、写作风格、安全拒答，标注者往往更擅长在两个回答里选一个更好的，而不一定能直接写出最优答案。OpenAI 在 TL;DR 总结工作中收集了大规模摘要比较数据，训练奖励模型，再用 RL 去最大化该奖励，结果是模型在人工偏好评价中明显优于纯监督学习基线，而且跨域迁移到 CNN/DM 新闻总结时也能保持较强效果。</p>

<h3 id="关键数学推导">关键数学推导</h3>

<p>RLHF 的核心通常分成两层数学对象：一个是<strong>偏好模型</strong>，一个是<strong>带 KL 约束的策略优化目标</strong>。最常见的偏好建模方式是 Bradley–Terry 形式：给定同一 prompt $x$ 下两个回答 $y^+$ 和 $y^-$，人类更喜欢 $y^+$ 的概率由两者奖励差决定：</p>

\[P(y^+ \succ y^- \mid x)=\sigma\big(r_\phi(x,y^+) - r_\phi(x,y^-)\big)\]

<p>于是奖励模型的训练就是一个二分类最大似然问题：</p>

\[\mathcal L_{\text{RM}}(\phi)=
-\mathbb E_{(x,y^+,y^-)\sim \mathcal D}
\left[
\log \sigma\big(r_\phi(x,y^+) - r_\phi(x,y^-)\big)
\right]\]

<p>这一步的假设是：人类偏好可以被一个潜在标量奖励函数近似，而且偏好数据主要以成对比较形式出现。这个写法在 DPO 论文回顾 RLHF 管线时写得非常清楚，也与 TL;DR 总结和 InstructGPT 采用的“比较数据 → 奖励模型”路线一致。</p>

<p>有了奖励模型后，经典 RLHF 的策略优化目标通常写成：</p>

\[\max_{\pi_\theta}\;
\mathbb E_{x \sim \mathcal D,\; y \sim \pi_\theta(\cdot\mid x)}
\big[r_\phi(x,y)\big]
\;-\;
\beta\, D_{\mathrm{KL}}\!\left(\pi_\theta(\cdot\mid x)\,\|\,\pi_{\text{ref}}(\cdot\mid x)\right)\]

<p>第一项是“追求更高奖励”，第二项是“别离参考模型太远”。初学者最容易漏掉的是第二项的重要性：<strong>如果没有 KL 约束，模型很容易为了讨好奖励模型而走到奖励模型并不可靠的区域，出现 reward hacking、模式崩塌或者语言质量劣化。</strong> DPO 论文把这个目标明确写为 prior RLHF 的标准形式；InstructGPT 和后续 TRL/PPO 实现则把它具体化成带参考模型的 PPO 训练。</p>

<p>从工程视角看，RLHF 与普通监督学习最大的不同是：<strong>训练目标取决于当前模型的采样结果</strong>。这意味着训练是“闭环”的，而不是固定数据上的静态拟合。也正因为如此，PPO 和 GRPO 这类在线方法会比 DPO 更重，但也更灵活。</p>

<h3 id="伪代码">伪代码</h3>

<p>下面这段伪代码概括的是经典 RLHF 管线，而不是任何一家公司的逐字实现。流程结构与 InstructGPT、TL;DR summarization 和 TRL/OpenRLHF 常见实现一致。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># Classical RLHF pipeline
</span>
<span class="c1"># Step 1: SFT
</span><span class="n">policy</span> <span class="o">=</span> <span class="n">init_from_pretrained_lm</span><span class="p">()</span>
<span class="n">policy</span> <span class="o">=</span> <span class="n">supervised_finetune</span><span class="p">(</span><span class="n">policy</span><span class="p">,</span> <span class="n">demonstration_dataset</span><span class="p">)</span>

<span class="c1"># Step 2: Reward Model
</span><span class="n">preference_pairs</span> <span class="o">=</span> <span class="n">collect_pairs</span><span class="p">(</span><span class="n">policy</span><span class="p">,</span> <span class="n">prompt_dataset</span><span class="p">)</span>   <span class="c1"># (x, y_plus, y_minus)
</span><span class="n">reward_model</span> <span class="o">=</span> <span class="n">init_reward_model</span><span class="p">(</span><span class="n">policy</span><span class="p">)</span>
<span class="k">for</span> <span class="n">batch</span> <span class="ow">in</span> <span class="n">preference_pairs</span><span class="p">:</span>
    <span class="n">loss_rm</span> <span class="o">=</span> <span class="o">-</span><span class="n">log_sigmoid</span><span class="p">(</span><span class="n">reward_model</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">y_plus</span><span class="p">)</span> <span class="o">-</span> <span class="n">reward_model</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">y_minus</span><span class="p">)).</span><span class="n">mean</span><span class="p">()</span>
    <span class="n">update</span><span class="p">(</span><span class="n">reward_model</span><span class="p">,</span> <span class="n">loss_rm</span><span class="p">)</span>

<span class="c1"># Step 3: RL optimization
</span><span class="n">ref_policy</span> <span class="o">=</span> <span class="n">copy</span><span class="p">(</span><span class="n">policy</span><span class="p">)</span>   <span class="c1"># frozen reference
</span><span class="n">value_model</span> <span class="o">=</span> <span class="n">init_value_model</span><span class="p">()</span>   <span class="c1"># PPO-style setups usually need it
</span><span class="k">for</span> <span class="n">iteration</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_iters</span><span class="p">):</span>
    <span class="n">trajectories</span> <span class="o">=</span> <span class="n">sample_responses</span><span class="p">(</span><span class="n">policy</span><span class="p">,</span> <span class="n">prompt_dataset</span><span class="p">)</span>
    <span class="n">rewards</span> <span class="o">=</span> <span class="n">reward_model</span><span class="p">.</span><span class="n">score</span><span class="p">(</span><span class="n">trajectories</span><span class="p">)</span> <span class="o">-</span> <span class="n">beta</span> <span class="o">*</span> <span class="n">kl</span><span class="p">(</span><span class="n">policy</span><span class="p">,</span> <span class="n">ref_policy</span><span class="p">)</span>
    <span class="n">advantages</span> <span class="o">=</span> <span class="n">estimate_advantages</span><span class="p">(</span><span class="n">rewards</span><span class="p">,</span> <span class="n">value_model</span><span class="p">)</span>
    <span class="n">update_policy_with_ppo</span><span class="p">(</span><span class="n">policy</span><span class="p">,</span> <span class="n">value_model</span><span class="p">,</span> <span class="n">trajectories</span><span class="p">,</span> <span class="n">advantages</span><span class="p">)</span>
</code></pre></div></div>

<h3 id="简短-python-示例片段">简短 Python 示例片段</h3>

<p>下面的片段只用于说明 <strong>TRL 当前 PPOTrainer 需要哪些核心对象</strong>。它仍然不是可直接运行的完整脚本，但字段名已经尽量对齐当前文档，避免把 <code class="language-plaintext highlighter-rouge">PPOTrainer</code> 误解成只要 <code class="language-plaintext highlighter-rouge">policy/ref/reward</code> 三件套就能启动。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 仅作结构示意：真实训练还需要数据预处理、生成配置、分布式训练等
</span><span class="kn">from</span> <span class="nn">transformers</span> <span class="kn">import</span> <span class="p">(</span>
    <span class="n">AutoModelForCausalLM</span><span class="p">,</span>
    <span class="n">AutoModelForSequenceClassification</span><span class="p">,</span>
    <span class="n">AutoTokenizer</span><span class="p">,</span>
<span class="p">)</span>
<span class="kn">from</span> <span class="nn">trl</span> <span class="kn">import</span> <span class="n">PPOTrainer</span><span class="p">,</span> <span class="n">PPOConfig</span>

<span class="n">policy</span> <span class="o">=</span> <span class="n">AutoModelForCausalLM</span><span class="p">.</span><span class="n">from_pretrained</span><span class="p">(</span><span class="s">"your-sft-model"</span><span class="p">)</span>
<span class="n">ref_policy</span> <span class="o">=</span> <span class="n">AutoModelForCausalLM</span><span class="p">.</span><span class="n">from_pretrained</span><span class="p">(</span><span class="s">"your-sft-model"</span><span class="p">)</span>
<span class="n">reward_model</span> <span class="o">=</span> <span class="n">AutoModelForSequenceClassification</span><span class="p">.</span><span class="n">from_pretrained</span><span class="p">(</span><span class="s">"your-rm"</span><span class="p">)</span>
<span class="n">value_model</span> <span class="o">=</span> <span class="n">AutoModelForSequenceClassification</span><span class="p">.</span><span class="n">from_pretrained</span><span class="p">(</span><span class="s">"your-value-model"</span><span class="p">)</span>
<span class="n">tokenizer</span> <span class="o">=</span> <span class="n">AutoTokenizer</span><span class="p">.</span><span class="n">from_pretrained</span><span class="p">(</span><span class="s">"your-sft-model"</span><span class="p">)</span>
<span class="n">train_dataset</span> <span class="o">=</span> <span class="n">load_your_prompt_dataset</span><span class="p">()</span>

<span class="n">args</span> <span class="o">=</span> <span class="n">PPOConfig</span><span class="p">(</span>
    <span class="n">learning_rate</span><span class="o">=</span><span class="mf">3e-6</span><span class="p">,</span>
    <span class="n">cliprange</span><span class="o">=</span><span class="mf">0.2</span><span class="p">,</span>
    <span class="n">kl_coef</span><span class="o">=</span><span class="mf">0.05</span><span class="p">,</span>
    <span class="n">num_ppo_epochs</span><span class="o">=</span><span class="mi">4</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">trainer</span> <span class="o">=</span> <span class="n">PPOTrainer</span><span class="p">(</span>
    <span class="n">args</span><span class="o">=</span><span class="n">args</span><span class="p">,</span>
    <span class="n">processing_class</span><span class="o">=</span><span class="n">tokenizer</span><span class="p">,</span>
    <span class="n">model</span><span class="o">=</span><span class="n">policy</span><span class="p">,</span>
    <span class="n">ref_model</span><span class="o">=</span><span class="n">ref_policy</span><span class="p">,</span>
    <span class="n">reward_model</span><span class="o">=</span><span class="n">reward_model</span><span class="p">,</span>
    <span class="n">value_model</span><span class="o">=</span><span class="n">value_model</span><span class="p">,</span>
    <span class="n">train_dataset</span><span class="o">=</span><span class="n">train_dataset</span><span class="p">,</span>
<span class="p">)</span>

<span class="c1"># 训练时的闭环仍然是：prompt -&gt; generate -&gt; reward -&gt; advantage/value -&gt; PPO update
</span></code></pre></div></div>

<h3 id="实践要点与超参数建议">实践要点与超参数建议</h3>

<p>对初学者最有价值的经验不是“某个神奇超参”，而是<strong>先把三阶段的职责分清</strong>。SFT 负责把分布收窄到“像样回答”的区域；奖励模型负责把“更好/更差”映射成标量；PPO 负责在 KL 约束下移动策略。如果你的 SFT 很差，后续 RM 和 PPO 往往也不会好，因为比较数据和采样数据都会变脏。OpenAI 的 InstructGPT 和 TRL 的 PPO 文档都把 SFT 作为 RLHF 的前置基础。</p>

<p>在工程细节上，Hugging Face 对 OpenAI 早期 RLHF 代码的复现实验总结出一批非常“接地气”的经验：<strong>关闭 dropout、奖励/值函数归一化、学习率退火、奖励白化、adaptive KL、必要时做 rejection sampling</strong>，这些都显著影响稳定性；而且他们在复现实验中还特别指出，PPO 训练里 PyTorch/TF 的 Adam 数值行为差异都可能导致更新激进。对初学者来说，这意味着：<strong>不要低估实现细节，RLHF 不是“把 PPO 套上去”就完事。</strong></p>

<p>如果你想做一个“能跑通”的入门版本，我的建议是：先不用追求 InstructGPT 规模，优先选择公开框架与小模型。奖励模型可以从序列分类头开始；如果数据很少，优先防止 RM 过拟合；PPO 阶段先严格盯住 <code class="language-plaintext highlighter-rouge">objective/rlhf_reward</code>、<code class="language-plaintext highlighter-rouge">objective/kl</code>、<code class="language-plaintext highlighter-rouge">val/ratio</code> 等指标，因为这些指标能直接告诉你训练是不是“往前走但没走飞”。TRL 文档特别提示，<code class="language-plaintext highlighter-rouge">val/ratio</code> 应围绕 1 左右波动，过大或过小都说明相邻策略之间变化过猛。</p>

<h3 id="典型实现与在线文档">典型实现与在线文档</h3>

<p>就“官方材料”而言，InstructGPT 的<strong>完整训练代码公开程度是未指定</strong>；但 OpenAI 已公开论文、博客、模型卡/评测仓库，以及更早的 <code class="language-plaintext highlighter-rouge">lm-human-preferences</code> 和 <code class="language-plaintext highlighter-rouge">summarize-from-feedback</code> 代码资源作为近似参考。对学习者而言，这些足以建立 RLHF 管线认知。</p>

<p>比较值得直接阅读的材料有：OpenAI 的 <strong>InstructGPT 论文与博客</strong>、OpenAI 的 <strong>Learning to Summarize from Human Feedback</strong> 论文与仓库、OpenAI 的 <strong>lm-human-preferences</strong> 代码库、Hugging Face 的 <strong>TRL Reward Modeling / PPOTrainer</strong> 文档、以及 <strong>OpenRLHF</strong> 这类偏工程化的高性能开源框架。中文方面，优先看 <strong>Hugging Face 中文博客/文档</strong>、<strong>OpenRLHF 中文文档</strong> 和 <strong>Spinning Up 中文版</strong>，因为它们更接近原始英文材料，而不是二次转述。</p>

<h3 id="常见问题与调试建议">常见问题与调试建议</h3>

<p>RLHF 最常见的问题，不是 loss 不下降，而是<strong>你根本不知道下降的是什么</strong>。奖励模型分数升高，不等于人类偏好一定同步升高；KL 很低，不等于模型真的学到了更好的行为；value loss 稳定，也不等于 policy 真在变好。这也是 OpenAI 后续研究 CriticGPT 的原因之一：随着模型越来越强，<strong>人类评审越来越难稳定地发现错误</strong>，而 RLHF 的标签质量会因此成为瓶颈。</p>

<p>如果你跑 PPO 式 RLHF，先排查四件事。第一，看回答是否很快塌缩成模板化短句；第二，看 KL 是否飞升；第三，看 <code class="language-plaintext highlighter-rouge">val/ratio</code> 是否经常远离 1；第四，看 RM 是否在训练集上好、在验证集上差。如果出现这些问题，优先从<strong>减小学习率、增大 KL 约束、缩短回复长度、检查 EOS/截断规则、复查 chosen/rejected 数据质量</strong>入手，而不是立刻改模型结构。</p>

<h3 id="小结与适用场景">小结与适用场景</h3>

<p>如果你想理解“现代大模型对齐最经典的一条路”，RLHF 必学；如果你想做一个<strong>工业上可解释、可扩展、可插入任意奖励函数</strong>的系统，RLHF 仍然很重要；但如果你只是想快速从偏好数据起步，RLHF 往往不是最容易的第一站，因为它的工程复杂度明显高于 DPO。这个判断并不是说 RLHF 过时，而是说它更像“全家桶”，适合在你已经会 SFT、懂 PPO、知道 RM 风险之后再系统掌握。</p>

<h2 id="ppo">PPO</h2>

<h3 id="通俗直观解释-1">通俗直观解释</h3>

<p>PPO 最容易理解的方式是：<strong>每次都朝“更高奖励”的方向走，但只允许走一小步安全步长</strong>。如果不加限制，策略梯度会鼓励模型把某些高回报动作的概率一路抬高，结果常常是“一步迈太大”，训练发散。TRPO 用 KL 约束来硬性限制步长，PPO 则用更简单的剪切目标来近似这种“别走太远”的思想，所以它在工程上更易实现。OpenAI 的 PPO 原始论文和 Spinning Up 教程都把它定位为在简单性、样本效率和稳定性之间取得平衡的方法。</p>

<p>放到 LLM 对齐里，PPO 的角色就变成：模型先根据当前策略回答一批 prompt，再根据奖励模型给分；如果某些回答分高，就提高这条回答路径上 token 的概率；但提高的幅度要被剪切项和 KL 项约束住，不然模型会迅速偏离语言质量良好的区域。DeepSeekMath 在介绍 GRPO 时，也先把 PPO 作为“当前 LLM RL 微调中广泛使用的 actor-critic 算法”来对照。</p>

<h3 id="关键数学推导-1">关键数学推导</h3>

<p>从策略梯度出发，我们希望最大化策略的期望回报 $J(\theta)$。在旧策略 $\pi_{\theta_{\text{old}}}$ 采样的数据上，可以把新旧策略的差异写成一个比值：</p>

\[r_t(\theta)=\frac{\pi_\theta(a_t\mid s_t)}{\pi_{\theta_{\text{old}}}(a_t\mid s_t)}\]

<p>于是得到经典的替代目标：</p>

\[L^{\text{CPI}}(\theta)=\hat{\mathbb E}_t\big[r_t(\theta)\hat A_t\big]\]

<p>这里 $\hat A_t$ 是 advantage，它表示在状态 $s_t$ 下选动作 $a_t$ 比平均动作好多少。这个目标的问题是：如果直接最大化，策略可能在少数状态上被推得太远。PPO 的关键改写就是加入剪切：</p>

\[L^{\text{CLIP}}(\theta)=
\hat{\mathbb E}_t
\Big[
\min\big(
r_t(\theta)\hat A_t,\;
\operatorname{clip}(r_t(\theta),1-\epsilon,1+\epsilon)\hat A_t
\big)
\Big]\]

<p>这就是 PPO 最核心的公式。论文给出的直觉是：如果 $\hat A_t&gt;0$，那就不希望 $r_t$ 大幅高于 $1+\epsilon$；如果 $\hat A_t&lt;0$，那就不希望 $r_t$ 大幅低于 $1-\epsilon$。这样做的结果不是“完全不让策略变”，而是“超过安全区后，继续变大对目标不再有额外好处”。</p>

<p>在 LLM RLHF 中，PPO 还经常配一个价值函数 $V_\psi$ 来降低方差，并配一个参考模型 KL 惩罚避免策略漂移。DeepSeekMath 对 PPO 的回顾特别指出，PPO 在 LLM 场景里通常需要训练 value function，并在 token 级奖励中加入来自参考模型的 KL 惩罚；这恰好也是后来 GRPO 要“去 critic”的原因。</p>

<h3 id="伪代码-1">伪代码</h3>

<p>下面的伪代码是 PPO 的最小骨架，结构上与 Schulman 的原始算法和 TRL 的 PPOTrainer 都一致。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">initialize</span> <span class="n">policy</span> <span class="n">πθ</span>
<span class="n">initialize</span> <span class="n">value</span> <span class="n">function</span> <span class="n">Vψ</span>

<span class="k">for</span> <span class="n">each</span> <span class="n">iteration</span><span class="p">:</span>
    <span class="n">rollouts</span> <span class="o">=</span> <span class="n">collect_trajectories</span><span class="p">(</span><span class="n">πθ_old</span><span class="p">)</span>
    <span class="n">rewards</span> <span class="o">=</span> <span class="n">compute_rewards</span><span class="p">(</span><span class="n">rollouts</span><span class="p">)</span>
    <span class="n">advantages</span> <span class="o">=</span> <span class="n">estimate_advantages</span><span class="p">(</span><span class="n">rewards</span><span class="p">,</span> <span class="n">Vψ</span><span class="p">)</span>

    <span class="k">for</span> <span class="n">epoch</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">K</span><span class="p">):</span>
        <span class="n">ratio</span> <span class="o">=</span> <span class="n">πθ</span><span class="p">(</span><span class="n">a</span><span class="o">|</span><span class="n">s</span><span class="p">)</span> <span class="o">/</span> <span class="n">πθ_old</span><span class="p">(</span><span class="n">a</span><span class="o">|</span><span class="n">s</span><span class="p">)</span>
        <span class="n">clipped_ratio</span> <span class="o">=</span> <span class="n">clip</span><span class="p">(</span><span class="n">ratio</span><span class="p">,</span> <span class="mi">1</span> <span class="o">-</span> <span class="n">eps</span><span class="p">,</span> <span class="mi">1</span> <span class="o">+</span> <span class="n">eps</span><span class="p">)</span>
        <span class="n">policy_loss</span> <span class="o">=</span> <span class="o">-</span><span class="n">mean</span><span class="p">(</span><span class="nb">min</span><span class="p">(</span><span class="n">ratio</span> <span class="o">*</span> <span class="n">advantages</span><span class="p">,</span>
                                <span class="n">clipped_ratio</span> <span class="o">*</span> <span class="n">advantages</span><span class="p">))</span>
        <span class="n">value_loss</span> <span class="o">=</span> <span class="n">mse</span><span class="p">(</span><span class="n">Vψ</span><span class="p">(</span><span class="n">s</span><span class="p">),</span> <span class="n">returns</span><span class="p">)</span>
        <span class="n">loss</span> <span class="o">=</span> <span class="n">policy_loss</span> <span class="o">+</span> <span class="n">c_v</span> <span class="o">*</span> <span class="n">value_loss</span> <span class="o">-</span> <span class="n">c_e</span> <span class="o">*</span> <span class="n">entropy</span><span class="p">(</span><span class="n">πθ</span><span class="p">)</span>
        <span class="n">update</span><span class="p">(</span><span class="n">θ</span><span class="p">,</span> <span class="n">ψ</span><span class="p">)</span>
</code></pre></div></div>

<h3 id="简短-python-示例片段-1">简短 Python 示例片段</h3>

<p>下面的片段强调的是 <strong>PPO 的对象依赖关系</strong>：当前策略、参考策略、奖励模型、价值模型、prompt 数据集和 tokenizer/processor 缺一不可。它是结构示意，不是最小可运行脚本。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="nn">transformers</span> <span class="kn">import</span> <span class="n">AutoTokenizer</span>
<span class="kn">from</span> <span class="nn">trl</span> <span class="kn">import</span> <span class="n">PPOConfig</span><span class="p">,</span> <span class="n">PPOTrainer</span>

<span class="n">tokenizer</span> <span class="o">=</span> <span class="n">AutoTokenizer</span><span class="p">.</span><span class="n">from_pretrained</span><span class="p">(</span><span class="s">"your-sft-model"</span><span class="p">)</span>
<span class="n">args</span> <span class="o">=</span> <span class="n">PPOConfig</span><span class="p">(</span>
    <span class="n">learning_rate</span><span class="o">=</span><span class="mf">3e-6</span><span class="p">,</span>
    <span class="n">num_ppo_epochs</span><span class="o">=</span><span class="mi">4</span><span class="p">,</span>
    <span class="n">cliprange</span><span class="o">=</span><span class="mf">0.2</span><span class="p">,</span>
    <span class="n">vf_coef</span><span class="o">=</span><span class="mf">0.1</span><span class="p">,</span>
    <span class="n">kl_coef</span><span class="o">=</span><span class="mf">0.05</span><span class="p">,</span>
    <span class="n">gamma</span><span class="o">=</span><span class="mf">1.0</span><span class="p">,</span>
    <span class="n">lam</span><span class="o">=</span><span class="mf">0.95</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">trainer</span> <span class="o">=</span> <span class="n">PPOTrainer</span><span class="p">(</span>
    <span class="n">args</span><span class="o">=</span><span class="n">args</span><span class="p">,</span>
    <span class="n">processing_class</span><span class="o">=</span><span class="n">tokenizer</span><span class="p">,</span>
    <span class="n">model</span><span class="o">=</span><span class="n">policy_model</span><span class="p">,</span>
    <span class="n">ref_model</span><span class="o">=</span><span class="n">ref_model</span><span class="p">,</span>
    <span class="n">reward_model</span><span class="o">=</span><span class="n">reward_model</span><span class="p">,</span>
    <span class="n">value_model</span><span class="o">=</span><span class="n">value_model</span><span class="p">,</span>
    <span class="n">train_dataset</span><span class="o">=</span><span class="n">prompt_dataset</span><span class="p">,</span>
<span class="p">)</span>

<span class="c1"># prompts -&gt; generate -&gt; reward -&gt; value/advantage -&gt; PPO update
</span></code></pre></div></div>

<h3 id="实践要点与超参数建议-1">实践要点与超参数建议</h3>

<p>对初学者最友好的 PPO 起点，往往不是原论文的通用设置，而是现成的 LLM 实现默认值。以 TRL 为例，PPOConfig 的关键默认值包括：<code class="language-plaintext highlighter-rouge">cliprange=0.2</code>、<code class="language-plaintext highlighter-rouge">vf_coef=0.1</code>、<code class="language-plaintext highlighter-rouge">cliprange_value=0.2</code>、<code class="language-plaintext highlighter-rouge">gamma=1.0</code>、<code class="language-plaintext highlighter-rouge">lam=0.95</code>、<code class="language-plaintext highlighter-rouge">kl_coef=0.05</code>、<code class="language-plaintext highlighter-rouge">num_ppo_epochs=4</code>。这组默认值反映了 LLM RLHF 的典型做法：<strong>把整段回答的质量看得比传统逐步折扣更重要，因此常把 $\gamma$ 设成 1；同时用 KL 和 value loss 控制训练稳定性。</strong></p>

<p>如果你要调 PPO，优先盯三件事。第一，<code class="language-plaintext highlighter-rouge">val/ratio</code> 是否明显脱离 1；第二，<code class="language-plaintext highlighter-rouge">objective/kl</code> 是否持续上冲；第三，<code class="language-plaintext highlighter-rouge">policy/clipfrac_avg</code> 是否过高。TRL 官方文档直接给了调试建议：<code class="language-plaintext highlighter-rouge">val/ratio</code> 应在 1 附近波动，若到 2、1000 或掉到 0.1 这类量级，说明连续两次策略更新差异太大，应先回头检查学习率、奖励尺度、KL 系数和 batch 构造。</p>

<p>再往前走一步，Hugging Face 对 OpenAI 早期 RLHF 的复现材料还给出一批很实用的 PPO 细节：训练中通常<strong>关闭 dropout</strong>；为 reward 与 policy 做<strong>学习率退火</strong>；在某些任务中做<strong>reward whitening</strong>；并通过 <strong>adaptive KL</strong> 动态调节 KL 系数。这些做法不一定对所有任务都最优，但对“为什么我的 PPO 跑不稳”这个初学者问题非常关键。</p>

<h3 id="典型实现与在线文档-1">典型实现与在线文档</h3>

<p>学习 PPO 最推荐的顺序是：先看 <strong>PPO 原始论文</strong>，理解剪切目标；再看 <strong>OpenAI Spinning Up</strong>，把公式和实现联系起来；如果你的目标是 LLM 对齐，再看 <strong>TRL PPOTrainer</strong> 和 <strong>OpenAI/复现项目的 RLHF with PPO 实现细节</strong>。这样你不会把“机器人控制 PPO”和“语言模型 PPO”混为一谈。</p>

<p>中文学习材料里，<strong>Spinning Up 中文版</strong>适合建立基础策略优化直觉，<strong>Hugging Face 中文博客的 RLHF with PPO 实现细节</strong>更适合理解 LLM 对齐中的工程坑。</p>

<h3 id="常见问题与调试建议-1">常见问题与调试建议</h3>

<p>PPO 在 LLM 场景最典型的失败模式是“<strong>奖励涨了，但文本坏了</strong>”。这常见于 KL 太弱、奖励模型可被钻空子、或者 EOS/截断规则处理不当。OpenAI 复现工作中特别提到 rejection sampling、截断、EOS 处理和固定低分惩罚，这些都不是数学主角，但它们会直接改变模型被奖励的输出分布。</p>

<p>第二类问题是“<strong>critic 学坏导致 advantage 失真</strong>”。如果 value loss 长期震荡、优势估计噪声很大，policy 就会跟着乱动。DeepSeekMath 在提出 GRPO 时正是明确指出：在 LLM 里，value model 常与 policy 模型同量级，内存和算力负担都很大，而且当奖励只在结尾给出时，训练一个逐 token 准确的 value function 还会变得更难。</p>

<h3 id="小结与适用场景-1">小结与适用场景</h3>

<p>PPO 适合你需要<strong>在线探索</strong>、奖励函数复杂、且希望保留“显式 RL 闭环”能力的时候。它是最经典、最通用、也最重的一条路。如果你的目标是“先把偏好学习跑起来”，DPO 通常更轻；如果你的目标是“在可验证推理任务里做在线强化学习，但不想背 critic 成本”，GRPO 往往更合适。</p>

<h2 id="dpo">DPO</h2>

<h3 id="通俗直观解释-2">通俗直观解释</h3>

<p>DPO 的核心直觉非常适合初学者：<strong>既然我们手里已经有“同一 prompt 下哪个回答更好”的配对数据，那为什么还要先学一个奖励模型，再跑一轮 RL，最后才能让策略更偏向好回答？</strong> DPO 的答案是：在特定假设下，可以直接把这个优化目标写成分类损失，于是训练就变成“提高优选回答的相对似然，降低弱选回答的相对似然”。</p>

<p>这也是为什么 DPO 论文会说它“stable, performant, and computationally lightweight”，并强调<strong>不需要在微调时从语言模型中在线采样，也不需要显式奖励模型</strong>。对于初学者，这意味着：比起 PPO，你可以更快地建立“偏好数据如何改变模型分布”的感觉，因为训练就像一个偏好版的二分类。</p>

<h3 id="关键数学推导-2">关键数学推导</h3>

<p>DPO 推导的起点并不是“完全抛弃 RLHF”，而是<strong>从 RLHF 的 KL 约束目标出发</strong>。假设我们要解的仍是下面这个问题：</p>

\[\max_{\pi}\;
\mathbb E_{x\sim \mathcal D,\; y\sim \pi(\cdot|x)}[r(x,y)]
-\beta D_{\mathrm{KL}}(\pi(\cdot|x)\|\pi_{\text{ref}}(\cdot|x))\]

<p>DPO 论文说明，这个目标的最优策略可以写成：</p>

\[\pi^*(y\mid x)=\frac{1}{Z(x)}\pi_{\text{ref}}(y\mid x)\exp\!\left(\frac{1}{\beta}r(x,y)\right)\]

<p>其中 $Z(x)$ 是归一化常数。把它取对数并整理，可得：</p>

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

<p>接下来引入 Bradley–Terry 偏好模型：</p>

\[P(y^+ \succ y^- \mid x)=\sigma\big(r(x,y^+) - r(x,y^-)\big)\]

<p>把上面的 reward 重参数化代进去，$Z(x)$ 会在成对比较中抵消，于是得到只与策略和参考策略有关的偏好概率。最终得到 DPO 损失：</p>

\[\mathcal L_{\text{DPO}}(\theta)=
-\mathbb E_{(x,y^+,y^-)}
\log \sigma\!\left(
\beta \log \frac{\pi_\theta(y^+\mid x)}{\pi_{\text{ref}}(y^+\mid x)}
-
\beta \log \frac{\pi_\theta(y^-\mid x)}{\pi_{\text{ref}}(y^-\mid x)}
\right)\]

<p>这就是 DPO 的核心：<strong>把“学奖励再做 RL”改成了“直接学一个更符合偏好的策略”</strong>。需要强调的是，这个结论依赖于明确的建模假设，尤其是 KL 约束形式和 Bradley–Terry 偏好建模；因此它不是“所有 RLHF 都被 DPO 完全取代”的数学结论，而是“在这类设定下，目标可被直接改写”的结论。</p>

<p>从实现角度看，DPO 还有一个非常实用的“隐式奖励”视角。TRL 文档把 chosen/rejected 的隐式奖励记成相对参考模型的对数概率比；如果你监控 <code class="language-plaintext highlighter-rouge">rewards/chosen</code>、<code class="language-plaintext highlighter-rouge">rewards/rejected</code> 和 <code class="language-plaintext highlighter-rouge">rewards/margins</code>，其实就在看模型相对参考模型把“好回答”和“差回答”分开得有多明显。</p>

<h3 id="伪代码-2">伪代码</h3>

<p>下面的伪代码几乎就是 DPO 的本质。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">initialize</span> <span class="n">policy</span> <span class="n">πθ</span> <span class="k">from</span> <span class="n">SFT</span> <span class="n">model</span>
<span class="n">freeze</span> <span class="n">reference</span> <span class="n">policy</span> <span class="n">πref</span>

<span class="k">for</span> <span class="n">batch</span> <span class="ow">in</span> <span class="n">preference_dataset</span><span class="p">:</span>   <span class="c1"># (x, y_plus, y_minus)
</span>    <span class="n">logp_plus</span> <span class="o">=</span> <span class="n">log</span> <span class="n">πθ</span><span class="p">(</span><span class="n">y_plus</span> <span class="o">|</span> <span class="n">x</span><span class="p">)</span>
    <span class="n">logp_minus</span> <span class="o">=</span> <span class="n">log</span> <span class="n">πθ</span><span class="p">(</span><span class="n">y_minus</span> <span class="o">|</span> <span class="n">x</span><span class="p">)</span>

    <span class="n">ref_logp_plus</span> <span class="o">=</span> <span class="n">log</span> <span class="n">πref</span><span class="p">(</span><span class="n">y_plus</span> <span class="o">|</span> <span class="n">x</span><span class="p">)</span>
    <span class="n">ref_logp_minus</span> <span class="o">=</span> <span class="n">log</span> <span class="n">πref</span><span class="p">(</span><span class="n">y_minus</span> <span class="o">|</span> <span class="n">x</span><span class="p">)</span>

    <span class="n">margin</span> <span class="o">=</span> <span class="n">beta</span> <span class="o">*</span> <span class="p">((</span><span class="n">logp_plus</span> <span class="o">-</span> <span class="n">ref_logp_plus</span><span class="p">)</span> <span class="o">-</span>
                     <span class="p">(</span><span class="n">logp_minus</span> <span class="o">-</span> <span class="n">ref_logp_minus</span><span class="p">))</span>

    <span class="n">loss</span> <span class="o">=</span> <span class="o">-</span><span class="n">mean</span><span class="p">(</span><span class="n">log_sigmoid</span><span class="p">(</span><span class="n">margin</span><span class="p">))</span>
    <span class="n">update</span><span class="p">(</span><span class="n">θ</span><span class="p">,</span> <span class="n">loss</span><span class="p">)</span>
</code></pre></div></div>

<h3 id="简短-python-示例片段-2">简短 Python 示例片段</h3>

<p>下面这段代码非常接近 TRL 官方文档中的最小示例。它足够短，适合你亲自替换模型与数据集跑一个 demo。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="nn">datasets</span> <span class="kn">import</span> <span class="n">load_dataset</span>
<span class="kn">from</span> <span class="nn">trl</span> <span class="kn">import</span> <span class="n">DPOTrainer</span>

<span class="n">trainer</span> <span class="o">=</span> <span class="n">DPOTrainer</span><span class="p">(</span>
    <span class="n">model</span><span class="o">=</span><span class="s">"Qwen/Qwen3-0.6B"</span><span class="p">,</span>
    <span class="n">train_dataset</span><span class="o">=</span><span class="n">load_dataset</span><span class="p">(</span><span class="s">"trl-lib/ultrafeedback_binarized"</span><span class="p">,</span> <span class="n">split</span><span class="o">=</span><span class="s">"train"</span><span class="p">),</span>
<span class="p">)</span>

<span class="n">trainer</span><span class="p">.</span><span class="n">train</span><span class="p">()</span>
</code></pre></div></div>

<h3 id="实践要点与超参数建议-2">实践要点与超参数建议</h3>

<p>DPO 的第一原则不是调 $\beta$，而是<strong>确保 preference pair 真的可信</strong>。如果 chosen / rejected 顺序经常反了，或者两条回答质量几乎没有可比差异，DPO 会学到很奇怪的边界。因为它没有单独的奖励模型做缓冲，所以数据质量问题会更直接地反映到策略里。DPO 论文与 TRL 文档都强调它处理的是“同一 prompt 的优选/弱选完成对”。</p>

<p>$\beta$ 是 DPO 里最值得优先理解的超参数。TRL 文档给出的默认值是 <strong>0.1</strong>，并把它解释为控制与参考模型偏离程度的关键参数；而后续 $\beta$-DPO 工作又专门指出，<strong>$\beta$ 对性能是敏感的</strong>。对于初学者，最稳妥的建议是：<strong>先用 0.1 跑通，再围绕它做少量扫描，不要一上来就扫特别大范围。</strong></p>

<p>在工程上，TRL 当前默认 <strong><code class="language-plaintext highlighter-rouge">disable_dropout=True</code></strong>，并提供 <code class="language-plaintext highlighter-rouge">precompute_ref_log_probs=True</code> 来节省显存；此外，<code class="language-plaintext highlighter-rouge">padding_free</code> 可以配合 FlashAttention 进一步减少 padding 开销。这些都很适合大模型或长序列场景。</p>

<p>如果你的偏好标签有噪声，可以关注 Robust DPO / EXO 一类扩展。TRL 文档中已经把 <code class="language-plaintext highlighter-rouge">label_smoothing</code> 暴露出来，并给出了一些典型值说明，例如 Robust DPO 中常把它解释成标签翻转概率，推荐值可取 0.1 的量级；但如果你只是入门，建议先不要同时引入太多 DPO 变体。</p>

<h3 id="典型实现与在线文档-2">典型实现与在线文档</h3>

<p>学习 DPO 的“黄金组合”通常是：先读 <strong>原始论文</strong>，再看 <strong>Stanford 的参考实现仓库</strong>，然后直接上 <strong>TRL DPOTrainer</strong>。这样做的好处是你既知道数学上为什么成立，又知道工程上怎么落地。</p>

<p>中文方面，优先读 <strong>Hugging Face 中文博客《从 RLHF 到 DPO》</strong>、<strong>Hugging Face 中文版 DPOTrainer 文档</strong>、以及 <strong>使用 DPO 微调 Llama 2</strong> 这类带脚本的实战教程。它们跟英文原始材料一致性较高。</p>

<h3 id="常见问题与调试建议-2">常见问题与调试建议</h3>

<p>DPO 第一类常见问题是“<strong>loss 在降，但模型回答更死板</strong>”。这通常意味着模型在过度贴近某类 chosen 模板，或者 $\beta$ / 数据分布让模型相对参考模型的偏移过于单一。此时优先看 <code class="language-plaintext highlighter-rouge">rewards/margins</code>、<code class="language-plaintext highlighter-rouge">rewards/accuracies</code>、<code class="language-plaintext highlighter-rouge">logps/chosen</code>、<code class="language-plaintext highlighter-rouge">logps/rejected</code> 是否同步改善，而不是只看总 loss。</p>

<p>第二类问题是“<strong>拿未做过 SFT 的底座直接做 DPO</strong>”。理论上不是绝对不行，但实践上很容易因为 pair 数据分布太窄而出问题。DPO 论文本身就是在 RLHF 的 SFT 起点上做改写；TRL 文档的示例也默认你已经有可用的 base/SFT 模型。对初学者来说，<strong>DPO 最好接在一个已经能正常完成任务的 SFT 模型之后</strong>。</p>

<h3 id="小结与适用场景-2">小结与适用场景</h3>

<p>当你手里有<strong>静态偏好数据</strong>，又希望用尽可能简单的训练方式完成偏好对齐时，DPO 几乎总是值得优先尝试。它比 PPO 更轻，训练与调试门槛更低，也更适合“先学会偏好优化，再回头学在线 RL”。但它的前提很明确：你要有质量不错的 chosen/rejected 数据，而且最好已经有一个 decent 的 SFT 起点。</p>

<h2 id="grpo">GRPO</h2>

<h3 id="通俗直观解释-3">通俗直观解释</h3>

<p>GRPO 可以把它想成“<strong>不再问一个回答绝对值多少，而是问它在同组候选里相对表现如何</strong>”。对同一个 prompt，你用旧策略一次采样出 $G$ 个回答；然后用奖励模型、规则验证器或准确率函数给它们打分；接着看每个回答相对这组回答的平均水平高多少、低多少；最后用这个“组内相对优势”更新策略。DeepSeekMath 论文把它定义为 PPO 的一个变体，最大卖点是：<strong>不再需要单独训练 critic/value model</strong>。</p>

<p>为什么这个思路在推理任务上特别流行？因为数学、代码、格式约束这些任务里，很多奖励都可以<strong>由规则直接验证</strong>，例如最终答案是否正确、输出是否满足模板、测试用例是否通过。DeepSeek-R1 进一步说明，他们在 reasoning 任务上更倾向于使用这类<strong>规则型奖励信号</strong>，而刻意避免大规模依赖神经奖励模型，因为后者在强 RL 下更容易被 reward hacking。</p>

<h3 id="关键数学推导-3">关键数学推导</h3>

<p>GRPO 这部分最容易让人混淆的地方在于：<strong>原始 DeepSeekMath 写法、后续 PPO-style 推广写法、以及 TRL 当前实现的工程修正式，并不是同一个公式。</strong> 把它们拆开看会清楚很多。</p>

<p>先看原始 DeepSeekMath 的核心直觉。它把 PPO 中“用 value function 做 baseline”的部分，替换成“<strong>用同组样本的平均奖励做 baseline</strong>”。设同一个问题 $q$ 采样出 $G$ 个回答 $o_1,\dots,o_G$，对应奖励为 $r_1,\dots,r_G$。最常见的组内优势写法是：</p>

\[A_i=\frac{r_i-\operatorname{mean}(r_{1:G})}{\operatorname{std}(r_{1:G})}\]

<p>如果是 outcome supervision，那么一条回答中所有 token 共用这个终局归一化奖励；如果是 process supervision，则可以给每个推理步骤打分，再把未来步骤的归一化奖励累加到每个 token 上。DeepSeekMath 对 outcome supervision 和 process supervision 都给了明确说明。到这一步为止，重点是“<strong>优势来自组内相对奖励，而不是 value model</strong>”。</p>

<p>接下来，为了和 PPO 对照，很多讲解或实现会再引入 PPO-style 的概率比率。令</p>

\[\begin{aligned}
\rho_{i,t}(\theta)
&amp;= \frac{\pi_\theta(o_{i,t}\mid q,o_{i,&lt;t})}
{\pi_{\theta_{\text{old}}}(o_{i,t}\mid q,o_{i,&lt;t})}
\end{aligned}\]

<p>则一种常见的 <strong>PPO-style surrogate 写法</strong> 可以写成：</p>

\[\begin{aligned}
\mathcal L_{\text{GRPO}}(\theta)
&amp;= -\mathbb E\Big[
\min\big(
\rho_{i,t}A_i,\;
\operatorname{clip}(\rho_{i,t},1-\epsilon,1+\epsilon)A_i
\big)
-\beta D_{\mathrm{KL}}(\pi_\theta\|\pi_{\text{ref}})
\Big]
\end{aligned}\]

<p>其中 $\beta$ 控制与参考策略的距离，$\epsilon$ 控制剪切范围。你会发现，它和 PPO 的结构非常像；区别在于 advantage 的来源：PPO 用 value/GAE，GRPO 用组内相对奖励。正因如此，你可以把 GRPO 记成“<strong>PPO 的 critic-free、group-relative 版本</strong>”。</p>

<p>最后要把“原始算法”和“工程变体”明确区分开。TRL 当前文档明确指出，组内标准差缩放可能引入 question-level difficulty bias；原始按回答长度做归一化的写法也被后续工作指出可能引入 length bias；因此当前工程实现会提供 <code class="language-plaintext highlighter-rouge">scale_rewards="group" / "batch" / "none"</code>、<code class="language-plaintext highlighter-rouge">loss_type="grpo" / "dapo" / "dr_grpo"</code>、<code class="language-plaintext highlighter-rouge">mask_truncated_completions</code> 等选项来缓解这些偏差。也就是说，网上你看到的 GRPO 公式之所以不完全一致，往往不是谁写错了，而是<strong>讨论的对象已经从“原始 DeepSeekMath 公式”切换成“后续工程修正版”</strong>。</p>

<h3 id="伪代码-3">伪代码</h3>

<p>下面的伪代码表达的是“对每个 prompt 生成多个回答，再用组内相对优势更新”的核心思想。它与 DeepSeekMath 的算法框架和 TRL 的 GRPOTrainer 设计是一致的。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">initialize</span> <span class="n">policy</span> <span class="n">πθ</span>
<span class="n">initialize</span> <span class="n">reference</span> <span class="n">policy</span> <span class="n">πref</span>   <span class="c1"># optional in some modern variants
</span>
<span class="k">for</span> <span class="n">each</span> <span class="n">iteration</span><span class="p">:</span>
    <span class="n">batch_prompts</span> <span class="o">=</span> <span class="n">sample_prompts</span><span class="p">()</span>

    <span class="n">groups</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="k">for</span> <span class="n">q</span> <span class="ow">in</span> <span class="n">batch_prompts</span><span class="p">:</span>
        <span class="n">completions</span> <span class="o">=</span> <span class="n">sample_G_responses</span><span class="p">(</span><span class="n">πθ_old</span><span class="p">,</span> <span class="n">q</span><span class="p">,</span> <span class="n">G</span><span class="p">)</span>
        <span class="n">rewards</span> <span class="o">=</span> <span class="n">reward_fn</span><span class="p">(</span><span class="n">q</span><span class="p">,</span> <span class="n">completions</span><span class="p">)</span>   <span class="c1"># rule-based or RM-based
</span>        <span class="n">advantages</span> <span class="o">=</span> <span class="n">normalize_within_group</span><span class="p">(</span><span class="n">rewards</span><span class="p">)</span>
        <span class="n">groups</span><span class="p">.</span><span class="n">append</span><span class="p">((</span><span class="n">q</span><span class="p">,</span> <span class="n">completions</span><span class="p">,</span> <span class="n">rewards</span><span class="p">,</span> <span class="n">advantages</span><span class="p">))</span>

    <span class="k">for</span> <span class="n">epoch</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_iterations</span><span class="p">):</span>
        <span class="k">for</span> <span class="n">q</span><span class="p">,</span> <span class="n">completions</span><span class="p">,</span> <span class="n">rewards</span><span class="p">,</span> <span class="n">advantages</span> <span class="ow">in</span> <span class="n">groups</span><span class="p">:</span>
            <span class="n">ratio</span> <span class="o">=</span> <span class="n">πθ</span><span class="p">(</span><span class="n">token</span> <span class="o">|</span> <span class="n">prefix</span><span class="p">)</span> <span class="o">/</span> <span class="n">πθ_old</span><span class="p">(</span><span class="n">token</span> <span class="o">|</span> <span class="n">prefix</span><span class="p">)</span>
            <span class="n">loss</span> <span class="o">=</span> <span class="o">-</span><span class="n">mean</span><span class="p">(</span><span class="nb">min</span><span class="p">(</span><span class="n">ratio</span> <span class="o">*</span> <span class="n">advantages</span><span class="p">,</span>
                             <span class="n">clip</span><span class="p">(</span><span class="n">ratio</span><span class="p">,</span> <span class="mi">1</span><span class="o">-</span><span class="n">eps</span><span class="p">,</span> <span class="mi">1</span><span class="o">+</span><span class="n">eps</span><span class="p">)</span> <span class="o">*</span> <span class="n">advantages</span><span class="p">)</span>
                         <span class="o">-</span> <span class="n">beta</span> <span class="o">*</span> <span class="n">kl</span><span class="p">(</span><span class="n">πθ</span><span class="p">,</span> <span class="n">πref</span><span class="p">))</span>
            <span class="n">update</span><span class="p">(</span><span class="n">θ</span><span class="p">,</span> <span class="n">loss</span><span class="p">)</span>
</code></pre></div></div>

<h3 id="简短-python-示例片段-3">简短 Python 示例片段</h3>

<p>下面这段代码来自 TRL 的最小 GRPO 用法思路：定义数据集，定义奖励函数，把它交给 <code class="language-plaintext highlighter-rouge">GRPOTrainer</code>。这对初学者非常友好，因为它把“奖励函数”显式暴露出来了。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="nn">datasets</span> <span class="kn">import</span> <span class="n">load_dataset</span>
<span class="kn">from</span> <span class="nn">trl</span> <span class="kn">import</span> <span class="n">GRPOTrainer</span>
<span class="kn">from</span> <span class="nn">trl.rewards</span> <span class="kn">import</span> <span class="n">accuracy_reward</span>

<span class="n">dataset</span> <span class="o">=</span> <span class="n">load_dataset</span><span class="p">(</span><span class="s">"trl-lib/DeepMath-103K"</span><span class="p">,</span> <span class="n">split</span><span class="o">=</span><span class="s">"train"</span><span class="p">)</span>

<span class="n">trainer</span> <span class="o">=</span> <span class="n">GRPOTrainer</span><span class="p">(</span>
    <span class="n">model</span><span class="o">=</span><span class="s">"Qwen/Qwen2-0.5B-Instruct"</span><span class="p">,</span>
    <span class="n">reward_funcs</span><span class="o">=</span><span class="n">accuracy_reward</span><span class="p">,</span>
    <span class="n">train_dataset</span><span class="o">=</span><span class="n">dataset</span><span class="p">,</span>
<span class="p">)</span>

<span class="n">trainer</span><span class="p">.</span><span class="n">train</span><span class="p">()</span>
</code></pre></div></div>

<h3 id="实践要点与超参数建议-3">实践要点与超参数建议</h3>

<p>如果你是初学者，我最推荐把 GRPO 用在<strong>可验证任务</strong>上，而不是开放式写作上。原因很简单：当 reward 是“答案是否正确”“测试是否通过”“格式是否满足模板”时，组内相对比较非常自然，而且不容易被一个脆弱的神经奖励模型带偏。DeepSeek-R1 也正是在 reasoning 任务上强调使用 rule-based reward 与 format reward，并明确写出他们避免大规模使用神经奖励模型，是因为担心 reward hacking 和额外复杂度。</p>

<p>从 TRL 当前默认值看，GRPO 的入门配置相对克制：<code class="language-plaintext highlighter-rouge">learning_rate</code> 默认 <strong>1e-6</strong>，<code class="language-plaintext highlighter-rouge">num_generations</code> 默认 <strong>8</strong>，<code class="language-plaintext highlighter-rouge">max_completion_length</code> 默认 <strong>256</strong>，<code class="language-plaintext highlighter-rouge">epsilon</code> 默认 <strong>0.2</strong>，<code class="language-plaintext highlighter-rouge">num_iterations</code> 默认 <strong>1</strong>，<code class="language-plaintext highlighter-rouge">beta</code> 默认 <strong>0.0</strong>。文档还特别解释了：当 <code class="language-plaintext highlighter-rouge">beta=0.0</code> 时，参考模型根本不加载，这可以进一步降低显存与训练开销。与此同时，DeepSeekMath 在原始论文里给出过一个更“经典 RL”风格的配置：policy learning rate 取 <strong>1e-6</strong>，KL coefficient 取 <strong>0.04</strong>。因此，一个很实用的理解是：<strong>原始 GRPO 更像“带 KL 的 PPO 变体”，而许多后续实现开始尝试把 KL 弱化甚至去掉。</strong></p>

<p>另一个必须注意的点是奖励归一化。TRL 文档直接提醒：按组标准差缩放虽然是原始直觉，但可能引入题目难度偏差；可选替代是 <code class="language-plaintext highlighter-rouge">scale_rewards=False</code> 或 <code class="language-plaintext highlighter-rouge">scale_rewards="batch"</code>。如果你在一个 batch 里混合了难度差异很大的题目，这个设置会明显影响训练信号。对刚开始实验的人，我建议先使用默认方式跑一次，再在一个小验证集上比较 <code class="language-plaintext highlighter-rouge">group</code>、<code class="language-plaintext highlighter-rouge">batch</code>、<code class="language-plaintext highlighter-rouge">none</code> 三种缩放策略，而不是一开始就盲信某一篇经验帖。</p>

<p>最后，GRPO 在长回答任务上要格外注意<strong>长度偏差与截断问题</strong>。TRL 当前已经把 <code class="language-plaintext highlighter-rouge">mask_truncated_completions=True</code> 明确标为一个有助于稳定性的好实践；同时也提供了针对长度偏差修正的 <code class="language-plaintext highlighter-rouge">dr_grpo</code> 等 loss 类型。你不必在第一次实验就把这些都用上，但必须知道：<strong>原始 GRPO 不是“公式一抄就稳”的算法，长 CoT 场景下很多偏差都是真问题。</strong></p>

<h3 id="典型实现与在线文档-3">典型实现与在线文档</h3>

<p>GRPO 最核心的原始材料是 <strong>DeepSeekMath 论文</strong>；如果你想看它如何进入更大规模 reasoning 体系，则读 <strong>DeepSeek-R1</strong>。工程上，<strong>TRL 的 GRPOTrainer</strong> 是最适合初学者的公开实现接口，而 <strong>Hugging Face open-r1</strong> 则更接近“开放复现 DeepSeek-R1 路线”的研究/工程仓库。</p>

<p>中文资源里，优先看 <strong>Hugging Face 中文课程中实现 GRPO 的章节</strong> 和 <strong>OpenRLHF 中文文档</strong>。前者偏教学，后者偏工程系统。</p>

<h3 id="常见问题与调试建议-3">常见问题与调试建议</h3>

<p>GRPO 的第一类问题是“<strong>组内奖励几乎一样，学不到东西</strong>”。从公式直接可见，当同组样本的奖励全接近相同，归一化后的 advantage 就会接近 0，更新信号自然变弱。这通常不是算法坏了，而是你的 reward 太粗糙、组内采样多样性太低，或者任务本身对当前模型来说已经过于简单。</p>

<p>第二类问题是“<strong>算力一下子爆炸</strong>”。GRPO 每个 prompt 要采样 $G$ 个 completion，训练成本天然受 <code class="language-plaintext highlighter-rouge">num_generations</code> 影响；TRL 的 quick start 甚至直接写出，8 GPU 下一个示例训练大约需要 1 天。OpenRLHF 也强调，在 RLHF 系统里，生成阶段常是主要开销来源。对入门实验来说，你完全可以先把 <code class="language-plaintext highlighter-rouge">num_generations</code>、<code class="language-plaintext highlighter-rouge">max_completion_length</code> 和 batch size 压小，把奖励函数和日志跑通，再考虑扩展。</p>

<p>第三类问题是“<strong>你学到的是格式，而不是能力</strong>”。如果奖励函数只奖励 <code class="language-plaintext highlighter-rouge">&lt;think&gt;...&lt;/think&gt;&lt;answer&gt;...&lt;/answer&gt;</code> 之类模板，模型很可能先学会的是“写得像被奖励的样子”，而不是更强的推理能力。因此 reward 设计最好至少同时覆盖<strong>正确性</strong>与<strong>格式性</strong>，这也是 DeepSeek-R1 中 accuracy reward 与 format reward 并用的原因。</p>

<h3 id="小结与适用场景-3">小结与适用场景</h3>

<p>如果你的任务是数学、代码、形式化推理、结构化输出，且有可靠 verifier，GRPO 非常值得学；如果你的任务是开放式闲聊、主观写作或通用帮助性，GRPO 也能用，但它的奖励设计会比 DPO 和经典 RM-PPO 更难做稳。从方法定位上看，GRPO 更接近“<strong>面向可验证推理任务的在线 RL 变体</strong>”。</p>

<h2 id="方法比较与流程概览">方法比较与流程概览</h2>

<h3 id="比较表">比较表</h3>

<table>
  <thead>
    <tr>
      <th>方法</th>
      <th>目标</th>
      <th style="text-align: right">是否需奖励模型</th>
      <th style="text-align: right">是否需价值模型</th>
      <th>训练复杂度</th>
      <th>优点</th>
      <th>缺点</th>
      <th>适用场景</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td>RLHF 管线</td>
      <td>用人类偏好把模型行为对齐到“更有帮助/更安全/更符合偏好”的方向</td>
      <td style="text-align: right">通常需要</td>
      <td style="text-align: right">取决于具体优化器；用 PPO 时通常需要</td>
      <td>高</td>
      <td>框架完整、可插入多种奖励、工业实践成熟</td>
      <td>数据与工程链路长，调试成本高</td>
      <td>通用大模型对齐、需要在线闭环时</td>
    </tr>
    <tr>
      <td>PPO</td>
      <td>稳定地在线优化策略，避免一步更新过大</td>
      <td style="text-align: right">在 LLM 对齐中通常需要奖励信号，常来自 RM 或规则</td>
      <td style="text-align: right">需要</td>
      <td>高</td>
      <td>通用、成熟、适合复杂奖励</td>
      <td>actor/ref/reward/value 四件套很重</td>
      <td>经典 RLHF、复杂环境、在线优化</td>
    </tr>
    <tr>
      <td>DPO</td>
      <td>直接从偏好对中优化策略</td>
      <td style="text-align: right">不需要显式 RM</td>
      <td style="text-align: right">不需要</td>
      <td>中低</td>
      <td>简洁、稳定、无需在线采样回路</td>
      <td>强依赖高质量 preference pairs；$\beta$ 敏感</td>
      <td>静态偏好数据、低门槛对齐</td>
    </tr>
    <tr>
      <td>GRPO</td>
      <td>用组内相对奖励做在线策略优化</td>
      <td style="text-align: right">可选：可用 RM，也可用规则/验证器</td>
      <td style="text-align: right">不需要</td>
      <td>中高</td>
      <td>去掉 critic，适合可验证推理任务</td>
      <td>多候选采样成本高；公式和工程变体较多</td>
      <td>数学、代码、结构化推理、规则奖励</td>
    </tr>
  </tbody>
</table>

<p>表里的“是否需要奖励模型/价值模型”都要按语境理解。<strong>RLHF 是管线，不是单算法</strong>；如果 RLHF 用 PPO，则通常需要 RM 和 value model。DPO 不需要显式 RM，也不需要 value model。GRPO 在原始 DeepSeekMath 中保留了可选 KL 正则，但核心思想就是用组内相对奖励替代 critic；在 TRL 当前实现里，<code class="language-plaintext highlighter-rouge">beta=0.0</code> 甚至是默认值。</p>

<h3 id="一张帮助记忆的流程对比图">一张帮助记忆的流程对比图</h3>

<pre><code class="language-mermaid">flowchart TB
    subgraph A[经典 RLHF]
        A1[SFT 模型] --&gt; A2[采样回答]
        A2 --&gt; A3[偏好比较]
        A3 --&gt; A4[奖励模型]
        A4 --&gt; A5[PPO 更新策略]
    end

    subgraph B[DPO]
        B1[SFT 模型] --&gt; B2[偏好比较数据]
        B2 --&gt; B3[直接优化 chosen/rejected 概率边界]
    end

    subgraph C[GRPO]
        C1[SFT 或基座模型] --&gt; C2[同一 prompt 采样 G 个回答]
        C2 --&gt; C3[规则/奖励函数打分]
        C3 --&gt; C4[组内相对优势]
        C4 --&gt; C5[更新策略]
    end
</code></pre>

<p>这张图对应的记忆口诀可以很简单：<strong>RLHF 是全流程，PPO 是重武器，DPO 是轻武器，GRPO 是推理场景里常见的 critic-free 在线武器。</strong>这种归纳是基于原始论文与官方文档做的教学性总结，不是任何单一论文的原话。</p>

<h2 id="参考链接">参考链接</h2>

<h3 id="原始论文">原始论文</h3>

<ul>
  <li><strong>Training language models to follow instructions with human feedback</strong>，InstructGPT 原始论文。</li>
  <li><strong>Learning to summarize from human feedback</strong>，OpenAI 在摘要任务上的经典 RLHF 论文。</li>
  <li><strong>Proximal Policy Optimization Algorithms</strong>，PPO 原始论文。</li>
  <li><strong>High-Dimensional Continuous Control Using Generalized Advantage Estimation</strong>，GAE 原始论文。</li>
  <li><strong>Direct Preference Optimization: Your Language Model is Secretly a Reward Model</strong>，DPO 原始论文。</li>
  <li><strong>DeepSeekMath: Pushing the Limits of Mathematical Reasoning in Open Language Models</strong>，GRPO 最核心的原始论文来源。</li>
  <li><strong>DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via Reinforcement Learning</strong>，展示 GRPO 在 R1 系列中的核心地位与规则奖励实践。</li>
  <li><strong>The N+ Implementation Details of RLHF with PPO: A Case Study on TL;DR Summarization</strong>，公开复现 RLHF-with-PPO 的工程论文。</li>
  <li><strong>$\beta$-DPO: Direct Preference Optimization with Dynamic $\beta$</strong>，讨论 DPO 中 $\beta$ 敏感性的后续工作。</li>
</ul>

<h3 id="官方博客与官方文档">官方博客与官方文档</h3>

<ul>
  <li><strong>OpenAI: Aligning language models to follow instructions</strong>，InstructGPT 官方博客。</li>
  <li><strong>OpenAI: Learning to summarize with human feedback</strong>，总结任务 RLHF 官方介绍。</li>
  <li><strong>OpenAI Spinning Up</strong>，策略优化和 PPO 的权威入门材料。</li>
  <li><strong>Hugging Face TRL 总览</strong>，统一了解 PPO / DPO / GRPO / Reward Modeling 的入口。</li>
  <li><strong>TRL PPOTrainer 文档</strong>。</li>
  <li><strong>TRL DPOTrainer 文档</strong>。</li>
  <li><strong>TRL GRPOTrainer 文档</strong>。</li>
  <li><strong>TRL Reward Modeling 文档</strong>。</li>
  <li><strong>The N Implementation Details of RLHF with PPO</strong>，Hugging Face 对 OpenAI 早期 RLHF 代码的工程级复现总结。</li>
</ul>

<h3 id="官方或常用开源仓库">官方或常用开源仓库</h3>

<ul>
  <li><strong>openai/following-instructions-human-feedback</strong>，OpenAI 的 InstructGPT 公开仓库，主要提供论文、模型卡和评测样本；完整训练代码未在该仓库中提供。</li>
  <li><strong>openai/lm-human-preferences</strong>，OpenAI 公开的早期“人类偏好微调语言模型”代码库。</li>
  <li><strong>openai/summarize-from-feedback</strong>，OpenAI 的摘要 RLHF 代码与数据资源。</li>
  <li><strong>openai/spinningup</strong>，OpenAI 的深度强化学习教学仓库。</li>
  <li><strong>eric-mitchell/direct-preference-optimization</strong>，DPO 参考实现。</li>
  <li><strong>huggingface/trl</strong>，当前最常见的 LLM 对齐训练工具库之一。</li>
  <li><strong>OpenRLHF/OpenRLHF</strong>，高性能 RLHF 工程框架。</li>
  <li><strong>deepseek-ai/DeepSeek-Math</strong>，DeepSeekMath 公开仓库。</li>
  <li><strong>huggingface/open-r1</strong>，Open-R1 开放复现项目。</li>
</ul>

<h3 id="中文高质量资源">中文高质量资源</h3>

<ul>
  <li><strong>Hugging Face 中文：从 RLHF 到 DPO</strong>。</li>
  <li><strong>Hugging Face 中文：DPOTrainer 文档</strong>。</li>
  <li><strong>Hugging Face 中文：使用 DPO 微调 Llama 2</strong>。</li>
  <li><strong>Hugging Face 中文课程：在 TRL 中实现 GRPO</strong>。</li>
  <li><strong>Spinning Up 中文版</strong>。</li>
  <li><strong>OpenRLHF 中文文档</strong>。</li>
</ul>

<h3 id="额外说明">额外说明</h3>

<ul>
  <li>对于 <strong>InstructGPT 的完整训练代码仓库</strong>，本文检索到的是 OpenAI 的公开说明性仓库与相关 RLHF 参考代码；若你需要“完全复刻 InstructGPT 原训练流水线”的逐步脚本，公开材料中并未给出一套与论文一一对应的官方训练仓库，可视为<strong>未指定</strong>。</li>
</ul>]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[强化学习初学者速通文档]]></summary></entry><entry><title type="html">JAX and TPU学习记录</title><link href="https://wqh011128.github.io/blog/2026/03/26/JAX_and_TPU%E8%AE%B0%E5%BD%95.html" rel="alternate" type="text/html" title="JAX and TPU学习记录" /><published>2026-03-26T00:00:00+00:00</published><updated>2026-03-26T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/03/26/JAX_and_TPU%E8%AE%B0%E5%BD%95</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/03/26/JAX_and_TPU%E8%AE%B0%E5%BD%95.html"><![CDATA[<h1 id="jax--tpu-两个重点问题">JAX + TPU 两个重点问题</h1>

<p>这次其实只需要抓住两个问题：</p>

<ol>
  <li><strong>如何监控：我到底有没有真正用上 TPU，运行时状态怎么看？</strong></li>
  <li><strong><code class="language-plaintext highlighter-rouge">XLA_FLAGS</code> 和 <code class="language-plaintext highlighter-rouge">LIBTPU_INIT_ARGS</code> 到底分别控制什么，会额外产出什么？</strong></li>
</ol>

<hr />

<h2 id="1-如何监控-jax--tpu">1. 如何监控 JAX + TPU</h2>

<p>我现在更倾向于把“监控”分成三层看，因为很多混淆都来自把这三层混在一起。</p>

<h3 id="11-第一层机器上有没有-tpu-设备">1.1 第一层：机器上有没有 TPU 设备</h3>

<p>这一层回答的是：</p>

<blockquote>
  <p><strong>机器有没有挂上 TPU 芯片。</strong></p>
</blockquote>

<p>常用命令：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>tpu-info
tpu-info <span class="nt">--streaming</span> <span class="nt">--rate</span> 2
</code></pre></div></div>

<p>如果 <code class="language-plaintext highlighter-rouge">tpu-info</code> 能看到型号、芯片数量、<code class="language-plaintext highlighter-rouge">/dev/vfio/*</code>，说明<strong>设备层面</strong>是存在的。</p>

<p>但这里要注意一个关键点：</p>

<ul>
  <li><strong>能看到 TPU 芯片，不等于能看到 TPU 利用率</strong></li>
  <li><code class="language-plaintext highlighter-rouge">HBM Usage</code>、<code class="language-plaintext highlighter-rouge">Duty cycle</code>、<code class="language-plaintext highlighter-rouge">TensorCore Utilization</code> 出现 <code class="language-plaintext highlighter-rouge">N/A</code>，通常说明 <strong>runtime 指标没有正确暴露出来</strong></li>
</ul>

<p>最常见原因有这些：</p>

<ul>
  <li>当前环境没有正确接上 <code class="language-plaintext highlighter-rouge">libtpu</code></li>
  <li><code class="language-plaintext highlighter-rouge">tpu-info</code> 和当前环境不兼容</li>
  <li>当前没有真正跑 TPU workload</li>
  <li>runtime metrics 本身还不可用</li>
</ul>

<p>可以先做一个最小检查：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>python - <span class="o">&lt;&lt;</span><span class="sh">'</span><span class="no">PY</span><span class="sh">'
try:
    import libtpu
    print("libtpu OK:", libtpu.__file__)
except Exception as e:
    print("libtpu import failed:", repr(e))
</span><span class="no">PY
</span></code></pre></div></div>

<p>如果这里连 <code class="language-plaintext highlighter-rouge">libtpu</code> 都导不进来，那么 <code class="language-plaintext highlighter-rouge">tpu-info</code> 读不到完整指标就不奇怪。</p>

<h3 id="12-第二层jax-有没有成功连到-tpu-backend">1.2 第二层：JAX 有没有成功连到 TPU backend</h3>

<p>这一层回答的是：</p>

<blockquote>
  <p><strong>JAX 这个进程，是否真的初始化了 TPU backend。</strong></p>
</blockquote>

<p>最小判断代码：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">jax</span>

<span class="k">print</span><span class="p">(</span><span class="s">"backend:"</span><span class="p">,</span> <span class="n">jax</span><span class="p">.</span><span class="n">default_backend</span><span class="p">())</span>
<span class="k">print</span><span class="p">(</span><span class="s">"devices:"</span><span class="p">,</span> <span class="n">jax</span><span class="p">.</span><span class="n">devices</span><span class="p">())</span>

<span class="k">assert</span> <span class="n">jax</span><span class="p">.</span><span class="n">devices</span><span class="p">()[</span><span class="mi">0</span><span class="p">].</span><span class="n">platform</span> <span class="o">==</span> <span class="s">"tpu"</span>
</code></pre></div></div>

<p>如果输出类似：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">backend</span><span class="p">:</span> <span class="n">tpu</span>
<span class="n">devices</span><span class="p">:</span> <span class="p">[</span><span class="n">TpuDevice</span><span class="p">(...)]</span>
</code></pre></div></div>

<p>说明的事情只有一件：</p>

<blockquote>
  <p><strong>JAX backend 已经接到 TPU 了。</strong></p>
</blockquote>

<p>但这还不等于“整段 Python 程序都在 TPU 上跑”。<br />
真正上 TPU 的，只是 JAX/XLA 编译后的数组计算部分；数据加载、Python 循环、日志打印仍然主要在 host 侧。</p>

<h3 id="13-第三层具体这一步计算是否真的在-tpu-上执行">1.3 第三层：具体这一步计算是否真的在 TPU 上执行</h3>

<p>这一层回答的是：</p>

<blockquote>
  <p><strong>不是环境认到了 TPU，而是这一步 <code class="language-plaintext highlighter-rouge">step()</code>、这次 matmul、这次训练迭代，是否真的在 TPU 上算了。</strong></p>
</blockquote>

<p>最实用的判断方式有三种。</p>

<h4 id="方法-1看输出数组落在哪个设备上">方法 1：看输出数组落在哪个设备上</h4>

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

<span class="o">@</span><span class="n">jax</span><span class="p">.</span><span class="n">jit</span>
<span class="k">def</span> <span class="nf">step</span><span class="p">(</span><span class="n">x</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">x</span> <span class="o">@</span> <span class="n">x</span> <span class="o">+</span> <span class="mi">1</span>

<span class="n">x</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">ones</span><span class="p">((</span><span class="mi">4096</span><span class="p">,</span> <span class="mi">4096</span><span class="p">))</span>
<span class="n">y</span> <span class="o">=</span> <span class="n">step</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="n">y</span><span class="p">.</span><span class="n">block_until_ready</span><span class="p">()</span>

<span class="k">print</span><span class="p">(</span><span class="s">"device:"</span><span class="p">,</span> <span class="n">y</span><span class="p">.</span><span class="n">device</span><span class="p">)</span>
<span class="k">print</span><span class="p">(</span><span class="s">"sharding:"</span><span class="p">,</span> <span class="n">y</span><span class="p">.</span><span class="n">sharding</span><span class="p">)</span>
<span class="k">print</span><span class="p">(</span><span class="s">"addressable shard devices:"</span><span class="p">,</span> <span class="p">[</span><span class="n">s</span><span class="p">.</span><span class="n">device</span> <span class="k">for</span> <span class="n">s</span> <span class="ow">in</span> <span class="n">y</span><span class="p">.</span><span class="n">addressable_shards</span><span class="p">])</span>
</code></pre></div></div>

<p>如果 <code class="language-plaintext highlighter-rouge">device</code> 或 <code class="language-plaintext highlighter-rouge">addressable_shards</code> 对应的是 <code class="language-plaintext highlighter-rouge">TpuDevice(...)</code>，说明这一步结果确实落在 TPU 上。</p>

<h4 id="方法-2显式同步-block_until_ready">方法 2：显式同步 <code class="language-plaintext highlighter-rouge">block_until_ready()</code></h4>

<p>这是判断 TPU 计算是否<strong>真的完成</strong>时最容易漏掉的一步。</p>

<p>JAX 默认是 <strong>asynchronous dispatch</strong>。也就是说：</p>

<ul>
  <li>Python 线程可能只是把任务提交给 TPU</li>
  <li>任务还在设备端排队或执行</li>
  <li>主线程已经继续往下跑了</li>
</ul>

<p>所以像下面这种代码：</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">step</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="s">"done"</span><span class="p">)</span>
</code></pre></div></div>

<p>很多时候只能说明 <strong>dispatch 完成了</strong>，不能说明 <strong>TPU 已经算完了</strong>。</p>

<p>如果要做准确判断或准确计时，应该这样写：</p>

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

<span class="o">@</span><span class="n">jax</span><span class="p">.</span><span class="n">jit</span>
<span class="k">def</span> <span class="nf">step</span><span class="p">(</span><span class="n">x</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">x</span> <span class="o">@</span> <span class="n">x</span>

<span class="n">x</span> <span class="o">=</span> <span class="n">jnp</span><span class="p">.</span><span class="n">ones</span><span class="p">((</span><span class="mi">4096</span><span class="p">,</span> <span class="mi">4096</span><span class="p">))</span>

<span class="n">t0</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span>
<span class="n">y</span> <span class="o">=</span> <span class="n">step</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="n">t1</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span>

<span class="n">y</span><span class="p">.</span><span class="n">block_until_ready</span><span class="p">()</span>
<span class="n">t2</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span>

<span class="k">print</span><span class="p">(</span><span class="s">"dispatch time:"</span><span class="p">,</span> <span class="n">t1</span> <span class="o">-</span> <span class="n">t0</span><span class="p">)</span>
<span class="k">print</span><span class="p">(</span><span class="s">"real execution time:"</span><span class="p">,</span> <span class="n">t2</span> <span class="o">-</span> <span class="n">t0</span><span class="p">)</span>
</code></pre></div></div>

<p>通常：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">dispatch time</code> 只是任务提交时间</li>
  <li><code class="language-plaintext highlighter-rouge">real execution time</code> 才更接近真实 TPU 执行时间</li>
</ul>

<h4 id="方法-3需要铁证时用-profiler">方法 3：需要铁证时，用 profiler</h4>

<p>如果要看更硬的证据，直接做 profile。</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">jax</span>

<span class="n">jax</span><span class="p">.</span><span class="n">profiler</span><span class="p">.</span><span class="n">start_server</span><span class="p">(</span><span class="mi">9999</span><span class="p">)</span>
</code></pre></div></div>

<p>或者：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">jax</span>

<span class="n">jax</span><span class="p">.</span><span class="n">profiler</span><span class="p">.</span><span class="n">start_trace</span><span class="p">(</span><span class="s">"/tmp/profile-data"</span><span class="p">)</span>
<span class="c1"># run workload
</span><span class="n">jax</span><span class="p">.</span><span class="n">profiler</span><span class="p">.</span><span class="n">stop_trace</span><span class="p">()</span>
</code></pre></div></div>

<p>这时候关心的就不是“数组在哪”，而是 trace 里有没有真正的 TPU activity。</p>

<h3 id="14-监控时最容易误判的三件事">1.4 监控时最容易误判的三件事</h3>

<h4 id="误判-1tpu-info-有芯片信息就以为-tpu-一定在正常工作">误判 1：<code class="language-plaintext highlighter-rouge">tpu-info</code> 有芯片信息，就以为 TPU 一定在正常工作</h4>

<p>不对。<br />
这最多说明<strong>设备存在</strong>，不说明 <strong>runtime metrics 可用</strong>，更不说明 <strong>JAX 已经在跑 TPU 计算</strong>。</p>

<h4 id="误判-2程序很快打印-done就以为-tpu-已经算完">误判 2：程序很快打印 <code class="language-plaintext highlighter-rouge">done</code>，就以为 TPU 已经算完</h4>

<p>不对。<br />
这通常只是 asynchronous dispatch 的表现，必须配合 <code class="language-plaintext highlighter-rouge">block_until_ready()</code> 看。</p>

<h4 id="误判-3报-the-tpu-is-already-in-use-by-process-with-pid-以为是-jax-坏了">误判 3：报 <code class="language-plaintext highlighter-rouge">The TPU is already in use by process with pid ...</code>，以为是 JAX 坏了</h4>

<p>这类报错的核心含义其实很简单：</p>

<blockquote>
  <p><strong>另一独立进程已经先初始化并占住了 TPU runtime。</strong></p>
</blockquote>

<p>常用检查命令：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>ps <span class="nt">-fp</span> 22589
<span class="nb">sudo </span>lsof <span class="nt">-w</span> /dev/vfio/<span class="k">*</span>
tpu-info
</code></pre></div></div>

<p>这通常是另一个脚本、另一个终端、<code class="language-plaintext highlighter-rouge">tmux</code>、notebook kernel 还没退出，不是 TPU “坏了”。</p>

<h3 id="15-一个最小但可靠的监控顺序">1.5 一个最小但可靠的监控顺序</h3>

<p>如果我只想快速判断 TPU 状态，我会按这个顺序看：</p>

<ol>
  <li><code class="language-plaintext highlighter-rouge">tpu-info</code>
先确认机器层面有没有 TPU 设备。</li>
  <li><code class="language-plaintext highlighter-rouge">jax.default_backend()</code> 和 <code class="language-plaintext highlighter-rouge">jax.devices()</code>
确认当前 JAX 进程有没有成功连到 TPU backend。</li>
  <li>输出数组的 <code class="language-plaintext highlighter-rouge">device / sharding / addressable_shards</code>
确认具体计算结果是不是落在 TPU 上。</li>
  <li><code class="language-plaintext highlighter-rouge">block_until_ready()</code>
确认这一步不是只完成 dispatch，而是真的执行完。</li>
  <li>profiler
需要最硬证据时再上。</li>
</ol>

<hr />

<h2 id="2-hlo--llo--compile--backend-到底是什么关系">2. HLO / LLO / compile / backend 到底是什么关系</h2>

<p>如果只想抓主线，可以先记这一条：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>trace -&gt; lower -&gt; HLO -&gt; XLA 优化 -&gt; backend lowering(LLO) -&gt; executable
</code></pre></div></div>

<p>这一节其实就是在解释这条链路，以及它在项目里意味着什么。</p>

<h3 id="21-四个词先分清">2.1 四个词先分清</h3>

<h4 id="hlo-是什么">HLO 是什么</h4>

<p>HLO 可以理解成 XLA 里的<strong>高层中间表示</strong>。<br />
它更接近张量计算图，也更适合看：</p>

<ul>
  <li>图优化</li>
  <li>fusion</li>
  <li>sharding</li>
  <li>layout 传播</li>
</ul>

<p>在项目里，下面两种写法都会碰到 HLO：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">lowered_func</span><span class="p">.</span><span class="n">as_text</span><span class="p">(</span><span class="s">"hlo"</span><span class="p">)</span>
<span class="n">lowered_func</span><span class="p">.</span><span class="nb">compile</span><span class="p">().</span><span class="n">as_text</span><span class="p">()</span>
</code></pre></div></div>

<p>区别是：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">as_text("hlo")</code> 更像是直接查看 lower 后的 HLO</li>
  <li><code class="language-plaintext highlighter-rouge">compile().as_text()</code> 则是在真正编译之后查看结果</li>
</ul>

<h4 id="compile-是什么">compile() 是什么</h4>

<p><code class="language-plaintext highlighter-rouge">compile()</code> 不是“把内容打印出来”，而是：</p>

<ul>
  <li>真正触发一次编译</li>
  <li>让 lowered computation 继续进入 XLA 和 backend 的编译链</li>
  <li>最终生成可执行对象</li>
</ul>

<p>所以这句代码：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">lowered_func</span><span class="p">.</span><span class="nb">compile</span><span class="p">().</span><span class="n">as_text</span><span class="p">()</span>
</code></pre></div></div>

<p>它的关键点不在 <code class="language-plaintext highlighter-rouge">as_text()</code>，而在 <strong><code class="language-plaintext highlighter-rouge">compile()</code> 已经把后面的编译过程跑起来了</strong>。</p>

<h4 id="backend-是什么">backend 是什么</h4>

<p>这里的 backend 可以理解为：</p>

<blockquote>
  <p><strong>负责把编译器中间表示继续变成目标设备可执行结果的那一层。</strong></p>
</blockquote>

<p>在这个问题里，主要就是 TPU backend / <code class="language-plaintext highlighter-rouge">libtpu</code>。</p>

<h4 id="llo-是什么">LLO 是什么</h4>

<p>LLO 可以粗略理解成<strong>更靠近后端和设备的一层表示</strong>。<br />
它的位置比 HLO 更低，更接近最终 executable。</p>

<p>所以可以把它粗略记成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>HLO -&gt; 更低层后端表示(LLO) -&gt; executable
</code></pre></div></div>

<p>LLO 通常不是项目自己构造的数据结构，而是 TPU backend / <code class="language-plaintext highlighter-rouge">libtpu</code> 在编译过程中额外吐出来的调试产物。</p>

<h3 id="22-这对当前项目意味着什么">2.2 这对当前项目意味着什么</h3>

<p>当前项目里，如果只是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">lowered_func</span><span class="p">.</span><span class="n">as_text</span><span class="p">(</span><span class="s">"hlo"</span><span class="p">)</span>
</code></pre></div></div>

<p>更偏向“看一份 HLO 文本”。</p>

<p>如果是：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">lowered_func</span><span class="p">.</span><span class="nb">compile</span><span class="p">().</span><span class="n">as_text</span><span class="p">()</span>
</code></pre></div></div>

<p>那就不只是看文本了，而是会真的触发一次编译过程。<br />
也正因为如此，编译链后面的很多东西都会被触发出来，比如：</p>

<ul>
  <li>优化后的 HLO</li>
  <li><code class="language-plaintext highlighter-rouge">XLA_FLAGS</code> 对应的 HLO pass dump</li>
  <li><code class="language-plaintext highlighter-rouge">LIBTPU_INIT_ARGS</code> 对应的 LLO dump</li>
</ul>

<p>所以项目里“导出优化后 HLO”这件事，本质上已经进入了<strong>真实编译</strong>，而不只是字符串导出。</p>

<h3 id="23-为什么会同时看到-xla_dump-和-llo_dump">2.3 为什么会同时看到 <code class="language-plaintext highlighter-rouge">xla_dump</code> 和 <code class="language-plaintext highlighter-rouge">llo_dump</code></h3>

<p>因为一次 <code class="language-plaintext highlighter-rouge">compile()</code> 本来就会经过多层编译链路。</p>

<p>如果设置：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nv">XLA_FLAGS</span><span class="o">=</span><span class="s2">"--xla_dump_to=/tmp/xla_dump --xla_dump_hlo_pass_re=.*"</span>
<span class="nv">LIBTPU_INIT_ARGS</span><span class="o">=</span><span class="s2">"--xla_jf_dump_to=/tmp/dump_llo"</span>
</code></pre></div></div>

<p>那么通常会同时出现两类 dump：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">xla_dump</code>
更偏 XLA / HLO pass，适合看编译器中高层是怎么改图的。</li>
  <li><code class="language-plaintext highlighter-rouge">llo_dump</code>
更偏 TPU backend / <code class="language-plaintext highlighter-rouge">libtpu</code>，适合看更靠近设备的一层 lowering。</li>
</ul>

<p>它们不是重复关系，而是编译链上<strong>不同阶段</strong>的调试产物。</p>

<h3 id="24-xla_flags-和-libtpu_init_args-分别控制什么">2.4 <code class="language-plaintext highlighter-rouge">XLA_FLAGS</code> 和 <code class="language-plaintext highlighter-rouge">LIBTPU_INIT_ARGS</code> 分别控制什么</h3>

<p>如果只看用途，可以直接记成：</p>

<ul>
  <li><strong><code class="language-plaintext highlighter-rouge">XLA_FLAGS</code></strong>：看 HLO / XLA pass</li>
  <li><strong><code class="language-plaintext highlighter-rouge">LIBTPU_INIT_ARGS</code></strong>：看 TPU backend / LLO lowering</li>
</ul>

<p>更重要的一点是：</p>

<blockquote>
  <p><strong>它们通常不是在改变最终 HLO 的语义，而是在让你看到更多编译过程中的中间状态。</strong></p>
</blockquote>

<p>所以同时打开时，通常会有三类结果：</p>

<ol>
  <li>项目自己主动导出的 HLO 文本</li>
  <li><code class="language-plaintext highlighter-rouge">/tmp/xla_dump</code> 里的 HLO pass dump</li>
  <li><code class="language-plaintext highlighter-rouge">/tmp/dump_llo</code> 里的 TPU/LLO dump</li>
</ol>

<p>前者适合做稳定分析，后两者适合排查“图到底是在哪一层被改掉的”。</p>

<h3 id="25-为什么不能在每个-layer-的-compile-前动态改-libtpu_init_args">2.5 为什么不能在每个 layer 的 <code class="language-plaintext highlighter-rouge">compile()</code> 前动态改 <code class="language-plaintext highlighter-rouge">LIBTPU_INIT_ARGS</code></h3>

<p>直觉上很容易想到这种写法：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">os</span><span class="p">.</span><span class="n">environ</span><span class="p">[</span><span class="s">"LIBTPU_INIT_ARGS"</span><span class="p">]</span> <span class="o">=</span> <span class="s">"..."</span>
<span class="n">lowered_func</span><span class="p">.</span><span class="nb">compile</span><span class="p">()</span>
<span class="n">os</span><span class="p">.</span><span class="n">environ</span><span class="p">.</span><span class="n">pop</span><span class="p">(</span><span class="s">"LIBTPU_INIT_ARGS"</span><span class="p">,</span> <span class="bp">None</span><span class="p">)</span>
</code></pre></div></div>

<p>但这通常不可靠，原因很简单：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">LIBTPU_INIT_ARGS</code> 属于 backend / runtime 初始化参数</li>
  <li>它往往在 backend 初始化时就已经读取了</li>
  <li>不是每次 <code class="language-plaintext highlighter-rouge">compile()</code> 都重新读一次</li>
</ul>

<p>所以程序跑到 <code class="language-plaintext highlighter-rouge">save_hlo_and_data()</code> 时，backend 很可能早就初始化完了。<br />
这时再临时改环境变量，大概率不会影响当前这次 <code class="language-plaintext highlighter-rouge">compile()</code>。</p>

<p>因此它更像是：</p>

<ul>
  <li><strong>进程启动前配置</strong></li>
  <li><strong>backend 初始化前配置</strong></li>
</ul>

<p>而不是 layer 级别的 Python 参数。</p>

<h3 id="26-这次修改真正解决的问题">2.6 这次修改真正解决的问题</h3>

<p>当前项目在导出优化后 HLO 时，会触发真实编译：</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">lowered_func</span><span class="p">.</span><span class="nb">compile</span><span class="p">().</span><span class="n">as_text</span><span class="p">()</span>
</code></pre></div></div>

<p>这次修改的目标不是改变编译行为本身，而是：</p>

<ul>
  <li>继续保留全局 <code class="language-plaintext highlighter-rouge">LIBTPU_INIT_ARGS="--xla_jf_dump_to=..."</code></li>
  <li>在每个 layer 的 <code class="language-plaintext highlighter-rouge">compile()</code> 前后观察全局 LLO dump 目录</li>
  <li>把本轮新增或变化的文件归档到当前 layer 目录</li>
</ul>

<p>也就是把“全局 dump”尽量整理成“按 layer 可读”。</p>

<h3 id="27-当前项目采用的最小方案">2.7 当前项目采用的最小方案</h3>

<p>因为 <code class="language-plaintext highlighter-rouge">LIBTPU_INIT_ARGS</code> 只能全局生效，所以当前最稳妥的办法不是“让 libtpu 直接按 layer dump”，而是：</p>

<blockquote>
  <p><strong>全局 dump，再由项目代码按每次 <code class="language-plaintext highlighter-rouge">compile()</code> 的时间窗口做归档。</strong></p>
</blockquote>

<p>具体流程是：</p>

<ol>
  <li>compile 前读取全局 LLO dump 根目录</li>
  <li>先做一次快照</li>
  <li>执行当前 layer 的 <code class="language-plaintext highlighter-rouge">compile()</code></li>
  <li>compile 后再做一次快照</li>
  <li>找出本轮新增或变化的文件</li>
  <li>复制到 <code class="language-plaintext highlighter-rouge">output/&lt;layer_name&gt;/llo_dump</code></li>
</ol>

<p>这个方案的优点是：</p>

<ul>
  <li>不改 backend 初始化逻辑</li>
  <li>不需要多进程</li>
  <li>改动面小</li>
</ul>

<p>它的限制也很明确：</p>

<ul>
  <li>这是按“时间窗口”归因，不是按文件内容精确识别 layer</li>
  <li>如果 backend 异步写文件，可能有少量归因误差</li>
  <li>如果多个重复 block 共用同一个名字，仍然不容易区分实例</li>
</ul>

<h3 id="28-运行时怎么用">2.8 运行时怎么用</h3>

<p>如果要让 LLO 归档生效，运行前需要先设置全局环境变量：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nv">LIBTPU_INIT_ARGS</span><span class="o">=</span><span class="s2">"--xla_jf_dump_to=/tmp/dump_llo"</span> <span class="se">\</span>
uv run entrypoints.py <span class="se">\</span>
  <span class="nt">--model_type</span> llama3 <span class="se">\</span>
  <span class="nt">--config-path</span> configs/llama/llama3-8B-sq4k-bf16-L1.ini <span class="se">\</span>
  <span class="nt">--output_dir</span> ./outputs/ <span class="se">\</span>
  <span class="nt">--export_after_optimize</span>
</code></pre></div></div>

<p>如果还想同时看 HLO pass dump，再加：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nv">XLA_FLAGS</span><span class="o">=</span><span class="s2">"--xla_dump_to=/tmp/xla_dump --xla_dump_hlo_pass_re=.*"</span>
</code></pre></div></div>

<p>运行后通常会有两类目录：</p>

<ul>
  <li>全局原始目录：<code class="language-plaintext highlighter-rouge">/tmp/dump_llo</code></li>
  <li>项目归档目录：<code class="language-plaintext highlighter-rouge">./outputs/&lt;config_name&gt;/&lt;layer_name&gt;/llo_dump</code></li>
</ul>

<hr />

<h2 id="3-最后压缩成几句话">3. 最后压缩成几句话</h2>

<ol>
  <li><strong>监控 TPU 要分三层：设备是否存在、JAX 是否连上、具体计算是否真的执行。</strong></li>
  <li><strong><code class="language-plaintext highlighter-rouge">compile()</code> 会真正触发编译，所以它不只是在“导出文本”，也会带出后端 dump。</strong></li>
  <li><strong>HLO 更偏 XLA 高层表示，LLO 更偏 TPU backend 更低层的 lowering。</strong></li>
  <li><strong><code class="language-plaintext highlighter-rouge">XLA_FLAGS</code> 主要看 HLO/XLA pass，<code class="language-plaintext highlighter-rouge">LIBTPU_INIT_ARGS</code> 主要看 TPU backend/LLO。</strong></li>
  <li><strong>当前项目不能让 libtpu 直接按 layer dump，所以采用的是“全局 dump + 每次 compile 后按 layer 归档”。</strong></li>
</ol>]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[JAX + TPU 两个重点问题]]></summary></entry><entry><title type="html">FLUX.1梳理</title><link href="https://wqh011128.github.io/blog/2026/03/25/FLUX.1.html" rel="alternate" type="text/html" title="FLUX.1梳理" /><published>2026-03-25T00:00:00+00:00</published><updated>2026-03-25T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/03/25/FLUX.1</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/03/25/FLUX.1.html"><![CDATA[<h1 id="flux1-梳理">FLUX.1 梳理</h1>

<h2 id="1-flow-matching">1. Flow Matching</h2>

<p>Flow matching的核心在于——<strong>跳过加噪的过程</strong>。</p>

<blockquote>
  <p>常识建立：对于扩散模型的生成来说，假设$x$ 是原始分布，$x_1$ 是目标分布，$t$是timestep。则要求$x$ 能在$t$个step内将分布转移到$x_1$内，$x$ $\rightarrow$ $x_1$。</p>
</blockquote>

<p>下面给出Flow matching的关键概念和直观理解。</p>

<h2 id="11-ode-与-sde">1.1 ODE 与 SDE</h2>

<p>ODE 的全称是 <strong>Ordinary Differential Equation</strong>，即<strong>常微分方程</strong>。<br />
它用于描述一个变量随时间变化的确定性规律。一般形式可以写为：
\(\frac{\mathrm{d}x}{\mathrm{d}t} = f(x,t)\)</p>

<p>其中：</p>

<ul>
  <li>$x$ 表示系统状态，在扩散模型中可以认为是某个时刻的数据分布；</li>
  <li>$t$ 表示时间；</li>
  <li>$f(x,t)$ 表示状态在当前时刻的变化率。</li>
</ul>

<p>ODE 的特点是：在给定初始条件后，系统的演化轨迹是确定的，<strong>不包含随机扰动</strong>。</p>

<p>SDE 的全称是 <strong>Stochastic Differential Equation</strong>，即<strong>随机微分方程</strong>。<br />
它用于描述既受到确定性规律影响、又受到随机噪声扰动的动态系统。一般形式可以写为：
\(\mathrm{d}x = f(x,t)\,\mathrm{d}t + g(x,t)\,\mathrm{d}w\)</p>

<p>其中：</p>

<ul>
  <li>$f(x,t)\,\mathrm{d}t$ 是确定性的漂移项（drift term）；</li>
  <li>$g(x,t)\,\mathrm{d}w$ 是随机扩散项（diffusion term）；</li>
  <li>$w$ 表示 Brownian motion（布朗运动）。</li>
</ul>

<p>SDE 的特点是：即使初始条件相同，由于<strong>存在随机项</strong>，系统轨迹通常也不是唯一的。</p>

<hr />

<h2 id="12-从ode-出发理解flow-matching">1.2 从ODE 出发，理解Flow Matching</h2>

<p>再来看到ODE的公式：
\(\frac{\mathrm{d}x}{\mathrm{d}t} = f(x,t)\)
那么，x对t的求导 $\frac{\mathrm{d}x}{\mathrm{d}t}$ 实际上表示状态变量 $x$ 关于时间 $t$ 的瞬时变化率，也就是$x$的变化方向，规定叫做Vector field。</p>

<blockquote>
  <p>在下图中，分布周围的蓝色箭头就是$f(x,t)$，可以直观理解为 “风吹着分布$x$从init状态$\rightarrow$target状态”。</p>
</blockquote>

<p><img src="/img/flow.png" alt="" /></p>

<p>规定$t\in [0,1]$，且将上面三张图分别定义init状态（t=0）、中间状态t时刻、和最终状态（t=1）。</p>

<p>则每一个step $t$时刻，数据的分布为$x_t$，vector field为$\frac{\mathrm{d}x_t}{\mathrm{d}t}$。</p>

<p>因此，可以得到一个初步结论：</p>

<p>在$t$个step内，从$x\rightarrow x_1$存在路径path，并且$f(x,t)$决定着$x_t$的走向。</p>

<blockquote>
  <p>[!Note]</p>

  <p>所以，在扩散模型中，如果我们要将$x$生成到某个分布（例如猫和狗），意味着$x$和$x_1$已经确定，我们需要求解出合适的vector field，使得x移动的路径满足要求。</p>
</blockquote>

<p>那么，$f(x,t)$该如何求解？先来看下面四张图：</p>

<blockquote>
  <p>图中的$\mu$就是$f$</p>
</blockquote>

<p><img src="/img/flow2.png" alt="" /></p>

<ul>
  <li>（a）给定目标$x_1$情况下的Path条件分布，也即x的移动路径；</li>
  <li>（b）不给定目标情况下，x移动路径的边缘分布，等于全部$x_1$下的期望；</li>
  <li>（c）给定目标$x_1$情况下的vector field条件分布；</li>
  <li>（d）不给定目标情况下，vector field的边缘分布；</li>
</ul>

<p>对于（a），分布变化的过程等价于求解一条从$N(0,\sigma^2I)\rightarrow N(x_1,0)$的条件概率路径，</p>

<p>把正态分布参数化，定义成下面的形式：
\(P_t(x|x_1):N(\alpha_{t}\cdot x_1,\beta_t^2\cdot I)\)
这样的好处是，可以将求解的过程转化成线性求解：</p>

<p><img src="/img/alphabeta.png" alt="" /></p>

<p>可以求出条件 vector field：</p>

\[\begin{aligned}
\frac{P_t(x\mid x_1)}{\mathrm{d}t}
&amp;= f_t(x\mid x_1) \\
&amp;= \frac{\beta_t^\prime}{\beta_t}(x-\alpha_t) + \alpha_t^\prime
\end{aligned}\]

<p>边缘分布等于条件分布的期望，先铺垫以下概念：</p>

<ol>
  <li>Divergence：在当前t step x的位置，流出的量减去流入的量（直观理解为当前位置，可以出去的方向数 减去 可以进来的方向数）</li>
</ol>

\[\operatorname{div}(v_t)(x)=\sum_{i=1}^{d}\frac{\partial}{\partial x_i}v_t(x)\]

<ol>
  <li>Continuity Equation：质量守恒。</li>
</ol>

\[\frac{\mathrm{d}}{\mathrm{d}t}p_t(x)+\operatorname{div}\bigl(p_t\cdot f_t\bigr)(x)=0\]

<blockquote>
  <p>直观理解：</p>

  <ul>
    <li>
      <p>当这个位置的散度（流出的量-流入的量）为0时，也就是$\operatorname{div}\bigl(p_t\cdot f_t\bigr)(x)$=0时，物体的质量随时间不会发生改变。</p>
    </li>
    <li>如果散度大于0，那么质量会随时间减小，$\frac{\mathrm{d}}{\mathrm{d}t}p_t(x)$&lt;0。</li>
    <li>如果散度小于0，那么质量会随时间变大，$\frac{\mathrm{d}}{\mathrm{d}t}p_t(x)$&gt;0。</li>
  </ul>
</blockquote>

<hr />

<p>Proof:</p>

\[\begin{aligned}
\frac{\mathrm{d}}{\mathrm{d}t}p_t(x)
&amp;=
\frac{\mathrm{d}}{\mathrm{d}t}
\int p_t(x\mid x_1)\,p(x_1)\,\mathrm{d}x_1 \\
&amp;=
\int
\frac{\mathrm{d}}{\mathrm{d}t}p_t(x\mid x_1)\,p(x_1)\,\mathrm{d}x_1 \\
&amp;=
\int
\Bigl[
-\operatorname{div}\bigl(p_t(x\mid x_1)\,f_t(x\mid x_1)\bigr)
\Bigr]
p(x_1)\,\mathrm{d}x_1 \\
&amp;=
-\operatorname{div}
\int
p_t(x\mid x_1)\,f_t(x\mid x_1)\,p(x_1)\,\mathrm{d}x_1 \\
&amp;=
-\operatorname{div}
\left[
p_t(x)
\int
f_t(x\mid x_1)\,
\frac{p_t(x\mid x_1)\,p(x_1)}{p_t(x)}
\,\mathrm{d}x_1
\right] \\
&amp;=
-\operatorname{div}\bigl(p_t\cdot f_t\bigr)(x)
\end{aligned}\]

<p>Therefore,</p>

\[\begin{aligned}
f_t(x)
&amp;= \int
f_t(x\mid x_1)\,
\frac{p_t(x\mid x_1)\,p(x_1)}{p_t(x)}
\,\mathrm{d}x_1
\end{aligned}\]

<hr />

<h2 id="13-how-to-train-a-flow-matching-model">1.3 How to train a Flow-Matching Model</h2>

<p>在上一部分中，我们已经得到边缘向量场的表达式：</p>

\[\begin{aligned}
f_t(x)
&amp;= \int f_t(x\mid x_1)\,
\frac{p_t(x\mid x_1)p(x_1)}{p_t(x)}
\,\mathrm{d}x_1 \\
&amp;= \mathbb{E}_{x_1\sim p(x_1\mid x,t)}\bigl[f_t(x\mid x_1)\bigr].
\end{aligned}\]

<p>这说明，边缘向量场 $f_t(x)$ 本质上是条件向量场 $f_t(x\mid x_1)$ 关于后验分布的条件期望。</p>

<p>我们的目标是训练一个<strong>参数化模型 $f_t^\theta(x)$</strong>，使其能够逼近真实的边缘向量场 $f_t(x)$。<br />
一个自然的训练目标是最小化如下损失：
\(\begin{aligned}
\mathcal{L}_{\mathrm{fm}}(\theta)
&amp;= \mathbb{E}_{t,x_1,x}
\left[
\left\|f_t^\theta(x)-f_t(x)\right\|^2
\right].
\end{aligned}\)</p>

<p>但是这个目标通常是 <strong>intractable</strong> 的，因为 $f_t(x)$ 本身往往无法直接计算。</p>

<hr />

<p>为了解决上述问题，考虑定义条件 Flow Matching 损失：
\(\begin{aligned}
\mathcal{L}_{\mathrm{cfm}}(\theta)
&amp;= \mathbb{E}_{t,x_1,x}
\left[
\left\|f_t^\theta(x)-f_t(x\mid x_1)\right\|^2
\right].
\end{aligned}\)</p>

<p>下面说明两个损失函数 <code class="language-plaintext highlighter-rouge">L_fm</code> 和 <code class="language-plaintext highlighter-rouge">L_cfm</code> 仅相差一个与参数 <code class="language-plaintext highlighter-rouge">theta</code> 无关的常数，因此二者具有相同的最优解。</p>

<p><strong>1. 展开 <code class="language-plaintext highlighter-rouge">L_fm</code></strong></p>

\[\begin{aligned}
\mathcal{L}_{\mathrm{fm}}(\theta)
&amp;=
\mathbb{E}_{t,x_1,x}
\left[
\|f_t^\theta(x)-f_t(x)\|^2
\right] \\
&amp;=
\mathbb{E}_{t,x_1,x}
\left[
\|f_t^\theta(x)\|^2
-2f_t^\theta(x)^\top f_t(x)
+\|f_t(x)\|^2
\right].
\end{aligned}\]

<p><strong>2. 展开 <code class="language-plaintext highlighter-rouge">L_cfm</code></strong></p>

\[\begin{aligned}
\mathcal{L}_{\mathrm{cfm}}(\theta)
&amp;=
\mathbb{E}_{t,x_1,x}
\left[
\|f_t^\theta(x)-f_t(x\mid x_1)\|^2
\right] \\
&amp;=
\mathbb{E}_{t,x_1,x}
\left[
\|f_t^\theta(x)\|^2
-2\,f_t^\theta(x)^\top f_t(x\mid x_1)
+\|f_t(x\mid x_1)\|^2
\right].
\end{aligned}\]

<p><strong>3. 利用边缘向量场的定义</strong></p>

<p>由上一节结论，</p>

\[\begin{aligned}
f_t(x)
&amp;= \int p(x_1\mid x,t)\,f_t(x\mid x_1)\,\mathrm{d}x_1 \\
&amp;= \mathbb{E}_{x_1\sim p(x_1\mid x,t)}\bigl[f_t(x\mid x_1)\bigr].
\end{aligned}\]

<p>因此，对于固定的 $x,t$，有</p>

\[\mathbb{E}_{x_1\sim p(x_1\mid x,t)}
\left[
f_t(x)-f_t(x\mid x_1)
\right]
=0.\]

<p>进一步可得</p>

\[\mathbb{E}_{t,x}
\left[
f_t^\theta(x)^\top
\mathbb{E}_{x_1\sim p(x_1\mid x,t)}
\bigl[f_t(x)-f_t(x\mid x_1)\bigr]
\right]
=0.\]

<p>也就是</p>

\[\begin{aligned}
\mathbb{E}_{t,x_1,x}
\left[
f_t^\theta(x)^\top f_t(x)
\right]
&amp;= \mathbb{E}_{t,x_1,x}
\left[
f_t^\theta(x)^\top f_t(x\mid x_1)
\right].
\end{aligned}\]

<p><strong>4. 两个损失之差</strong></p>

<p>于是，
\(\begin{aligned}
\mathcal{L}_{\mathrm{fm}}(\theta)-\mathcal{L}_{\mathrm{cfm}}(\theta)
&amp;= \mathbb{E}_{t,x_1,x}
\left[
\|f_t(x)\|^2-\|f_t(x\mid x_1)\|^2
\right].
\end{aligned}\)</p>

<p>右边不含参数 $\theta$，因此它只是一个常数。于是有</p>

\[\begin{aligned}
\mathcal{L}_{\mathrm{fm}}(\theta)
&amp;= \mathcal{L}_{\mathrm{cfm}}(\theta)+\text{constant}.
\end{aligned}\]

<p>因此，</p>

\[\begin{aligned}
\arg\min_\theta \mathcal{L}_{\mathrm{fm}}(\theta)
&amp;= \arg\min_\theta \mathcal{L}_{\mathrm{cfm}}(\theta).
\end{aligned}\]

<hr />

<p>所以，虽然我们原本想训练的是边缘向量场 $u_t(x)$，但实际中可以转而优化如下可计算的目标：</p>

\[\begin{aligned}
\mathcal{L}(\theta)
&amp;= \mathbb{E}_{t,x_1,x}
\left[
\left\|f_t^\theta(x)-f_t(x\mid x_1)\right\|^2
\right].
\end{aligned}\]

<p>这就是 Flow Matching / Conditional Flow Matching 的核心训练形式。</p>

<p>它的含义是：</p>

<ul>
  <li>我们真正想拟合的是边缘向量场 $f_t(x)$；</li>
  <li>但由于 $f_t(x)$ 不易直接计算，于是改为拟合条件向量场 $f_t(x\mid x_1)$；</li>
  <li>由于二者对应的损失只差一个与参数无关的常数，因此训练条件目标就等价于训练边缘目标。</li>
</ul>

<h2 id="14-pseudo-codes-of-training-and-inference">1.4 Pseudo codes of training and inference</h2>

<blockquote>
  <p>线性情况下，
$\frac{P_t(x|x_1)}{\mathrm{d}t}=f_t(x|x_1)=\frac{\beta_t^\prime}{\beta_t}(x-\alpha_t) + \alpha_t^\prime$</p>
</blockquote>

<p><strong>Training</strong></p>

<p><strong>1.</strong> 从数据集中采样一个图像 $x_1$。</p>

<p><strong>2.</strong> 随机采样一个 timestep $t\in[0,1]$，以及一个 noise $\sigma\in N(0,I)$。</p>

<p><strong>3.</strong> 给定 $x_1$ 下的条件路径分布为</p>

\[p_t(x\mid x_1)=t x_1 + (1-t)\sigma\]

<p>它服从分布</p>

\[N(\alpha_t\cdot x_1,\beta_t^2\cdot I),\]

<p>所以 $\alpha_t=t,\beta_t=1-t$。</p>

<p><strong>4.</strong> 因此</p>

\[f_t(x\mid x_1)=x_1-\sigma\]

<p>或者你直接对 $t$ 求导，也是这个结果，不用管服从什么分布。</p>

<p><strong>5.</strong> 所以损失函数为</p>

\[\mathcal{L}(\theta)=\left\|f_t^\theta(x)-(x_1-\sigma)\right\|^2.\]

<p><strong>6.</strong> BP优化，最终模型能够拟合vector field。可以注意到不论 $t$ 采样多少，监督信号都是不变的。</p>

<p><strong>Inference</strong></p>

<p><strong>1.</strong> 初始化 $t=0$，step size 为 $h=\frac{1}{n}$，其中 $n$ 为离散步数。</p>

<p><strong>2.</strong> 令 $x_0$ 为初始状态。</p>

<p><strong>3.</strong> 对 $i=1,\dots,n-1$ 循环执行：</p>

\[x_{t+h}=x_t+f_t^\theta(x_t)\cdot h\]

\[t\leftarrow t+h\]

<p><strong>4.</strong> 返回 $x_1$。</p>

<h2 id="2-dit">2. DiT</h2>

<p><img src="/img/dit.png" alt="" /></p>

<p>主要从这张图看一下时间步这种条件是如何影响模型的生成的，也就是$f_t^\theta(x_t)$怎么接受$x$和$t$。</p>

<p>在 DiT 中，若采用 <strong>adaLN-Zero</strong> 机制，则模型虽然形式上可写为
$f_t^\theta(x_t)$，
但更准确地说，它实际还依赖于时间步条件以及其他外部条件，因此应理解为
$f_\theta(x_t, t, c)$。</p>

<p>其中：</p>

<ul>
  <li>$x_t$ 表示时刻 $t$ 的带噪 latent；</li>
  <li>$t$ 表示时间步或噪声强度；</li>
  <li>$c$ 表示附加条件，例如类别标签或文本条件。</li>
</ul>

<ol>
  <li>$x_t$ 如何进入模型</li>
</ol>

<p>输入的带噪 latent $x_t$ 先经过 patchify，被划分为一系列 token，作为 Transformer 主干网络的输入。<br />
因此，$x_t$ 提供的是当前样本的空间内容信息，是模型生成更新的基础。</p>

<ol>
  <li>$t$ 和条件 $c$ 如何进入模型</li>
</ol>

<p>在 adaLN-Zero 中，时间步 $t$ 以及其他条件 $c$ 不直接作为 token 与输入序列拼接，也不通过 cross-attention 进入主干，而是先经过 embedding 和 MLP，生成一组用于调制 Transformer block 的参数。</p>

<p>设条件向量记为</p>

\[e = \phi(t, c),\]

<p>其中 $\phi(\cdot)$ 表示由时间步嵌入、条件嵌入及后续 MLP 构成的映射。<br />
随后，模型根据该条件向量生成每个 block 所需的调制参数，例如：</p>

\[(\alpha_1,\beta_1,\gamma_1,\alpha_2,\beta_2,\gamma_2)=\mathrm{MLP}(e).\]

<p>这些参数分别作用于 block 中的 attention 分支和 feedforward 分支。</p>

<ol>
  <li>adaLN-Zero 的具体作用方式</li>
</ol>

<p>对于输入特征 $h$，普通 LayerNorm 可写为</p>

\[\mathrm{LN}(h).\]

<p>在 adaLN 中，LayerNorm 的输出会被条件相关参数进一步调制：</p>

\[\mathrm{adaLN}(h,e)=\gamma(e)\odot \mathrm{LN}(h)+\beta(e),\]

<p>其中：</p>

<ul>
  <li>$\gamma(e)$ 表示缩放参数（scale）；</li>
  <li>$\beta(e)$ 表示平移参数（shift）；</li>
  <li>$\odot$ 表示逐元素乘法。</li>
</ul>

<p>此外，残差分支还会通过门控参数 $\alpha(e)$ 进行控制，因此对应的更新形式可以写为</p>

\[h+\alpha(e)\cdot F\bigl(\mathrm{adaLN}(h,e)\bigr),\]

<p>其中 $F(\cdot)$ 表示 self-attention 或 pointwise feedforward 等子模块。</p>

<ol>
  <li>Zero 的含义</li>
</ol>

<p>adaLN-Zero 中的 “Zero” 表示：在初始化时，调制相关的残差门控通常被初始化为接近零，使得网络初始状态更接近恒等映射。<br />
这样做的作用是：</p>

<ul>
  <li>降低训练初期条件注入对主干特征的剧烈扰动；</li>
  <li>提高深层 Transformer 训练的稳定性；</li>
  <li>使模型逐步学会如何利用时间步和条件信息调制生成过程。</li>
</ul>

<ol>
  <li>对生成过程的理解</li>
</ol>

<p>因此，在 adaLN-Zero 中：</p>

<ul>
  <li>$x_t$ 决定当前 noisy latent 的内容；</li>
  <li>$t$ 告诉模型当前处于去噪过程的哪个阶段；</li>
  <li>条件 $c$（例如 text）提供目标语义信息。</li>
</ul>

<p>但这些条件并不是直接替代输入特征，而是通过生成 scale、shift 和 gate 等参数，逐层调制 Transformer block 对 $x_t$ 的处理方式。<br />
因此，$f_t^\theta(x_t)$ 的本质是一个在条件控制下对输入 latent 进行变换的函数，其条件依赖性由 adaLN-Zero 注入到每一层的归一化与残差更新过程中。</p>

<h2 id="3-flux1-model">3. FLUX.1 Model</h2>

<p><img src="/img/flux.1.png" alt="" /></p>

<p>前面已经把 <strong>Flow Matching</strong> 和 <strong>DiT 中条件如何注入</strong> 两件事讲清楚了。</p>

<p>接下来就要回答真正和 FLUX.1 本身有关的问题：</p>

<ul>
  <li>输入文本是怎么进入模型的？</li>
  <li>图像为什么不是直接在像素空间生成，而是在 latent 空间生成？</li>
  <li>FLUX.1 和普通 DiT 到底有什么结构差别？</li>
  <li>模型每一步预测的到底是什么量？</li>
</ul>

<p>如果把这些问题串起来，FLUX.1 的整体逻辑就会变得非常清楚。</p>

<h2 id="31-总体架构总览">3.1 总体架构总览</h2>

<p>先给出一句最核心的概括：</p>

<blockquote>
  <p>FLUX.1 本质上是一个 <strong>latent rectified-flow transformer</strong>。</p>
</blockquote>

<p>这句话可以拆成三层含义：</p>

<ol>
  <li>
    <p><strong>latent</strong></p>

    <p>模型不是直接在原始像素空间上生成图像，而是先在压缩后的 latent 空间中做生成，最后再 decode 回图像。</p>
  </li>
  <li>
    <p><strong>rectified flow</strong></p>

    <p>模型学习的是从噪声分布流向目标图像分布的 <strong>vector field</strong>，因此训练目标更接近前面讲的 Flow Matching，而不是传统 DDPM 里一步步预测噪声的写法。</p>
  </li>
  <li>
    <p><strong>transformer</strong></p>

    <p>主干网络不再是 UNet，而是将图像 latent 划分成 token，再和文本 token 一起送入 Transformer 中建模。</p>
  </li>
</ol>

<p>从 diffusers 的参考实现来看，一个完整的 FLUX pipeline 主要包括下面几个部分：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">VAE</code>：负责图像和 latent 之间的编码/解码；</li>
  <li><code class="language-plaintext highlighter-rouge">text_encoder</code>：<code class="language-plaintext highlighter-rouge">CLIPTextModel</code>，可以理解为提供更偏全局的文本语义；</li>
  <li><code class="language-plaintext highlighter-rouge">text_encoder_2</code>：<code class="language-plaintext highlighter-rouge">T5EncoderModel</code>，可以理解为提供更细粒度、更长序列的文本条件；</li>
  <li><code class="language-plaintext highlighter-rouge">transformer</code>：真正负责预测 latent 更新方向的 <code class="language-plaintext highlighter-rouge">FluxTransformer2DModel</code>；</li>
  <li><code class="language-plaintext highlighter-rouge">scheduler</code>：<code class="language-plaintext highlighter-rouge">FlowMatchEulerDiscreteScheduler</code>，负责按照 flow matching 的离散更新方式推进采样过程。</li>
</ul>

<p>因此，FLUX.1 并不是“只换了个 loss 的 DiT”。</p>

<p>更准确地说，它是：</p>

<ol>
  <li>在 <strong>latent 空间</strong>中进行生成；</li>
  <li>用 <strong>文本-图像联合 Transformer</strong> 作为主干；</li>
  <li>用 <strong>flow / rectified flow</strong> 的目标来学习生成轨迹。</li>
</ol>

<hr />

<h2 id="32-输入表示textlatent-与-token">3.2 输入表示：text、latent 与 token</h2>

<p>在真正进入 Transformer 之前，FLUX.1 需要先把“文本”和“图像”都表示成 token 形式。</p>

<h3 id="321-图像侧先变成-latent再变成-image-tokens">3.2.1 图像侧：先变成 latent，再变成 image tokens</h3>

<p>图像侧的流程可以理解为：</p>

<p><code class="language-plaintext highlighter-rouge">image -&gt; VAE encoder -&gt; latent grid -&gt; image tokens</code></p>

<p>为什么不直接在像素空间生成？</p>

<p>原因很简单：像素空间太大，直接建模的计算量和显存压力都非常高。<br />
因此和 Stable Diffusion 系列一样，FLUX.1 也是先在更低维的 latent 空间中完成主要生成过程，再由 VAE decoder 还原出最终图像。</p>

<p>在 diffusers 的 <code class="language-plaintext highlighter-rouge">FluxTransformer2DModel</code> 中，图像输入 <code class="language-plaintext highlighter-rouge">hidden_states</code> 的形状是：</p>

\[(\text{batch},\ \text{image\_sequence\_length},\ \text{in\_channels})\]

<p>并且参考实现中的 <code class="language-plaintext highlighter-rouge">in_channels=64</code>。<br />
这说明进入 Transformer 时，图像已经不再是二维 feature map，而是被整理成一串 image tokens，每个 token 对应一个 latent 网格位置上的特征。</p>

<p>所以可以把 FLUX.1 的图像输入理解为：</p>

<ul>
  <li>空间结构仍然保留在 token 的排列关系中；</li>
  <li>但网络实际处理的是一串序列，而不是传统卷积网络里的 feature map。</li>
</ul>

<h3 id="322-文本侧不是一条文本通路而是两条">3.2.2 文本侧：不是一条文本通路，而是两条</h3>

<p>FLUX pipeline 的另一个关键点在于：<strong>它并不是只用一个 text encoder。</strong></p>

<p>在 diffusers 参考实现中，它同时使用：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">CLIPTextModel</code></li>
  <li><code class="language-plaintext highlighter-rouge">T5EncoderModel</code></li>
</ul>

<p>可以把这两条文本通路粗略理解为：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">CLIP</code> 更偏向于提供一个全局、压缩后的语义摘要；</li>
  <li><code class="language-plaintext highlighter-rouge">T5</code> 更偏向于提供细粒度的文本序列信息，保留 prompt 中更多 token 级别的内容。</li>
</ul>

<p>这两部分信息在后续主干中扮演的角色并不完全一样：</p>

<ul>
  <li>一部分作为 <strong>序列条件</strong> 参与 joint attention；</li>
  <li>一部分作为 <strong>pooled condition</strong>，进入时间步/条件调制通路，用于控制 block 的行为。</li>
</ul>

<p>这也是为什么 FLUX.1 的 prompt adherence 很强。<br />
因为它并不是简单地把一句文本编码成一个向量，而是同时保留了：</p>

<ul>
  <li>全局语义；</li>
  <li>长文本细节；</li>
  <li>token 级别的结构信息。</li>
</ul>

<hr />

<h2 id="33-flux1-主干先双流再单流">3.3 FLUX.1 主干：先双流，再单流</h2>

<p>如果只看名字，很多人会觉得 FLUX.1 就是一个普通 DiT。<br />
但从 diffusers 中 <code class="language-plaintext highlighter-rouge">FluxTransformer2DModel</code> 的结构参数来看，它其实有一个非常鲜明的组织方式：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">num_layers = 19</code></li>
  <li><code class="language-plaintext highlighter-rouge">num_single_layers = 38</code></li>
</ul>

<p>文档中把前者称作 <strong>dual stream DiT blocks</strong>，后者称作 <strong>single stream DiT blocks</strong>。</p>

<p>这件事非常重要，因为它揭示了 FLUX.1 的核心结构逻辑。</p>

<h3 id="331-dual-stream先分别处理-text-和-image">3.3.1 dual stream：先分别处理 text 和 image</h3>

<p>所谓 dual stream，可以把它理解成：</p>

<ul>
  <li>文本 token 有自己的流；</li>
  <li>图像 token 也有自己的流；</li>
  <li>但二者并不是完全隔离，而是在 attention 中交换信息。</li>
</ul>

<p>为什么要这样设计？</p>

<p>因为 text token 和 image token 的统计特性并不一样：</p>

<ul>
  <li>text token 更偏离散语义；</li>
  <li>image token 更偏连续空间特征。</li>
</ul>

<p>如果一上来就把两者完全混在一起做统一建模，模型需要在很浅层就同时解决：</p>

<ul>
  <li>文本语义理解；</li>
  <li>图像结构建模；</li>
  <li>跨模态对齐；</li>
</ul>

<p>这会让优化变得更困难。</p>

<p>因此 FLUX.1 的策略是：</p>

<ol>
  <li>先保留 text/image 各自的表示流；</li>
  <li>在较浅层完成跨模态对齐；</li>
  <li>等两种 token 已经“知道彼此是谁”之后，再进入更深的统一建模。</li>
</ol>

<h3 id="332-single-stream后面再做深度融合">3.3.2 single stream：后面再做深度融合</h3>

<p>在 dual stream blocks 之后，FLUX.1 还有更多的 single stream blocks。</p>

<p>这可以理解为：</p>

<ul>
  <li>前面解决“对齐”问题；</li>
  <li>后面解决“联合生成”问题。</li>
</ul>

<p>一旦 text token 和 image token 已经在浅层对齐完成，后面的单流建模就可以更专注于：</p>

<ul>
  <li>细化空间结构；</li>
  <li>强化 prompt 对局部内容的控制；</li>
  <li>输出更准确的 vector field。</li>
</ul>

<p>因此，FLUX.1 的主干不是“从头到尾都一个结构”，而是分成两个阶段：</p>

<ol>
  <li><strong>dual stream 阶段</strong>：做跨模态对齐；</li>
  <li><strong>single stream 阶段</strong>：做联合生成与精细更新。</li>
</ol>

<p>这也是它和“普通一串统一 Transformer block”的一个重要差别。</p>

<hr />

<h2 id="34-flux1-预测的到底是什么">3.4 FLUX.1 预测的到底是什么</h2>

<p>这一点必须和传统 diffusion 区分开。</p>

<p>在很多 DDPM / latent diffusion 模型中，网络常见的预测目标是：</p>

<ul>
  <li>噪声 $\epsilon$；</li>
  <li>或者某种等价的重参数化目标，例如 $v$。</li>
</ul>

<p>但在 FLUX.1 里，更自然的理解方式是：</p>

<blockquote>
  <p>模型预测的是当前 latent 状态应该往哪个方向走，也就是 <strong>vector field / velocity</strong>。</p>
</blockquote>

<p>因此，输入当前时刻的 latent $x_t$ 后，模型输出可以理解为：</p>

\[f_t^\theta(x_t,\text{text})\]

<p>它描述的是：<strong>在当前时间步，latent 应该沿哪个方向更新，才能逐渐流向目标图像分布。</strong></p>

<p>所以推理时的核心更新形式就是：</p>

\[x_{t+h}=x_t+h\cdot f_t^\theta(x_t,\text{text})\]

<p>这和前面讲的 ODE 完全对应。</p>

<p>也就是说，FLUX.1 的生成过程更像是：</p>

<ul>
  <li>在 latent 空间里定义一条从 noise 到 image 的轨迹；</li>
  <li>Transformer 在每个时刻给出局部方向；</li>
  <li>scheduler 按照这个方向做离散积分。</li>
</ul>

<p>从这个角度看，FLUX.1 和“逐步去噪”的直觉并不完全一样。<br />
它更像是在学习一张“速度场地图”，然后沿着这张地图把噪声流到图像。</p>

<hr />

<h2 id="35-条件信息到底怎么进入-flux1">3.5 条件信息到底怎么进入 FLUX.1</h2>

<p>前面在 DiT 一节已经介绍过 <code class="language-plaintext highlighter-rouge">adaLN-Zero</code> 的基本思想。<br />
到了 FLUX.1 这里，可以把条件注入方式理解得更具体一些。</p>

<p>对于 FLUX.1 来说，至少有三类条件会参与计算：</p>

<ol>
  <li>当前 latent 状态 $x_t$；</li>
  <li>时间步 <code class="language-plaintext highlighter-rouge">timestep</code>；</li>
  <li>文本条件（包括序列级文本表示和 pooled 文本表示）。</li>
</ol>

<p>从 <code class="language-plaintext highlighter-rouge">FluxTransformer2DModel.forward</code> 的输入可以看到，主干会接收：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">hidden_states</code></li>
  <li><code class="language-plaintext highlighter-rouge">encoder_hidden_states</code></li>
  <li><code class="language-plaintext highlighter-rouge">pooled_projections</code></li>
  <li><code class="language-plaintext highlighter-rouge">timestep</code></li>
  <li>以及和位置编码有关的 <code class="language-plaintext highlighter-rouge">img_ids</code>、<code class="language-plaintext highlighter-rouge">txt_ids</code></li>
</ul>

<p>这说明 FLUX.1 的条件注入并不是“只有一条 text embedding”这么简单，而是至少分成了两条路径：</p>

<h3 id="351-一条进入-attention">3.5.1 一条进入 attention</h3>

<p><code class="language-plaintext highlighter-rouge">encoder_hidden_states</code> 可以理解为文本 token 序列对应的条件表示。<br />
这部分信息参与 joint attention，使图像 token 在更新时能够感知 prompt 的细粒度语义。</p>

<h3 id="352-一条进入-block-调制">3.5.2 一条进入 block 调制</h3>

<p><code class="language-plaintext highlighter-rouge">pooled_projections</code> 和 <code class="language-plaintext highlighter-rouge">timestep</code> 则更像全局条件。<br />
它们不会只告诉模型“要画什么”，还会告诉模型：</p>

<ul>
  <li>当前处于采样轨迹的哪个阶段；</li>
  <li>当前 block 应该更偏向语义布局，还是更偏向细节修正。</li>
</ul>

<p>因此，如果把 FLUX.1 看成一个受条件控制的动力系统，那么：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">encoder_hidden_states</code> 更像局部语义约束；</li>
  <li><code class="language-plaintext highlighter-rouge">pooled_projections</code> 更像全局控制信号；</li>
  <li><code class="language-plaintext highlighter-rouge">timestep</code> 更像系统当前所处的时间位置。</li>
</ul>

<p>这三者一起决定了最终的 vector field。</p>

<hr />

<h2 id="36-dev-和-schnell-的差别">3.6 <code class="language-plaintext highlighter-rouge">dev</code> 和 <code class="language-plaintext highlighter-rouge">schnell</code> 的差别</h2>

<p>FLUX.1 常见的开源版本里，最常被拿来比较的就是：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">FLUX.1-dev</code></li>
  <li><code class="language-plaintext highlighter-rouge">FLUX.1-schnell</code></li>
</ul>

<p>二者都基于 FLUX.1 主干，但推理习惯和蒸馏方式不同。</p>

<h3 id="361-dev">3.6.1 <code class="language-plaintext highlighter-rouge">dev</code></h3>

<p>根据 diffusers 文档，<code class="language-plaintext highlighter-rouge">dev</code> 是 <strong>guidance-distilled</strong> 版本。</p>

<p>它的典型特点是：</p>

<ul>
  <li>生成质量更稳；</li>
  <li>常见设置下大约需要 <strong>50 steps</strong> 才能比较充分地发挥效果；</li>
  <li>没有 <code class="language-plaintext highlighter-rouge">max_sequence_length=256</code> 这类限制；</li>
  <li>推理时通常会配合非零 <code class="language-plaintext highlighter-rouge">guidance_scale</code> 使用。</li>
</ul>

<p>因此，<code class="language-plaintext highlighter-rouge">dev</code> 更像是：</p>

<ul>
  <li>速度不是第一优先级；</li>
  <li>更注重质量、细节和 prompt 跟随能力。</li>
</ul>

<h3 id="362-schnell">3.6.2 <code class="language-plaintext highlighter-rouge">schnell</code></h3>

<p>根据 diffusers 文档，<code class="language-plaintext highlighter-rouge">schnell</code> 是 <strong>timestep-distilled</strong> 版本。</p>

<p>它的特点更激进：</p>

<ul>
  <li>通常 <strong>1 到 4 步</strong> 就能出图；</li>
  <li><code class="language-plaintext highlighter-rouge">guidance_scale</code> 需要设为 <code class="language-plaintext highlighter-rouge">0</code>；</li>
  <li><code class="language-plaintext highlighter-rouge">max_sequence_length</code> 不能超过 <code class="language-plaintext highlighter-rouge">256</code>。</li>
</ul>

<p>因此，<code class="language-plaintext highlighter-rouge">schnell</code> 的本质是：<br />
用额外蒸馏把原本多步的采样轨迹压缩到极少数步骤里，从而换取极快的生成速度。</p>

<p>可以简单理解为：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">dev</code>：更偏质量；</li>
  <li><code class="language-plaintext highlighter-rouge">schnell</code>：更偏速度。</li>
</ul>

<hr />

<h2 id="37-flux1-和-sdxl--普通-dit-的差别">3.7 FLUX.1 和 SDXL / 普通 DiT 的差别</h2>

<p>如果要快速定位 FLUX.1 在整个生成模型谱系里的位置，可以这样看：</p>

<h3 id="371-和-sdxl-的差别">3.7.1 和 SDXL 的差别</h3>

<p>SDXL 的主干仍然是 <strong>UNet</strong> 思路：</p>

<ul>
  <li>图像特征在多尺度卷积结构里流动；</li>
  <li>条件通常通过 cross-attention 等方式注入。</li>
</ul>

<p>而 FLUX.1 的主干已经变成：</p>

<ul>
  <li>图像先转成 token；</li>
  <li>文本也转成 token；</li>
  <li>再用 Transformer 做统一建模。</li>
</ul>

<p>所以二者最大的差别不是“loss 不一样”，而是：</p>

<ul>
  <li>SDXL 更像卷积式的 latent diffusion；</li>
  <li>FLUX.1 更像 token-based 的 latent flow transformer。</li>
</ul>

<h3 id="372-和普通-dit-的差别">3.7.2 和普通 DiT 的差别</h3>

<p>普通 DiT 的核心思想是：<br />
用 Transformer 代替 UNet 来处理图像 token。</p>

<p>而 FLUX.1 在此基础上又往前走了一步：</p>

<ul>
  <li>不只是“图像 token + 条件”；</li>
  <li>而是显式地做 <strong>文本流 / 图像流的分阶段融合</strong>；</li>
  <li>再配合 rectified flow / flow matching 目标来学习生成轨迹。</li>
</ul>

<p>所以从直观上说：</p>

<ul>
  <li>DiT 解决的是“能不能用 Transformer 生成图像”；</li>
  <li>FLUX.1 更进一步解决的是“怎么让文本条件和图像 latent 在 Transformer 里更高效地协同工作”。</li>
</ul>

<hr />

<h2 id="38-一个完整的-forward-pass">3.8 一个完整的 forward pass</h2>

<p>最后，把前面的所有内容串起来，看一遍 FLUX.1 的完整前向流程。</p>

<p>假设输入是一句 prompt，例如：</p>

<p><code class="language-plaintext highlighter-rouge">"a tiny astronaut hatching from an egg on the moon"</code></p>

<p>则一次典型生成可以理解为下面几步：</p>

<h3 id="381-文本编码">3.8.1 文本编码</h3>

<p>prompt 先分别送入文本编码器，得到两类文本条件：</p>

<ul>
  <li>序列级文本 token 表示；</li>
  <li>pooled 的全局文本表示。</li>
</ul>

<h3 id="382-初始化-latent">3.8.2 初始化 latent</h3>

<p>在图像侧，先采样一个初始噪声 latent，记作 $x_0$。<br />
此时它还不对应任何真实图像，只是 latent 空间中的随机起点。</p>

<h3 id="383-整理成-image-tokens">3.8.3 整理成 image tokens</h3>

<p>这个 latent 会被整理成 image token 序列，作为 Transformer 的 <code class="language-plaintext highlighter-rouge">hidden_states</code> 输入。</p>

<h3 id="384-注入时间步和条件">3.8.4 注入时间步和条件</h3>

<p>当前的 <code class="language-plaintext highlighter-rouge">timestep</code>、pooled 文本条件、文本 token 条件一起进入主干网络：</p>

<ul>
  <li>一部分调制 block；</li>
  <li>一部分进入 attention；</li>
  <li>一部分决定 text/image 的对齐关系与融合方式。</li>
</ul>

<h3 id="385-transformer-预测-vector-field">3.8.5 Transformer 预测 vector field</h3>

<p>经过 dual stream blocks 和 single stream blocks 后，模型输出当前时刻的更新方向，也就是：</p>

\[f_t^\theta(x_t,\text{text})\]

<h3 id="386-scheduler-更新-latent">3.8.6 scheduler 更新 latent</h3>

<p>随后，<code class="language-plaintext highlighter-rouge">FlowMatchEulerDiscreteScheduler</code> 根据这个方向，把 latent 从当前时刻推进到下一个时刻：</p>

\[x_{t+h}=x_t+h\cdot f_t^\theta(x_t,\text{text})\]

<p>重复多次之后，latent 会逐渐从随机噪声流向目标图像分布。</p>

<h3 id="387-decode-成最终图像">3.8.7 decode 成最终图像</h3>

<p>当采样结束后，得到最终 latent $x_1$，再送入 VAE decoder，输出最终图像。</p>

<hr />

<h2 id="39-从工程角度如何理解-flux1">3.9 从工程角度如何理解 FLUX.1</h2>

<p>如果从工程实现的角度看，FLUX.1 之所以强，主要是因为它同时把三件事做到了：</p>

<ol>
  <li>
    <p><strong>文本理解强</strong></p>

    <p>双文本编码器 + 联合 Transformer，使 prompt 中的复杂语义能够更充分地注入图像生成过程。</p>
  </li>
  <li>
    <p><strong>主干表达能力强</strong></p>

    <p>12B 级别的 Transformer 主干，加上先 dual-stream 再 single-stream 的结构，使它在跨模态对齐和图像细节生成上都有很高容量。</p>
  </li>
  <li>
    <p><strong>采样路径更直接</strong></p>

    <p>flow matching / rectified flow 的目标让模型更自然地学习“从噪声走向图像”的方向场，因此在推理时可以用更直接的轨迹更新 latent。</p>
  </li>
</ol>

<p>当然，它的代价也很明显：</p>

<ul>
  <li>模型非常大；</li>
  <li>显存压力高；</li>
  <li>推理成本高；</li>
  <li>consumer GPU 上部署并不轻松。</li>
</ul>

<p>所以可以把 FLUX.1 看成一个典型的“拿更大算力和更复杂结构，换更强 prompt adherence 与更高图像质量”的模型。</p>

<hr />

<p>到这里再回头看，你会发现 FLUX.1 的逻辑其实很统一：</p>

<ul>
  <li><strong>理论层面</strong>：用 Flow Matching 学 vector field；</li>
  <li><strong>网络层面</strong>：用 Transformer 建模 text/image token；</li>
  <li><strong>条件层面</strong>：用时间步和文本条件调制主干；</li>
  <li><strong>推理层面</strong>：在 latent 空间中沿着 learned flow 逐步更新。</li>
</ul>

<p>这四件事合在一起，才构成了完整的 FLUX.1。</p>]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[FLUX.1 梳理]]></summary></entry><entry><title type="html">Git 常见问题与操作笔记</title><link href="https://wqh011128.github.io/blog/2026/03/11/GIT.html" rel="alternate" type="text/html" title="Git 常见问题与操作笔记" /><published>2026-03-11T00:00:00+00:00</published><updated>2026-03-11T00:00:00+00:00</updated><id>https://wqh011128.github.io/blog/2026/03/11/GIT</id><content type="html" xml:base="https://wqh011128.github.io/blog/2026/03/11/GIT.html"><![CDATA[<p>这篇笔记按“遇到的问题”整理 Git 日常操作。命令默认在项目根目录执行；如果涉及已经推送到远端的历史改写，优先用 <code class="language-plaintext highlighter-rouge">--force-with-lease</code>，不要直接用 <code class="language-plaintext highlighter-rouge">--force</code>。</p>

<h2 id="github-recovery-codes-应该写进笔记里吗">GitHub recovery codes 应该写进笔记里吗？</h2>

<p>不要。GitHub recovery codes 等同于备用登录凭证，不应该放在仓库、博客、截图或公开笔记里。</p>

<p>如果 recovery codes 曾经被推送到公开仓库或 GitHub Pages，建议在 GitHub 的 2FA 设置里重新生成 recovery codes，让旧 codes 失效。</p>

<h2 id="怎么把-remote-的-https-地址换成-ssh">怎么把 remote 的 HTTPS 地址换成 SSH？</h2>

<p>查看当前 remote：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git remote <span class="nt">-v</span>
</code></pre></div></div>

<p>把 <code class="language-plaintext highlighter-rouge">origin</code> 从 HTTPS 改成 SSH：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git remote set-url origin git@github.com:&lt;yourname&gt;/&lt;repo&gt;.git
</code></pre></div></div>

<p>再次确认：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git remote <span class="nt">-v</span>
</code></pre></div></div>

<h2 id="git-的全局身份提交模板和-ssh-key-怎么配置">Git 的全局身份、提交模板和 SSH key 怎么配置？</h2>

<p>身份信息：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git config <span class="nt">--global</span> user.name <span class="s2">"wqh011128"</span>
git config <span class="nt">--global</span> user.email <span class="s2">"wqh011128@163.com"</span>
</code></pre></div></div>

<p>提交信息模板：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nb">cat</span> <span class="o">&lt;&lt;</span><span class="sh">'</span><span class="no">EOF</span><span class="sh">' &gt; ~/.gitmessage.txt
&lt;type&gt;[&lt;scope&gt;]: &lt;short-summary&gt;

Problem:
&lt;description of the problem being solved&gt;

Solution:
&lt;description of the solution implemented&gt;

Test:
&lt;description of how the change was tested&gt;

JIRA: ISSUE-&lt;Number&gt;
</span><span class="no">EOF

</span>git config <span class="nt">--global</span> commit.template ~/.gitmessage.txt
</code></pre></div></div>

<p>Windows 上如果经常遇到 CRLF/LF 问题，可以按项目规范决定是否设置：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git config <span class="nt">--global</span> core.autocrlf <span class="nb">true</span>
</code></pre></div></div>

<p>生成 SSH key 并查看公钥：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>ssh-keygen <span class="nt">-t</span> ed25519 <span class="nt">-C</span> <span class="s2">"wqh011128@163.com"</span>
<span class="nb">cat</span> ~/.ssh/id_ed25519.pub
</code></pre></div></div>

<p>如果环境不支持 <code class="language-plaintext highlighter-rouge">ed25519</code>，再考虑 RSA：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>ssh-keygen <span class="nt">-t</span> rsa <span class="nt">-b</span> 2048 <span class="nt">-C</span> <span class="s2">"wqh011128@163.com"</span>
<span class="nb">cat</span> ~/.ssh/id_rsa.pub
</code></pre></div></div>

<h2 id="git-fetchgit-pull-和-git-pull---rebase-有什么区别"><code class="language-plaintext highlighter-rouge">git fetch</code>、<code class="language-plaintext highlighter-rouge">git pull</code> 和 <code class="language-plaintext highlighter-rouge">git pull --rebase</code> 有什么区别？</h2>

<p><code class="language-plaintext highlighter-rouge">git fetch</code> 只更新远端分支信息，不会改动当前工作区和当前分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git fetch origin
</code></pre></div></div>

<p><code class="language-plaintext highlighter-rouge">git pull</code> 等价于先 <code class="language-plaintext highlighter-rouge">fetch</code>，再把远端内容整合到当前分支。默认通常是 merge：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git pull origin main
</code></pre></div></div>

<p>这条命令的意思是：把远端 <code class="language-plaintext highlighter-rouge">origin/main</code> 整合到你当前 checkout 的分支里。它不是“切到 main 再拉取”，当前分支是谁就影响谁。</p>

<p>如果想让历史更线性，常用 rebase：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git pull <span class="nt">--rebase</span> origin main
</code></pre></div></div>

<p>概念上可以理解成：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git fetch origin main
git rebase origin/main
</code></pre></div></div>

<h2 id="git-pull---rebase-origin-devvl-会改哪个分支"><code class="language-plaintext highlighter-rouge">git pull --rebase origin dev/vl</code> 会改哪个分支？</h2>

<p>会改你当前 checkout 的分支。</p>

<p>例如你当前在 <code class="language-plaintext highlighter-rouge">feat</code> 分支，执行：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git pull <span class="nt">--rebase</span> origin dev/vl
</code></pre></div></div>

<p>含义是：</p>

<ol>
  <li>从 <code class="language-plaintext highlighter-rouge">origin</code> 拉取 <code class="language-plaintext highlighter-rouge">dev/vl</code> 的最新提交。</li>
  <li>把当前分支 <code class="language-plaintext highlighter-rouge">feat</code> 上你自己的提交 rebase 到 <code class="language-plaintext highlighter-rouge">origin/dev/vl</code> 之后。</li>
</ol>

<p>所以它不是“把 <code class="language-plaintext highlighter-rouge">dev/vl</code> 拉到 <code class="language-plaintext highlighter-rouge">dev/vl</code>”，而是“用 <code class="language-plaintext highlighter-rouge">origin/dev/vl</code> 作为当前分支的新基底”。</p>

<p>例子：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>origin/dev/vl: A-B-C
local feat:    x-y
</code></pre></div></div>

<p>执行后大致变成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>A-B-C-x'-y'
</code></pre></div></div>

<h2 id="怎么复制一个-branch-继续开发">怎么复制一个 branch 继续开发？</h2>

<p>如果想基于最新 <code class="language-plaintext highlighter-rouge">main</code> 开一个新分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git switch main
git pull <span class="nt">--rebase</span> origin main
git switch <span class="nt">-c</span> feat/example
</code></pre></div></div>

<p>老版本 Git 可以用：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git checkout <span class="nt">-b</span> feat/example
</code></pre></div></div>

<h2 id="怎么查看添加和删除-remote">怎么查看、添加和删除 remote？</h2>

<p>查看远程连接：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git remote <span class="nt">-v</span>
</code></pre></div></div>

<p>添加远程连接：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git remote add origin &lt;remote-url&gt;
</code></pre></div></div>

<p>删除某个远程连接：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git remote remove origin
</code></pre></div></div>

<p>修改远程连接：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git remote set-url origin &lt;remote-url&gt;
</code></pre></div></div>

<h2 id="怎么从本地项目创建-github-仓库并推上去">怎么从本地项目创建 GitHub 仓库并推上去？</h2>

<p>如果本地还没有 Git 历史：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nb">cd</span> /path/to/your-project
git init
git add <span class="nb">.</span>
git commit <span class="nt">-m</span> <span class="s2">"Initial commit"</span>
</code></pre></div></div>

<p>如果本地已经有提交历史，从绑定远端开始即可。</p>

<p>在 GitHub 上创建仓库时，建议创建空仓库：</p>

<ul>
  <li>填仓库名。</li>
  <li>先选 Private 或 Public 都可以，按需求决定。</li>
  <li>不要勾选 README、<code class="language-plaintext highlighter-rouge">.gitignore</code>、license，避免和本地历史冲突。</li>
</ul>

<p>使用 HTTPS：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git remote add origin https://github.com/&lt;yourname&gt;/&lt;repo&gt;.git
git branch <span class="nt">-M</span> main
git push <span class="nt">-u</span> origin main
</code></pre></div></div>

<p>使用 SSH：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git remote add origin git@github.com:&lt;yourname&gt;/&lt;repo&gt;.git
git branch <span class="nt">-M</span> main
git push <span class="nt">-u</span> origin main
</code></pre></div></div>

<p>常见坑：</p>

<ul>
  <li>HTTPS push 被拒绝或要求登录：GitHub 通常需要 Personal Access Token，不是账号密码。</li>
  <li>远端已有 README 导致冲突：可以重建空仓库，或先 <code class="language-plaintext highlighter-rouge">git pull --rebase</code> 解决冲突后再 push。</li>
  <li>默认分支不是 <code class="language-plaintext highlighter-rouge">main</code>：用 <code class="language-plaintext highlighter-rouge">git branch -M main</code> 统一。</li>
</ul>

<h2 id="怎么创建新分支并推送到远端">怎么创建新分支并推送到远端？</h2>

<p>先确认当前状态：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git status
git branch <span class="nt">--show-current</span>
</code></pre></div></div>

<p>基于当前分支创建并切到新分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git switch <span class="nt">-c</span> feat/eau-opt
</code></pre></div></div>

<p>提交修改：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git add <span class="nt">-A</span>
pre-commit run <span class="nt">--all-files</span>
git commit <span class="nt">-m</span> <span class="s2">"Describe your change"</span>
</code></pre></div></div>

<p>推送新分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git push <span class="nt">-u</span> origin feat/eau-opt
</code></pre></div></div>

<p><code class="language-plaintext highlighter-rouge">-u</code> 会建立 upstream。以后在这个分支上直接 <code class="language-plaintext highlighter-rouge">git push</code> 或 <code class="language-plaintext highlighter-rouge">git pull</code> 就行。</p>

<h2 id="当前-feat-分支怎么更新到最新-originmain">当前 <code class="language-plaintext highlighter-rouge">feat</code> 分支怎么更新到最新 <code class="language-plaintext highlighter-rouge">origin/main</code>？</h2>

<p>推荐把当前 feature 分支 rebase 到最新 <code class="language-plaintext highlighter-rouge">origin/main</code> 上：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git status
git stash push <span class="nt">-u</span> <span class="nt">-m</span> <span class="s2">"wip before rebase"</span>
git fetch origin
git rebase origin/main
git stash pop
</code></pre></div></div>

<p>如果想让 Git 自动暂存未提交修改，可以用：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git pull <span class="nt">--rebase</span> <span class="nt">--autostash</span> origin main
</code></pre></div></div>

<p>注意：这条命令会影响当前 checkout 的分支。先确认自己在目标 feature 分支上。</p>

<h2 id="push-时远程比当前要领先怎么办">push 时远程比当前要领先怎么办？</h2>

<p>先区分清楚：准备开 PR 前，优先同步的是最新 <code class="language-plaintext highlighter-rouge">main</code>，不是盲目 <code class="language-plaintext highlighter-rouge">git pull</code> 当前分支。</p>

<p>推荐流程：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git status
git fetch origin
git rebase origin/main
</code></pre></div></div>

<p>这个习惯很重要。别人从其他分支合 PR 之后，<code class="language-plaintext highlighter-rouge">main</code> 可能已经领先；提前 rebase 到最新 <code class="language-plaintext highlighter-rouge">main</code>，通常能更早发现冲突，也能避免 PR 里混入过期历史。</p>

<p>如果已经在自己的分支上提交了 commit，再执行：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git rebase origin/main
</code></pre></div></div>

<p>Git 会把你的 commit 临时拿下来，再放到最新 <code class="language-plaintext highlighter-rouge">main</code> 后面重新应用。遇到冲突时按下面处理：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git status
<span class="c"># 手动解决冲突文件</span>
git add &lt;conflict-file&gt;
git rebase <span class="nt">--continue</span>
</code></pre></div></div>

<p>如果发现 rebase 方向不对，直接取消：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git rebase <span class="nt">--abort</span>
</code></pre></div></div>

<p>当 <code class="language-plaintext highlighter-rouge">rebase origin/main</code> 之后，<code class="language-plaintext highlighter-rouge">git status</code> 仍然可能显示：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Your branch and 'origin/xxx' have diverged,
and have N and M different commits each, respectively.
</code></pre></div></div>

<p>这通常不是说你落后于 <code class="language-plaintext highlighter-rouge">origin/main</code>，而是说本地分支和远程同名分支 <code class="language-plaintext highlighter-rouge">origin/xxx</code> 分叉了。比如之前用同一个分支开过 PR，PR 已经合进 <code class="language-plaintext highlighter-rouge">main</code>，但远程功能分支还停在旧位置。</p>

<p>这时不要直接 <code class="language-plaintext highlighter-rouge">git pull</code>。<code class="language-plaintext highlighter-rouge">git pull</code> 会把 <code class="language-plaintext highlighter-rouge">origin/xxx</code> 拉回来整合，可能产生重复提交或无意义 merge。先确认当前分支已经基于最新 <code class="language-plaintext highlighter-rouge">main</code>：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git log <span class="nt">--oneline</span> <span class="nt">--graph</span> <span class="nt">--decorate</span> <span class="nt">--left-right</span> origin/main...HEAD
</code></pre></div></div>

<p>如果确认当前分支只是 <code class="language-plaintext highlighter-rouge">origin/main + 你的新 commit</code>，就更新远程功能分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git push <span class="nt">--force-with-lease</span> origin xxx
</code></pre></div></div>

<p>一句话记忆：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>rebase origin/main：同步主线。
push --force-with-lease：更新远程功能分支。
不要用 git pull 去解决已经 rebase 到 main 后的同名分支分叉。
</code></pre></div></div>

<h2 id="flowra-工作流里-git-怎么走">Flowra 工作流里 Git 怎么走？</h2>

<p>这部分更像项目工具链笔记，主线可以按下面走：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>flowra ws myWorkspace
git clone &lt;repo-url&gt;
</code></pre></div></div>

<p>进入 workspace 后修改代码，然后格式化：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>flowra format your_project
</code></pre></div></div>

<p>再走正常 Git 流程：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git add <span class="nt">-A</span>
git commit <span class="nt">-m</span> <span class="s2">"xxx"</span>
git push
</code></pre></div></div>

<p>如果需要修改最后一次提交：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git commit <span class="nt">--amend</span>
git push <span class="nt">--force-with-lease</span>
</code></pre></div></div>

<h2 id="pre-commit-应该什么时候跑">pre-commit 应该什么时候跑？</h2>

<p>安装：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>pip <span class="nb">install </span>pre-commit
</code></pre></div></div>

<p>通常在 <code class="language-plaintext highlighter-rouge">git add</code> 后、<code class="language-plaintext highlighter-rouge">git commit</code> 前跑：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git add <span class="nt">-A</span>
pre-commit run <span class="nt">--all-files</span>
git commit <span class="nt">-m</span> <span class="s2">"Describe your change"</span>
</code></pre></div></div>

<p>如果 hook 自动修了文件，需要重新暂存：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git add <span class="nt">-A</span>
git commit <span class="nt">-m</span> <span class="s2">"Describe your change"</span>
</code></pre></div></div>

<h2 id="怎么修改过去的-commit-message">怎么修改过去的 commit message？</h2>

<p>如果历史 commit message 写错，并且 CI 或规范检查依赖 commit message，可以用 interactive rebase 的 <code class="language-plaintext highlighter-rouge">reword</code>。</p>

<p>选择要修改的最近 N 个 commit：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git rebase <span class="nt">-i</span> HEAD~N
</code></pre></div></div>

<p>在编辑器里，把需要修改 message 的那几行从 <code class="language-plaintext highlighter-rouge">pick</code> 改成 <code class="language-plaintext highlighter-rouge">reword</code>：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>reword a1b2c3d feat: old message
pick   b2c3d4e fix: another message
</code></pre></div></div>

<p>保存退出后，Git 会逐个打开 commit message 编辑器。</p>

<p>如果这些 commit 已经推送到远端，修改完成后需要：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git push <span class="nt">--force-with-lease</span>
</code></pre></div></div>

<h2 id="怎么合并过去的多个-commits">怎么合并过去的多个 commits？</h2>

<p>先更新当前分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git checkout feat/op_matcher
git fetch origin
git pull <span class="nt">--rebase</span>
</code></pre></div></div>

<p>把最近 6 个提交放进 interactive rebase：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git rebase <span class="nt">-i</span> HEAD~6
</code></pre></div></div>

<p>如果历史里包含 merge commit，并且需要保留 merge 结构，可以改用：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git rebase <span class="nt">-i</span> <span class="nt">--rebase-merges</span> HEAD~6
</code></pre></div></div>

<p>编辑器里会看到类似：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>pick a1b2c3d feat(op_matcher): reconstruct...
pick b2c3d4e feat(op_matcher): reconstruct...
pick c3d4e5f feat(op_matcher): reconstruct...
pick d4e5f6g feat(op_matcher): support bf16
pick e5f6g7h feat(op_matcher): reconstruct...
pick f6g7h8i feat(op_matcher): optimize compare
</code></pre></div></div>

<p>如果想把前 3 个合成 1 个，把后两行改成 <code class="language-plaintext highlighter-rouge">squash</code> 或 <code class="language-plaintext highlighter-rouge">s</code>：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>pick a1b2c3d feat(op_matcher): reconstruct...
s    b2c3d4e feat(op_matcher): reconstruct...
s    c3d4e5f feat(op_matcher): reconstruct...
pick d4e5f6g feat(op_matcher): support bf16
pick e5f6g7h feat(op_matcher): reconstruct...
pick f6g7h8i feat(op_matcher): optimize compare
</code></pre></div></div>

<p>保存退出后，Git 会让你编辑新的 commit message。删掉重复内容，保留一个清晰的 summary 和必要 bullet 即可。</p>

<p>如果 rebase 过程中遇到冲突：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git add &lt;conflict-file&gt;
git rebase <span class="nt">--continue</span>
</code></pre></div></div>

<p>如果这些 commit 已经推送过，最后同步远端：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git push <span class="nt">--force-with-lease</span>
</code></pre></div></div>

<h2 id="git-commit--c-head-是什么"><code class="language-plaintext highlighter-rouge">git commit -c HEAD</code> 是什么？</h2>

<p><code class="language-plaintext highlighter-rouge">git commit -c HEAD</code> 会复用当前分支最后一次提交的 message 作为模板，并打开编辑器让你修改，然后创建一个新提交。</p>

<p>基本流程：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git add <span class="nb">.</span>
git commit <span class="nt">-c</span> HEAD
</code></pre></div></div>

<p>它和 <code class="language-plaintext highlighter-rouge">git commit --amend</code> 的区别：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">git commit --amend</code>：修改上一次提交，不产生额外的新提交，而是用新 commit 替换旧 commit。</li>
  <li><code class="language-plaintext highlighter-rouge">git commit -c HEAD</code>：创建一个全新的提交，只是借用上一次提交信息作为草稿。</li>
</ul>

<p><code class="language-plaintext highlighter-rouge">-c</code> 和 <code class="language-plaintext highlighter-rouge">-C</code> 的区别：</p>

<ul>
  <li><code class="language-plaintext highlighter-rouge">git commit -c HEAD</code>：复用信息，并打开编辑器供你修改。</li>
  <li><code class="language-plaintext highlighter-rouge">git commit -C HEAD</code>：直接复用信息，不打开编辑器。</li>
</ul>

<p>如果想“复制上个提交的消息并改一改，然后发一个新提交”，用 <code class="language-plaintext highlighter-rouge">git commit -c HEAD</code>。</p>

<h2 id="合并冲突时怎么保留某一边">合并冲突时怎么保留某一边？</h2>

<p>如果你明确知道要保留当前分支这一边：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git checkout <span class="nt">--ours</span> op_matcher/verify_eau_opcode.py
git add op_matcher/verify_eau_opcode.py
</code></pre></div></div>

<p>如果你明确知道要保留另一边：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git checkout <span class="nt">--theirs</span> op_matcher/verify_eau_opcode.py
git add op_matcher/verify_eau_opcode.py
</code></pre></div></div>

<p>注意：在 rebase 场景里，<code class="language-plaintext highlighter-rouge">ours</code> 和 <code class="language-plaintext highlighter-rouge">theirs</code> 的直觉可能和 merge 场景不同。执行前最好用 <code class="language-plaintext highlighter-rouge">git status</code> 和冲突内容确认清楚。</p>

<h2 id="怎么修正最近一次-commit-的-author">怎么修正最近一次 commit 的 author？</h2>

<p>如果最近一次提交的 author 信息不对：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git commit <span class="nt">--amend</span> <span class="nt">--reset-author</span> <span class="nt">--no-edit</span>
</code></pre></div></div>

<p>如果该提交已经推送过：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git push <span class="nt">--force-with-lease</span>
</code></pre></div></div>

<h2 id="wsl-里-clone-项目失败可以先看什么">WSL 里 clone 项目失败可以先看什么？</h2>

<p>如果是 DNS 或网络解析问题，可以先检查 <code class="language-plaintext highlighter-rouge">/etc/resolv.conf</code>：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nb">sudo </span>vim /etc/resolv.conf
</code></pre></div></div>

<p>有些网络环境里，需要把 <code class="language-plaintext highlighter-rouge">nameserver</code> 换成本机或可用 DNS 的 IP。这个属于环境问题，不一定是 Git 本身的问题。</p>

<h2 id="amend-之后为什么普通-git-push-会让我先-pull">amend 之后为什么普通 <code class="language-plaintext highlighter-rouge">git push</code> 会让我先 pull？</h2>

<p>假设最开始本地和远端都是提交 <code class="language-plaintext highlighter-rouge">A</code>：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>remote: A
local:  A
</code></pre></div></div>

<p>执行：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git commit <span class="nt">--amend</span>
</code></pre></div></div>

<p>本地提交会从 <code class="language-plaintext highlighter-rouge">A</code> 变成一个新的提交，比如 <code class="language-plaintext highlighter-rouge">B</code>：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>remote: A
local:  B
</code></pre></div></div>

<p><code class="language-plaintext highlighter-rouge">B</code> 不是追加在 <code class="language-plaintext highlighter-rouge">A</code> 后面的提交，而是替换了 <code class="language-plaintext highlighter-rouge">A</code>。这不是 fast-forward 关系，所以普通 <code class="language-plaintext highlighter-rouge">git push</code> 会拒绝，并提示你先 pull。</p>

<p>但这个场景通常不应该先 pull。你真正想做的是用本地新的 <code class="language-plaintext highlighter-rouge">B</code> 替换远端旧的 <code class="language-plaintext highlighter-rouge">A</code>，因此应该：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git push <span class="nt">--force-with-lease</span>
</code></pre></div></div>

<h2 id="--force-和---force-with-lease-有什么区别"><code class="language-plaintext highlighter-rouge">--force</code> 和 <code class="language-plaintext highlighter-rouge">--force-with-lease</code> 有什么区别？</h2>

<p><code class="language-plaintext highlighter-rouge">--force</code> 会直接强推，不太管远端是否已经被别人更新。</p>

<p><code class="language-plaintext highlighter-rouge">--force-with-lease</code> 会先检查远端是否还是你本地记录里的状态。如果远端在这期间被别人更新过，它会拒绝推送，避免你覆盖别人的提交。</p>

<p>所以平时更推荐：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git push <span class="nt">--force-with-lease</span>
</code></pre></div></div>

<h2 id="怎么退回以前的提交">怎么退回以前的提交？</h2>

<p>如果提交已经推送到远端，最保守的撤销方式是 <code class="language-plaintext highlighter-rouge">revert</code>。它会新增一个“反向提交”，保留历史：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git revert HEAD
git push
</code></pre></div></div>

<p>如果后来决定把“原 commit + revert commit”都从历史中删掉，可以用 interactive rebase：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git rebase <span class="nt">-i</span> HEAD~3
</code></pre></div></div>

<p>在打开的编辑器里，把这两条从：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>pick &lt;sha1&gt; feat: original change
pick &lt;sha2&gt; revert: revert original change
</code></pre></div></div>

<p>改成：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>drop &lt;sha1&gt; feat: original change
drop &lt;sha2&gt; revert: revert original change
</code></pre></div></div>

<p>如果 rebase 过程中状态混乱，先退出：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git rebase <span class="nt">--abort</span>
</code></pre></div></div>

<p>rebase 成功后检查：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git status
git log <span class="nt">--oneline</span> <span class="nt">-5</span>
</code></pre></div></div>

<p>如果此时看到类似：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Your branch is behind origin/... by 2 commits
</code></pre></div></div>

<p>并且确认这 2 个 commit 正是刚刚 drop 掉的那两个，不要 <code class="language-plaintext highlighter-rouge">git pull</code>，否则会把它们又拉回来。应该执行：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git push <span class="nt">--force-with-lease</span>
</code></pre></div></div>

<p>完整链路：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git revert HEAD
git push

git rebase <span class="nt">-i</span> HEAD~3
<span class="c"># 把原 commit 和 revert commit 改成 drop</span>

git status
git log <span class="nt">--oneline</span> <span class="nt">-5</span>

git push <span class="nt">--force-with-lease</span>
</code></pre></div></div>

<h2 id="两个本地-clone-同一个分支一个-force-push-后另一个怎么同步">两个本地 clone 同一个分支，一个 force push 后另一个怎么同步？</h2>

<p>场景：</p>

<ul>
  <li>本地 1 对 <code class="language-plaintext highlighter-rouge">feat/dump_llo</code> 做了 rebase。</li>
  <li>本地 1 执行了 <code class="language-plaintext highlighter-rouge">git push --force-with-lease</code>。</li>
  <li>远端 <code class="language-plaintext highlighter-rouge">origin/feat/dump_llo</code> 变成一条新的提交历史。</li>
  <li>本地 2 仍然保留旧历史。</li>
</ul>

<p>这时本地 2 执行：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git pull
</code></pre></div></div>

<p>可能会看到：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>+ 815a613...2889e70 feat/dump_llo -&gt; origin/feat/dump_llo  (forced update)
hint: You have divergent branches and need to specify how to reconcile them.
fatal: Need to specify how to reconcile divergent branches.
</code></pre></div></div>

<p>原因是本地旧历史和远端新历史已经不是 fast-forward 关系。<code class="language-plaintext highlighter-rouge">git pull</code> 不知道你想 merge、rebase，还是只接受 fast-forward。</p>

<p>如果本地 2 没有需要保留的提交，最直接的做法是让它完全对齐远端：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git fetch origin
git switch feat/dump_llo
git reset <span class="nt">--hard</span> origin/feat/dump_llo
</code></pre></div></div>

<p>如果担心本地 2 上有内容需要找回，先备份：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git branch backup/feat-dump-llo-local2
git fetch origin
git reset <span class="nt">--hard</span> origin/feat/dump_llo
</code></pre></div></div>

<p>这个场景不推荐直接 <code class="language-plaintext highlighter-rouge">git pull</code>，因为它会尝试整合两条不同历史，而你的目标通常只是“本地 2 跟远端同步”。</p>

<h2 id="git-worktree-和-branch-有什么区别">Git worktree 和 branch 有什么区别？</h2>

<p><code class="language-plaintext highlighter-rouge">branch</code> 是 Git 历史上的提交指针，例如 <code class="language-plaintext highlighter-rouge">main</code>、<code class="language-plaintext highlighter-rouge">codex/reconstruct</code>。它记录代码历史走到哪里。</p>

<p><code class="language-plaintext highlighter-rouge">worktree</code> 是磁盘上的实际工作目录。一个仓库可以有多个 worktree，每个 worktree 都是一份可以编辑、运行、提交代码的工作现场。</p>

<p>可以这样理解：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>repo
├─ main worktree -&gt; main
├─ worktree A    -&gt; feature/login
└─ worktree B    -&gt; codex/reconstruct
</code></pre></div></div>

<p>关键点：</p>

<ul>
  <li>一个 worktree 同一时间通常只能 checkout 一个 branch，或者处于 detached HEAD。</li>
  <li>一个 branch 同一时间通常只能被一个 worktree 占用。</li>
  <li>如果某个 branch 已经被别的 worktree 使用，新的 worktree 不能直接 <code class="language-plaintext highlighter-rouge">git switch</code> 到这个 branch。</li>
  <li>删除 Codex 对话不等于删除 Git worktree。worktree 是 Git 记录的本地目录，需要用 Git 命令清理。</li>
</ul>

<h2 id="worktree-常用命令有哪些">worktree 常用命令有哪些？</h2>

<p>查看当前有哪些 worktree：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git worktree list
</code></pre></div></div>

<p>查看某个 worktree 的状态：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git <span class="nt">-C</span> <span class="s2">"C:/path/to/worktree"</span> status <span class="nt">--short</span> <span class="nt">--branch</span>
</code></pre></div></div>

<p>新建一个 worktree，并让它绑定新分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git worktree add <span class="s2">"C:/path/to/new-worktree"</span> <span class="nt">-b</span> feat/example main
</code></pre></div></div>

<p>新建一个 worktree，并让它 checkout 已有分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git worktree add <span class="s2">"C:/path/to/new-worktree"</span> feat/example
</code></pre></div></div>

<p>删除不再使用的普通 worktree：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git worktree remove <span class="s2">"C:/path/to/old-worktree"</span>
</code></pre></div></div>

<p>如果目录已经被手动删掉，但 Git 记录还在：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git worktree prune
</code></pre></div></div>

<p>查看当前 worktree 绑定的分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git branch <span class="nt">--show-current</span>
</code></pre></div></div>

<p>如果显示不出分支名，可能是 detached HEAD：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git status <span class="nt">--short</span> <span class="nt">--branch</span>
</code></pre></div></div>

<p>切换分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git switch main
git switch codex/reconstruct
</code></pre></div></div>

<h2 id="worktree-里目标分支被占用怎么办">worktree 里目标分支被占用怎么办？</h2>

<p>先查看 worktree 结构：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git worktree list
</code></pre></div></div>

<p>可能会看到类似：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>E:/my_github_page/wqh011128.github.io                         [codex/reconstruct]
C:/Users/吴启航/.codex/worktrees/4855/wqh011128.github.io     (detached HEAD)
</code></pre></div></div>

<p>这说明 <code class="language-plaintext highlighter-rouge">codex/reconstruct</code> 分支已经被 <code class="language-plaintext highlighter-rouge">E:/my_github_page/wqh011128.github.io</code> 这个 worktree 占用，当前 Codex worktree 不能直接接管这个分支。</p>

<p>如果被占用的是主工作目录，不能用下面这条命令删除：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git worktree remove <span class="s2">"E:/my_github_page/wqh011128.github.io"</span>
</code></pre></div></div>

<p>否则会看到：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>fatal: 'E:/my_github_page/wqh011128.github.io' is a main working tree
</code></pre></div></div>

<p>正确思路是让占用分支的 worktree 先切到别的分支，例如 <code class="language-plaintext highlighter-rouge">main</code>，从而释放目标分支。</p>

<p>推荐步骤：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git worktree list
git <span class="nt">-C</span> <span class="s2">"E:/my_github_page/wqh011128.github.io"</span> status <span class="nt">--short</span> <span class="nt">--branch</span>
git <span class="nt">-C</span> <span class="s2">"C:/Users/吴启航/.codex/worktrees/4855/wqh011128.github.io"</span> status <span class="nt">--short</span> <span class="nt">--branch</span>
</code></pre></div></div>

<p>确认没有未提交内容后，让主工作目录切回 <code class="language-plaintext highlighter-rouge">main</code>：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git <span class="nt">-C</span> <span class="s2">"E:/my_github_page/wqh011128.github.io"</span> switch main
</code></pre></div></div>

<p>然后在当前 worktree 接管目标分支：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git switch codex/reconstruct
</code></pre></div></div>

<p>最后确认：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git worktree list
git status <span class="nt">--short</span> <span class="nt">--branch</span>
</code></pre></div></div>

<p>期望结构类似：</p>

<div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>E:/my_github_page/wqh011128.github.io                         [main]
C:/Users/吴启航/.codex/worktrees/4855/wqh011128.github.io     [codex/reconstruct]
</code></pre></div></div>

<h2 id="worktree-上有本地修改时怎么处理">worktree 上有本地修改时怎么处理？</h2>

<p>如果目标分支或目标 worktree 上有本地修改，先决定是否保留。不要直接切分支、删除 worktree 或 reset。</p>

<p>如果确认不保留已跟踪文件的修改：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git restore <span class="nt">--staged</span> <span class="nt">--worktree</span> <span class="nb">.</span>
</code></pre></div></div>

<p>如果还存在未跟踪文件或目录，例如 <code class="language-plaintext highlighter-rouge">skills/</code>：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git clean <span class="nt">-fd</span> <span class="nt">--</span> <span class="s2">"skills/"</span>
</code></pre></div></div>

<p>然后再按需要 rebase 到 <code class="language-plaintext highlighter-rouge">main</code>：</p>

<div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>git rebase main
</code></pre></div></div>

<h2 id="日常使用-worktree-的习惯是什么">日常使用 worktree 的习惯是什么？</h2>

<ul>
  <li>主目录长期放 <code class="language-plaintext highlighter-rouge">main</code>。</li>
  <li>每个 Codex 任务使用单独 worktree 和单独分支。</li>
  <li>一个长期开发分支尽量固定在一个 worktree 使用。</li>
  <li>用完的 worktree 先检查 <code class="language-plaintext highlighter-rouge">git status --short --branch</code>，确认干净后再 <code class="language-plaintext highlighter-rouge">git worktree remove</code>。</li>
  <li>看到 detached HEAD 先不要慌，它只是说明当前 worktree 没有绑定分支；先用 <code class="language-plaintext highlighter-rouge">git worktree list</code> 判断目标分支是否被别的 worktree 占用。</li>
</ul>]]></content><author><name></name></author><category term="blog" /><summary type="html"><![CDATA[这篇笔记按“遇到的问题”整理 Git 日常操作。命令默认在项目根目录执行；如果涉及已经推送到远端的历史改写，优先用 --force-with-lease，不要直接用 --force。]]></summary></entry></feed>