从稳定 Softmax 到在线归约:TileLang 原理详解

本文面向已经完成 Task3 的 TileLang Softmax 实现、希望理解“为什么要这样写”的读者。

如果你还在填写 solution.py,请先阅读Task3 TileLang Softmax 零基础完成指南。本文重点解释算法、数据划分和 TileLang 原语之间的关系,而不是按步骤给出答案。


1. 从一组分数到一组权重

Softmax 的输入通常是一组分数。例如模型面对三个候选项时,可能输出:

[0, 1, 2]

分数 2 最大,说明模型最倾向第三项;但 [0, 1, 2] 不是一组概率,因为数值不在 0 到 1 的范围内,和也不等于 1。

Softmax 将它转换为:

[约 0.090, 约 0.245, 约 0.665]

现在每个值都为正,且总和为 1。它们可以被解释为权重或概率:第三项仍最受偏好,但前两项没有被完全丢弃。

它做了哪两件事

第一步是 expexp(x) 也写作 e^x,其中 e≈2.718。不需要手算,只需知道:

exp(0) = 1
x 越大,exp(x) 越大,且差距会扩大

第二步是归一化:将每个 exp 结果除以所有结果的总和。

分数:       [0,     1,     2]
计算 exp:   [1,  2.718, 7.389]
总和:       11.107
归一化结果: [0.090, 0.245, 0.665]

因此,对一行向量 x,定义为:

softmax(xᵢ) = exp(xᵢ) / Σⱼ exp(xⱼ)

xᵢ 表示第 i 个元素;Σⱼ 表示将这一行的全部元素相加。

注意分母包含整行所有列,所以同一行中任意一个输入值改变,所有输出都可能改变:

向量加法:C[i] 只依赖 A[i] 和 B[i]。
Softmax:B[i, j] 依赖 A[i, :] 整行。

前者是逐元素计算;后者同时包含逐元素计算和归约(最大值、求和)。


2. 数值稳定性:为什么要减去行最大值

直接计算 exp(x) 并不总是安全。对于测试中的一行:

[-1000, 0, 1000, 1]

因为 IEEE 754 的 float32 能表示的最大有限数大约只有 3.4 × 10^38 , exp(1000)float32 中会溢出为无穷大;之后的“无穷大除以无穷大”不再是有效的数值计算。

令:

m = max(x)

则可改写为:

softmax(xᵢ)
= exp(xᵢ - m) / Σⱼ exp(xⱼ - m)

这是严格相等的,因为分子和分母同时除以 exp(m)

以该例为例,m=1000

x - m = [-2000, -1000, 0, -999]

最大元素变为 0,其他元素都不大于 0,因此指数值在 (0, 1] 范围内,不会发生正向溢出。

这也解释了一个调试信号:若普通随机输入通过、极端值测试失败,优先检查“是否在指数前减去了每行最大值”。


3. 任务如何把二维矩阵映射成工作

任务的参数有两个分块维度:

BLOCK_N:一次处理多少行
BLOCK_M:一次处理每行多少列

对应的 kernel 起点是:

with T.Kernel(T.ceildiv(N, BLOCK_N), threads=256) as pid:

pid 选择一个行块,行起点是:

row_start = pid × BLOCK_N

例如 N=17、BLOCK_N=16 时:

pid = 0 → 行 0 至 15
pid = 1 → 行 16 和尾部位置

因此需要 T.ceildiv(N, BLOCK_N),而不是 N // BLOCK_N。否则最后一行没有任何工作块负责。

在每一个行块内部,列也按 BLOCK_M 分段:

col_start = m_blk_id × BLOCK_M

m_blk_id 遍历:

T.Serial(T.ceildiv(M, BLOCK_M))

例如 M=513、BLOCK_M=256,每一行分成三个 tile:0~255、256~511、512~尾部。

为什么 threads=256 不等于“只处理 256 个元素”

threads=256 是此 kernel 的线程配置;BLOCK_N × BLOCK_M 是一个数据 tile 的逻辑形状。测试默认是 16 × 256 = 4096 个逻辑元素。

T.Parallel(BLOCK_N, BLOCK_M) 表达的是可并行处理这 4096 个逻辑位置。编译器会把这些迭代映射到 256 个线程,而不是要求你手动决定每个物理线程处理哪一个元素。

因此需要分开记忆:

  • BLOCK_NBLOCK_M:数据如何分块;
  • threads:GPU 工作块的线程配置;
  • T.Parallel:一批逻辑迭代可以并行。

4. Fragment:当前工作块的临时 tile

典型实现会分配:

a_tile = T.alloc_fragment((BLOCK_N, BLOCK_M), T.float32)
exp_tile = T.alloc_fragment((BLOCK_N, BLOCK_M), T.float32)
tile_max = T.alloc_fragment((BLOCK_N,), T.float32)
tile_sum = T.alloc_fragment((BLOCK_N,), T.float32)
lse = T.alloc_fragment((BLOCK_N,), T.float32)

T.alloc_fragment 分配 kernel 内的临时缓冲。对当前任务,最重要的是它们的逻辑角色:

a_tile    :当前读取的输入行块 × 列块
exp_tile  :当前 tile 的稳定指数项
tile_max  :当前 tile 中每一行的最大值
tile_sum  :当前 tile 中每一行的指数和
lse       :扫描到当前位置时,整段前缀的归一化信息

因为 Softmax 的分母依赖整行,a_tile 只容纳一部分列还不够;我们需要用 lse 把多次列 tile 扫描的结果合并起来。


5. 尾块填充为什么使用负无穷

本任务为每个 tile 元素显式计算全局坐标,并在越界时写入负无穷:

row = pid * BLOCK_N + i
col = m_blk_id * BLOCK_M + j
if row < N:
    if col < M:
        a_tile[i, j] = A[row, col]
    else:
        a_tile[i, j] = -T.infinity(T.float32)
else:
    a_tile[i, j] = -T.infinity(T.float32)

这是 Softmax 特别合适的填充值:

max(x, -∞) = max(x)
exp(-∞) = 0
sum(x, 0) = sum(x)

因此,填充值不会改变有效列的最大值和分母。

输出阶段也使用同一组坐标,只在 row < Ncol < M 时写入 B[row, col]。这里不应将 -inf 写回输出;它只是在计算阶段表示无效输入位置。

为什么不依赖 annotate_safe_value 和整块 T.copy

当前课程环境中,T.annotate_safe_value({A: ...}) 可以通过编译,但整块 T.copy 的越界通道仍可能按 0 填充。0 会参与 Softmax:

max(x, 0) 可能改变行最大值
exp(0) = 1 会增大分母

这正是只有不整除尺寸失败、整除尺寸通过的典型原因。显式掩码虽然更长,但不依赖该版本的隐式填充路径,因此是本任务可靠的写法。


6. 为什么 dim=1 才是按行归约

a_tile 的形状为:

(BLOCK_N, BLOCK_M)

第一维是行,第二维是列。对 Softmax 而言,要保留行、消去列:

T.reduce_max(a_tile, tile_max, dim=1, clear=True)
T.reduce_sum(exp_tile, tile_sum, dim=1, clear=True)

输出缓冲的形状是 (BLOCK_N,),每一行对应一个值。

若误用 dim=0,就会跨不同行的元素求最大值或求和。这会把多行混在一起,破坏“每行输出和为 1”的定义。

clear=True 让每次归约从恰当的初始值开始:

  • 最大值归约从负无穷开始;
  • 求和归约从 0 开始。

这里的 tile_maxtile_sum 只服务于“当前列 tile”,每次循环都应重新计算,所以使用 clear=True


7. 为什么用 exp2,以及 log2_e 的作用

任务实现使用:

log2_e = 1.44269504
T.exp2(x * log2_e)

因为:

exp(x) = 2 ^ (x × log₂(e))

所以 T.exp2(x * log2_e) 与自然指数 exp(x) 数学等价。该写法便于与 T.log2 一起构成在线 LogSumExp 公式。

计算当前 tile 时:

exp_tile[i, j] = T.exp2(
    a_tile[i, j] * log2_e - tile_max[i] * log2_e
)

它对应:

exp(a_tile[i, j] - tile_max[i])

减去 tile_max[i] 使当前 tile 的指数计算稳定。随后:

T.reduce_sum(exp_tile, tile_sum, dim=1, clear=True)

得到当前 tile 每行的局部指数和。


8. 在线 LogSumExp 如何合并多段列 tile

M > BLOCK_M,整行被切成多个列 tile。最朴素的稳定 Softmax 需要至少三次扫描:

第 1 次:找整行最大值
第 2 次:计算指数并累加整行分母
第 3 次:计算并写出最终输出

在线 Softmax 将“找最大值”和“累加分母”合并为第一次扫描,只剩两次完整读取。

lse 表示什么

用以 2 为底的表示,令前面已经扫描过的所有元素的状态为:

L_old = log₂(Σ已扫描元素 2 ^ (x × log₂(e)))

对当前 tile,先求其行最大值 m,再求:

s = Σ当前 tile 2 ^ (x × log₂(e) - m × log₂(e))

那么合并后的状态是:

L_new = m × log₂(e)
        + log₂(2 ^ (L_old - m × log₂(e)) + s)

这正是代码中的更新:

lse[i] = tile_max[i] * log2_e + T.log2(
    T.exp2(lse[i] - tile_max[i] * log2_e) + tile_sum[i]
)

第一轮前 lse=-inf,于是:

2 ^ (-∞) = 0

公式自然退化为当前 tile 的结果,不需要专门为第一轮写分支。

为什么列 tile 的循环必须是 T.Serial

新的 lse 依赖前一个列 tile 留下的旧 lse。这是严格的数据依赖:

tile 0 的结果 → tile 1 的输入 → tile 2 的输入

所以列 tile 之间必须串行扫描:

for m_blk_id in T.Serial(...):

但每一轮 tile 内部的元素计算彼此独立,仍可以用 T.Parallel


9. 为什么最后仍需第二次扫描

第一次扫描结束后,lse[i] 已包含对应整行的最终归一化信息。但此前计算的 exp_tile 只保留在当前 tile 的临时 fragment 中,扫描到下一列 tile 时就会被覆盖。

若要在第一次扫描后直接写回,必须保存整行的中间指数值;对任意长行而言,这需要更多临时存储。

更实用的做法是重新读取每个 tile:

row = pid * BLOCK_N + i
col = m_blk_id * BLOCK_M + j
if row < N:
    if col < M:
        B[row, col] = T.exp2(A[row, col] * log2_e - lse[i])

因为:

2 ^ (xᵢ × log₂(e) - L)
= exp(xᵢ) / Σⱼ exp(xⱼ)

这就是最终 Softmax 值。

所以“两次扫描”不是重复劳动,而是在额外临时存储与再次读取输入之间作出的权衡。在线算法避免了第三次扫描,同时不必存下整行中间结果。


10. 测试案例怎样验证设计

测试输入覆盖:

(1, 1)      :最小矩阵
(3, 7)      :行、列都小于默认 tile
(16, 256)   :恰好填满一个 tile
(17, 513)   :行尾和列尾同时存在
(64, 4096)  :每行包含多个完整列 tile
[-1000, 0, 1000, 1]:数值稳定性

每个案例都针对一种可能的错误假设:

  • 只用整除:(17, 513) 会漏掉尾部;
  • 未填充负无穷:尾块可能污染最大值或指数和;
  • dim 写错:多行会被混合;
  • 忘记第二次扫描:无法得到最终归一化输出;
  • 直接算 exp(x):极端值可能溢出。

测试使用容差比较:

torch.testing.assert_close(out, torch.softmax(x, dim=1), atol=2e-3, rtol=2e-3)

因为浮点计算存在舍入误差。通过测试的标准是数值足够接近 PyTorch 的参考实现,而不是逐位完全相同。


11. benchmark 应当怎样解读

benchmark_softmax.py 输出两种不同性质的数据:

  • compile=...ms:为当前 NMBLOCK_NBLOCK_M 生成 kernel 的一次性成本;
  • TileLang(us):编译后的 kernel 多次运行的平均耗时。

二者不要混为一谈。部署或重复调用时,通常更关心运行耗时;首次运行或输入形状频繁变化时,编译成本也值得考虑。

Softmax 需要两次读取输入、至少一次写回输出,并执行两类归约。因此它比简单逐元素加法更复杂,性能会受到内存访问、归约通信、tile 形状和线程配置共同影响。

本任务不要求结果一定超过 PyTorch。一个有效 benchmark 结论应限定在当前 GPU、当前输入规模、当前参数和当前测量方法下。


12. 从 Softmax 走向 Attention

Softmax 是 Transformer Attention 的核心步骤之一。Attention 先得到一行分数,再对该行做 Softmax,最后用权重加权 Value。

FlashAttention 的关键思路正是把在线 Softmax 的状态与后续矩阵乘法累加结合起来:

读取一块 Key/Value

计算当前分数 tile

在线更新最大值、归一化状态和输出累加器

继续下一块,而不是物化整个分数矩阵

Task3 不要求实现这些融合优化,但已经练习了最关键的基础:分块扫描、稳定归约、尾块处理和在线状态更新。


13. 总结

这份 Softmax kernel 的关键逻辑是:

  1. T.ceildiv(N, BLOCK_N) 覆盖所有行块;
  2. T.Serial(T.ceildiv(M, BLOCK_M)) 扫描每一行的列 tile;
  3. 显式将尾块无效位置填充为负无穷,并只写回有效输出位置;
  4. dim=1 对每行做 max 与 sum 归约;
  5. 用在线 LogSumExp 在第一次扫描合并整行状态;
  6. 第二次扫描重新读取输入,并使用最终 lse 写出 Softmax。

继续学习时,可以查看 TileLang Instructionsreduce_max / reduce_sum API。建议先尝试独立回答:

  1. 为什么无效列填充 -inf 而不是 0?
  2. 为什么 T.Serial 只用于列 tile 之间,而 T.Parallel 用于 tile 内?
  3. 为什么在线 Softmax 仍需要第二次扫描?
  4. dim=1 改为 dim=0,计算的含义会怎样改变?