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 NM 决定输出有多少行,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_M、BLOCK_SIZE_N、BLOCK_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: OK 和 GPU: 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) 时,lhs、rhs 才是 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.float32。output = 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每行都应打印 PASS。max_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_gemm 和 torch.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 官方矩阵乘法教程。