Skip to content

线性代数与 GEMM ​

标签
AI/ml/数学
字数
3597 字
阅读时间
15 分钟

AI Infra 对数学的要求是看到公式能立刻回答三个工程问题:

  1. 形状是否匹配? 内维对不对,输出是什么 shape
  2. 代价是多少? 多少 FLOPs、多少访存
  3. 实现风险是什么? 大数据类型、内存布局、分块、数值稳定性

以语言模型的输出投影为例:

Y=XW,X∈R(BS)×H,W∈RH×V
  • 形状:内维 H 相同,输出 (BS)×V
  • 代价:每个输出元素做 H 次乘加,共约 2BSHV FLOPs
  • 风险:权重很大、输出 logits 很大 —— 要考虑 dtype、内存布局、分块、并行策略,以及 Softmax 的数值稳定性

这三问贯穿整个技术栈,后面每一层优化都是在回答其中某一个。

AI Infra 对数学的要求是「看到公式能立刻回答三个工程问题」。

   ┌──────────────────────────────────────────────────────────┐
   │ ① 形状是否匹配?                                           │
   │   内维对不对,输出是什么 shape                              │
   ├──────────────────────────────────────────────────────────┤
   │ ② 代价是多少?                                             │
   │   多少 FLOPs、多少访存                                     │
   ├──────────────────────────────────────────────────────────┤
   │ ③ 实现风险是什么?                                          │
   │   大数据类型、内存布局、分块、数值稳定性                     │
   └──────────────────────────────────────────────────────────┘

   以语言模型的输出投影为例:Y = X W
     X ∈ R^{(BS)×H}、W ∈ R^{H×V}
     形状:内维 H 相同,输出 (BS) × V
     代价:每个输出元素做 H 次乘加,共约 2BSHV FLOPs
     风险:权重很大、输出 logits 很大 —— 要考虑 dtype、内存布局、
           分块、并行策略,以及 Softmax 的数值稳定性
   └─ 这三问贯穿整个技术栈,后面每一层优化都是在回答其中某一个。

符号与约定 ​

符号含义
x标量
x向量
X矩阵或高阶张量
Xij矩阵第 i 行第 j 列
X⊤转置
‖x‖2L2 范数
⊙逐元素乘法

深度学习代码用 batch-first 记法:

符号含义
Bbatch size
Ssequence length
Hhidden size
Vvocabulary size
Nhattention head 数
Dh=H/Nh每个 head 的维度

张量不只是 shape ​

工程里的「张量」是这几项信息的组合:

工程里的「张量」是六项信息的组合

   数据地址 + shape + stride + dtype + device + layout
        │
        └─ 两个张量 shape 相同,不代表内存布局相同;
           数值相同,也不代表 dtype 与误差特性相同

   例子:(2,3) 的行主序矩阵   ┌ 1 2 3 ┐
                              └ 4 5 6 ┘
     底层连续存储   [1, 2, 3, 4, 5, 6]
     stride        (3, 1)
     元素偏移      offset(i, j) = 3i + j
        │
        ▼ 转置视图 X^T
     shape         变成 (3, 2)
     stride        变成 (1, 3)
     数据          一个字节都不用搬
     └─ 代价是:按转置后的最后一维扫描时访问不再连续

   看到这几个函数要同时想逻辑维度与物理布局
     view        只改 stride(要求能表示),不拷贝
     reshape     必要时会拷贝一次
     contiguous  显式把数据重排成连续

两个张量 shape 相同,不代表内存布局相同;数值相同,也不代表 dtype 与误差特性相同。

stride 与偏移 ​

一个 shape 为 (2,3) 的行主序矩阵:

X=[123456]

底层连续存储为 [1,2,3,4,5,6],以元素为单位的 stride 是 (3,1)。元素 Xij 的线性偏移:

offset(i,j)=3i+j

转置视图 X⊤ 的 shape 变成 (3,2),但 stride 变成 (1,3),数据一个字节都不用搬。代价是:按转置后的最后一维扫描时访问不再连续。

看到 view / reshape / transpose / contiguous,要同时想逻辑维度和物理布局。 view 与 reshape 的差别就在这 —— view 只改 stride(要求能表示),reshape 必要时会拷贝一次。contiguous() 则是显式把数据重排成连续。

四类运算的形状规则,广播最容易出错。

   逐元素         要求对应元素能配对
   广播           从尾部维度向前对齐,每对维度
                  「相等 / 有一个是 1 / 有一方不存在」
   点积           n × n → 标量
   矩阵乘         (M, K) @ (K, N) → (M, N),内维 K 归约
   batch matmul   前若干维是 batch 维,只对最后两维做矩阵乘

   例:Q K^T ∈ R^{B × N_h × S_q × S_k}
     前两维 (B, N_h) 是 batch 维
     每个 (batch, head) 独立做一次矩阵乘
     └─ 「高阶张量乘法」不神秘:最后一维做矩阵乘、前面维度负责批处理

   广播的危险在于「不报错但算错」
     广播是逻辑视图,不会先复制出完整张量。
     但错误 shape 也可能恰好能广播 ——
     代码不报错,语义已经错了。
     └─ 这类 bug 在自定义算子里很常见,形状检查要写严格。

逐元素、广播、点积、矩阵乘 ​

运算形状规则
逐元素要求对应元素能配对
广播从尾部维度向前对齐,每对维度「相等 / 有一个是 1 / 有一方不存在」
点积n×n→ 标量
矩阵乘(M,K)@(K,N)→(M,N),内维 K 归约
batch matmul前若干维是 batch 维,只对最后两维做矩阵乘
QK⊤∈RB×Nh×Sq×Sk

前两维 (B,Nh) 是 batch 维,每个 (batch, head) 独立做一次矩阵乘。「高阶张量乘法」不神秘,就是最后一维做矩阵乘、前面维度负责批处理。

广播的危险在于「不报错但算错」

广播是逻辑视图,不会先复制出完整张量。但错误 shape 也可能恰好能广播 —— 代码不报错,语义已经错了。这类 bug 在自定义算子里很常见,形状检查要写严格。

范数与容差 ​

‖x‖1=∑i|xi|,‖x‖2=∑ixi2,‖x‖∞=maxi|xi|

范数在工程里用于:衡量参数/激活/梯度规模、梯度裁剪、比较参考实现与优化实现的误差、正则化、向量归一化。

比较浮点结果时,测试库通常组合绝对与相对容差:

|x^−x|≤atol+rtol⋅|x|

为什么不能只用相对容差:参考值 x 接近 0 时相对误差会被放大。这也是 torch.allclose / numpy.testing 的默认行为。

线性层其实是仿射变换 ​

严格的线性变换满足 f(ax+by)=af(x)+bf(y)。神经网络里的「Linear 层」还带 bias:

y=Wx+b

这是仿射变换。一条推论值得记住:多个没有激活函数的线性/仿射层可以合并成一个 —— 所以非线性激活是深层网络能表达复杂函数的关键,不是可选项。

秩、低秩分解与 LoRA ​

矩阵的秩可以理解为「独立方向的数量」:

rank(A)≤min(M,N)

秩低意味着信息冗余,可用两个小矩阵近似:

A≈UV,U∈RM×r, V∈Rr×N, r≪min(M,N)

参数量从 MN 降到 r(M+N) —— 这是低秩适配(LoRA)与模型压缩的数学基础。一个 4096×4096 的权重,rank 8 的 LoRA 只增加 8×8192=65,536 个参数,不到原矩阵的 0.4%。

GEMM 与分块 ​

分块矩阵乘法是恒等式 ​

把 A、B 按块划分,则 AB 的每个块可以由小块的乘加得到:

[A11A12A21A22][B11B12B21B22]=[A11B11+A12B21⋯⋯⋯]

这是 tiling 的数学依据 —— 它改变计算顺序与数据搬运,不改变目标公式。 理解这一点就不会把 tiling 当成「工程近似」。

朴素实现浪费在哪 ​

对 (M,K)@(K,N):输出元素 MN 个,每个做 K 次乘加,FLOPs 约 2MKN。

朴素实现若每次乘加都从全局内存加载 Aik 与 Bkj,同一个元素会被反复读取:

数据复用说明
A 的一个元素被同一输出行的多个列复用
B 的一个元素被同一输出列的多个行复用
累加器在寄存器中被复用 K 次

GPU kernel 的 tiling 就是把这件事变成:取一小块 A 与一小块 B → 搬到片上共享内存/寄存器 → 复用这些数据算出多个输出元素 → 沿 K 方向迭代累加 → 写回 HBM。

分块矩阵乘是恒等式,不是工程近似。

   ┌ A₁₁  A₁₂ ┐ ┌ B₁₁  B₁₂ ┐   ┌ A₁₁B₁₁ + A₁₂B₂₁   … ┐
   └ A₂₁  A₂₂ ┘ └ B₂₁  B₂₂ ┘ = └      …             … ┘

   └─ 每个块可以由小块的乘加得到 ⟹ 这是 tiling 的数学依据
      它改变计算顺序与数据搬运,不改变目标公式

   朴素实现浪费在哪:(M,K) @ (K,N) 的输出有 MN 个元素,
   每个做 K 次乘加,FLOPs 约 2MKN
   若每次乘加都从全局内存加载 A_ik 与 B_kj,同一个元素会被反复读取
     A 的一个元素   被同一输出行的多个列复用
     B 的一个元素   被同一输出列的多个行复用
     累加器         在寄存器中被复用 K 次
        │
        ▼ GPU kernel 的 tiling 就是把这件事实体化
     取一小块 A 与一小块 B ──▶ 搬到片上共享内存 / 寄存器
                          ──▶ 复用这些数据算出多个输出元素
                          ──▶ 沿 K 方向迭代累加 ──▶ 写回 HBM

算术强度可以估 ​

理想情况下每个输入矩阵只从 HBM 读一次、输出写一次,粗略字节数 2(MK+KN+MN):

AI≈2MKN2(MK+KN+MN)

这不是精确性能模型,但足以判断一个 GEMM 更可能受算力还是带宽限制 —— 判据是把它和硬件的 ops:byte 比(H100 约 295 FLOP/Byte)比较,见 01-GPU 硬件架构与存储层次。

tile 不是越大越好 ​

更大的 tile 提高复用,但会消耗更多:

  • 共享内存
  • 寄存器
  • 每 block 的线程数与同步开销
  • 边界处理成本

资源占用过高会降低一个 SM 上同时驻留的 block/warp 数,影响延迟隐藏。tile 的选择是在数据复用、占用率、指令效率、形状适配之间的折中。

尾块 ​

M,N,K 不是 tile size 整数倍时,最后一块越界。常见处理:加边界判断/掩码;padding 到对齐尺寸;为常见整齐 shape 提供快路径。三种方式的开销结构不同 —— padding 增加计算和存储,分支增加控制开销,哪个更快取决于具体 shape 与硬件。

浮点分块结果可能不同 ​

实数加法满足结合律,浮点加法不严格满足:

(a+b)+c≠a+(b+c)

不同 tile 划分、不同线程归约树、是否走 Tensor Core 路径,都会改变累加顺序,所以优化前后结果可能不逐位一致。正确性验证应使用适合 dtype 与问题规模的误差容限,并关注误差是否随归约长度系统性放大。机制见 01-数值计算与精度。

从数学式到 Kernel 的检查表 ​

看到一个算子,按这个顺序拆:

  1. 输入、输出、中间量的 shape 是什么?
  2. 哪些维度保留,哪些维度归约?
  3. 能否分块?分块状态如何合并?
  4. FLOPs 与理论最小访存量是多少?
  5. 是否存在广播、转置或非连续 stride?
  6. 哪些操作对精度敏感,需要高精度累加?
  7. 是否会出现 exp 溢出、除零、消减或长归约误差?
  8. 中间张量能否融合消除?
  9. 边界 shape 与尾块如何处理?
  10. 用什么参考实现、dtype 容差和极端输入验证?

这份清单把抽象数学变成可执行的工程动作。第 3、6、7 条分别指向后面三块内容:分块状态的合并方式决定能否写成分块 kernel(Online Softmax 就是一例),精度敏感的归约决定 accumulate dtype,数值风险决定要不要用稳定形式。

十问可以分成三组,各自对应一类风险。

   ┌─ 形状组 ──────────────────────────────────────────────┐
   │ ① 输入、输出、中间量的 shape 是什么?                    │
   │ ② 哪些维度保留,哪些维度归约?                           │
   │ ⑤ 是否存在广播、转置或非连续 stride?                     │
   │ ⑨ 边界 shape 与尾块如何处理?                            │
   └───────────────────────────────────────────────────────┘
   ┌─ 代价组 ──────────────────────────────────────────────┐
   │ ③ 能否分块?分块状态如何合并?                           │
   │ ④ FLOPs 与理论最小访存量是多少?                         │
   │ ⑧ 中间张量能否融合消除?                                 │
   └───────────────────────────────────────────────────────┘
   ┌─ 数值组 ──────────────────────────────────────────────┐
   │ ⑥ 哪些操作对精度敏感,需要高精度累加?                    │
   │ ⑦ 是否会出现 exp 溢出、除零、消减或长归约误差?            │
   │ ⑩ 用什么参考实现、dtype 容差和极端输入验证?              │
   └───────────────────────────────────────────────────────┘

   第 3、6、7 条分别指向后面三块内容
     分块状态的合并方式   决定能否写成分块 kernel(Online Softmax 就是一例)
     精度敏感的归约       决定 accumulate dtype
     数值风险             决定要不要用稳定形式

   这份清单把抽象数学变成可执行的工程动作。

相关 ​

参考 ​

贡献者 ​

文件历史 ​