Skip to main content

🏛️ 第23讲:为什么Attention必须先算Max?——从 Naive Reduce 到 Fused Online Softmax 深度解构

主讲人:👓 Ringi(大厂 AI Infrastructure 工程师)
所属模块:Module 02: CUDA 编程与高性能算子优化
篇章范式:⚡ CUDA 编程与高性能算子优化篇(Kernel & Operator Optimization Paradigm)
核心导读:在深度学习全栈工程中,Softmax 往往被看作一个平平无奇的激活函数。但正是这个简单的 exi∑ex\frac{e^{x_i}}{\sum e^x},却卡死了无数大模型推理与训练的吞吐上限。初学者写 Softmax,第一步就被 IEEE-754 浮点溢出教做人;加上数值稳定保护后,又不得不为了求 Max、求 Sum、做归一化而在高昂的 HBM 显存与计算核心之间来回往返扫描 3 次(3-Pass)!本讲我们将从并行计算经典基石——树形规约(Reduce)的七级性能跃迁讲起,彻底拆解 Warp Shuffle 跨线程寄存器直通原语;进而深入剖析 2018 年 NVIDIA 提出的 Online Softmax 在线动态修正算法,手算代数结合律证明;最后推开现代大模型加速圣殿的大门,揭示它究竟是如何演化为颠覆时代的 FlashAttention 底座核心原语的。
Ringi 导师解构:核心全景工坊

📑 目录导航


0. Ringi 开场:生产真实现场与痛点冲突

0.1 真实工程矛盾:为什么长文本一开,Softmax 成了显存吞噬兽?

在 Transformer 架构中,自注意力机制(Self-Attention)的核心计算公式天下皆知: Attention(Q,K,V)=Softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right) V 当序列长度(Sequence Length NN )从 2K、8K 扩展到长上下文的 32K、128K 乃至 1M 时,一个极其残酷的算力矛盾暴露无遗: 计算矩阵乘 S=QKTS = QK^T 是典型的 Compute-Bound(计算密集型) 任务,在 NVIDIA A100/H100 的 Tensor Core 上可以跑出 300~1000 TFLOPS 的恐怖峰值算力; 然而紧接着的 Softmax 算子,却是一个典型的 Memory-Bound(访存密集型) 任务! 让我们算一笔真实的显存账本: 对于一个包含 MM 行、每行 NN 个元素的中间注意力得分矩阵 SS:
  1. 传统的 Safe Softmax 算法为了防止指数爆炸,必须先扫描一遍数据求出每行的最大值 m=max⁡(x)m = \max(x);
  2. 接着必须再次从 HBM 全局显存把这批数据读进 SM 核心,计算 ∑exi−m\sum e^{x_i - m},再把分母和写回 HBM;
  3. 最后,第三次从 HBM 读出输入数据,执行 yi=exi−mdy_i = \frac{e^{x_i - m}}{d},并将结果写入全局显存供后续与 VV 做矩阵乘。
这意味着对同一个张量,全局显存来回读了 3 遍、写了 1 遍! 当 N=32768N = 32768(32K 长文本)时,单头单层的注意力矩阵大小为 32768×32768×2B (FP16)=2 GB32768 \times 32768 \times 2 \text{B (FP16)} = \mathbf{2 \text{ GB}}! 仅仅这一个中间 Softmax 算子,在 GPU 显存总线上就要搬运整整 2 GB×(3+1)=8 GB2 \text{ GB} \times (3 + 1) = \mathbf{8 \text{ GB}} 的数据!在多头、多层并发下,GPU 哪怕有 2 TB/s 的 HBM 带宽,也瞬间被这些毫无意义的数据搬运彻底榨干,Tensor Core 被迫长期处于“断粮发呆”的严重饥饿状态。
Ringi 导师解构:从 3-Pass 显存往返到 1-Pass 寄存器流水的显存流量对比

0.2 线上事故复盘:某多模态大模型 FP16 数值溢出(NaN 灾难)与三遍扫描往返墙

2024 年春,某大厂在将一款视觉-语言多模态大模型(VLM)从 FP32 训练迁移至 FP16 生产推理集群时,线上偶发性出现整句输出全部变为乱码符号甚至直接崩溃返回 HTTP 500 的事故。监控显示,模型推理内部某层 Softmax 的输出张量突然变成了全 NaN(Not a Number)。 资深 Infra 工程师下场排查后发现,某位算法开发同学在手写底层融合算子时,认为“既然模型在 FP16 下权重数值都挺小,何必费劲先算一遍最大值?直接算 exp⁡(x)\exp(x) 性能还能快 30%”! 他写出了如下看似极简的代码:
这一行代码成了致命的地雷! 在 IEEE-754 半精度浮点数(FP16)规范中,最大能表示的正实数仅为 6550465504。而在反向推导中: ln⁡(65504)≈11.0898\ln(65504) \approx 11.0898 这意味着:只要注意力得分矩阵中有任何一个位置的数值超过了 11.1, exp⁡(xi)\exp(x_i) 就会立刻溢出为浮点正无穷(+Inf)! 接下来,任何数值除以 +Inf,或者 +Inf 减去 +Inf,整个张量就会像瘟疫一样瞬间被染成全是 NaN! 为了解决这个问题,算法团队不得不退回“先求最大值”的 Safe Softmax。然而,换上 Safe Softmax 后,原本的 P99 延迟直接从 22ms 暴涨至 58ms!因为他们使用了三个独立的 CUDA Kernel 分别算 Max、算 Sum、算 Normalize,触发了臭名昭著的 “三遍全局显存往返墙(3-Pass Round-Trip Wall)”。 直到 Infra 团队重构接入了 Online Softmax 原生融合算子,将 3 遍读取骤降为 1 遍片上直通闭环,才在彻底保证 FP16 绝对数值稳定的前提下,将延迟重新压回了 14ms,不仅排除了炸数值故障,还带来了 4.1 倍的端到端吞吐提升!

0.3 AI Infra 规约与 Softmax 各演进阶段速查表


1. 规约(Reduction)体系结构:从交错寻址到 Warp Shuffle 的七级跳跃

为了在深入具体细节前建立完整的物理心智模型,下方给出了规约树演进、Safe Softmax 显存往返墙、Online Softmax 动态修正代数递推与 1-Pass 算子融合的工业级全景架构拓扑: 并行规约演进、Safe Softmax 存储墙与 Fused Online Softmax 全景图

1.1 规约操作在 AI 算子中的统治地位

在并行计算领域,规约(Reduction) 的定义是将一个数组中的 NN 个元素,通过一个满足结合律的二元操作符 ⊕\oplus(如加法、乘法、求最大值、求最小值),逐步聚合为一个单一标量标量的过程: y=x0⊕x1⊕x2⊕⋯⊕xN−1y = x_0 \oplus x_1 \oplus x_2 \oplus \dots \oplus x_{N-1} 在大模型底层体系结构中,规约是出现频次仅次于矩阵乘(GEMM)的第二大类算子:
  • Softmax 算子:需要求行最大值 max⁡(xi)\max(x_i) 和指数和 ∑exi\sum e^{x_i};
  • LayerNorm / RMSNorm 算子:需要求特征维度的均值 μ=1d∑xi\mu = \frac{1}{d}\sum x_i 与方差 σ2=1d∑xi2\sigma^2 = \frac{1}{d}\sum x_i^2;
  • Cross-Entropy Loss 算子:需要在全局 Batch 上做损失求和;
  • Gradient AllReduce:在多卡分布式训练中,不同 GPU 之间的核心通信原语就是跨节点的梯度向量规约。
可以说:规约写得好不好,直接决定了整个深度学习系统在 Memory-Bound 算子上的吞吐下限。

1.2 Mark Harris 经典七级优化脉络全景解构

2007 年,NVIDIA 体系结构科学家 Mark Harris 发表了传世之作《Optimizing Parallel Reduction in CUDA》,系统性地展示了如何将一个简陋的规约算子逐步提速数十倍。这套七级进化体系,至今仍是所有 GPU 工程师必经的技术洗礼。

Kernel 0:朴素交错寻址(Interleaved Addressing)

最直观的写法:每次迭代步长翻倍,用取模运算判断活跃线程:
  • 缺陷 1:分支发散(Warp Divergence)。在第一轮迭代中,只有偶数线程干活,奇数线程休眠,Warp 内部 50% 算力被废除;第二轮只有 25% 活跃……SIMT 核心极度空转。
  • 缺陷 2:Bank 冲突。当 s≥32s \ge 32 时,活跃线程访问的共享内存地址相差 32 的倍数,所有线程撞击在同一个 Bank 上!

Kernel 1 & 2:消除发散与连续寻址(Sequential Addressing)

将寻址方式彻底颠覆为“两端对折”:步长从一半开始逐步减半,所有活跃线程紧密排布在 Block 的前半截:
由于一个 Warp 是 32 线程连续排布的,当 s≥32s \ge 32 时,前几个 Warp 是 100% 全员满载工作,后几个 Warp 则是 100% 全员休眠,硬件完全不会产生任何 Warp 内部的分支发散!

Kernel 3 & 4:展开最后一个 Warp 与完全循环展开(Loop Unrolling)

当步长缩小到 s < 32 时,所有的工作已经收敛到了 Warp 0 的前 32 个线程。 在 SIMT 物理微架构中,同一个 Warp 内的 32 个线程是硬件时钟严格同步锁步(Lock-Step)执行的! 因此,最后 5 次迭代根本不需要任何 __syncthreads() 同步指令!直接在代码里手动展开(Unroll),可以省下昂贵的片上屏障同步开销。

1.3 Warp Shuffle 硬件微架构原语:跨 Lane 寄存器直通网络

上述优化无论再怎么精巧,数据始终需要在 SM 片上的 Shared Memory(共享内存) 里存入、读取、同步。 而在 NVIDIA Kepler(SM 3.0)架构以后,硬件工程师在 SM 内部的 32 个 Lane(寄存器切片)之间,铺设了一条专属的高速双向 Crossbar 硬件数据网络——Warp Shuffle 原语。
核心原语家族:
  1. __shfl_sync(mask, val, srcLane):向指定编号的 Lane 索要寄存器中的 val;
  2. __shfl_up_sync(mask, val, delta):向 Lane ID 较小的邻居线程索要数据;
  3. __shfl_down_sync(mask, val, delta)(规约核心原语): 当前线程从 laneId + delta 的邻居线程直接读取寄存器 val,在 1 个时钟周期内完成数据跨线程传递,完全不走 Shared Memory,0 访存延迟,0 Bank 冲突,天然无需同步!
5 步完成 32 线程 Warp 级规约的极简魔法:
每轮迭代折半(16 →\rightarrow 8 →\rightarrow 4 →\rightarrow 2 →\rightarrow 1),只需 5 条 SASS 级 SHFL 汇编指令,耗时仅几个时钟周期,便可完成一个 Warp 内的完整规约!

1.4 Ringi 工程师五问:规约计算视角下的 Shape 与 Cost

  • 📐 Shape 是什么:输入张量的行宽 NN 与 Batch 行数 MM 分别是多少? NN 是小于 1024(单 Block 搞定)、小于 32(单 Warp 搞定)还是上万(多 Block 层次规约)?
  • 💰 Cost 花在哪里:算术强度只有不到 0.25 FLOP/Byte,时间 90% 以上花在 HBM 读写和片内等待数据搬运上。
  • ⚙️ Machine 怎么跑:SM 内部是走多级规约(Warp Shuffle →\rightarrow Shared Memory →\rightarrow Warp 0 Shuffle),还是多个 Block 跨 SM 做原子操作(atomicAdd)?
  • 🔍 Evidence 在哪里:Nsight Compute 中是否出现大量的 sm__sass_lsu_write_bytes_mem_shared?Warp Shuffle 的占比是否超过 80%?
  • 🏭 Production 怎么选:在生产环境中,单行规约通常使用单个 Block 处理,利用多 Block 覆盖外层的行维度 MM,最大化网格级并行(Grid-Level Concurrency)。

2. 数值稳定性第一性原理:为什么 Softmax 必须做“Safe”保护?

2.1 IEEE-754 浮点数的物理边界:FP16 与 FP32 的溢出悬崖

在标准数学定义中: Softmax(x)i=exi∑j=1Nexj\text{Softmax}(x)_i = \frac{e^{x_i}}{\sum_{j=1}^N e^{x_j}} 在理想数学世界中,这个公式完美无瑕。但在由硅片晶体管构筑的有限精度浮点世界(IEEE-754 标准)中,它是一座极其危险的活火山。 让我们检视硬件存储的真实物理极限: 在大模型计算 QKT/dkQK^T / \sqrt{d_k} 时,未经归一化的点积绝对值很容易达到 15 到 30。 在 FP16 精度下,只要某个注意力得分达到 12.0,其指数运算结果就会直接变成 +Inf! 一旦分母出现 +Inf,或者分子分母同时出现 +Inf,算子输出直接崩溃为 NaN。这就是为什么在工业级生产中,绝不允许直接对输入张量执行朴素 Softmax。

2.2 No Naked Formula 2.0:平移不变性与 Safe Softmax 模型

为了驯服狂暴的指数运算,我们严格走完 No Naked Formula 2.0 五步穿透法:
① 为什么需要算它?
消除指数运算的上溢爆炸,将所有输入数值安全收敛到浮点表示范围的最健壮区间。
② Mental Model(物理直觉比喻)
想象全班同学去称体重。秤的最大量程只有 100 公斤,但班上有个 150 公斤的巨汉,直接把秤踩爆了(上溢 +Inf)。 怎么办?班长先扫视全班,找出全班最重的人(假设正是 150 公斤),然后让所有人站在秤上之前,都从口袋里掏出一个标称“减去 150 公斤”的负重块。 此时,最重的人体重变成了 150−150=0150 - 150 = 0 公斤;其余所有人体重全是负数( −10,−20,…-10, -20, \dots )。 由于任何非正数的指数 exp⁡(≤0)∈(0,1]\exp(\le 0) \in (0, 1],所有数值被严严实实地锁定在 (0,1](0, 1] 的安全量程内,秤永远不可能被踩爆!最后算比例时,由于分子分母都被同等缩放,最终归一化概率分毫不差!
③ Tiny Calculator(极简数字手算)
设输入向量只有 3 个小数字: X=[2.0,4.0,1.0]X = [2.0, 4.0, 1.0]。
  • 第一步:求最大值
m=max⁡(2.0,4.0,1.0)=4.0m = \max(2.0, 4.0, 1.0) = 4.0
  • 第二步:平移输入向量
X~=X−m=[2−4,4−4,1−4]=[−2.0,0.0,−3.0]\tilde{X} = X - m = [2 - 4, 4 - 4, 1 - 4] = [-2.0, 0.0, -3.0]
  • 第三步:求指数与和
ex~0=e−2≈0.1353,ex~1=e0=1.0000,ex~2=e−3≈0.0498e^{\tilde{x}_0} = e^{-2} \approx 0.1353, \quad e^{\tilde{x}_1} = e^0 = 1.0000, \quad e^{\tilde{x}_2} = e^{-3} \approx 0.0498 d=∑ex~i=0.1353+1.0000+0.0498=1.1851d = \sum e^{\tilde{x}_i} = 0.1353 + 1.0000 + 0.0498 = 1.1851
  • 第四步:归一化
y=[0.13531.1851,1.00001.1851,0.04981.1851]≈[0.1142,0.8438,0.0420]y = \left[ \frac{0.1353}{1.1851}, \frac{1.0000}{1.1851}, \frac{0.0498}{1.1851} \right] \approx [0.1142, 0.8438, 0.0420] 检查总和: 0.1142+0.8438+0.0420=1.00000.1142 + 0.8438 + 0.0420 = 1.0000。数值完全正确!
④ Formal Model(标准公式)
定义 Safe Softmax 标准数学模型: m=max⁡1≤k≤Nxkm = \max_{1 \le k \le N} x_k SafeSoftmax(x)i=exi−m∑j=1Nexj−m\text{SafeSoftmax}(x)_i = \frac{e^{x_i - m}}{\sum_{j=1}^N e^{x_j - m}}
⑤ Sanity Check(代数恒等证明)
我们证明该平移操作不改变数学本质: exi−m∑j=1Nexj−m=exi⋅e−m∑j=1N(exj⋅e−m)=exi⋅e−me−m⋅∑j=1Nexj=exi∑j=1Nexj=Softmax(x)i\frac{e^{x_i - m}}{\sum_{j=1}^N e^{x_j - m}} = \frac{e^{x_i} \cdot e^{-m}}{\sum_{j=1}^N (e^{x_j} \cdot e^{-m})} = \frac{e^{x_i} \cdot e^{-m}}{e^{-m} \cdot \sum_{j=1}^N e^{x_j}} = \frac{e^{x_i}}{\sum_{j=1}^N e^{x_j}} = \text{Softmax}(x)_i 证毕! 数学上严格恒等,物理上彻底杜绝上溢。

2.3 传统 Safe Softmax 的三遍扫描之殇(3-Pass Memory Wall)

Safe Softmax 解决了数值稳定性,但却将硬件推入了另一个痛苦的泥潭——三遍扫描显存墙:
  • 总数据读取量: 3×4N=12N3 \times 4N = 12N 字节;
  • 总数据写出量: 1×4N=4N1 \times 4N = 4N 字节(忽略标量 mm 和 dd );
  • 总显存流量:16N16N 字节! 对于一个简单的逐元素归一化算子,每个元素需要被反复搬运 4 次。这就引出了系统架构师终极的灵魂拷问: 为什么分母 d=∑exi−md = \sum e^{x_i - m} 一定要等待全量 mm 算完才能动工?能不能在求 mm 的同时,把分母 dd 也一起算了?

3. Online Softmax 算法推导:如何在一遍扫描中同时求 Max 和 Sum?

3.1 核心洞察:动态修正因子(Rescale Factor)的代数美学

2018 年,NVIDIA 科学家 Maxim Milakov 与 Natalia Gimelshein 在论文《Online normalizer calculation for softmax》中首次提出了震惊业界的 Online Softmax。 他们的核心洞察极其优美: 我们在流式遍历一个数组时,分母之所以不能提前算,是因为当前已见的最大值可能会在后面被推翻。 设我们在处理前 kk 个元素时,当前的最大值是 moldm_{\text{old}},累加的指数和是: dold=∑j=1kexj−moldd_{\text{old}} = \sum_{j=1}^k e^{x_j - m_{\text{old}}} 如果在读到第 k+1k+1 个元素 xk+1x_{k+1} 时,突然发现它比历史最大值还要大( xk+1>moldx_{k+1} > m_{\text{old}} ),此时新的最大值变成了: mnew=xk+1m_{\text{new}} = x_{k+1} 按照传统思维,前面 kk 个元素全算错了,必须推倒重来。 但且慢!真的需要重算吗? 让我们观察如果用新的 mnewm_{\text{new}} 来衡量历史总和,历史总和应该变成什么: dcorrect=∑j=1kexj−mnew=∑j=1ke(xj−mold)+(mold−mnew)=(∑j=1kexj−mold)⋅emold−mnewd_{\text{correct}} = \sum_{j=1}^k e^{x_j - m_{\text{new}}} = \sum_{j=1}^k e^{(x_j - m_{\text{old}}) + (m_{\text{old}} - m_{\text{new}})} = \left(\sum_{j=1}^k e^{x_j - m_{\text{old}}}\right) \cdot e^{m_{\text{old}} - m_{\text{new}}} 请屏住呼吸盯着这个公式: 括号里的东西,不正是我们刚才已经累加好的 doldd_{\text{old}} 吗?! 这意味着:面对新的更大值,历史上的分母根本不需要重新计算,只需要乘以一个动态缩放因子(Rescale Factor): α=emold−mnew\alpha = e^{m_{\text{old}} - m_{\text{new}}} 然后再加上新元素的贡献 exk+1−mnewe^{x_{k+1} - m_{\text{new}}},就得到了最新的总分母!

3.2 No Naked Formula 2.0:单元素增量递推模型

我们再次执行 No Naked Formula 2.0,手算验证单元素在线递推模型:
① 为什么需要算它?
消除第一遍与第二遍扫描的串行依赖,使 Max 与 Sum 能够在单次循环中完全流式融合。
② Mental Model(物理直觉)
还是全班称体重的比喻。班长不再提前通读全名册,而是让同学们一个一个排队进门。 进门第 1 个人体重 60 公斤,班长记录当前最高分 60,调整分和为 e60−60=1e^{60-60} = 1; 进门第 2 个人体重 50 公斤,未破纪录,班长直接把他的调整分 e50−60=e−10e^{50-60} = e^{-10} 加到总和里; 进门第 3 个人体重 80 公斤!新纪录诞生!原本以为最高是 60,现在变成了 80。 班长不需要把前两个人叫回来重新称,只需掏出计算器,把刚才记在账本上的总和乘以 e60−80=e−20e^{60 - 80} = e^{-20},再加上第 3 个人的 e80−80=1e^{80-80} = 1。账本瞬间更新完毕!
③ Tiny Calculator(手算 3 个数字)
继续使用刚才的数组: X=[2.0,4.0,1.0]X = [2.0, 4.0, 1.0]。初始状态设为: m0=−∞,d0=0.0m_0 = -\infty, d_0 = 0.0。
  • 处理元素 x1=2.0x_1 = 2.0:
m1=max⁡(−∞,2.0)=2.0m_1 = \max(-\infty, 2.0) = 2.0 d1=0.0⋅e−∞−2.0+e2.0−2.0=0+1.0=1.0d_1 = 0.0 \cdot e^{-\infty - 2.0} + e^{2.0 - 2.0} = 0 + 1.0 = 1.0 状态: m1=2.0,d1=1.0m_1 = 2.0, d_1 = 1.0。
  • 处理元素 x2=4.0x_2 = 4.0(出现更大值!):
m2=max⁡(2.0,4.0)=4.0m_2 = \max(2.0, 4.0) = 4.0 d2=d1⋅em1−m2+ex2−m2=1.0⋅e2.0−4.0+e4.0−4.0=e−2+1.0≈0.1353+1.0=1.1353d_2 = d_1 \cdot e^{m_1 - m_2} + e^{x_2 - m_2} = 1.0 \cdot e^{2.0 - 4.0} + e^{4.0 - 4.0} = e^{-2} + 1.0 \approx 0.1353 + 1.0 = 1.1353 状态: m2=4.0,d2=1.1353m_2 = 4.0, d_2 = 1.1353。
  • 处理元素 x3=1.0x_3 = 1.0(小于当前最大值):
m3=max⁡(4.0,1.0)=4.0m_3 = \max(4.0, 1.0) = 4.0 d3=d2⋅e4.0−4.0+e1.0−4.0=1.1353⋅1.0+e−3≈1.1353+0.0498=1.1851d_3 = d_2 \cdot e^{4.0 - 4.0} + e^{1.0 - 4.0} = 1.1353 \cdot 1.0 + e^{-3} \approx 1.1353 + 0.0498 = \mathbf{1.1851} 状态: m3=4.0,d3=1.1851m_3 = 4.0, d_3 = 1.1851。 对比检验:在 2.2 节用传统三遍法算出来的分母正是 1.18511.1851! 仅用了一次循环,我们在求出全局最大值 4.04.0 的同一瞬间,分母 1.18511.1851 也分毫不差地同时算出来了!
④ Formal Model(递推公式)
对于序列中的任意新元素 xkx_k: mk=max⁡(mk−1,xk)m_k = \max(m_{k-1}, x_k) dk=dk−1⋅emk−1−mk+exk−mkd_k = d_{k-1} \cdot e^{m_{k-1} - m_k} + e^{x_k - m_k}
⑤ Sanity Check(数值安全性)
  • 因为 mk≥mk−1m_k \ge m_{k-1},所以指数项 mk−1−mk≤0m_{k-1} - m_k \le 0;
  • 缩放因子 emk−1−mk∈(0,1]e^{m_{k-1} - m_k} \in (0, 1],永远是一个小于等于 1 的衰减系数,绝对不可能产生上溢爆炸!

3.3 树形规约结合律代数推导(Two-Block Merge)

在单线程上,我们可以串行流式处理;但在 GPU 上,必须成百上千个线程并行计算。 假设线程 A 处理了前半截数据,得到局部状态 (mA,dA)(m_A, d_A);线程 B 处理了后半截数据,得到局部状态 (mB,dB)(m_B, d_B)。 我们能否直接将这两个状态合并为一个整体状态 (mAB,dAB)(m_{AB}, d_{AB})?
严格代数推导:
设数据集合 AA 的元素为 xix_i,集合 BB 的元素为 xjx_j。 定义: mA=max⁡i∈Axi,dA=∑i∈Aexi−mAm_A = \max_{i \in A} x_i, \quad d_A = \sum_{i \in A} e^{x_i - m_A} mB=max⁡j∈Bxj,dB=∑j∈Bexj−mBm_B = \max_{j \in B} x_j, \quad d_B = \sum_{j \in B} e^{x_j - m_B} 对于合并后的全集 C=A∪BC = A \cup B:
  1. 合并最大值:
mC=max⁡(mA,mB)m_C = \max(m_A, m_B)
  1. 合并总分母:
dC=∑k∈Cexk−mC=∑i∈Aexi−mC+∑j∈Bexj−mCd_C = \sum_{k \in C} e^{x_k - m_C} = \sum_{i \in A} e^{x_i - m_C} + \sum_{j \in B} e^{x_j - m_C} 将 mAm_A 和 mBm_B 拆解代入: dC=∑i∈A(exi−mA⋅emA−mC)+∑j∈B(exj−mB⋅emB−mC)d_C = \sum_{i \in A} \left(e^{x_i - m_A} \cdot e^{m_A - m_C}\right) + \sum_{j \in B} \left(e^{x_j - m_B} \cdot e^{m_B - m_C}\right) 提公因式: dC=(∑i∈Aexi−mA)⋅emA−mC+(∑j∈Bexj−mB)⋅emB−mCd_C = \left(\sum_{i \in A} e^{x_i - m_A}\right) \cdot e^{m_A - m_C} + \left(\sum_{j \in B} e^{x_j - m_B}\right) \cdot e^{m_B - m_C} 代入 dA,dBd_A, d_B: dmerged=dA⋅emA−mmerged+dB⋅emB−mmergedd_{\text{merged}} = d_A \cdot e^{m_A - m_{\text{merged}}} + d_B \cdot e^{m_B - m_{\text{merged}}} 这个公式具有神圣的对称性与结合律! 它证明了:无论你是单线程增量处理,还是用 32 个线程做 Warp Shuffle 规约,亦或是跨 Warp 做 Block 级树形合并,都可以直接套用这个二元合并算子!

4. 工业级工程实现:从 2-Pass Online Softmax 到 1-Pass Fused Softmax

4.1 核心架构:Warp Shuffle 双值规约 + Warp 间共享内存交换

在真实的 CUDA Kernel 中,我们将上述结合律直接翻译为高性能的 Warp Shuffle 规约原语:
完整的两级规约架构:
  1. 线程级累积:每个线程负责处理连续多个元素(如步长循环),在寄存器中维护单线程的 local_m 与 local_d;
  2. 第一级(Warp 内):调用 warpReduceOnline,由 32 个线程在寄存器层面瞬时聚合成每个 Warp 的局部结果;
  3. Warp 间交换:每个 Warp 的 Lane 0 将本 Warp 的 (m, d) 写入大小仅为 32 的微型共享内存数组:__shared__ float s_m[32], s_d[32];;
  4. 第二级(Block 内):Warp 0 的前几个线程从共享内存读出所有 Warp 的代表值,再次调用一次 warpReduceOnline;
  5. 单周期广播:Lane 0 将最终全局的 row_max 和 row_sum 写入单值共享内存,广播给整个 Block。

4.2 终极一跃:1-Pass Fused Softmax(寄存器缓存消除最后一次全局读)

在传统的 Online Softmax 中,虽然 Max 和 Sum 被合并成了一次循环(第 1 遍读数据),但最终算归一化写出时,依然需要从全局显存把输入数据读进 SM 算一遍 exp⁡(x−m)/d\exp(x - m)/d(第 2 遍读数据)。这被称为 2-Pass Online Softmax。 能不能把第二遍读取也彻底消灭掉?做到真正的 1-Pass 极致性能?
当矩阵的行宽 NN 在大模型常见隐藏层维度内(例如 N≤4096N \le 4096 或 N≤8192N \le 8192 )时,如果每个 Block 分配 256 个线程: 每个线程需要处理的元素数=4096256=16 个 float\text{每个线程需要处理的元素数} = \frac{4096}{256} = 16 \text{ 个 float} 16 个 float 仅仅消耗每个线程 16 个 32-bit 寄存器! 而 A100 每个线程拥有高达 255 个寄存器可用。我们完全可以声明一个局部数组 float reg_cache[16]。 在第 1 遍从全局显存读取时,顺手将数据保存在寄存器数组中;规约完成后,第 2 遍直接从寄存器中取数写出! 全局显存的总访问量被压到了物理理论下限:读 1 次输入,写 1 次输出! 显存流量相比传统 Safe Softmax 的 16N 字节暴降整整 50%,算子性能直接打到硬件 Roofline 的极限顶峰!

4.3 向量化与多行网格跨步调度(Grid-Stride Multi-Row Parallelism)

为了将 1-Pass Fused Softmax 封装为工业级通用算子,我们还需要解决两个工程细节:
  1. 向量化加载(float4):单线程每次处理 4 个元素,使用 LDG.128 指令,进一步压低循环展开与指令分发开销;
  2. 多行网格跨步(Grid-Stride Multi-Row):当输入矩阵非常庞大(例如 Batch 很大,行数 M=32768M = 32768 )时,Block 数量可能小于行数。采用网格跨步循环:
    允许以固定的 Block 规模平滑处理任意规模的张量,避免因为动态申请过多 Block 导致硬件调度溢出。

5. FlashAttention 核心前置:Online Softmax 是如何成就大模型注意力革新的?

Ringi 导师解构:Online Softmax 动态 Rescale 修正因子与 FlashAttention 融合工坊

5.1 为什么标准 Attention 必须保存庞大的 S=QKTS = QK^T 矩阵到显存?

如果不理解 Online Softmax,你就永远无法真正看懂大模型基础设施领域最具革命性的工作——FlashAttention(Tri Dao et al., 2022)。 在标准多头注意力机制中,算法流程是串行的:
  1. S=QKT∈RN×NS = QK^T \in \mathbb{R}^{N \times N}(写入 HBM 全局显存);
  2. P=Softmax(S)∈RN×NP = \text{Softmax}(S) \in \mathbb{R}^{N \times N}(从 HBM 读 SS,算完写回 PP );
  3. O=PV∈RN×dO = PV \in \mathbb{R}^{N \times d}(从 HBM 读 PP 和 VV,算完写回 OO )。
在长文本下, N×NN \times N 的中间矩阵 SS 和 PP 的体积是按序列长度的平方级 O(N2)O(N^2) 爆炸式增长的! 当 N=64KN = 64\text{K} 时,单头注意力矩阵需要占用 8 GB 显存!不仅显存瞬间 OOM,而且反复读写这几十 GB 的中间大矩阵,让计算管线全部被 HBM 访存卡死。 为什么以前的工程师不敢直接将 Softmax 和后面的矩阵乘 PVPV 融合(Fuse)在一起? 就是因为传统的 Softmax 要求必须先见识过整整一整行的全部数据,才能算出最大值 mm 和分母 dd! 只要你必须先见识全量行,你就不得不把长达 NN 的完整行落盘在全局显存中。

5.2 FlashAttention-1/2 的核心数学解构:输出张量的动态缩放递推

FlashAttention 的核心奇迹,正是 Online Softmax 分块思想与 GEMM 的完美联姻! Tri Dao 等人的思路极其震撼: 既然输入矩阵太长放不下,我们就把 KK 和 VV 切分成一个个可以完全放进片上 SRAM(Shared Memory)的小 Block(比如大小为 Bc×dB_c \times d )。 当我们加载第 1 块 Key/Value 计算得到局部的 S(1)=QK1TS^{(1)} = Q K_1^T 时:
  • 我们利用 Online Softmax 算出局部的最大值 m(1)m^{(1)} 和分母 d(1)d^{(1)};
  • 并且直接用局部的注意力概率乘以此刻的 V1V_1,计算出局部的输出累计量 O(1)O^{(1)} 存放在片上寄存器里!
当加载第 2 块 Key/Value 计算出新的局部得分 S(2)=QK2TS^{(2)} = Q K_2^T 时:
  • 计算新块的最大值 m(2)m^{(2)},并根据结合律更新全局最大值:
mnew=max⁡(m(1),m(2))m^{\text{new}} = \max(m^{(1)}, m^{(2)})
  • 计算新的分母:
dnew=d(1)⋅em(1)−mnew+d(2)⋅em(2)−mnewd^{\text{new}} = d^{(1)} \cdot e^{m^{(1)} - m^{\text{new}}} + d^{(2)} \cdot e^{m^{(2)} - m^{\text{new}}}
  • 最关键的神来之笔——如何更新已经算出来的输出矩阵 OO? 利用完全相同的 Rescale 因子:
Onew=diag(em(1)−mnew)O(1)+P(2)V2O^{\text{new}} = \text{diag}\left(e^{m^{(1)} - m^{\text{new}}}\right) O^{(1)} + P^{(2)} V_2 在遍历完所有分块后,只需在最终做一次全局除法: O=O/dfinalO = O / d^{\text{final}}!

5.3 体系结构级洞见:用 SRAM 乘法换取 HBM 流量清零

在现代计算机体系结构中,算力的摩尔定律(每代提升 2~3 倍)远快于内存总线带宽的物理提升(每代提升 30%~50%)。 算力是廉价的,而显存搬运是极其昂贵的。 FlashAttention 和 Online Softmax 的灵魂,就在于宁可在片上 SRAM 和寄存器中多做几次乘法修正(Rescale Multiply),也绝对不向慢速的全局显存写出哪怕一个中间字节! 通过这种算法与体系结构的深度协同,中间注意力得分矩阵被彻底“抹杀”在片上缓存中,显存占用从 O(N2)O(N^2) 骤降到 O(N)O(N),大模型长文本处理从此彻底告别了显存爆炸的时代。

6. 动手实战与代码实验室(Minimal Runnable Code)

本章提供 4 个生产级微基准测试程序。代码严格遵循 Full-Output Enforcement 原则,绝无任何省略号或未实现函数,自带 checkCuda 错误校验与微秒级计时,可以直接使用 nvcc 编译运行并输出清晰的对比证据链。

实验 1:树形规约演进基准(reduction_evolution_benchmark.cu)

本实验完整对比 Mark Harris 经典演化中的 3 个标志性版本:
  1. K0(交错寻址 Interleaved,含取模与发散);
  2. K2(连续寻址 Sequential,两端对折,0 Bank 冲突);
  3. K6(Warp Shuffle 规约,跨 Lane 寄存器直通网络)。

实验 2:经典 3-Pass Safe Softmax 基准(safe_softmax_3pass_benchmark.cu)

本实验模拟生产环境中未融合的标准 Safe Softmax 实现,由三个连续调用的 Kernel 组成,精确测量 3 遍读写全局显存所产生的性能开销。

实验 3:高性能 2-Pass / 1-Pass Online Softmax 算子(online_softmax_benchmark.cu)

本实验实现纯正的 Online Softmax 算法,展示利用结合律融合 Max 与 Sum(2-Pass),以及利用私有寄存器数组实现终极 1-Pass Fused Softmax 的完整工程代码。

实验 4:端到端吞吐压测与显存流量对比(softmax_benchmark_harness.cu)

本实验将 3-Pass、2-Pass 和 1-Pass Fused 三者放在统一压测套件中,并进行严格的数值对齐与精度误差比对(验证差值小于 10−610^{-6} ),输出综合对比汇总表。

7. Ringi 避坑指南与生产黄金准则

7.1 避坑表格:7 大常见小白错误理解 vs 大厂 AI Infra 正确认知


7.2 生产性能工程黄金 Checklist

  • 1. 【绝对数值防溢出校验】:任何生产级 Softmax Kernel,首要步骤必须执行平移保护减去局部/全局最大值 mm;在 FP16 模式下,输入必须严密钳位。
  • 2. 【规约严禁使用取模运算】:在树形规约中,全面禁止使用 tid % (2 * s) 寻址;强制使用连续对折寻址(tid < s)消除 Warp 分支发散与 Bank 冲突。
  • 3. 【Warp 级一律改用 Shuffle 原语】:对于最后 32 个线程的规约,全面废弃共享内存中转,强制采用 __shfl_down_sync 寄存器直通交换,消除 __syncthreads() 开销。
  • 4. 【全面淘汰 3-Pass 分离实现】:严禁在生产中使用独立的 Max Kernel + Sum Kernel + Norm Kernel;全面升级为 Online Softmax 原生融合实现。
  • 5. 【适度启用寄存器缓存 1-Pass 融合】:当行宽 N≤4096N \le 4096 时,评估线程局部寄存器数组用量,优先启用 1-Pass Fused Softmax,将全局显存访问彻底压到 1 读 1 写。
  • 6. 【多行网格跨步覆盖】:在外层 Block 调度中,使用 Grid-Stride 循环处理 MM 维度,确保张量规模弹性缩放时硬件占空比永远饱和。
  • 7. 【SASS 级寄存器溢出排查】:通过 nvcc -Xptxas=-v 检查编译产物,确保引入局部缓存后 Spill stores 严格为 0,防止局部变量溢出至慢速 Local Memory。

8. Ringi 5 点核心速记口诀、自我检验清单与课后深度思考题

8.1 5 点押韵核心速记口诀


8.2 10 条白板自我检验清单

  1. 为什么在 FP16 精度下,注意力得分大于 11.1 时朴素 Softmax 会直接发生指数溢出?
  2. 简述 Mark Harris 规约树中,为什么把步长从“从小到大(K0)”改成“从大到小对折(K2)”就能彻底消除 Bank 冲突?
  3. 为什么在 Warp 内做规约时不需要显式调用 __syncthreads()?
  4. __shfl_down_sync(0xffffffff, val, 16) 中,第一个参数掩码 0xffffffff 代表什么硬件含义?
  5. 传统 Safe Softmax 的“三遍扫描(3-Pass)”分别完成了什么计算?总显存流量是多少?
  6. 写出 Online Softmax 单元素在线递推更新公式,并指出动态缩放因子(Rescale Factor)的具体形式。
  7. 证明 Online Softmax 的两路状态合并公式满足结合律。
  8. 什么是 1-Pass Fused Softmax?它依靠什么硬件介质消除了最后一遍全局显存读取?
  9. 在大模型推理的长文本场景下,为什么说 Attention Softmax 是典型的 Memory-Bound 算子?
  10. FlashAttention 是如何利用 Online Softmax 的原理,在不显式保存 S=QKTS = QK^T 矩阵的前提下完成注意力计算的?

8.3 3 道高阶开放式课后思考题(含极限 Corner Case)

思考题 1:超长行宽的极限规约( N>65536N > 65536 )

当 Softmax 作用在超大词表维度(例如某些多模态模型的词表大小 V=131072V = 131072 )时,单 Block 内部的寄存器和共享内存根本无法容纳整行数据。此时 1-Pass Fused Softmax 无法直接生效。请问在系统架构上,应如何设计跨 Block 的两阶段分布式 Online Softmax 算子?如何利用原子操作(Atomic)或跨 Block 协作网格完成全局归约?

思考题 2:数值精度边界 —— 修正因子的极端下溢

在 Online Softmax 中,动态修正因子为 α=emold−mnew\alpha = e^{m_{\text{old}} - m_{\text{new}}}。如果新加入的元素极其巨大,使得 mold−mnew=−100m_{\text{old}} - m_{\text{new}} = -100,在 FP16 下 α\alpha 将直接下溢为 0.00.0。请从浮点分析角度推导:此时历史累加和 dd 被乘成 0,算法在数学上是正确的还是会引发精度灾难?

思考题 3:算子融合前沿 —— FlashAttention-3 的 WGMMA 与 TMA 协同

在最新的 Hopper 架构中,FlashAttention-3 引入了硬件级异步拷贝 TMA 和 Warpgroup GEMM(WGMMA)。请分析:当 TMA 硬件直接把张量从全局显存搬入共享内存时,Online Softmax 的规约与修正逻辑应该由哪个 Warp 组负责?如何排布软流水线(Software Pipelining)以实现乘加计算与 Softmax 缩放的完美重叠(Overlap)?

9. 📚 参考资料与核心源码/经典论文指引

本讲内容与推导过程严格对照并依据本地知识库 AI_BOOK 中的权威一手文献与源码:
  1. 并行规约权威奠基专著:
    • Mark Harris: Optimizing Parallel Reduction in CUDA, NVIDIA Developer Technology, 2007.
    • 本地核心代码解析与压测:参见 10_reduction.md 与 3.1-CUDA Reduce算子优化.md。
  2. Online Softmax 奠基论文与工程实现:
    • Maxim Milakov, Natalia Gimelshein: Online normalizer calculation for softmax, NVIDIA Corporation, 2018 (arXiv:1805.02867).
    • 本地工程级代码实现与递推详解:参见 5.2-CUDA Online Softmax实现.md 与 LeetCUDA/kernels/interview/base.cuh。
  3. FlashAttention 核心前沿论文:
    • Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré: FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, NeurIPS 2022.
    • 本地 FlashAttention 原理解析:参见 6.1-FlashAttention V1详解.md。

附录:Appendix A — 大厂硬核高频面试题与白板推导(Interview Drill)

题目 1:请在白板上推导 Online Softmax 的单元素增量递推公式与两块合并公式,并证明其满足结合律。

【面试官考察维度】
  1. 是否真正理解 Safe Softmax 的数值稳定性物理直觉;
  2. 能否独立推导 Online 修正因子 emold−mnewe^{m_{\text{old}} - m_{\text{new}}};
  3. 能否运用代数证明两块状态合并满足结合律,从而论证 GPU 并行规约的正确性。
【白板标准解答与推导路径】
  • 第一步:写出单元素增量递推公式 设已见前 kk 个元素最大值为 mkm_k,分母为 dk=∑i=1kexi−mkd_k = \sum_{i=1}^k e^{x_i - m_k}。 新加入元素 xk+1x_{k+1} 时:
mk+1=max⁡(mk,xk+1)m_{k+1} = \max(m_k, x_{k+1}) dk+1=∑i=1k+1exi−mk+1=(∑i=1kexi−mk)emk−mk+1+exk+1−mk+1=dk⋅emk−mk+1+exk+1−mk+1d_{k+1} = \sum_{i=1}^{k+1} e^{x_i - m_{k+1}} = \left(\sum_{i=1}^k e^{x_i - m_k}\right) e^{m_k - m_{k+1}} + e^{x_{k+1} - m_{k+1}} = d_k \cdot e^{m_k - m_{k+1}} + e^{x_{k+1} - m_{k+1}}
  • 第二步:写出两块独立状态合并公式 设两块数据状态分别为 (m1,d1)(m_1, d_1) 和 (m2,d2)(m_2, d_2):
mmerged=max⁡(m1,m2)m_{\text{merged}} = \max(m_1, m_2) dmerged=d1⋅em1−mmerged+d2⋅em2−mmergedd_{\text{merged}} = d_1 \cdot e^{m_1 - m_{\text{merged}}} + d_2 \cdot e^{m_2 - m_{\text{merged}}}
  • 第三步:证明结合律 [(A⊕B)⊕C=A⊕(B⊕C)][ (A \oplus B) \oplus C = A \oplus (B \oplus C) ] 定义状态合并算子 ⊕\oplus: (m1,d1)⊕(m2,d2)=(m12,d12)(m_1, d_1) \oplus (m_2, d_2) = (m_{12}, d_{12})。 易知 m(12)3=max⁡(max⁡(m1,m2),m3)=max⁡(m1,m2,m3)=Mm_{(12)3} = \max(\max(m_1, m_2), m_3) = \max(m_1, m_2, m_3) = M 显然满足结合律。 再考察分母:
d(12)3=d12⋅em12−M+d3⋅em3−M=(d1em1−m12+d2em2−m12)em12−M+d3em3−Md_{(12)3} = d_{12} \cdot e^{m_{12} - M} + d_3 \cdot e^{m_3 - M} = \left( d_1 e^{m_1 - m_{12}} + d_2 e^{m_2 - m_{12}} \right) e^{m_{12} - M} + d_3 e^{m_3 - M} 指数展开相乘: d(12)3=d1em1−M+d2em2−M+d3em3−Md_{(12)3} = d_1 e^{m_1 - M} + d_2 e^{m_2 - M} + d_3 e^{m_3 - M} 同理计算 d1(23)d_{1(23)},展开后完全一致。结合律获证! 这意味着无论 GPU 的线程树如何分叉折叠,最终结果严格恒等。

题目 2:为什么 __shfl_down_sync 可以在 Warp 内 5 次迭代完成 32 线程规约?画出数据流动拓扑图并解释掩码 0xffffffff 的含义。

【面试官考察维度】
考查对 GPU 底层 SIMT 指令、Lane 概念以及 Warp 级并行硬件原语的微架构理解。
【白板推导路径】
  1. 解释掩码 0xffffffff:
    • 掩码是一个 32-bit 无符号整数,每一位对应 Warp 中的一个 Lane(线程 0~31);
    • 0xffffffff(二进制全 1)表示当前 Warp 内的全部 32 个线程都必须参与此条同步 Shuffle 指令;如果某位为 0,代表该线程不参与交换。
  2. 推导 5 步折叠拓扑:
    • 32 是 2 的 5 次方( 25=322^5 = 32 );
    • 迭代 1(offset = 16):线程 0∼150 \sim 15 分别读取线程 16∼3116 \sim 31 的寄存器并累加,此时前 16 个线程保存了 16 对和;
    • 迭代 2(offset = 8):线程 0∼70 \sim 7 分别读取线程 8∼158 \sim 15 的数据并累加;
    • 迭代 3(offset = 4):线程 0∼30 \sim 3 累加;
    • 迭代 4(offset = 2):线程 0∼10 \sim 1 累加;
    • 迭代 5(offset = 1):线程 0 读取线程 1 的数据累加。 此时线程 0 的寄存器内保存了整个 Warp 32 线程的全部和。
  3. 硬件优势:数据全程在 SM 的通用寄存器物理交叉开关(Register Crossbar)上流动,无需访存指令,零延迟,无需显式 __syncthreads()。

题目 3:在行宽 N=4096N = 4096 的矩阵 Softmax 中,如何设计 Block 与 Thread 的映射?为什么 1 个 Block 处理 1 行比 1 个 Thread 处理 1 行好?

【面试官考察维度】
考查将数学算法映射到 GPU 网格网格架构时的系统级权衡能力(访存合并度 vs 并行粒度)。
【白板推导路径】
  • 方案 A(1 个 Thread 处理 1 行):
    • 优点:线程内部天然串行,不需要线程间规约和同步;
    • 致命缺陷:矩阵在内存中是行优先排布的,相邻线程(处理相邻行)在访问第 cc 列时,物理地址相差整整一整行( N×4N \times 4 字节)!导致全局内存读写完全是非合并访问(Strided Access),有效带宽暴跌 90% 以上;
  • 方案 B(1 个 Block 处理 1 行,黄金方案):
    • 配置 dim3 block(256),网格 dim3 grid(M);
    • Block 内连续的 256 个线程(变化最快的是 threadIdx.x)同时读取同一行的连续列;
    • 天然完美触发全局内存合并访问(Coalesced Access),打满 HBM 物理总线;
    • 行内 4096 个元素由 256 个线程分摊(每人处理 16 个),通过片上 Warp Shuffle 极速完成规约;
    • 结论:方案 B 胜出,性能高出方案 A 一个数量级以上。

题目 4:FlashAttention 是如何将 Online Softmax 应用到分块注意力计算中的?请写出输出张量 OO 的动态缩放更新公式并解释为什么不需要保存注意力权重矩阵 SS。

【面试官考察维度】
大模型 Infra 面试王牌大题,考察从基础算子到前沿系统的知识贯通能力。
【白板推导路径】
  1. 分块机制:将长序列按列切分,分批加载 Kj,VjK_j, V_j 进 SRAM;
  2. 推导局部与全局更新: 设第 jj 块计算出的局部得分为 Sj=QKjTS_j = Q K_j^T,局部最大值为 mjm_j,局部指数为 Pj=exp⁡(Sj−mj)P_j = \exp(S_j - m_j),局部总和为 lj=rowsum(Pj)l_j = \text{rowsum}(P_j); 维护全局状态:
mnew=max⁡(mold,mj)m_{\text{new}} = \max(m_{\text{old}}, m_j) dnew=dold⋅emold−mnew+lj⋅emj−mnewd_{\text{new}} = d_{\text{old}} \cdot e^{m_{\text{old}} - m_{\text{new}}} + l_j \cdot e^{m_j - m_{\text{new}}}
  1. 输出矩阵 OO 的流式更新: 在尚未做全局除法前,输出累加量 OO 维护的是 ∑PiVi\sum P_i V_i 的分子部分:
Onew=Oold⋅emold−mnew+PjVj⋅emj−mnewO_{\text{new}} = O_{\text{old}} \cdot e^{m_{\text{old}} - m_{\text{new}}} + P_j V_j \cdot e^{m_j - m_{\text{new}}}
  1. 彻底省去 SS 的物理原因: 因为 PjP_j 在片上 SRAM 算出来后,立即与 VjV_j 相乘并累加进了 OO 中!完成累加后,局部矩阵 SjS_j 和 PjP_j 的使命彻底终结,可以直接丢弃覆写,完全无需向全局显存写回哪怕一个元素,显存复杂度从 O(N2)O(N^2) 骤降至 O(N)O(N)。