Task3 NineToothed GEMM:完成矩阵乘法任务

Task2 中,你已经用 arrangement() 将向量切块,再用 application() 完成加法。这次仍然补全这两个函数,计算对象变成了两个矩阵:C = A @ B

本题的输入和输出使用 float16,中间用 float32 累加。完成后要提交 ninetoothed_gemm.py、测试输出和基线性能截图。下面从一个可以手算的例子开始,逐步写出实现。

1. 从一个元素算到一个矩阵块

GEMM 是通用矩阵乘法的缩写。本题只计算 C = A @ B,没有额外的缩放或偏置。文件中的 lhs 对应 A,rhs 对应 B,output 对应 C。

假设 A 有 4 行、6 列,B 有 6 行、8 列,结果 C 就有 4 行、8 列:

A: (4, 6)    @    B: (6, 8)    →    C: (4, 8)
    M  K              K  N              M  N

M 决定输出有多少行,N 决定输出有多少列。K 是 A 每行与 B 每列对上的长度。一项输出要沿这个长度乘一次、加一次,直到 K 个位置全部用完。

例如 C[0, 2] 使用 A 的第 0 行和 B 的第 2 列。假设这两段数据是:

A[0, :] = [1, 2, 3, 4, 5, 6]
B[:, 2] = [1, 0, 1, 0, 1, 1]
 
C[0, 2] = 1×1 + 2×0 + 3×1 + 4×0 + 5×1 + 6×1 = 15

这里的 : 表示沿该维取全部元素;下标从 0 开始。

GPU 一次可以合作计算一小片输出。先在纸面上把每个矩阵切成 2 × 2 的块,称为 tile:

A 的块网格:2 行 × 3 列       B 的块网格:3 行 × 4 列
 
        K0    K1    K2                N0    N1    N2    N3
M0     A00   A01   A02        K0     B00   B01   B02   B03
M1     A10   A11   A12        K1     B10   B11   B12   B13
                             K2     B20   B21   B22   B23
 
C 的块网格:2 行 × 4 列
 
        N0    N1    N2    N3
M0     C00   C01   C02   C03
M1     C10   C11   C12   C13

图中的 A00 等名字都代表一个 2 × 2 矩阵。下面将 C 的第 0 行块、第 1 列块写成 C_tile[0, 1],它覆盖原矩阵的 C[0:2, 2:4]。这个记法将块下标与原始元素下标区分开。

为了算出这一整块,还要取 A 的前两行、B 的第 2、3 列。沿用刚才的数据,并补上另一行、另一列:

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

K 的长度为 6,每次取 2 个位置,因此分三轮。每一轮都是 (2, 2) @ (2, 2)

K 段取 A 的哪些列、B 的哪些行本轮乘法结果
第 0 段0:2[[1, 2], [0, 1]]
第 1 段2:4[[3, 7], [0, 1]]
第 2 段4:6[[11, 6], [1, 1]]

把三个结果逐元素相加:

C_tile[0, 1] = A00 @ B01 + A01 @ B11 + A02 @ B21
            = [[15, 15],
               [ 1,  3]]

左上角的 15,正是刚才单独算出的 C[0, 2]。分块改变了组织计算的方式,仍然包含原来全部的乘积与加法。

这些 2 × 2 小块用于手算。后面的 GPU 实现保留课程默认的 64 × 64 × 64 配置;ntl.dot 的硬件约束由该配置满足。

2. 确定每份 GPU 工作负责什么

本题让一份 GPU 程序负责一个 C tile。这里的“程序”指一次处理一组 tile 的 GPU 工作实例,内部由线程协作执行。8 个 C tile 对应 8 份工作,各自写入不同的输出区域。

C_tile[0, 1] 为例,它需要 A 的一行块 [A00, A01, A02],以及 B 的一列块 [B01, B11, B21]。两串块按相同的 K 段编号配对,计算结果留在同一个累加器中,最后写回 C。

这样分工有两个直接好处。不同 C tile 可以独立计算;同一个 C tile 的三个部分结果在程序内完成相加,不需要让多个程序协调写入同一片输出。

接下来切换到实际的 64 分块,同时保留同样的块网格:

M = 128,K = 192,N = 256
BLOCK_SIZE_M = BLOCK_SIZE_N = BLOCK_SIZE_K = 64
 
A: (128, 192) → 2 × 3 个块
B: (192, 256) → 3 × 4 个块
C: (128, 256) → 2 × 4 个块

文中会将三个分块大小简写为 BM、BN、BK。一次局部乘法的形状是:

A tile: (BM, BK)    @    B tile: (BK, BN)    →    C tile: (BM, BN)

代码始终使用文件里的完整名称 BLOCK_SIZE_MBLOCK_SIZE_NBLOCK_SIZE_K

3. 打开任务文件,认识两段代码所处的位置

运行命令使用课程 GPU 环境中的 Bash 终端。先进入项目根目录;下面的 /data/gollamago 按你的实际项目路径替换。

cd /data/gollamago
source ./setup_env.sh
cd assignment/task3
python -m pip install ninetoothed
python -c "import ninetoothed, torch; print('NineToothed: OK'); print('GPU:', torch.cuda.is_available())"

确认出现 NineToothed: OKGPU: True 后,打开 ninetoothed_gemm.py。需要填写的是 arrangement()application() 中的两处占位异常。

文件下方的这段代码连接了两个函数:

_KERNEL = ninetoothed.make(
    arrangement,
    application,
    (Tensor(2), Tensor(2), Tensor(2)),
)

模块加载时,make 用三个二维符号张量分析布局与计算。Tensor(2) 中的 2 是维数;真实的行列数来自之后传入的张量。符号张量保存形状和访问方式等描述,里面没有输入矩阵的数值。

调用 nt_gemm(lhs, rhs) 时,lhsrhs 才是 GPU 上的 PyTorch 数据。封装函数创建 C,再调用 _KERNEL(lhs, rhs, output)。首次调用还可能触发对应配置的后端编译。

两个函数中的同名参数也有一处需要留意:arrangement() 从整个矩阵的符号描述开始;application() 描述一个程序取得局部数据后的计算。在我们的例子中,后者的 lhs 是 3 个 A tile 的序列,rhs 是 3 个 B tile 的序列,output 是一个 64 × 64 的 C tile。

4. 把需要的数据组织到同一个程序

arrangement() 要得到的对应关系已经确定:

程序坐标 (i, j)
    lhs:A 第 i 行的全部 K tile
    rhs:B 第 j 列的全部 K tile
    output:C_tile[i, j]

NineToothed 使用安排结果的最外层形状确定程序网格。因此,三者最外层都应为 C 的块网格 (2, 4)。A、B 的 K 序列放在里面,供各自程序遍历。

先处理 C:output.tile((BM, BN)) 得到外层 (2, 4),其中每个元素描述一个 (64, 64) 的输出块。

对 A 直接做 tile((BM, BK)) 后,外层是 (2, 3),第二维仍然是 K 块数。需要再把一整行 K 块收在一起,然后沿 N 方向映射到 C 的各列:

操作操作后的外层形状外层一个位置包含什么
lhs.tile((BM, BK))(2, 3)一个 64 × 64 的 A tile
再做 .tile((1, -1))(2, 1)形状为 (1, 3) 的一组 A tile
再做 .expand((-1, 4))(2, 4)同一组 A tile,供不同输出列使用
.dtype.squeeze(0)(2, 4)形状为 (3,) 的 A tile 序列

两种操作中 -1 的含义各有约定:在 tile 中表示取满当前维度,在 expand 中表示保持该维长度。于是 tile((1, -1)) 取一行、取满这行的 3 个 K 块;expand((-1, 4)) 保留 2 行,将长度为 1 的第二维扩展到 4。

最后一行中的 .dtype 表示“外层一个元素的类型”。经过嵌套切块,这个元素本身是一组 tile,具有自己的形状 (1, 3)squeeze(0) 去掉它长度为 1 的第 0 维,得到 (3,),计算时就能用 lhs[k] 取块。

B 沿列收集,过程如下:

操作操作后的外层形状外层一个位置包含什么
rhs.tile((BK, BN))(3, 4)一个 64 × 64 的 B tile
再做 .tile((-1, 1))(1, 4)形状为 (3, 1) 的一组 B tile
再做 .expand((2, -1))(2, 4)同一组 B tile,供不同输出行使用
.dtype.squeeze(1)(2, 4)形状为 (3,) 的 B tile 序列

这里收集的是 B 的块列,每一个块内部仍保持 (BK, BN) 的行列方向。expand 只描述哪些程序访问同一组数据,不会创建一份更大的 B。

代码中的 2 和 4 要从输出布局取得:output_arranged.shape[0] 是 M 块数,output_arranged.shape[1] 是 N 块数。这样换成其他矩阵尺寸时,程序网格会随之改变。

接着写 application()。当前程序的 output.shape(BM, BN),因此可以据此建立全零累加器;lhs.shape[0] 是 K 块数,本例为 3。循环中用 ntl.dot(lhs[k], rhs[k]) 计算一个部分结果并累加,循环完成后写回 output。

将两个函数填写为:

def arrangement(lhs, rhs, output):
    output_arranged = output.tile((BLOCK_SIZE_M, BLOCK_SIZE_N))
 
    lhs_arranged = lhs.tile((BLOCK_SIZE_M, BLOCK_SIZE_K))
    lhs_arranged = lhs_arranged.tile((1, -1))
    lhs_arranged = lhs_arranged.expand((-1, output_arranged.shape[1]))
    lhs_arranged.dtype = lhs_arranged.dtype.squeeze(0)
 
    rhs_arranged = rhs.tile((BLOCK_SIZE_K, BLOCK_SIZE_N))
    rhs_arranged = rhs_arranged.tile((-1, 1))
    rhs_arranged = rhs_arranged.expand((output_arranged.shape[0], -1))
    rhs_arranged.dtype = rhs_arranged.dtype.squeeze(1)
 
    return lhs_arranged, rhs_arranged, output_arranged
 
 
def application(lhs, rhs, output):
    accumulator = ntl.zeros(output.shape, dtype=ntl.float32)
 
    for k in range(lhs.shape[0]):
        accumulator += ntl.dot(lhs[k], rhs[k])
 
    output = accumulator

累加器建立在循环外,类型为 ntl.float32output = accumulator 是 NineToothed 编译器识别的输出赋值,会生成当前 tile 的写回操作。外部的 nt_gemm() 再把真实 PyTorch 输出返回给调用者。这里沿用 Task2 的 application 写法,无需额外返回值。

填写后检查两处 raise NotImplementedError(...) 都已删除。若它们还在,make 可能在模块导入、pytest 收集测试时就遇到异常。

5. 用测试确认实现覆盖了哪些情况

先运行作业自带的测试:

python -m pytest -q test_ninetoothed_gemm.py

正常情况下会看到 1 passed。这个测试使用 M=N=K=512,默认分块为 64,即 8×8 个输出块,每个程序遍历 8 段 K。它检查基本计算与多轮累加,所有维度都恰好整除。

方阵的几个维度数值相同,容易掩盖 M/N 写反的错误。再补一个非方阵、几个单独的尾块,以及三维同时不整除的情况。下面这段可以完整复制到同一个 Bash 终端中运行:

python - <<'PY'
import torch
from ninetoothed_gemm import nt_gemm
 
torch.manual_seed(0)
cases = [
    (128, 256, 192),  # M、N、K 各不相同,均能整除 64
    (65, 128, 128),   # 只有 M 有尾块
    (128, 97, 128),   # 只有 N 有尾块
    (128, 128, 130),  # K 有两个完整块和一个尾块
    (65, 97, 130),    # M、N、K 同时有尾块
]
for m, n, k in cases:
    a = torch.randn((m, k), device='cuda', dtype=torch.float16)
    b = torch.randn((k, n), device='cuda', dtype=torch.float16)
    out = nt_gemm(a, b)
    ref = torch.mm(a, b)
    assert out.shape == (m, n)
    assert out.dtype == torch.float16
    assert torch.allclose(out, ref, atol=0.025, rtol=0.025), (m, n, k)
    max_error = (out.float() - ref.float()).abs().max().item()
    print(f'M={m}, N={n}, K={k}: PASS, max_abs_error={max_error:.6f}')
PY

每行都应打印 PASSmax_abs_error 用来帮助观察;通过与否由 torch.allclose 判断。它同时使用绝对容差和相对容差,逐元素条件为:

|out - ref| <= 0.025 + 0.025 × |ref|

因此,单看最大绝对误差是否超过 0.025,无法代替整个判断。

尾块的计算仍然使用固定 tile 形状。对本题的生成代码,越界输入在带掩码的读取中取 0,越界输出位置不写入。K=130 时,第三段只有两个有效位置,其余位置的零不会增加乘积和。复盘文档会展开这段索引关系。

遇到失败时,可以沿着具体症状检查:

现象检查方向
512 方阵通过,非方阵失败A 是否沿输出 N 维 expand,B 是否沿输出 M 维 expand
K 只需一轮时正确,多轮时错误累加器是否在循环外初始化,循环里是否用 +=
dot 输入维度不匹配A 内块是否为 (BM, BK),B 内块是否为 (BK, BN);两处 squeeze 是否完整
只在 K 尾块失败K 块数是否保留向上取整,是否改变过默认布局或读取方式
导入时报错,测试尚未执行确认占位异常已删除,再区分依赖导入失败与 kernel 构建失败

不要通过扩大容差掩盖布局或漏算问题。先保留任务的容差,把失败范围缩小到某一类输入。

6. 保存性能基线

正确性检查通过后运行:

python benchmark_ninetoothed_gemm.py

脚本测试边长为 8、16、32、64、128、256、512、1024、2048、4096 的方阵。每个尺寸先和 torch.mm 比较,再计时;若比较失败,脚本会抛出异常,该尺寸不会打印 PASS

输出中的 NineToothed(ms)PyTorch(ms) 都以毫秒为单位。计时函数重复调用 nt_gemmtorch.mm;前者包含任务封装中的输出分配与 kernel 调用。因此这份数据应按脚本提供的调用方式比较,不宜当作纯设备指令的耗时。

保存默认 64×64×64 配置的结果作为基线。若要进一步尝试自动调优,在启动 Python 前设置环境变量:

NINETOOTHED_AUTOTUNE=1 python benchmark_ninetoothed_gemm.py

三个 block size 会进入指定范围的配置搜索。首次等待时间可能增加,脚本中的毫秒数不单独报告这段搜索成本。记录结果时同时记下 GPU 型号、是否开启调优和软件版本,便于之后比较。

提交前应有完成的 ninetoothed_gemm.py、自带测试的通过输出,以及默认配置下 benchmark 的完整结果截图。补充用例的输出也可以保留,帮助说明尾块和非方阵确实检查过。

接着阅读《Task3 NineToothed GEMM:理解布局、索引与执行》,自己打印嵌套张量的各层形状,并尝试修改 BK、预测循环次数。

本文实现对应任务 README 引用的 NineToothed 官方矩阵乘法教程