Task3 TileLang Softmax 零基础完成指南

这份指南写给已经了解 Task2 向量加法、但第一次接触 Softmax 的读者。

目标是完成 solution.pytl_softmax(),通过测试,并运行 benchmark。

完成任务后,请阅读从稳定 Softmax 到在线归约:TileLang 原理详解,理解两次扫描、在线 LogSumExp 和性能取舍。


0. 先理解 Softmax:把分数变成权重

模型经常会为多个候选项给出一组分数。例如:

三个候选项的分数:[1, 2, 3]

分数越大,表示模型越倾向该项;但它们不是概率,因为它们不必在 0 到 1 之间,三个数的和也不等于 1。

Softmax 的作用是把这组分数转换为一组可比较的权重:

输入分数:[-1, 0, 2]
Softmax :[约 0.042, 约 0.114, 约 0.844]

输出具有两个特征:

每个值都大于 0
一行中所有值之和约等于 1

因此,它常被解释为“模型分给每个候选项的概率”或“权重”。例如在 Transformer 的 Attention 中,Softmax 将注意力分数变成对不同位置的权重。

Softmax 如何做到这一点

它分两步:

  1. 对每个分数计算 exp
  2. 用每个结果除以全部结果的总和。

exp(x) 是一个数学函数,也可写成 e^x,其中 e 约为 2.718。你不需要手算它;只需记住:

exp(0) = 1
输入越大,exp 的结果越大,而且差距会被拉开

例如:

分数:       [0,     1,     2]
计算 exp:   [1,  2.718, 7.389]
总和:       11.107
除以总和后: [0.090, 0.245, 0.665]

这里 Σ 是“把所有项相加”的符号,xᵢ 表示第 i 个元素。因此最普通的 Softmax 公式是:

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

输入 A 是一个形状为 (N, M)float32 矩阵。本任务对每一行独立执行上述两步,输出 B 与输入形状相同:

A = [[1, 2, 3],
     [4, 5, 6]]
 
B[0, :] = softmax([1, 2, 3])
B[1, :] = softmax([4, 5, 6])

任务实际使用一个数值更安全、但数学结果相同的公式:

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

其中 max(x) 就是一行中最大的分数。为什么需要它,会在第 3 节解释。

你需要:

  1. 将行分组,使任意 N 都能覆盖;
  2. 按块扫描每一行,使任意 M 都能覆盖;
  3. 对每行求最大值和指数和;
  4. 用稳定方式得到每个输出;
  5. 正确处理行尾、列尾和极端数值;
  6. 通过 test_softmax.py 和运行 benchmark_softmax.py

1. 先确认运行环境

以下命令在课程提供的 JupyterLab 终端中运行:

cd /gollamago
source ./setup_env.sh
cd assignment/task3

确认当前 Python 能导入 TileLang 并识别 GPU:

python -c "import tilelang, torch; print('TileLang: OK'); print('GPU:', torch.cuda.is_available())"

预期包含:

TileLang: OK
GPU: True

若导入失败,先解决环境问题。source ./setup_env.sh 只对当前终端有效;新开 JupyterLab 终端后需要重新执行。


2. 认识要补全的骨架

打开 solution.py

import tilelang
import tilelang.language as T
 
 
@tilelang.jit
def tl_softmax(A, BLOCK_N: int, BLOCK_M: int):
    N, M = T.const("N, M")
    A: T.Tensor((N, M), T.float32)
    B = T.empty((N, M), T.float32)
 
    # 在这里实现 Softmax
    raise NotImplementedError("请根据步骤实现 Softmax")

这里:

  • A:输入矩阵;
  • B:输出矩阵;
  • N:行数;
  • M:每行的列数;
  • BLOCK_N:一个 GPU 工作块一次处理的行数;
  • BLOCK_M:每次从一行读取并处理的列数;
  • Ttilelang.language 的简称,其中有 T.KernelT.ParallelT.reduce_max 等 TileLang 工具。

测试固定使用:

BLOCK_N = 16
BLOCK_M = 256

因此,一个工作块最多处理 16 行;每一行每次处理 256 列。对于 M=4096,同一行要分 16 次读取。


3. 两个必须先懂的 Softmax 概念

Softmax 必须按行计算

给定:

A = [[1, 2, 3],
     [4, 5, 6]]

正确做法是分别处理两行:

B[0, :] = softmax([1, 2, 3])
B[1, :] = softmax([4, 5, 6])

不能把整个 N × M 矩阵的所有元素一起求一个最大值和一个总和。

在 TileLang 归约中:

T.reduce_max(tile, row_max, dim=1)
T.reduce_sum(tile, row_sum, dim=1)

dim=1 表示沿 tile 的列方向归约,得到“每一行一个结果”。

为什么必须减去最大值

数学上,下面两式相等:

exp(xᵢ) / Σⱼ exp(xⱼ)
= exp(xᵢ - max(x)) / Σⱼ exp(xⱼ - max(x))

但直接计算前者可能溢出。例如测试包含:

[-1000, 0, 1000, 1]

exp(1000) 超出 float32 的可表示范围;而减去最大值 1000 后:

[-2000, -1000, 0, -999]

所有指数都不会上溢。这个“先减最大值”的步骤不能省略。


4. 为什么要扫描一行两次

如果整行只有 BLOCK_M=256 列,读一次即可。但测试包含 M=513M=4096,不能把整行都放进一个 (BLOCK_N, BLOCK_M) 临时 tile。

所以把每行拆成多个列 tile:

M = 513,BLOCK_M = 256
 
第 0 个列 tile:列 0 至 255
第 1 个列 tile:列 256 至 511
第 2 个列 tile:列 512,以及尾部位置

Softmax 的分母依赖整行。第一个 tile 读完时,尚不知道后面是否有更大的值或更多指数项,因此不能立刻写出最终结果。

本任务采用两次扫描

第一次扫描每行所有列 tile:得到整行的稳定归一化信息 lse
第二次扫描每行所有列 tile:用 lse 计算并写出最终 B

第一次扫描采用在线更新,所以不需要额外保存整行的中间指数结果。


5. 第一步:启动覆盖全部行的 kernel

B = T.empty(...) 后添加:

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

这里:

  • T.ceildiv(N, BLOCK_N) 保证最后不足 BLOCK_N 行的尾块也能被处理;
  • pid 是当前行块的编号;
  • 当前块第一行的下标是 pid * BLOCK_N
  • threads=256 是这份练习使用的线程配置。它不等于“每个元素必有一个固定线程”;TileLang 会映射后续 T.Parallel 的工作。

log2_e 约等于 log₂(e)。后续使用 T.exp2 时:

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

这与自然指数 exp(x) 数学等价。


6. 第二步:分配当前工作块需要的临时空间

with T.Kernel(...) 内添加:

    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.fill(lse, -T.infinity(T.float32))

T.alloc_fragment 分配当前 GPU 工作块使用的临时数据。这里不需要记忆“fragment 位于哪一级硬件存储”,只需理解形状和用途:

a_tile    :当前读入的 BLOCK_N 行 × BLOCK_M 列输入
exp_tile  :当前 tile 的指数值
tile_max  :当前 tile 每行的最大值
tile_sum  :当前 tile 每行的指数和
lse       :已经扫描过的所有列 tile 的每行归一化信息

lse 初始为负无穷,表示“当前行还没有读到任何元素”。这让第一次在线更新可以与后续更新使用同一个公式。


7. 第三步:第一次扫描,逐块更新 lse

T.fill(...) 后添加:

    for m_blk_id in T.Serial(T.ceildiv(M, BLOCK_M)):
        for i, j in T.Parallel(BLOCK_N, BLOCK_M):
            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)
 
        T.reduce_max(a_tile, tile_max, dim=1, clear=True)
 
        for i, j in T.Parallel(BLOCK_N, BLOCK_M):
            exp_tile[i, j] = T.exp2(
                a_tile[i, j] * log2_e - tile_max[i] * log2_e
            )
 
        T.reduce_sum(exp_tile, tile_sum, dim=1, clear=True)
 
        for i in T.Parallel(BLOCK_N):
            lse[i] = tile_max[i] * log2_e + T.log2(
                T.exp2(lse[i] - tile_max[i] * log2_e) + tile_sum[i]
            )

这一段看起来长,但可以按四件事理解。

3.1 T.Serial:按列 tile 的顺序扫描

for m_blk_id in T.Serial(T.ceildiv(M, BLOCK_M)):

T.ceildiv(M, BLOCK_M) 计算一行需要多少个列 tile。这里必须是 T.Serial,因为本轮得到的 lse 会成为下一轮的输入;这些轮次不能互相独立并行。

3.2 显式判断:读取并安全填充尾块

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)

这段代码位于 T.Parallel(BLOCK_N, BLOCK_M) 内。它先计算当前 tile 元素在原矩阵中的行、列坐标;只有坐标有效时才读取 A[row, col],否则向 a_tile 写入负无穷。

对于 M=513 的最后一块,大多数位置已经超过真实列数。越界位置会被显式填充为负无穷:

max(真实值, -inf) 不改变最大值
exp(-inf) = 0,不改变指数和

因此尾块可以参与后续同一套计算,又不会影响真实结果。

3.3 reduce_maxreduce_sum:每行各得到一个数

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

计算当前 tile 每一行的最大值;tile_max 的形状为 (BLOCK_N,),所以每行对应一个最大值。

随后计算:

exp(a_tile - 当前 tile 的行最大值)

并通过:

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

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

3.4 lse:把当前 tile 合并到此前扫描的结果

最后一段公式把“此前所有列 tile 的结果”和“当前 tile 的结果”稳定地合并为新的 lse。你现在不需要推导它;只需遵守:

  • lse 初值是 -T.infinity(T.float32)
  • 更新公式逐字保留;
  • 它必须位于 T.Serial 循环内;
  • 它按行执行,所以循环是 T.Parallel(BLOCK_N)

完成第一个循环后,lse[i] 已经包含对应整行的归一化信息。


8. 第四步:第二次扫描并写出最终结果

在第一次扫描之后、离开 with 之前添加:

    for m_blk_id in T.Serial(T.ceildiv(M, BLOCK_M)):
        for i, j in T.Parallel(BLOCK_N, BLOCK_M):
            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])

现在 lse 已经知道整行分母,因此每个元素都能计算:

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

这里再次读取 A[row, col],而不是保留上一轮的 a_tile,因为第一次扫描结束后 tile 临时空间已经被后续列块覆盖。写入 B 前的两层判断确保最后一个不完整行块或列块只会访问真实存在的位置。


9. 拼成完整核心实现

将原来的 raise NotImplementedError(...) 替换为下面这段;保持它位于 B = T.empty(...)return B 之间:

log2_e = 1.44269504
num_row_blocks = T.ceildiv(N, BLOCK_N)
 
with T.Kernel(num_row_blocks, threads=256) as pid:
    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.fill(lse, -T.infinity(T.float32))
 
    for m_blk_id in T.Serial(T.ceildiv(M, BLOCK_M)):
        for i, j in T.Parallel(BLOCK_N, BLOCK_M):
            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)
 
        T.reduce_max(a_tile, tile_max, dim=1, clear=True)
 
        for i, j in T.Parallel(BLOCK_N, BLOCK_M):
            exp_tile[i, j] = T.exp2(
                a_tile[i, j] * log2_e - tile_max[i] * log2_e
            )
 
        T.reduce_sum(exp_tile, tile_sum, dim=1, clear=True)
 
        for i in T.Parallel(BLOCK_N):
            lse[i] = tile_max[i] * log2_e + T.log2(
                T.exp2(lse[i] - tile_max[i] * log2_e) + tile_sum[i]
            )
 
    for m_blk_id in T.Serial(T.ceildiv(M, BLOCK_M)):
        for i, j in T.Parallel(BLOCK_N, BLOCK_M):
            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])
 
return B

不要忘记删除原来的 raise NotImplementedError(...)。如果你的文件已有部分尝试,保留已有正确的导入和张量声明,按本指南核对核心结构即可。


10. 运行测试

assignment/task3 目录执行:

python -m pytest -q test_softmax.py

预期结果类似:

6 passed

测试覆盖:

(N, M) = (1, 1)、(3, 7)、(16, 256)、(17, 513)、(64, 4096)
以及 [-1000, 0, 1000, 1] 的极端数值行

它们分别检查单元素、行列尾块、刚好整除、多列 tile 和数值稳定性。

常见错误

NotImplementedError

确认已经删除骨架中的占位异常。

只有 (17, 513) 一类用例失败

检查两处 T.ceildiv:行块应使用 T.ceildiv(N, BLOCK_N),列块应使用 T.ceildiv(M, BLOCK_M);同时确认越界的行或列都显式写入了 -T.infinity(T.float32)

结果每行和不接近 1,或极端数值测试失败

确认指数计算前减的是每行最大值,且 lse 更新公式没有被省略。

不同的行结果互相影响

检查两次归约是否都使用 dim=1dim=0 会沿行方向归约,不符合按行 Softmax。

编译或缩进错误

确认两个 T.Serial 循环与 T.fill 都在 with T.Kernel(...) 内;第二次扫描必须在第一次扫描结束后执行。


11. 运行 benchmark

测试通过后运行:

python benchmark_softmax.py

输出会显示每种 (N, M) 的正确性、TileLang 平均耗时、PyTorch 平均耗时、最大误差和首次编译时间。确认每一行的 correct 均为 PASS,然后保存终端截图。

不要用首次编译时间判断 kernel 性能。编译只需发生一次;反复调用时更应关注 TileLang(us)。本任务不要求你的实现一定快于 PyTorch,正确性和稳定性优先。


12. 提交前检查

  • 使用 T.ceildiv(N, BLOCK_N) 覆盖行尾块;
  • 使用 T.Serial(T.ceildiv(M, BLOCK_M)) 扫描列 tile;
  • 无效行、列均显式写入 -T.infinity(T.float32)
  • T.reduce_max(..., dim=1)T.reduce_sum(..., dim=1) 均按行归约;
  • lse 初始化为负无穷并在第一次扫描中更新;
  • 第二次扫描使用 lse 计算并写回 B
  • 已删除 NotImplementedError
  • python -m pytest -q test_softmax.py 通过;
  • python benchmark_softmax.py 的每行均为 PASS
  • 已保存 benchmark 终端截图。

完成后,建议阅读从稳定 Softmax 到在线归约:TileLang 原理详解,再深入理解 lse 合并公式为何正确,以及为什么在线 Softmax 需要两次扫描而不是三次。