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 输入方便定位漏算,随机输入还能揭露一些块错配问题。
继续阅读源码时,可以从这三个位置对应回本文:
- 官方 Basics:矩阵乘法:安排输入块行、块列与输出块。
- Tensor 的实现:
tile、expand、squeeze的形状与索引变化。 - 后端读写生成:读取填充值和访问掩码。