从稳定 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。它们可以被解释为权重或概率:第三项仍最受偏好,但前两项没有被完全丢弃。
它做了哪两件事
第一步是 exp。exp(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_Mm_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_N、BLOCK_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 < N 且 col < 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_max 和 tile_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:为当前N、M、BLOCK_N、BLOCK_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 的关键逻辑是:
- 用
T.ceildiv(N, BLOCK_N)覆盖所有行块; - 用
T.Serial(T.ceildiv(M, BLOCK_M))扫描每一行的列 tile; - 显式将尾块无效位置填充为负无穷,并只写回有效输出位置;
- 用
dim=1对每行做 max 与 sum 归约; - 用在线 LogSumExp 在第一次扫描合并整行状态;
- 第二次扫描重新读取输入,并使用最终
lse写出 Softmax。
继续学习时,可以查看 TileLang Instructions 与 reduce_max / reduce_sum API。建议先尝试独立回答:
- 为什么无效列填充
-inf而不是 0? - 为什么
T.Serial只用于列 tile 之间,而T.Parallel用于 tile 内? - 为什么在线 Softmax 仍需要第二次扫描?
- 若
dim=1改为dim=0,计算的含义会怎样改变?