Task3 NineToothed GEMM:理解布局、索引与执行

完成任务后,可以先留住代码里的这行:

lhs_arranged.dtype = lhs_arranged.dtype.squeeze(0)

它修改了哪一层?为什么这一层会多出一个 1?删掉它之后,lhs[k] 为什么就能取得正确的小矩阵?这些问题都能通过查看形状、追踪一次读取来回答。

本文沿用完成指南中的实际尺寸:A 为 (128, 192),B 为 (192, 256),C 为 (128, 256),BM、BN、BK 都为 64。BM、BN、BK 分别是三个 BLOCK_SIZE_* 的简写。输出有 2×4 个 tile,每个输出 tile 需要沿 K 方向累加 3 次。

1. 把嵌套张量的每一层打印出来

普通 PyTorch 张量中的一个元素通常是标量,因此 dtype 常见的值是 torch.float16。NineToothed 的符号张量允许一个元素本身也是张量。执行一次 tile 后,外层元素的类型就是一个 tile 的描述,所以能继续访问 .dtype.shape

用 A 举例:

a.shape                 → 当前这层有多少个元素
若 a.dtype 也是 Tensor:
a.dtype.shape           → 每个外层元素内部的形状
若还能再往里一层:
a.dtype.dtype.shape     → 最内层小矩阵的形状

这里每一层的“元素”含义不同。最外层可以是一组 tile,中间层的元素可以是一个小矩阵,最内层才走到矩阵内的数值。只看 .shape 的两个数字,很容易漏掉它们正在描述哪一层。

A 的全部变化如下。表中的每一列对应一层,空白表示没有这一层 Tensor:

A 经历的操作.shape.dtype.shape.dtype.dtype.shape
原矩阵描述(128, 192)
tile((64, 64))(2, 3)(64, 64)
tile((1, -1))(2, 1)(1, 3)(64, 64)
expand((-1, 4))(2, 4)(1, 3)(64, 64)
.dtype 换成 .dtype.squeeze(0)(2, 4)(3,)(64, 64)

第一次 tile 将标量矩阵分成小矩阵。第二次 tile 的输入已经是 (2, 3) 的块网格,(1, -1) 在这一层取一行、取满 3 列。因此新外层有 (2, 1) 组,每组中排列着 (1, 3) 个小矩阵。原来每块 (64, 64) 的形状保持不变,只是深入了一层。

expand 将外层第二维从 1 扩展为 4,让一行 A tile 对应 4 个输出列。它保持组内 (1, 3) 的形状。随后 .dtype.squeeze(0) 把组内从 (1, 3) 改成 (3,),外层 (2, 4) 也保持不变。

对 B 做同样的追踪:

B 经历的操作.shape.dtype.shape.dtype.dtype.shape
原矩阵描述(192, 256)
tile((64, 64))(3, 4)(64, 64)
tile((-1, 1))(1, 4)(3, 1)(64, 64)
expand((2, -1))(2, 4)(3, 1)(64, 64)
.dtype 换成 .dtype.squeeze(1)(2, 4)(3,)(64, 64)

B 收集的是块列,所以组内形状为 (3, 1),要去掉的维度编号是 1。这里 A、B 的最终层次相同,里面的块仍然分别沿 A 的行、B 的列排列。

C 只做了一次 tile,其层次为:

C 的输出块网格 (2, 4) → 一个输出块 (64, 64)

A、B 多出来的那层 (3,),正是一个输出块所依赖的 K 块序列。

这些形状可以直接打印。下面代码单独运行,只构造符号描述,不调用 GPU kernel;环境需要能导入 NineToothed。

from ninetoothed import Tensor
 
 
def show(label, tensor):
    shapes = []
    while isinstance(tensor, Tensor):
        shapes.append(str(tensor.shape))
        tensor = tensor.dtype
    print(label + ': ' + ' -> '.join(shapes))
 
 
BM = BN = BK = 64
c = Tensor(shape=(128, 256)).tile((BM, BN))
 
a = Tensor(shape=(128, 192))
show('A original', a)
a = a.tile((BM, BK))
show('A tiled', a)
a = a.tile((1, -1))
show('A grouped', a)
a = a.expand((-1, c.shape[1]))
show('A expanded', a)
a.dtype = a.dtype.squeeze(0)
show('A squeezed', a)
 
b = Tensor(shape=(192, 256))
b = b.tile((BK, BN))
b = b.tile((-1, 1))
b = b.expand((c.shape[0], -1))
b.dtype = b.dtype.squeeze(1)
show('B final', b)
show('C final', c)

预期输出:

A original: (128, 192)
A tiled: (2, 3) -> (64, 64)
A grouped: (2, 1) -> (1, 3) -> (64, 64)
A expanded: (2, 4) -> (1, 3) -> (64, 64)
A squeezed: (2, 4) -> (3,) -> (64, 64)
B final: (2, 4) -> (3,) -> (64, 64)
C final: (2, 4) -> (64, 64)

调试布局时,这样逐层打印能定位很多问题。例如,漏掉 squeeze 时 A 仍然是 (1, 3)lhs.shape[0] 就会得到 1;遍历次数和索引方式都随之改变。

2. 从符号布局走到 application

任务文件在模块级调用:

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

这里的三个 Tensor(2) 只有二维参数的描述,具体 M、N、K 仍然是符号。前面打印形状的代码使用具体数字,是为了方便观察;实际 kernel 的布局规则会随着传入矩阵的尺寸计算块数。

make 在构建阶段调用 arrangement,得到三个嵌套布局,并分析 application 的代码。等外部调用 nt_gemm(a, b),真实数据、尺寸和步长才参与对应 kernel 的运行。application 中的循环会进入生成的 GPU 程序,执行时不会由 Python 逐轮调用 ntl.dot

编译器根据最外层网格安排程序,然后把该位置的内层对象映射为 application 的参数。所以同一个名字在两个函数中有不同的形状:

名字arrangement 返回时的最外层当前 application 中的 .shape再索引一次
lhs(2, 4)(3,)lhs[k](64, 64)
rhs(2, 4)(3,)rhs[k](64, 64)
output(2, 4)(64, 64)对应当前 C tile 的元素

于是这段计算就有了明确的形状来源:

accumulator = ntl.zeros(output.shape, dtype=ntl.float32)
for k in range(lhs.shape[0]):
    accumulator += ntl.dot(lhs[k], rhs[k])

累加器为 64 × 64,循环 3 次。输出网格的 (2, 4) 已经用于选择当前程序,因此不会再出现在 lhs.shape 中。

同样,output = accumulator 会被编译器识别为向输出参数赋值并生成写回操作。只把这个函数作为普通 Python 函数单独调用,无法获得 NineToothed kernel 的语义;它需要与 arrangement 一起交给 make

3. 追踪一次读取,确认数据确实对齐

现在选定程序坐标 (i, j) = (1, 2)。它负责第二行块、第三列块:

C_tile[1, 2] 对应 C[64:128, 128:192]

当 K 循环走到第 p=1 段时:

lhs[p] 对应 A[64:128, 64:128]
rhs[p] 对应 B[64:128, 128:192]

用块内下标 r、s、t 表示局部行、局部列和局部 K 位置,原矩阵索引为:

lhs[p][r, t]  → A[i × BM + r, p × BK + t]
rhs[p][t, s]  → B[p × BK + t, j × BN + s]
output[r, s]  → C[i × BM + r, j × BN + s]

A 的索引不需要 j,因为同一行输出的所有列块都读取这行 A;B 的索引不需要 i,因为同一列输出的所有行块都读取这列 B。这正是两次 expand 建立的对应关系。

可以据此检查两个容易混淆的概念。B 的块按列组成 K 序列,但单个 B tile 的行列仍是 (BK, BN),满足 dot 的输入要求。不同程序也可能读到相同的 A 或 B 数据;expand 描述这种共享访问,不保证这些数据只从显存读取一次,实际复用还取决于缓存和编译器安排。

把局部索引代回矩阵乘法,一轮 dot 完成:

partial[r, s] = Σt A[i×BM+r, p×BK+t] × B[p×BK+t, j×BN+s]

Σt 表示把当前 BK 段的乘积加起来。外面的 K 块循环再把各个 p 的 partial 相加,于是覆盖整个原始 K 维。

4. 为什么 M、N 放在外层,K 留在循环中

M 和 N 确定“结果写到哪里”。两个不同的输出 tile 写入不同的矩阵区域,各自计算即可。K 确定“同一个结果包含哪些加数”,这些加数最终必须合并。

如果把 K 块也直接当成互相独立的程序,负责 (i, j, 0)(i, j, 1)(i, j, 2) 的程序都会产生 C_tile[i, j] 的一部分。这时还需要额外的合并机制,直接赋值会覆盖其他部分。当前实现把 K 循环留在一个程序中,让合并过程由本地累加器完成。

在默认例子中,输出网格为 (2, 4),总共有 8 份程序工作,每份循环 3 次。GPU 可以调度这些独立工作;具体有多少份同时驻留,取决于硬件资源。

改变分块参数时,也可以按这个分工推算:

M 块数 = ceil(M / BM)
N 块数 = ceil(N / BN)
K 块数 = ceil(K / BK)

ceil 表示向上取整。BM、BN 改变输出块大小和输出网格;BK 改变每一轮覆盖的 K 长度和循环次数。无论如何划分,矩阵乘法仍须包含完整 K 范围内的乘积。

5. 尾块怎样参与计算

取完成指南中的边界用例:

M = 65,N = 97,K = 130
BM = BN = BK = 64
 
输出块网格:(2, 2)
K 循环次数:3

右下角输出 tile 只有 1 行、33 列有效数据;K 的最后一个 tile 只有 2 个有效位置。计算所用 tile 的计划形状仍然是 64 × 64,有效位置由访问掩码决定。掩码就是针对每个位置的有效性条件。

按上一节的索引公式,读取 A 时检查:

i×BM+r < M  并且  p×BK+t < K

读取 B 时检查:

p×BK+t < K  并且  j×BN+s < N

在本题的默认生成路径中,有效位置读入真实数据,越界位置的读取值为 0。最后一轮 p=2 对应原始 K 下标 128 到 191,其中只有 128、129 有效;其余位置按零参与点积,乘积和不会因此增加。

写回 C 时只允许:

i×BM+r < M  并且  j×BN+s < N

因此右下角 tile 只写回那 1×33 个真实存在的输出元素。输入的零填充值与输出的写入掩码分别解决计算和存储两件事。

这些行为来自布局生成的边界条件和后端的带掩码读写。本题无需手写它们,但测试仍应检查相应路径:M 尾块、N 尾块和 K 尾块各有不同的访问位置。自带的 512 方阵测试只覆盖整除情况;benchmark 的 8、16、32 虽然包含不满一个 tile 的输入,也无法替代“多个完整 K 块之后再接一个尾块”的检查。

6. float32 累加保留了什么

矩阵乘法的一项输出需要把很多乘积相加。浮点数能表示的精度有限,两个相邻可表示数之间的间距会随数值大小增加。

例如,float16 在 2048 附近已经无法表示每一个整数。考虑从 2048 开始,连续加两次 1:

每次都舍入回 float16:2048 → 2048 → 2048
用 float32 累加:     2048 → 2049 → 2050
最终转为 float16:                 2050

第二种方式保留了中途的 2049,所以后一个 1 还能继续加上去。虽然最终仍然写入 float16,较高精度的中间状态依然有价值。

本题用 ntl.zeros(output.shape, dtype=ntl.float32) 指定累加器精度,输入数据及输出存储为 float16。输入在生成时已经受 float16 精度限制,写回时也会发生舍入;float32 累加无法消除所有误差。

因此测试使用 torch.allclose,逐元素比较:

|out - ref| <= atol + rtol × |ref|

任务里 atol=rtol=0.025。参考值为 0 时,容差为 0.025;参考值为 100 时,容差为 2.525。这是当前作业选择的验收标准。检查输出形状、尾块和累加完整性时,应保持这个标准,并用针对性输入帮助定位原因。

7. 分块与性能的关系

计算一个 (BM, BN) 输出块的某一段 K 时,要读入 (BM, BK) 的 A 和 (BK, BN) 的 B。A tile 中的一个元素会参与 BN 个输出列的计算,B tile 中的一个元素会参与 BM 个输出行的计算。把这些计算安排在同一程序内,为输入复用提供了条件,也是分块矩阵乘法常用的组织方式。

分块变大,同时会增加单个程序需要处理的数据和临时结果。特别是累加器的大小为 BM × BN,它需要占用相应的 GPU 资源。程序数量、资源占用和数据复用会共同影响速度,因此从代码长度或块大小本身推不出性能优劣。

任务的自动调优将三个 block size 设为搜索参数。在参考版本中,这些参数采用给定范围内的 2 的幂,范围 32 到 128 对应 32、64、128;实际候选还受后端配置约束。环境变量要在启动 Python 前设置,因为文件导入时就决定使用固定值还是搜索参数。

阅读 benchmark 时可以先看两端:很小的矩阵计算量少,启动和调用等固定成本更突出;较大的矩阵能提供更多计算工作,分块及底层矩阵乘法实现的影响也更明显。脚本使用 torch.mm 作为对照,结果只对应当前设备、输入尺寸和调用方式。

首次调用某个尺寸发生在计时前,也可能触发编译或调优。终端等了多久和表格打印的毫秒数描述不同的成本,记录实验时应分别观察。

8. 改一个条件,再核对自己的推算

下面三个练习都可以从已经完成的代码出发。先写下预测,再展开答案;需要运行时,在课程环境执行。

练习一:只把 BK 从 64 改为 32。 保持 M=128、N=256、K=192,BM、BN 仍为 64。输出有多少个程序?每个程序循环几次?lhs[k]rhs[k] 的形状分别是什么?可以将第 1 节的形状观察代码改为 BM = BN = 64; BK = 32 来核对,观察过程无需运行 kernel。

展开核对

输出网格仍为 (2, 4),共有 8 个程序。K 块数变为 6,故每个程序循环 6 次。A 内块为 (64, 32),B 内块为 (32, 64),单次 dot 的结果仍为 (64, 64)

最终层次应为:

A: (2, 4) -> (6,) -> (64, 32)
B: (2, 4) -> (6,) -> (32, 64)
C: (2, 4) -> (64, 64)

练习二:保留二次 tile 和 expand,删掉两处 squeeze。 为了让计算仍然正确,application 的循环范围与索引需要怎样改变?

展开核对

当前 application 收到的 A 组形状为 (1, 3),B 组形状为 (3, 1)。因此使用 A 的第二维作为循环范围,再分别从块行和块列中取出小矩阵:

def application(lhs, rhs, output):
    accumulator = ntl.zeros(output.shape, dtype=ntl.float32)
    for k in range(lhs.shape[1]):
        accumulator += ntl.dot(lhs[0, k], rhs[k, 0])
    output = accumulator

这说明 squeeze 负责简化访问形式。若恢复原版 application,也需要恢复两处 squeeze,使循环与布局继续匹配。

练习三:用全 1 输入检查最后一段 K。M=65、N=97、K=130,A、B 所有元素都为 1。正确输出中的每个元素应该是多少?如果实现遗漏最后一个 K tile,会得到什么?

展开核对与运行代码

每个输出元素包含 130 个 1 × 1,因此应为 130。遗漏最后一段后只累计 128 项,结果会变成 128。这个输入有助于直接识别“少算最后一段”的问题。

assignment/task3 的 Bash 终端运行:

python - <<'PY'
import torch
from ninetoothed_gemm import nt_gemm
 
a = torch.ones((65, 130), device='cuda', dtype=torch.float16)
b = torch.ones((130, 97), device='cuda', dtype=torch.float16)
out = nt_gemm(a, b)
assert out.shape == (65, 97)
assert torch.equal(out, torch.full_like(out, 130))
print('PASS: every output element is 130')
PY

通过后,再运行完成指南中的随机非方阵和尾块用例。全 1 输入方便定位漏算,随机输入还能揭露一些块错配问题。

继续阅读源码时,可以从这三个位置对应回本文: