CUDA 算子工程:手写 FlashAttention v2 之路

第 6 章 Softmax 与 Online Softmax

作者 杨艺韬 · 4,230 字 · 发布于 · 更新于

Online softmax 不是 FlashAttention 发明的:它出自 Milakov & Gimelshein 的 Online normalizer calculation for softmax(NVIDIA,2018,arXiv:1805.02867)。 FlashAttention 做的事,是把这套一遍式的归一化和 attention 的分块矩阵乘缝到了一起。

6.1 Softmax 的标准定义

给定一个向量 x=(x1,x2,…,xN)\mathbf{x} = (x_1, x_2, \ldots, x_N),softmax 定义为:

softmax(xi)=exi∑j=1Nexj\text{softmax}(x_i) = \frac{e^{x_i}}{\sum_{j=1}^{N} e^{x_j}}

直接照定义计算需要两遍:

// 第一遍: 算分母 (求和)
float sum = 0;
for (int i = 0; i < N; ++i) sum += expf(x[i]);

// 第二遍: 算每个输出
for (int i = 0; i < N; ++i) y[i] = expf(x[i]) / sum;

但这段代码有一个致命问题:数值溢出。

6.2 数值稳定的 Safe Softmax

如果 xix_i 中有比较大的数(比如 100),那 e100e^{100} 大约是 2.7×10432.7 \times 10^{43}。FP32 的上限大约是 3.4×10383.4 \times 10^{38}——直接溢出成 +inf。如果 xix_i 全是 1000,e1000e^{1000} 直接溢出,所有项都变 inf,分子分母 inf/inf 输出 NaN。

LLM 训练和推理中,attention 的 logits(QKT/dQK^T/\sqrt{d})没有取值上界的保证。若在 FP16 里算 exp,风险更高:FP16 上限只有 65504,x>ln⁡65504≈11.09x > \ln 65504 \approx 11.09 就溢出。BF16 的指数位与 FP32 同为 8 位,表示范围和 FP32 相当,但尾数只有 7 位,所以实践中 softmax 一般在 FP32 里算。

解决方法是经典的 safe softmax 技巧——同时减去最大值:

softmax(xi)=exi−m∑j=1Nexj−m,m=max⁡jxj\text{softmax}(x_i) = \frac{e^{x_i - m}}{\sum_{j=1}^{N} e^{x_j - m}}, \quad m = \max_j x_j

数学上完全等价(分子分母都乘以 e−me^{-m} 抵消),但所有 exi−m≤1e^{x_i - m} \le 1,绝不溢出。

但 safe softmax 现在需要 3 遍:

// Pass 1: 找 max
float m = -INFINITY;
for (int i = 0; i < N; ++i) m = fmaxf(m, x[i]);

// Pass 2: 算分母
float sum = 0;
for (int i = 0; i < N; ++i) sum += expf(x[i] - m);

// Pass 3: 算每个输出
for (int i = 0; i < N; ++i) y[i] = expf(x[i] - m) / sum;

3 遍意味着 x 数组要从 HBM 读 3 次(一行大到缓存装不下时)。这是大问题——softmax 的算术强度本身就低(每个元素几次浮点操作 vs 4 字节读),读 3 遍再加写回 yy,每个元素 4 次访存,而下限是 1 读 1 写共 2 次。

6.3 Online Softmax:1 遍数学

Online Softmax 的核心思想:在遍历数据的过程中同时维护当前的 max 和 sum,不需要预先知道全局 max。

6.3.1 数学推导

定义两个状态量:

  • mnm_n = 截至第 nn 个元素时的最大值
  • ℓn=∑j=1nexj−mn\ell_n = \sum_{j=1}^{n} e^{x_j - m_n} = 截至第 nn 个元素的"调整后 exp 和"

初始状态 m0=−∞m_0 = -\infty,ℓ0=0\ell_0 = 0。

当看到新元素 xn+1x_{n+1} 时,更新规则:

mn+1=max⁡(mn,xn+1)ℓn+1=ℓn⋅emn−mn+1+exn+1−mn+1\begin{aligned} m_{n+1} &= \max(m_n, x_{n+1}) \\ \ell_{n+1} &= \ell_n \cdot e^{m_n - m_{n+1}} + e^{x_{n+1} - m_{n+1}} \end{aligned}

关键的"修正项"是 emn−mn+1e^{m_n - m_{n+1}}:当 max 被更新(mn+1>mnm_{n+1} > m_n 时这个值小于 1),把之前累积的 sum 也"按比例缩小",保证它仍然是相对于新 max 的 sum。

把这个递推走完,ℓN\ell_N 就是 safe softmax 的分母(原论文把它记作 dd,本专栏沿用 FlashAttention 论文和第 14–16 章的记号 ℓ\ell)。然后再走一遍计算输出(这一遍可以和下游计算 fuse 在一起)。

如果是 attention 这种"sum 之后还要点积"的场景,可以做到真正的 1 pass——这是 FA 的核心。

6.3.2 验证一个简单例子

考虑 x=(1,5,3)\mathbf{x} = (1, 5, 3)。

朴素方式:

  • m=5m = 5
  • ∑=e1−5+e5−5+e3−5=e−4+1+e−2≈0.0183+1+0.1353=1.1536\sum = e^{1-5} + e^{5-5} + e^{3-5} = e^{-4} + 1 + e^{-2} \approx 0.0183 + 1 + 0.1353 = 1.1536
  • y=(0.0183/1.1536,1/1.1536,0.1353/1.1536)≈(0.0159,0.8668,0.1173)y = (0.0183/1.1536, 1/1.1536, 0.1353/1.1536) \approx (0.0159, 0.8668, 0.1173)

Online:

步骤 xnx_n mnm_n ℓn\ell_n
初始 - −∞-\infty 0
n=1 1 1 0⋅e−∞−1+e0=10 \cdot e^{-\infty - 1} + e^{0} = 1
n=2 5 5 1⋅e1−5+e0=e−4+1≈1.01831 \cdot e^{1-5} + e^{0} = e^{-4} + 1 \approx 1.0183
n=3 3 5 1.0183⋅e5−5+e3−5=1.0183+e−2≈1.15361.0183 \cdot e^{5-5} + e^{3-5} = 1.0183 + e^{-2} \approx 1.1536

ℓ3=1.1536\ell_3 = 1.1536,和朴素方式得到的 ∑\sum 完全一致。

6.3.3 推导的几何理解

为什么 ℓn⋅emn−mn+1\ell_n \cdot e^{m_n - m_{n+1}} 是正确的修正?

设 mn+1=max⁡(mn,xn+1)m_{n+1} = \max(m_n, x_{n+1})。我们希望 ℓn+1=∑j=1n+1exj−mn+1\ell_{n+1} = \sum_{j=1}^{n+1} e^{x_j - m_{n+1}}。

把它拆开:

∑j=1n+1exj−mn+1=exn+1−mn+1+∑j=1nexj−mn+1\sum_{j=1}^{n+1} e^{x_j - m_{n+1}} = e^{x_{n+1} - m_{n+1}} + \sum_{j=1}^{n} e^{x_j - m_{n+1}}

第二项可以改写:

∑j=1nexj−mn+1=∑j=1nexj−mn⋅emn−mn+1=ℓn⋅emn−mn+1\sum_{j=1}^{n} e^{x_j - m_{n+1}} = \sum_{j=1}^{n} e^{x_j - m_n} \cdot e^{m_n - m_{n+1}} = \ell_n \cdot e^{m_n - m_{n+1}}

所以:

ℓn+1=ℓn⋅emn−mn+1+exn+1−mn+1\ell_{n+1} = \ell_n \cdot e^{m_n - m_{n+1}} + e^{x_{n+1} - m_{n+1}}

干净的代数变换。这就是为什么 online softmax 数学上严格等价于 safe softmax——它只是把"先扫一遍找 max 再扫一遍累加"改写成"边扫边更新"。

6.4 Online Softmax 的 GPU 实现

把 online softmax 实现成 GPU kernel。基本模板:

__global__ void softmax_online(const float* in, float* out, int N) {
    int tid = threadIdx.x;
    extern __shared__ float smem[];

    // 约定: 一个 block 处理一行, blockDim.x == 256 (8 个 warp),
    //       启动时动态 SMEM 至少 16 个 float (64 字节)

    // ====== Phase 1: Online sweep (1 pass) ======
    float m = -INFINITY;  // 当前 max
    float l = 0.0f;       // 当前调整后 sum

    for (int i = tid; i < N; i += blockDim.x) {
        float x = in[i];
        float m_new = fmaxf(m, x);
        if (m_new == -INFINITY) continue;  // 目前为止全是 -INF: 否则 expf(-INF - (-INF)) = NaN
        l = l * expf(m - m_new) + expf(x - m_new);
        m = m_new;
    }
    // 此时每个 thread 持有局部的 (m, l)

    // ====== Phase 2: Block-level reduce ======
    // Warp 内合并 (m, l)
    auto warp_combine = [](float& m, float& l, float m2, float l2) {
        float m_new = fmaxf(m, m2);
        if (m_new == -INFINITY) return;  // 两边都是空状态 (-INF, 0), 保持不动
        l = l * expf(m - m_new) + l2 * expf(m2 - m_new);
        m = m_new;
    };

    for (int offset = 16; offset > 0; offset >>= 1) {
        float m2 = __shfl_xor_sync(0xFFFFFFFF, m, offset);
        float l2 = __shfl_xor_sync(0xFFFFFFFF, l, offset);
        warp_combine(m, l, m2, l2);
    }
    // 此时 warp 内所有 lane 都持有 warp 的 (m, l)

    // 写入 SMEM, block 内再 reduce 一次
    int warp_id = tid / 32;
    int lane_id = tid % 32;
    if (lane_id == 0) {
        smem[warp_id * 2 + 0] = m;
        smem[warp_id * 2 + 1] = l;
    }
    __syncthreads();

    // Warp 0 收集 8 个 warp 的 (m, l)
    if (warp_id == 0) {
        m = (lane_id < 8) ? smem[lane_id * 2 + 0] : -INFINITY;
        l = (lane_id < 8) ? smem[lane_id * 2 + 1] : 0.0f;
        // 注意 mask 必须是 0xFFFFFFFF 而不是 0xFF:这里的守卫是
        // `warp_id == 0`,warp 0 的 32 个 lane **全都**在执行这条指令,
        // 而 `*_sync` 系列要求"执行它的线程必须被 mask 点名",
        // 写 0xFF 会让 lane 8..31 落在 mask 之外,行为未定义。
        // lane 8..31 持的是 (-INF, 0) 这个单位元(combine 里的
        // -INF 判断保证它真是单位元),跟着算也不会污染结果。
        for (int offset = 4; offset > 0; offset >>= 1) {
            float m2 = __shfl_xor_sync(0xFFFFFFFF, m, offset);
            float l2 = __shfl_xor_sync(0xFFFFFFFF, l, offset);
            warp_combine(m, l, m2, l2);
        }
        if (lane_id == 0) {
            smem[0] = m;
            smem[1] = l;
        }
    }
    __syncthreads();

    float final_m = smem[0];
    float final_l = smem[1];

    // ====== Phase 3: Normalize ======
    for (int i = tid; i < N; i += blockDim.x) {
        out[i] = expf(in[i] - final_m) / final_l;
    }
}

关键点:

  1. Phase 1 每个 thread 维护自己的 (m, l) 状态。
  2. Phase 2 用 warp_combine 合并不同 thread 的 (m, l)——这是 online softmax 的核心组合规则。
  3. Phase 3 用最终的全局 (m, l) 做归一化。
  4. 两处 m_new == -INFINITY 判断不能省:线程分到的元素全是 -INF(被 mask 掉),或者 N 小于 blockDim.x、有线程一个元素都没分到,它的状态就停在 (−∞,0)(-\infty, 0);两个这样的状态相遇时 expf(-INF - (-INF)) 是 expf(NaN),NaN 会一路传进最终结果。用 numpy 按本 kernel 的线程划分和蝶形归约逐步模拟,去掉这两行时 N=100(< 256)或前 300 个元素为 -INF 的行输出全是 NaN,加上后与 torch.softmax 的最大误差在 10−810^{-8} 量级。FA2 用的是同一类防护:csrc/flash_attn/src/softmax.h 在 row max 为 -INFINITY 时改用 0 去算缩放因子。

本 kernel 用 nvcc 13.4 以 -arch=sm_90a 编译通过,21 个寄存器、无 spill(nvcc 13.4 编译所得,不同 CUDA 版本可能不同)。

这里 Phase 1 和 Phase 3 都需要读 in[i]——所以严格说还是 2 pass。但 Phase 1 不再需要 max、sum 两次单独的 pass——它把 max 和 sum 合并到一遍里了。这对 SRAM 受限的 GEMM/attention 场景非常关键,因为可以让 in 数据只在 SMEM/寄存器里待一次。

6.4.1 Combine 函数的对称性

注意 warp_combine 函数:

auto warp_combine = [](float& m, float& l, float m2, float l2) {
    float m_new = fmaxf(m, m2);
    if (m_new == -INFINITY) return;  // 两边都是空状态 (-INF, 0), 保持不动
    l = l * expf(m - m_new) + l2 * expf(m2 - m_new);
    m = m_new;
};

这个函数是对称、可结合的:

  • combine((m1,l1),(m2,l2))=combine((m2,l2),(m1,l1))\text{combine}((m_1, l_1), (m_2, l_2)) = \text{combine}((m_2, l_2), (m_1, l_1))
  • combine(combine(a,b),c)=combine(a,combine(b,c))\text{combine}(\text{combine}(a, b), c) = \text{combine}(a, \text{combine}(b, c))

证明只需一步:把状态 (m,ℓ)(m, \ell) 看作它代表的量 ℓ⋅em\ell \cdot e^{m},combine 就是把两者代表的量相加,再改用新的 max 作基准表示。无论先合并哪两个,(a,b,c)(a, b, c) 合并的结果都是 (M, ℓ1em1−M+ℓ2em2−M+ℓ3em3−M)\bigl(M,\ \ell_1 e^{m_1 - M} + \ell_2 e^{m_2 - M} + \ell_3 e^{m_3 - M}\bigr),其中 M=max⁡(m1,m2,m3)M = \max(m_1, m_2, m_3);交换律从公式的对称性直接可见。原论文(式 4)给出了同一个算子,并声明它可结合、可交换,证明从略。注意这是实数意义上的等式:浮点下不同的合并顺序会带来末位舍入差异,结果不保证逐位相同。

结合律是并行归约的必要条件——没有它,改变合并顺序就可能改变结果;交换律则保证上面 __shfl_xor_sync 蝶形归约里配对的两个 lane 以相反顺序合并,也得到同一个值。再加上单位元 (−∞,0)(-\infty, 0)(前提是上面那个 -INFINITY 判断),combine 构成一个幺半群(monoid),这正是 online softmax 能套进任何并行归约框架的原因。

6.5 LSE 形式:Log-Sum-Exp

attention 中常用 softmax 的 log 版本(log-softmax),定义为:

log⁡softmax(xi)=xi−log⁡∑jexj=xi−LSE(x)\log \text{softmax}(x_i) = x_i - \log\sum_{j} e^{x_j} = x_i - \text{LSE}(\mathbf{x})

其中 LSE(x)=log⁡∑jexj\text{LSE}(\mathbf{x}) = \log \sum_j e^{x_j} 是 log-sum-exp。

数值稳定的 LSE:

LSE(x)=m+log⁡∑jexj−m,m=max⁡jxj\text{LSE}(\mathbf{x}) = m + \log \sum_j e^{x_j - m}, \quad m = \max_j x_j

在 online 形式中,存 (mn,ℓn)(m_n, \ell_n) 时实际上已经在算 LSE:

LSE=m+log⁡ℓ\text{LSE} = m + \log \ell

FA2 反向传播时需要存 LSE(一个标量/行)来支持梯度计算。FA2 前向除了输出 O,只额外存下 LSE,不存 N×NN\times N 的 attention 矩阵——这是它额外显存从 O(N2)O(N^2) 降到 O(N)O(N) 的关键。

6.6 Online Softmax 应用到 FA:1 Pass 真的成立

FA 的核心创新就是把 attention 改写成 online 形式:

朴素 attention:

O=softmax(QKT)⋅VO = \text{softmax}(QK^T) \cdot V

需要 3 个完整的 pass:

  1. 算 S=QKTS = QK^T(HBM 写 N×N 矩阵)。
  2. 算 P=softmax(S)P = \text{softmax}(S) 行级(HBM 写 N×N 矩阵)。
  3. 算 O=PVO = PV(HBM 写 N×d 矩阵)。

中间结果 SS 和 PP 都是 O(N2)O(N^2) 大小:N=4096 时单个 head 就是 40962≈1.68×1074096^2 \approx 1.68\times10^7 个元素,FP16 下 32 MB;乘上 head 数和 batch,写一次读一次的 HBM 流量非常可观。

FA 用 online softmax 把它压成 1 pass:

# 伪代码: 单个 query 行 q, K/V 按 B 行一块
m, l = -inf, 0
O = zeros(d)  # output accumulator

for j in range(0, N, B):  # block-by-block
    K_block = K[j:j+B]
    V_block = V[j:j+B]
    S_block = q @ K_block.T              # B 个分数
    m_block = max(S_block)               # block 局部 max
    P_block = exp(S_block - m_block)     # 局部 exp
    l_block = sum(P_block)

    # Combine 到 (m, l)
    m_new = max(m, m_block)
    alpha = exp(m - m_new)
    beta = exp(m_block - m_new)

    # 重要: 之前累积的 O 也要按 alpha 缩放
    O = O * alpha + beta * (P_block @ V_block)
    l = l * alpha + beta * l_block
    m = m_new

O = O / l  # 最终归一化

注意几点:

  1. m, l, O 三个状态量同步演化。每来一个 K/V block,都更新这三个。
  2. O 的缩放因子 α=em−mnew\alpha = e^{m - m_{\text{new}}} 来自 online softmax 的修正——之前累积的 attention 输出也要按新 max 缩放。
  3. 完全不写中间 S 或 P 矩阵到 HBM——只在寄存器/SMEM 中流过。
  4. 这里沿用了 6.4 节 combine 的写法:块内先减局部 max,再用 β\beta 对齐到新 max。第 14、15 章的 FA2 写法更省一步:直接用 mnewm_{\text{new}} 算 P=eS−mnewP = e^{S - m_{\text{new}}},于是不再需要 β\beta,两者数学上等价(用 numpy 对单个 query 行分块复算,与 softmax(qK^T)V 的差在 10−1510^{-15} 量级)。

这就是 FA 的本质:用 online softmax 让 attention 变成可流式的算法。Q/K/V 可以按 tile 喂给 kernel,算完就丢,不需要存中间结果。

第 14-15 章会把这个伪代码落到具体的 CUDA kernel 上。

6.7 Softmax 的访存优化

回到 standalone softmax 的 kernel 实现。性能优化要点:

6.7.1 整行处理 vs 分块处理

LLM 中 softmax 是按行操作的(attention 的每一行独立 softmax)。两种 block 配置:

整行处理:一个 block 处理一行的所有 N 个元素。

适合: 行长适中 (一个 block 就能处理完, 最好整行放得进 SMEM 或寄存器)
优势: 不需要跨 block 通信

分块处理:多个 block 协作处理一行。

适合: N 极大 (比如长序列 attention 的中间矩阵)
代价: 需要两阶段 + atomic 或 cluster reduce

一行交给一个 block 能处理完时,整行处理更简单:不需要跨 block 通信,也不需要第二个 kernel 或 atomic。

6.7.2 Vectorized I/O

和 reduce 类似,softmax 的访存也应该 vectorized:

// in 为 const half*; 每线程一次读 4 个 fp16 = 8 字节 (i 须是 4 的倍数, 地址 8 字节对齐)
uint2 packed = *reinterpret_cast<const uint2*>(&in[i]);
half2 v0 = *reinterpret_cast<half2*>(&packed.x);
half2 v1 = *reinterpret_cast<half2*>(&packed.y);

别写成两次 half2 读:*reinterpret_cast<const half2*>(&in[i]) 和 &in[i + 2] 各读 4 字节,编译器并不会替你合并——nvcc 13.4 在 sm_90a 上编出的是两条 32 位的 LDG.E,而上面的 uint2 写法是一条 LDG.E.64(nvcc 13.4 编译所得,不同 CUDA 版本可能不同)。

或者一次读 16 字节,用 int4 装 8 个 fp16(这时每个线程负责下标 i * 8 起的 8 个元素,地址须 16 字节对齐,nvcc 13.4 编出一条 LDG.E.128):

int4 packed = *reinterpret_cast<const int4*>(&in[i * 8]);
half h[8];
memcpy(h, &packed, 16);

6.7.3 Fused Softmax + Mask

attention 中 softmax 之前通常有 mask(causal mask、padding mask)。fused 写法:

float x = in[i] + mask[i];  // mask 通常是 0 或 -INFINITY
m = fmaxf(m, x);

把 mask add 直接 fuse 到 softmax 的 max/sum 阶段,避免单独走一遍 mask kernel。用 -INFINITY 做 mask 时,online 更新里 6.4 节那个 m_new == -INFINITY 判断就是必需的,否则被整段 mask 掉的线程会产生 NaN。整行都被 mask 时,结果本来就没有定义:torch.softmax 对全 -inf 的行输出 NaN,加了判断的 kernel 也输出 NaN(0/00/0)。

6.7.4 三种写法的访存账

原论文按"每个元素访问几次内存"给三种写法记账(Milakov & Gimelshein 2018,第 2–4 节):

实现 每元素访存次数 构成
朴素 softmax(不减 max,不安全) 3 读 2 遍 + 写 1 遍
Safe softmax(3-pass) 4 读 3 遍 + 写 1 遍
Online softmax 3 读 2 遍(max 与 sum 合成一遍)+ 写 1 遍

Online 相对 safe 省下的访存是 4/3≈1.334/3 \approx 1.33 倍。论文在 Tesla V100(fp32,每批 4000 个向量)上实测:向量长度到 1000 左右时三种写法差不多(数据还在 L1/L2 里),之后被 DRAM 带宽卡住,online 相对 safe 很快达到约 1.3 倍,与 1.33 倍的访存比吻合。对 standalone softmax 来说,这差不多就是这项改写的上限——它本来就是带宽 bound 的算子,每元素至少要 1 读 1 写。论文里真正大的收益来自融合:Softmax + TopK 融合后每元素只剩 1 次访存,实测最高约 5 倍。FA 把 softmax 嵌进 attention 也是同一个道理,收益来自根本不写 SS、PP 这两个 O(N2)O(N^2) 中间矩阵,那是另一个量级的事。

6.8 这一章给我们的"内核数学"

Online softmax 是这个专栏第一次正式接触"用算法重写让 GPU 友好"的思路。它带给读者两个核心 insight:

  1. 数学等价不等于 GPU 等价。同一个 softmax 公式有 3 遍和 2 遍两种数学等价的算法(与下游融合后还能压到 1 遍),访存次数不同,GPU 上的性能也就不同:standalone 时约 1.3 倍,融合进下游后差距可以到数倍。算法重写和"低层 kernel 优化"是 GPU 性能工程的两个独立维度,online softmax 是前者的典范。

  2. Streaming(流式)算法在 GPU 上是黄金。能够用单 pass 维护少量状态量来累积结果的算法,天然适合 GPU——因为它把"中间结果"留在寄存器/SMEM,避免 HBM 往返。Online softmax、Welford 在线方差(下一章)、prefix sum——这些都是流式算法的代表。

第 7 章我们把 online 思路用到另一个 LLM 高频算子上:LayerNorm 与 RMSNorm。LayerNorm 需要算均值和方差——传统写法要么先算均值再算方差、多读一遍,要么用 E[X2]−(E[X])2E[X^2]-(E[X])^2 一遍算完却有数值稳定问题,Welford 算法(online softmax 论文自述,其灵感正来自 Welford 的数值稳定在线方差算法)可以一遍同时算出均值和方差。读完第 7 章,读者会发现"online 思维"在 LLM 算子里几乎无处不在。

本章动手练习:

  1. 实现一个 N=4096 的整行 softmax kernel,先用 3-pass,再改成 online,对比性能。
  2. 用 PyTorch 的 torch.nn.functional.softmax 跑一遍,用 Nsight Compute 看它实际落到了哪个 kernel 上——PyTorch v2.11 对最后一维做 softmax 时,短行(dim_size <= 2048 且整行不超过 8 KB,见 pytorch-v2.11.0/aten/src/ATen/native/cuda/SoftMax.cu:1097)走 pytorch-v2.11.0/aten/src/ATen/native/cuda/PersistentSoftmax.cuh 的 softmax_warp_forward(:68);更长的行(比如 N=4096)在 pytorch-v2.11.0/aten/src/ATen/native/cuda/SoftMax.cu 的五个前向 kernel 里选:cunn_SoftMaxForwardFast(:702)、cunn_SoftMaxForward(:734)、cunn_SoftMaxForwardReg(:770)、cunn_SoftMaxForwardGmem(:827)、cunn_SoftMaxForwardSmem(:888)。分派逻辑在同文件 :1108 起:use_fast_softmax 为真时按 dim_size % ILP 在 Gmem / Fast 之间选;否则先算 potential_register_count 决定能不能整行塞进寄存器(Reg),塞不下就看 can_use_smem(同时要求整行放得进 sharedMemPerBlock、输入输出都 16 字节对齐、dim_size 是 ILP 的整数倍)走 Smem,再不行才退回最通用的 cunn_SoftMaxForward。
  3. 思考:如果 softmax 的 N 极大(比如 N=1M),整行处理放不进单 block 的 SMEM,怎么用 online softmax + cluster reduce 解决?