Same Algebra ⇏ \not\Rightarrow ⇒ Same Floating-Point Evaluation
错误版本·
在练习题目 LeetGPU | Matrix Multiplication 时,首先写了一个非标准的 Triton Kernel
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 @triton.jit def matrix_multiplication_kernel_err ( a, b, c, M, N, K, stride_am, stride_an, stride_bn, stride_bk, stride_cm, stride_ck ): m = tl.program_id(axis=0 ) k = tl.program_id(axis=1 ) BLOCK_N: tl.constexpr = 256 offsets = tl.arange(0 , BLOCK_N) acc = tl.zeros((BLOCK_N,), dtype=tl.float32) for n_start in tl.range (0 , N, BLOCK_N): n = n_start + offsets mask = n < N a_value = tl.load(a + m * stride_am + n * stride_an, mask=mask, other=0.0 ).to(tl.float32) b_value = tl.load(b + n * stride_bn + k * stride_bk, mask=mask, other=0.0 ).to(tl.float32) acc += a_value * b_value result = tl.sum (acc, axis=0 ) tl.store(c + m * stride_cm + k * stride_ck, result)
虽然数学上是在计算 C m , k = ∑ n A m , n B n , k C_{m, k} = \sum_{n} A_{m, n}B_{n, k} C m , k = ∑ n A m , n B n , k 。令 p n = A m , n B n , k p_n = A_{m, n}B_{n, k} p n = A m , n B n , k ,目标就是求 ∑ p n \sum p_n ∑ p n 。但是上面代码并没有按顺序累加,实际计算结构是
1 2 3 4 5 6 acc[0] = p0 + p256 + p512 + ... acc[1] = p1 + p257 + p513 + ... ... acc[255] = p255 + p511 + p767 + ... result = tl.sum(acc)
它被拆成了 256 条独立的跨步累加链,最后再做一次规约。但是对于 FP32 数据类型来说,有限精度浮点运算实际是
fl ( x ∘ y ) = ( x ∘ y ) ( 1 + δ ) , ∣ δ ∣ ≤ u \operatorname{fl}(x \circ y) = (x \circ y)(1+\delta), \quad |\delta|\le u
fl ( x ∘ y ) = ( x ∘ y ) ( 1 + δ ) , ∣ δ ∣ ≤ u
因此,每做一次加法,就可能发生一次舍入,则 fl ( fl ( a + b ) + c ) = = fl ( a + fl ( b + c ) ) \operatorname{fl}(\operatorname{fl}(a+b)+c) == \operatorname{fl}(a+\operatorname{fl}(b+c)) fl ( fl ( a + b ) + c ) == fl ( a + fl ( b + c )) 未必成立,不同的分组顺序会得到不同的局部和,不同的舍入时机就会导致最终结果误差不同。
本题的测试用例为 Matrix Multiplication — Triton 测试用例 ,测试 matrix_multiplication_kernel_err 会发现下面测试用例不通过
1 functional_14 max_dimensions 8,192 × 6,144 × 4,096 uniform(-1.0, 1.0)
该用例的特点是矩阵维度大,元素数值接近 0 0 0 。准确地说,这个用例真正敏感的并不是元素都很接近 0,而是 输入以 0 为中心、正负号混合,并且内积维度 N = 6144 N=6144 N = 6144 很大 。对任意一个输出元素,仍然记
s = ∑ n = 0 N − 1 p n , p n = A m , n B n , k . s = \sum_{n = 0}^{N-1} p_n, \qquad p_n = A_{m, n}B_{n, k}.
s = n = 0 ∑ N − 1 p n , p n = A m , n B n , k .
衡量求和问题对扰动敏感程度的一个常用量是求和条件数
κ s u m = ∑ n ∣ p n ∣ ∣ ∑ n p n ∣ = ∑ n ∣ p n ∣ ∣ s ∣ . \kappa_{\mathrm{sum}}
=\frac{\sum_n |p_n|}{\left|\sum_n p_n\right|}
=\frac{\sum_n |p_n|}{|s|}.
κ sum = ∣ ∑ n p n ∣ ∑ n ∣ p n ∣ = ∣ s ∣ ∑ n ∣ p n ∣ .
如果所有 p n p_n p n 基本同号,那么分子和分母量级接近,κ s u m \kappa_{\mathrm{sum}} κ sum 通常不会很大;但这里 A , B ∼ U ( − 1 , 1 ) A,B\sim U(-1,1) A , B ∼ U ( − 1 , 1 ) ,所以 p n p_n p n 也是关于 0 0 0 对称的随机变量,正负项会大量互相抵消。此时 ∑ ∣ p n ∣ \sum |p_n| ∑ ∣ p n ∣ 仍然很大,而最终的 ∣ s ∣ |s| ∣ s ∣ 却可能很小,于是 κ s u m \kappa_{\mathrm{sum}} κ sum 会迅速增大。极端情况下若真实和恰好为 0 0 0 ,相对条件数甚至可以视为无穷大。
对独立的 A m , n , B n , k ∼ U ( − 1 , 1 ) A_{m,n},B_{n,k}\sim U(-1,1) A m , n , B n , k ∼ U ( − 1 , 1 ) ,有
E [ p n ] = 0 , E ∣ p n ∣ = 1 4 , Var ( p n ) = 1 9 . \mathbb E [p_n] = 0,\qquad
\mathbb E|p_n|=\frac14,\qquad
\operatorname{Var}(p_n)=\frac19.
E [ p n ] = 0 , E ∣ p n ∣ = 4 1 , Var ( p n ) = 9 1 .
因此在 N = 6144 N=6144 N = 6144 时,可以粗略估计
∑ n ∣ p n ∣ ≈ N 4 = 1536 , std ( s ) = N 9 ≈ 26.1. \sum_n |p_n|\approx \frac{N}{4}= 1536,
\qquad
\operatorname{std}(s)=\sqrt{\frac{N}{9}}\approx 26.1.
n ∑ ∣ p n ∣ ≈ 4 N = 1536 , std ( s ) = 9 N ≈ 26.1.
如果某个输出元素的 ∣ s ∣ |s| ∣ s ∣ 恰好处在一个标准差附近,那么条件数已经约为
κ s u m ≈ 1536 26.1 ≈ 59. \kappa_{\mathrm{sum}}\approx \frac{1536}{26.1}\approx 59.
κ sum ≈ 26.1 1536 ≈ 59.
而输出矩阵共有 8192 × 4096 8192\times4096 8192 × 4096 个元素,其中可能会出现一些抵消更严重、∣ s ∣ |s| ∣ s ∣ 更接近 0 0 0 的位置;这些位置的 κ s u m \kappa_{\mathrm{sum}} κ sum 可以达到更高。这时浮点误差就会被条件数放大。对长度为 N N N 的浮点点积,在经典舍入模型下常见的前向误差估计具有下面的形式
∣ s ^ − s ∣ ≲ γ N ∑ n ∣ p n ∣ , γ N = N u 1 − N u , |\widehat{s}-s|
\lesssim
\gamma_N\sum_n |p_n|,
\qquad
\gamma_N =\frac{Nu}{1-Nu},
∣ s − s ∣ ≲ γ N n ∑ ∣ p n ∣ , γ N = 1 − N u N u ,
因此相对误差满足
∣ s ^ − s ∣ ∣ s ∣ ≲ γ N κ s u m . \frac{|\widehat{s}-s|}{|s|}
\lesssim
\gamma_N\kappa_{\mathrm{sum}}.
∣ s ∣ ∣ s − s ∣ ≲ γ N κ sum .
这里精确的常数会随乘法是否融合为 FMA、规约树的形状等实现细节变化,但核心项始终是 ∑ ∣ p n ∣ / ∣ s ∣ \sum|p_n|/|s| ∑ ∣ p n ∣/∣ s ∣ 。FP32 的单位舍入误差约为 u = 2 − 24 u=2^{-24} u = 2 − 24 ,当 N = 6144 N=6144 N = 6144 时 γ N ≈ 3.66 × 10 − 4 \gamma_N\approx 3.66\times10^{-4} γ N ≈ 3.66 × 1 0 − 4 ;一旦 κ s u m \kappa_{\mathrm{sum}} κ sum 很大,相对误差就可能被明显放大。这也解释了 灾难性抵消 为什么重要。抵消本身并不一定额外制造舍入误差,但它会把最终结果压到很小的量级,从而暴露并放大此前乘法、局部累加中已经产生的舍入误差。前面的错误 Kernel 先按 n m o d 256 n\bmod 256 n mod 256 把乘积拆成 256 条跨步累加链,再对 256 个局部和做 tl.sum:
q r = fl ( ∑ j p r + 256 j ) , s ^ = fl ( ∑ r = 0 255 q r ) . q_r =\operatorname{fl}\left(\sum_j p_{r+256j}\right),\qquad
\widehat{s}=\operatorname{fl}\left(\sum_{r = 0}^{255}q_r\right).
q r = fl ( j ∑ p r + 256 j ) , s = fl ( r = 0 ∑ 255 q r ) .
这种求和树与 PyTorch/CUDA GEMM 参考实现采用的计算路径并不相同。由于浮点加法不满足结合律,不同分组会在不同位置进行舍入;当某个输出元素又恰好具有很大的 κ s u m \kappa_{\mathrm{sum}} κ sum 时,本来只有几个 ULP(Unit in the Last Place) 的局部差异就可能被放大成明显的相对误差,最终越过测试的允许误差。
所以,这个大尺寸随机用例容易暴露问题,本质上是三个因素叠加:内积很长、数据正负混合导致强抵消、实现改变了求和树 。
修正版本·
参考官方的 Triton GEMM 写法,我又写了一个 分块版本的二维矩阵乘法核并且通过了 Matrix Multiplication — Triton 测试用例
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 @triton.jit def matrix_multiplication_kernel ( a, b, c, M, N, K, stride_am, stride_an, stride_bn, stride_bk, stride_cm, stride_ck, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, ): pid = tl.program_id(axis=0 ) num_pid_k = tl.cdiv(K, BLOCK_K) pid_m = pid // num_pid_k pid_k = pid % num_pid_k offsets_m = pid_m * BLOCK_M + tl.arange(0 , BLOCK_M) offsets_k = pid_k * BLOCK_K + tl.arange(0 , BLOCK_K) offsets_n = tl.arange(0 , BLOCK_N) a_ptrs = a + offsets_m[:, None ] * stride_am + offsets_n[None , :] * stride_an b_ptrs = b + offsets_n[:, None ] * stride_bn + offsets_k[None , :] * stride_bk acc = tl.zeros((BLOCK_M, BLOCK_K), dtype=tl.float32) for n_start in range (0 , N, BLOCK_N): a_values = tl.load( a_ptrs, mask=(offsets_m[:, None ] < M) & (offsets_n[None , :] + n_start < N), other=0.0 , ) b_values = tl.load( b_ptrs, mask=(offsets_n[:, None ] + n_start < N) & (offsets_k[None , :] < K), other=0.0 , ) acc += tl.dot(a_values, b_values, input_precision="ieee" ) a_ptrs += BLOCK_N * stride_an b_ptrs += BLOCK_N * stride_bn c_ptrs = c + offsets_m[:, None ] * stride_cm + offsets_k[None , :] * stride_ck tl.store(c_ptrs, acc, mask=(offsets_m[:, None ] < M) & (offsets_k[None , :] < K))
计算图是 C tile = D 0 + D 1 + D 2 + ⋯ C_{\text{tile}} = D_0 + D_1 + D_2 + \cdots C tile = D 0 + D 1 + D 2 + ⋯ ,其中 D i = dot ( A tile i , B tile i ) D_i = \operatorname{dot}(A_{\text{tile}_i}, B_{\text{tile}_i}) D i = dot ( A tile i , B tile i ) .
这里的 tl.dot 不是对两个向量做一次标量点积,而是对两个 二维 block 做小型矩阵乘法 。对当前 Kernel 的记号,设第 i i i 次循环覆盖归约维度上的区间
I i = [ i ⋅ B L O C K N , ( i + 1 ) ⋅ B L O C K N ) , I_i =[i\cdot BLOCK_N,(i+1)\cdot BLOCK_N),
I i = [ i ⋅ B L O C K N , ( i + 1 ) ⋅ B L O C K N ) ,
那么加载到 SRAM/寄存器语义中的两个 block 可以写成
A i ∈ R B L O C K M × B L O C K N , B i ∈ R B L O C K N × B L O C K K . A_i\in\mathbb R^{BLOCK_M\times BLOCK_N},\qquad
B_i\in\mathbb R^{BLOCK_N\times BLOCK_K}.
A i ∈ R B L O C K M × B L O C K N , B i ∈ R B L O C K N × B L O C K K .
tl.dot(a_values, b_values, input_precision="ieee") 返回一个
D i = A i B i ∈ R B L O C K M × B L O C K K D_i = A_iB_i\in\mathbb R^{BLOCK_M\times BLOCK_K}
D i = A i B i ∈ R B L O C K M × B L O C K K
的 block。展开到单个输出位置 ( r , c ) (r,c) ( r , c ) ,它所做的工作就是
( D i ) r , c = ∑ j = 0 B L O C K N − 1 ( A i ) r , j ( B i ) j , c . (D_i)_{r, c}
=\sum_{j = 0}^{BLOCK_N-1}(A_i)_{r, j}(B_i)_{j, c}.
( D i ) r , c = j = 0 ∑ B L O C K N − 1 ( A i ) r , j ( B i ) j , c .
因此整个 Kernel 的计算图可以写成
1 2 3 4 5 6 7 8 9 10 11 12 13 14 A_0 [BM×BN] ─┐ ├─ tl.dot ─> D_0 [BM×BK] ─┐ B_0 [BN×BK] ─┘ │ ├─ + ─> acc_1 A_1 [BM×BN] ─┐ │ ├─ tl.dot ─> D_1 [BM×BK] ─┘ B_1 [BN×BK] ─┘ acc_0 = 0 acc_1 = fl(acc_0 + D_0) acc_2 = fl(acc_1 + D_1) ... acc_T = fl(acc_{T-1} + D_{T-1}) C_tile = acc_T
其中
T = ⌈ N B L O C K N ⌉ . T =\left\lceil\frac{N}{BLOCK_N}\right\rceil.
T = ⌈ B L O C K N N ⌉ .
也就是说,一个 Triton program instance 不再只计算一个标量 C m , k C_{m,k} C m , k ,而是一次负责输出矩阵中的一个 B L O C K M × B L O C K K BLOCK_M\times BLOCK_K B L O C K M × B L O C K K tile。每轮循环沿着归约维 N N N 取一块连续的 A i A_i A i 和 B i B_i B i ,tl.dot 同时产生这一整块输出的部分和,然后外层循环继续把各个 D i D_i D i 累加到 FP32 的 acc 中。这正是典型 blocked GEMM 的结构:
C t i l e = A 0 B 0 + A 1 B 1 + ⋯ + A T − 1 B T − 1 . C_{\mathrm{tile}}
= A_0B_0+A_1B_1+\cdots+A_{T-1}B_{T-1}.
C tile = A 0 B 0 + A 1 B 1 + ⋯ + A T − 1 B T − 1 .
从数值角度看,这个版本仍然没有摆脱浮点非结合性:tl.dot 内部的乘加顺序以及不同 D i D_i D i 之间的累加仍然会产生舍入。但是它有两个重要区别。
第一,acc 明确使用 tl.float32,所以各个 block 的结果在整个归约过程中都保留在 FP32 累加器中。Triton 官方 GEMM 教程采用的也是“加载 A / B A/B A / B block → tl.dot → 累加到 FP32 accumulator”的结构。
第二,当前输入本身就是 FP32,而 Triton 在 NVIDIA GPU 上对 f32×f32 的 tl.dot 默认可能采用 TF32 输入精度。这里显式指定
表示要求 f32 点积使用 IEEE 精度路径,而不是默认的 TF32 输入精度,从而避免先把 FP32 输入有效尾数压缩到 TF32 精度后再做矩阵乘。需要注意,这个参数描述的是 tl.dot 的输入计算精度;具体最终 lower 成什么指令序列仍取决于 Triton 版本和目标硬件,不能仅凭 Python 代码把它理解成某一种固定的机器指令。
把两个版本放在一起看,它们的差别可以概括为:
1 2 3 4 5 6 7 8 9 10 11 12 13 错误版本: p_n │ ├─ 按 n mod 256 分成 256 条跨步链 │ q_0, q_1, ..., q_255 └─ tl.sum(q) -> C[m, k] 分块 GEMM: 连续的归约区间 I_0, I_1, ... │ ├─ A_i × B_i --tl.dot--> D_i(一个输出 tile 的部分和) │ └─ FP32 acc = D_0 + D_1 + ... -> C_tile
后者使用了 Triton 专门为矩阵乘提供的 block-level tl.dot ,计算组织方式也与常规高性能 GEMM 更一致。
参考资料·
LeetGPU | Matrix Multiplication
Triton’s documentation | Group GEMM
Triton’s documentation | Matrix Multiplication
Triton 中文站 | 分组 GEMM
Triton’s documentation | triton.language.dot
NVIDIA | Floating Point and IEEE 754
讨论
评论