Lecture 9: 在 GPU 上高效评估 DNN(Efficiently Evaluating DNNs on GPUs: Transformers and ConvNets)(日期:Oct 21, 2025)

目录 · ← l8 · l10 →

Lecture 9: 在 GPU 上高效评估 DNN(Efficiently Evaluating DNNs on GPUs: Transformers and ConvNets)(日期:Oct 21, 2025)

概述:本讲把前几讲的所有工具(数据并行、共享内存、SIMD、带宽分析)汇聚到现代 AI 推理上:如何高效调度 DNN 的每一层。核心主线有三条:① 算术强度(arithmetic intensity)与 roofline 模型——程序要么是 compute bound 要么是 bandwidth bound,而”更快更宽的硬件”与”提高算术强度的程序改造”会改变这个平衡;② 把卷积映射为矩阵乘法——explicit GEMM(im2col 物化矩阵)与 implicit GEMM(不物化、在 shared memory 里按块构造),以及分块(blocked/tiled)GEMM 如何通过缓存/共享内存复用提升算术强度;③ 层融合(layer fusion)——conv+scale/bias+maxpool、softmax、transformer 注意力(FlashAttention 式分块 softmax)如何通过融合减少中间数据的 DRAM 往返。最后展望:为什么 GPU 是 DNN 的好平台(张量核、高算术强度、cuDNN),以及为什么它可能不是最优平台(专用硬件 TPU/NPU/Neural Engine)。

注意:本讲对应 Assignment 4: Fused Conv+MaxPool on the Trainium2 Accelerator(Trainium2 卷积+池化融合)——你需要在 AWS Trainium2 上把卷积层与 max pooling 融合实现,正是本讲”层融合”主题的实战。


一、核心概念与定义

1. 算术强度(Arithmetic Intensity)与 Roofline 模型

  • 定义:算术强度 = 每搬运 1 字节数据执行多少次算术操作(Ops/BW)。Roofline 把”机器峰值吞吐(ops/sec)”对”算术强度”画成曲线:算术强度低时吞吐被带宽压住(bandwidth bound 区,斜率为带宽的直线);算术强度高时达到峰值算力(compute bound 区,水平线)。两个推论:① 同样内存系统下把峰值算力提高 → 程序更容易落入带宽受限区;② 提高程序算术强度(程序改造)→ 更容易达到 compute bound。
  • 现实类比:算术强度像”每趟卡车运的料能做出多少件成品”。料运得少(低强度)时,工厂速度取决于卡车(带宽);料运得足(高强度)时,工厂速度取决于机器(算力)。
  • 公式/图示
    Throughput (Ops/sec)
        │                ┌──── compute bound(水平线 = 峰值算力)
        │               ╱
        │              ╱
        │             ╱  bandwidth bound(斜线 = 带宽)
        │            ╱
        └───────────╱───────────────────► Arithmetic Intensity (Ops/BW)
                 (1/4 1/2 1 2 4 8 16)
    程序时间 ≈ max(总操作数/算力, 总字节数/带宽)
    

2. 流水线重叠与双缓冲(Pipelining / Double Buffering)

  • 定义:把”数据加载(load)→ 计算(arithmetic)→ 写回(store)”三个阶段重叠执行:上一块数据在计算时,下一块数据同时在加载。代价是片上存储成本——必须同时持有”正在处理的数据”和”正在传输的数据”两份缓冲,即 double buffering(双缓冲)
  • 现实类比:餐厅后厨”流水线备餐”——A 桌在炒菜时,B 桌的食材已经在切配(切配与炒菜并行),而不是炒完 A 再切 B。代价是需要两张案板(两份缓冲)。
  • 公式/图示
    无重叠:  [load][compute][store] [load][compute][store] ...
    有重叠:  [load0][compute0][store0]
                  [load1][compute1][store1]
                        [load2][compute2][store2] ...
    片上存储 = 正在算的块 + 正在传的块(双缓冲)
    

3. Loop Fusion(循环融合)

  • 定义:把多个遍历同一数据的循环合并成一个循环,提高算术强度——中间结果不再落 DRAM,而是留在寄存器/片上。例:计算 E = D + (A+B)*C,朴素写法是三个独立循环(add、mul、add),每个循环”2 次 load + 1 次 store 对应 1 次运算”(算术强度 1/3);融合成一个循环后是”4 次 load + 1 次 store 对应 3 次运算”(算术强度 3/5)。整体算术强度从 1/3 提升到 3/5。
  • 现实类比:洗菜、切菜、炒菜三件事如果每件都先把菜搬回仓库再取出来做(中间结果落 DRAM),不如在一个厨房里一气呵成——省掉两次”仓库往返”。
  • 公式/图示
    程序1(3 个循环):add(A,B)→tmp1; mul(tmp1,C)→tmp2; add(tmp2,D)→E
         每循环:2 load + 1 store / 1 op  →  算术强度 1/3(总)
    程序2(融合):E[i] = D[i] + (A[i]+B[i])*C[i]
         4 load + 1 store / 3 ops          →  算术强度 3/5
    

4. 卷积层(Convolutional Layer)

  • 定义:全连接层是”每个输出连所有输入”;卷积层是局部连接(每个输出只看输入的一个小窗口,如 3×3)且同一层所有单元共享同一组参数(weights + bias)。卷积核(filter)可看作一个”模式检测器”:输出像素的幅度 = 滤波器对输入局部区域的”响应”(如 Sobel 梯度检测核:水平梯度核 [[-1,0,1],[-2,0,2],[-1,0,1]] 响应水平梯度)。现代 CNN(如 Inception、MobileNet)由大量 Conv + Pool + ReLU 层堆叠。
  • 现实类比:卷积核像”印章”——同一个印章盖到图像的每个位置(共享参数),盖出来的”印痕”(响应图)标出哪些位置出现了印章上的图案。
  • 公式/图示
    output[j][i] = Σ_jj Σ_ii input[j+jj][i+ii] * weights[jj][ii]   (单通道 3×3 卷积)
    多滤波器:输出 W × H × num_filters(每滤波器一张响应图)
    卷积 + ReLU + Pool:W×H×C → W×H×K → W/2×H/2×K(pooling 减半空间分辨率)
    

5. GEMM(General Matrix Multiply,稠密矩阵乘)与 im2col(Explicit GEMM)

  • 定义:GEMM 是 C += A × B 的稠密矩阵乘法,是现代 AI 的”内核中的内核”——全连接层、卷积层、transformer 的注意力块都归结为 GEMM。im2col(explicit GEMM):把卷积”展开”成矩阵乘法——为每个输出位置构造一行输入窗口(拉平成向量),得到 (W×H) 行 × (R×S×C) 列的”卷积矩阵”,再与滤波器矩阵相乘。代价:矩阵带 0-padding、存储开销 O(N)(滤波器有 N 个元素时),且物化该矩阵使 DRAM 流量增加 R×S 倍(3×3 卷积就是 9 倍)——读激活张量来”拼矩阵”本身就要多读 9 遍数据。
  • 现实类比:im2col 像”把滑动窗口的每一帧都截图存档”——每个窗口的照片(矩阵行)都要实际写出来才能交给 GEMM 库;照片数量巨大(每个输出位置一张),占地方、费流量。
  • 公式/图示
    3×3 卷积 = 矩阵乘:
         [w0 w1 ... w8]          [窗口0拉平: 0 0 0 0 x00 x01 0 x10 x11]
         [ ... ]      ×          [窗口1拉平: 0 0 0 x00 x01 x02 x10 x11 x12]
         [ ... ]                 [窗口2拉平: ...                        ]
       num_filters×9              (W×H) 行 × 9 列(每列一个输入通道×空间位置)
    =  (W×H) × num_filters 的响应图
    

6. 分块 GEMM(Blocked / Tiled GEMM)与层次化分块

  • 定义:朴素三重循环 GEMM 算术强度极低(不利用 A、B 的时间局部性);分块(blocking/tiling)把计算组织成 C 的小块:算 C 的一个 BLOCKSIZE_J×BLOCKSIZE_I 子块时,让所需 A、B 子块驻留在缓存里反复复用(假设 BLOCKSIZE 选得足够小)。进一步层次化分块:L2 级块 → L1 级块 → 寄存器级块,逐级匹配内存层级。自检问题:BLOCKSIZE 是不是越大越好?(不是——块要能装进缓存/寄存器文件,太大就装不下、反而驱逐自己。)
  • 现实类比:盖一堵大墙(C),与其每次搬一块砖(一个标量)来回跑仓库,不如把一大片砖(A、B 子块)一次搬到脚边(缓存/shared memory),在脚边砌完一片再搬下一片。
  • 公式/图示
    朴素:  for j for i for k:  C[j][i] += A[j][k]*B[k][i]      ← 每步都从 DRAM 取 A、B
    分块:  for jblock for iblock for kblock:
              for j for i for k(块内): C[jb+j][ib+i] += A[jb+j][kb+k]*B[kb+k][ib+i]
                                        ← A、B 子块在缓存/片上驻留期间被复用 BLOCKSIZE 次
    层次:  jblock2/iblock2/kblock2(L2 级)→ jblock1/...(L1 级)→ 寄存器级(未画出)
    

7. Implicit GEMM(隐式 GEMM,不物化矩阵)

  • 定义:explicit im2col 要物化整个卷积矩阵(DRAM 流量 × R×S、额外存储);implicit GEMM 的改进实现是只把卷积矩阵的一个子块物化在 GPU 片上 shared memory 里,用调好的 shared-memory GEMM 例程(如 CUTLASS)做子块乘法——不需要额外片外存储,也不增加 DRAM 流量。索引仍直接指向原始权重张量和激活张量。
  • 现实类比:不再”每帧截图存档”,而是”投影仪边放边截取当前一帧”——需要用哪一帧,现场从原始视频(激活张量)里截哪一帧,看完即弃,不占相册(DRAM)。
  • 公式/图示
    explicit GEMM:  构造完整卷积矩阵(DRAM 流量 ×9,3×3 卷积)→ 调 GEMM 库
    implicit GEMM:  按块从激活张量现场构造子矩阵 → shared memory 中的子块 GEMM → 写回
                    (无额外片外存储、无额外 DRAM 流量)
    

8. Attention(注意力)与 Softmax

  • 定义:transformer 的注意力块:设 Q、K、V 都是 N×d 矩阵(N = 序列长度,d = 嵌入维度),计算 S = Q·Kᵀ(N×N 分数矩阵),对 S 的每一行做 softmax 得 P,再 O = P·V(N×d 输出)。注意 N 可达数千 → N² 矩阵太大,朴素实现需要 N² 空间(”Trouble!!!”)。softmax 行向量 x:m(x)=max_i x_if(x)=[e^{x1-m},...,e^{xB-m}]l(x)=Σf(x)_isoftmax(x)=f(x)/l(x)(减 max 保证数值稳定)。
  • 现实类比:注意力像”信息检索打分”——Q 是查询、K 是键、V 是值;S 是查询与键的匹配分数,softmax 把分数变成”权重”,O 是值的加权平均。N² 矩阵就像”每对词都写一张评分卡”——长句子时卡片堆满仓库。
  • 公式/图示
    S = Q Kᵀ   (N×d × d×N = N×N)
    P = softmax(S)  (逐行 softmax)
    O = P V     (N×N × N×d = N×d)
    

9. 分块 softmax 与 Fused Attention(FlashAttention 思路)

  • 定义:softmax 可以分块计算:把行向量 x 分成块 x⁽¹⁾、x⁽²⁾,则 m(x) = max(m(x⁽¹⁾), m(x⁽²⁾))f(x) = [e^{m(x⁽¹⁾)−m(x)}·f(x⁽¹⁾), e^{m(x⁽²⁾)−m(x)}·f(x⁽²⁾)]l(x) = e^{m(x⁽¹⁾)−m(x)}·l(x⁽¹⁾) + e^{m(x⁽²⁾)−m(x)}·l(x⁽²⁾)Fused attention(FlashAttention):for each j(Q 块):for each i(K/V 块):加载 Qᵢ、Kⱼᵀ、Vⱼ、Oᵢ 块 → 算 Sᵢⱼ = QᵢKⱼᵀ → 行方向算 Mᵢⱼ、Pᵢⱼ、lᵢⱼ → 把 PᵢⱼVⱼ 按缩放累加进 Oᵢ。效果:从不物化 N² 矩阵(省内存),算术强度高(读 3 个块做 2 次矩阵乘 + 若干行求和,O 块常驻缓存);代价:每步 i 循环要重新缩放之前累加的 O(额外计算)。
  • 现实类比:不再把全部评分卡堆满仓库再统一折算,而是”算一批、折算一批、累加进结果”,仓库里永远只有当前一批卡片——空间 O(N),时间上多花一点”重新折算”的功夫。
  • 公式/图示
    for each j(外层,Q 块):
        for each i(内层,Kᵀ/V 块):
            加载 Q_i, K_jᵀ, V_j, O_i(4 个块)
            S_ij = Q_i K_jᵀ
            M_ij, P_ij, l_ij = 行方向 max/exp/sum(分块 softmax)
            O_i = 缩放后的 O_i + P_ij V_j(按 m/l 重新缩放)
    

10. 层融合(Layer Fusion)

  • 定义:把相邻层的计算融合进同一个循环/kernel,避免中间结果落 DRAM。例:Conv → Scale/Bias → MaxPool 序列:如果分开跑,conv 输出(可高达 1 GB)要先写 DRAM、scale/bias 再读一遍、pool 再读一遍——带宽灾难。融合方案:scale/bias 是逐元素操作,可以在 conv 每算出一个元素后立即执行;max pool 可以在每算完一个 2×2 输出区域后立即取最大值。融合后 DRAM 流量从”N×H×W×K 写 + 2 次读”变成只写 pool 输出 N×H/2×W/2×K。
  • 现实类比:工厂流水线不再把每道工序的半成品全部运回仓库,而是”边加工边传”——每件工件做完前一道立即进下一道,仓库只存最终成品。
  • 公式/图示
    未融合:Conv →(写 N×H×W×K)→ Scale/Bias →(写/读 N×H×W×K)→ MaxPool →(写 N×H/2×W/2×K)
    融合后:Conv + Scale/Bias + MaxPool 一个 kernel:
            每个输出元素算出后立即 scale+bias;每算完 2×2 区域立即取 max
            只写 N×H/2×W/2×K(最终结果)
    

11. 张量核(Tensor Cores)

  • 定义:NVIDIA SM 内的专用矩阵乘单元(第 7 讲结尾预告”数百 TFLOPs 的张量核”,本讲展开其用途):为低精度(如 FP16)矩阵乘提供远超普通 fp32 ALU 的吞吐。注意:本讲主要把它们当作”为什么 GPU 是 DNN 好平台”的证据(V100 有大量 tensor core,需要大量并行工作才能喂饱——N=1、P=Q=64 时输出 524K 元素=2 MB;N=32、P=Q=256 时输出 256M 元素=1 GB)。
  • 现实类比:张量核是”印钞机”(专印矩阵乘这张钞票),普通 ALU 是”点钞机”——印得快,但一次要印一大批才划算(需要大 batch 大矩阵才能喂饱)。
  • 公式/图示D = A×B + C,单条指令完成一个小矩阵块(如 4×4×4 或 8×8×4)的乘加。

12. 低精度(Low Precision)与 DNN 优化三路径

  • 定义:DNN 权重与中间激活常用 16-bit、8-bit 值,正在向 4-bit 推进(极端是 1-bit)。幻灯片总结的优化技术分三类:① 更好的算法(手工设计模型:深度、宽度、滤波器数、stride;以及自动搜索高效拓扑 NAS);② 软件优化(对性能关键操作做好的调度:loop blocking/tiling、fusion——通常人工调优,研究界在努力自动化,如 torch.compile、cuDNN backend、Triton);③ 近似手段(模型压缩:低比特精度)。
  • 现实类比:三条路分别是”换更省的配方(模型)”、”改进做菜工序(调度)”、”用更便宜的食材(低精度)”——可以同时用。
  • 公式/图示:精度 32-bit → 16-bit → 8-bit → 4-bit → 1-bit(带宽与存储需求随位数线性下降)。

二、代码示例与详细解说

示例 1:Loop Fusion 提升算术强度(幻灯片原例)

// program1.c —— 三个独立循环,中间结果落内存
void add(int n, float* A, float* B, float* C) {
    for (int i = 0; i < n; i++)
        C[i] = A[i] + B[i];
}
void mul(int n, float* A, float* B, float* C) {
    for (int i = 0; i < n; i++)
        C[i] = A[i] * B[i];
}
// 计算 E = D + ((A + B) * C)
add(n, A, B, tmp1);      // 每个循环:2 load + 1 store 对应 1 op(算术强度 1/3)
mul(n, tmp1, C, tmp2);   // 每个循环:2 load + 1 store 对应 1 op(算术强度 1/3)
add(n, tmp2, D, E);      // 整体算术强度 = 1/3
// program2.c —— 融合成一个循环
void fused(int n, float* A, float* B, float* C, float* D, float* E) {
    for (int i = 0; i < n; i++)
        E[i] = D[i] + (A[i] + B[i]) * C[i];   // 4 load + 1 store 对应 3 op(算术强度 3/5)
}
// 程序 1 → 程序 2 的变换就叫 loop fusion

【代码做了什么?】

  • 程序 1:三个独立循环各遍历一次数组,tmp1tmp2 两个中间数组完整落内存,每个循环都是”2 次 load + 1 次 store 换 1 次算术”。
  • 程序 2:一个循环内完成 E[i] = D[i] + (A[i]+B[i])*C[i],中间值 A[i]+B[i] 留在寄存器里,不再写/读内存。

【并行机制解说】

  • 算术强度:1/3(程序 1)→ 3/5(程序 2),提升 80%。在 roofline 图上,程序 1 很可能落在带宽受限区,程序 2 可能够到 compute bound 区。
  • 融合的两个前提:① 循环有相同遍历结构(同一 i 域);② 中间结果无跨元素依赖tmp1[i] 只被 multmp1[i] 使用——逐元素独立)。这正是第 8 讲”理解依赖”的直接应用。
  • 融合的代价:无额外片上存储需求(中间值就在寄存器里),比双缓冲还便宜;但对”循环体变复杂”的 kernel 要小心寄存器压力。
  • 对应概念:loop fusion、算术强度、roofline、bandwidth bound vs compute bound

示例 2:im2col——把卷积展开成 GEMM(伪代码)

// im2col_pseudo.cpp —— 伪代码:把 3×3 卷积映射为矩阵乘(explicit GEMM)
// 输入:input  W×H 单通道图像;weights  num_filters×9(每个滤波器 9 个权重)
// 输出:output W×H×num_filters
// 设 im2col_matrix 为 (W*H) 行 × 9 列的矩阵,每行是"以某输出像素为中心的 3×3 窗口拉平"

// Step 1: 构造 im2col 矩阵(关键:这是"物化"步骤,需要 0-padding 越界窗口)
for (int outIdx = 0; outIdx < W * H; outIdx++) {
    int j = outIdx / W, i = outIdx % W;          // 输出像素 (i, j)
    for (int jj = 0; jj < 3; jj++)
        for (int ii = 0; ii < 3; ii++) {
            int sy = j + jj - 1, sx = i + ii - 1; // 输入坐标(-1 表示 padding 边界)
            im2col_matrix[outIdx][jj * 3 + ii] =
                (sy >= 0 && sy < H && sx >= 0 && sx < W) ? input[sy * W + sx] : 0.f;
        }
}
// Step 2: 一次 GEMM 完成所有输出:output_matrix = im2col_matrix × weightsᵀ
//         ((W*H)×9  ×  9×num_filters  =  (W*H)×num_filters)
//         可以直接调用任何调优 GEMM 库(BLAS/cuBLAS)

// 多输入通道版:im2col 每行长度变为 9×C(3×3×C),weights 变为 num_filters×(9×C)
// 批量版:batch 内每张图都构造一份 im2col 矩阵

【代码做了什么?】

  • Step 1 把”滑动窗口”物化为行向量:每个输出像素对应一行 9 个(或 9×C 个)元素,越界位置填 0(padding)。这是 im2col 的全部”魔法”——卷积的窗口结构被编码进矩阵布局里。
  • Step 2 把”9 次乘加”变成”9 维行向量 × 权重矩阵列”的点积——整层卷积变成一次标准 GEMM,直接复用调优矩阵乘库。

【并行机制解说】

  • 并行性:GEMM 有海量并行度(每个输出元素独立),这正是 GPU 需要的(对应第 8 讲”暴露大量并行度”)。
  • 代价(本讲强调):物化矩阵的 DRAM 流量是输入的 R×S 倍(3×3 → 9 倍)——构造 im2col 矩阵时要反复读激活张量的重叠窗口;还占用大量额外存储。这就是示例 3 引入 implicit GEMM 的动机。
  • 对应概念:im2col / explicit GEMM、卷积→矩阵乘映射、带宽开销

示例 3:分块(tiled)GEMM——CUDA shared memory 实现

// tiled_gemm.cu —— 编译:nvcc -o tiled_gemm tiled_gemm.cu
// 计算 C = A*B(M×N = M×K · K×N),用 shared memory 分块复用
#include <cuda_runtime.h>

#define TILE 16   // 每个 block 计算 C 的一个 TILE×TILE 子块

__global__ void tiledGemm(const float* A, const float* B, float* C,
                          int M, int N, int K)
{
    __shared__ float As[TILE][TILE];   // A 子块(片上)
    __shared__ float Bs[TILE][TILE];   // B 子块(片上)

    int row = blockIdx.y * TILE + threadIdx.y;   // C 子块内的行
    int col = blockIdx.x * TILE + threadIdx.x;   // C 子块内的列

    float acc = 0.0f;                            // 寄存器累加器(私有)

    for (int kk = 0; kk < K; kk += TILE) {
        // 1) 协作加载 A、B 子块到 shared memory
        As[threadIdx.y][threadIdx.x] = A[(blockIdx.y * TILE + threadIdx.y) * K + kk + threadIdx.x];
        Bs[threadIdx.y][threadIdx.x] = B[(kk + threadIdx.y) * N + blockIdx.x * TILE + threadIdx.x];
        __syncthreads();                         // 2) 屏障:确保子块加载完成

        // 3) 块内小 GEMM:累加 As×Bs
        for (int k = 0; k < TILE; k++)
            acc += As[threadIdx.y][k] * Bs[k][threadIdx.x];
        __syncthreads();                         // 4) 屏障:防止覆盖 As/Bs 前有人还在读
    }

    C[row * N + col] = acc;                      // 写回全局
}
// host 端:dim3 block(TILE,TILE); dim3 grid(N/TILE, M/TILE);
//          tiledGemm<<<grid, block>>>(dA, dB, dC, M, N, K);

【代码做了什么?】

  • 每个 block 负责 C 的一个 TILE×TILE 子块;外层 kk 循环沿 K 维滑动,每轮把 A、B 的相应子块协作装入 shared memory,块内线程各自累加一个输出元素(acc 在寄存器里)。
  • 两次 __syncthreads():装完子块后(防止读到旧数据)与子块用完后(防止下一轮覆盖时还有人没读完)。

【并行机制解说】

  • 算术强度:朴素三重循环每个 C[j][i] 都从 DRAM 取 A[j][k]、B[k][i](每 1 次乘加 ~2 次 DRAM 访问);分块后,每块 As、Bs 在 shared memory 驻留期间被复用 TILE 次——DRAM 访问量 ÷ TILE。这就是”compute partial result for block of C while required blocks of A and B remain in cache”(幻灯片语)在 GPU 上的直译:shared memory 就是 GPU 的”cache”。
  • 自检题:BLOCKSIZE(TILE)越大越好吗?不是——块必须装得进 shared memory(还有寄存器压力),太大放不下,且会降低每 SM 可驻留的 block 数。
  • 完整工程里还有层次化分块(L2 → L1 → 寄存器)与向量化变体(splat + muladd 向量化 i 循环;预转置 B 以便向量化最内层 k 循环;以及寄存器块 C_accum[SIMD_WIDTH] 同时向量化 j、i 两维——幻灯片 37-39 页三种方案)。当 i 维很小时方案 1 不好用,需要预转置(方案 2);方案 3 假设 A、C 也预转置,在 SIMD_WIDTH×SIMD_WIDTH 的寄存器块上做乘加。
  • 对应概念:分块/tiling、算术强度、shared memory、多级内存层级

示例 4:注意力分数计算——朴素版与融合(分块 softmax)版

# attention_naive.py —— 朴素注意力:物化 N×N 矩阵(N 大时内存爆炸)
import numpy as np
N, d = 1024, 64
Q = np.random.randn(N, d); K = np.random.randn(N, d); V = np.random.randn(N, d)

S = Q @ K.T                    # N×N 分数矩阵(N=1024 时 4 MB;N=10万时 40 GB!)
m = S.max(axis=1, keepdims=True)
P = np.exp(S - m)              # f(x):减行最大值(数值稳定)
l = P.sum(axis=1, keepdims=True)
P = P / l                      # softmax 归一化
O = P @ V                      # N×d 输出
# attention_fused.py —— 融合(FlashAttention 思路):按块扫描,从不物化 N×N
import numpy as np
N, d, BLOCK = 1024, 64, 128
Q = np.random.randn(N, d); K = np.random.randn(N, d); V = np.random.randn(N, d)

O = np.zeros((N, d))
m_prev = np.full(N, -np.inf)   # 每行的 running max
l_prev = np.zeros(N)           # 每行的 running sum of exp

for j in range(0, N, BLOCK):           # Q 块
    Qb = Q[j:j+BLOCK]                  # BLOCK×d
    for i in range(0, N, BLOCK):       # K/V 块
        Kb, Vb = K[i:i+BLOCK], V[i:i+BLOCK]
        Sij = Qb @ Kb.T                # BLOCK×BLOCK 分数块
        m_ij = Sij.max(axis=1)         # 块的 row-max
        Pij = np.exp(Sij - m_ij[:, None])          # 块内 f(x)
        l_ij = Pij.sum(axis=1)                     # 块内 l(x)
        # 用新的 m 重新缩放已累加的 O 与 l,再并入本块
        m_new = np.maximum(m_prev[j:j+BLOCK], m_ij)
        alpha = np.exp(m_prev[j:j+BLOCK] - m_new)  # 旧部分的缩放
        beta  = np.exp(m_ij - m_new)               # 新块的缩放
        O[j:j+BLOCK] = O[j:j+BLOCK] * alpha[:, None] + (Pij * beta[:, None]) @ Vb
        l_prev[j:j+BLOCK] = l_prev[j:j+BLOCK] * alpha + l_ij * beta
        m_prev[j:j+BLOCK] = m_new
O = O / l_prev[:, None]                # 最终归一化

【代码做了什么?】

  • 朴素版:一次性算出 N×N 的 S,逐行 softmax,再乘 V。三步各读写一次 N×N 矩阵——内存 O(N²) 且算术强度低(每步整个矩阵从 DRAM 进出)。
  • 融合版:外层 j 循环遍历 Q 块,内层 i 循环遍历 K/V 块;每步只物化 BLOCK×BLOCK 的 Sᵢⱼ 块,立刻算 M/P/l 并按新 max 重新缩放已累加的 O 与 l,然后并入 PᵢⱼVⱼ。循环结束用 running l 归一化。

【并行机制解说】

  • 分块 softmax 的数学依据(幻灯片 63 页):m(x)=max(m(x⁽¹⁾),m(x⁽²⁾))f(x)=[e^{m(x⁽¹⁾)−m(x)}f(x⁽¹⁾), e^{m(x⁽²⁾)−m(x)}f(x⁽²⁾)]l(x)=e^{m(x⁽¹⁾)−m(x)}l(x⁽¹⁾)+e^{m(x⁽²⁾)−m(x)}l(x⁽²⁾)——所以 softmax 可以在块间”增量”计算。
  • 收益:① 内存:从不物化 N² 矩阵(省到 O(N));② 带宽/算术强度:每个 i 步读 3 个块(Q、K、V)、做 2 次矩阵乘 + 若干行求和,O 块常驻缓存——高算术强度、低 DRAM 往返。
  • 代价:每步 i 循环必须重新缩放之前累加的 O(比朴素版多的计算量);以及需要保存每行的 running m/l。
  • GPU 上这正是 FlashAttention 的核心结构(幻灯片以 Thunderkittens 实现为例)。
  • 对应概念:attention、分块 softmax、fused attention / FlashAttention、算术强度

示例 5:融合 Conv + Scale/Bias(+ MaxPool 思路)

// fused_conv_scalebias.c —— 把 scale/bias 融合进卷积循环(幻灯片 57 页)
// 未融合时:Conv 输出整层后,Scale/Bias 再遍历一遍(中间结果落内存)
float input[IMAGE_BATCH_SIZE][INPUT_HEIGHT][INPUT_WIDTH][INPUT_DEPTH];
float output[IMAGE_BATCH_SIZE][INPUT_HEIGHT][INPUT_WIDTH][LAYER_NUM_FILTERS];
float layer_weights[LAYER_NUM_FILTERS][LAYER_CONVY][LAYER_CONVX][INPUT_DEPTH];
float scale[LAYER_NUM_FILTERS], bias[LAYER_NUM_FILTERS];

// 假设卷积 stride = 1
for (int img = 0; img < IMAGE_BATCH_SIZE; img++)        // 所有 batch 图像
    for (int j = 0; j < INPUT_HEIGHT; j++)
        for (int i = 0; i < INPUT_WIDTH; i++)           // 所有输出像素
            for (int f = 0; f < LAYER_NUM_FILTERS; f++) {  // 所有输出通道
                float tmp = 0.0f;
                for (int kk = 0; kk < INPUT_DEPTH; kk++)    // 累加所有输入通道响应
                    for (int jj = 0; jj < LAYER_FILTER_Y; jj++)
                        for (int ii = 0; ii < LAYER_FILTER_X; ii++)
                            tmp += layer_weights[f][jj][ii][kk]
                                 * input[img][j+jj][i+ii][kk];
                output[img][j][i][f] = tmp * scale[f] + bias[f];  // ← 融合:算完即 scale+bias
            }

// 课堂练习(幻灯片):如何再融合紧随其后的 max pool(取输出矩阵 2×2 块的最大值)?
// 提示:把黄色循环(i、j)按 2×2 分块——每个 2×2 输出块算完后立即取 max 写 pool 结果,
//       而不是先把整层输出写内存再读回来 pool。

【代码做了什么?】

  • 七重循环的”直接实现”(batched conv):每个输出元素 = 所有输入通道上 3×3 窗口的加权和(共享滤波器权重,局部连接)。
  • 融合点:output[...][f] = tmp * scale[f] + bias[f]——scale/bias 是逐元素操作,紧跟在每个 tmp 算完之后,tmp 还在寄存器里就完成,根本不写中间层

【并行机制解说】

  • 为什么必须融合:幻灯片算过账——conv 输出可达 1 GB(N=32、P=Q=256 情形),”dumping 1 GB to memory and reading it back just to scale, then rereading to pool”是灾难。融合后 DRAM 只写最终 pool 输出(N×H/2×W/2×K),读写量减少约一个数量级。
  • maxpool 融合方法(课堂练习答案):把 i、j 循环按 2×2 分块——每算完一个 2×2 的 conv 输出区域,立即在寄存器里取 max 写入 pool 输出;这样中间 conv 输出永远不出片上(这正是 Assignment 4 的任务)。
  • 与示例 1 的 loop fusion 一脉相承,只是融合对象从”数组级”变成”张量层级”。现代框架里:TensorFlow 曾把少数融合算子写死;cuDNN backend 由编译器现场生成融合实现(无运行时开销、中间结果不经内存);torch.compile 等编译器在自动做类似调度。
  • 对应概念:layer fusion、带宽、convolutional layer、算术强度

三、关键要点

  1. “上一张幻灯片你就知道了软件侧性能优化的几乎一切”:程序是 compute bound 还是 bandwidth bound,取决于机器算力、带宽与程序算术强度的对比;重叠通信与计算需要双缓冲(片上存储成本);数据移动耗能、片上存储资源与计算资源此消彼长——现代 AI 优化的软件侧核心就是提高算术强度、减少数据移动
  2. 卷积的两种 GEMM 化路线:explicit GEMM(im2col)物化矩阵、直接调库,但 DRAM 流量 ×R×S、存储爆炸;implicit GEMM 只把子块物化在 shared memory,用 CUTLASS 等调优子块 GEMM——不增加 DRAM 流量。DNN 层尺寸千差万别(MobileNet 的 1×1/3×3 dw/逐点卷积、Inception 多分支),没有银弹库,所以 CUTLASS/Triton/Thunderkittens/NKI 这类”可编程原语”层至关重要。
  3. 分块(blocking)是提升算术强度的通用手术刀:朴素 GEMM 每步都访问 DRAM;分块让 A、B 子块在缓存/shared memory 里复用 BLOCKSIZE 次;层次化分块逐级匹配 L2 → L1 → 寄存器。自检:BLOCKSIZE 不是越大越好——块必须能驻留(cache 容量、shared memory 容量、寄存器压力)。
  4. 融合 = 消除中间数据的 DRAM 往返:conv+scale/bias+maxpool、softmax 逐行、FlashAttention 分块 softmax——共同点都是”中间结果留在片上,算完即用”;FlashAttention 还展示了分块数学(增量 max/exp/sum 重缩放)如何让”不物化 N² 矩阵”成为可能,代价是额外的重缩放计算。
  5. GPU 是 DNN 的好平台但未必最优:高算术强度的矩阵乘正对 GPU 的”flop 富矿”(5120 个 fp32 ALU + 张量核),且有 cuDNN 等高度优化的 kernel 库;但”通用处理器真的需要吗?”——TPU、NPU、Neural Engine、IPU 等专用硬件正是下节课(专门化加速)的主题。低精度(16/8/4/1-bit)是另一条通用优化路径。

四、常见陷阱与注意事项

  1. 只优化算力不优化带宽(或反之):在带宽受限区盲目加计算资源毫无收益——roofline 告诉你要先看程序在哪个区。幻灯片反问:”这是 compute bound 还是 BW bound?”每次优化前先回答这个问题。
  2. BLOCKSIZE 贪大:分块 GEMM 的块必须能放进 cache/shared memory/寄存器文件;太大放不下、太小复用不足。层次化分块要逐级匹配内存层级,且寄存器级的分块(最内层)往往被忽略。
  3. im2col 的隐性成本:物化卷积矩阵使 DRAM 流量 ×R×S(3×3 就是 9 倍)并需要大额存储——如果直接拿 explicit GEMM 实现卷积而不考虑这一点,性能会被带宽拖垮;implicit GEMM 才是生产级选择。
  4. softmax/attention 的数值与空间陷阱:① 不减行 max 直接 exp 会溢出(数值稳定性);② 朴素实现物化 N×N 矩阵,长序列时空间 O(N²) 直接爆内存;③ 融合实现里”重缩放旧 O”这步最容易写错(必须用新 m 重缩放,而不是直接用旧 m)。
  5. 忘记融合的带宽账:分开实现 Conv → Scale/Bias → MaxPool 时,1 GB 级中间张量被写了又读、读了又写——看似每层都”简单”,实际带宽开销是数量级的。作业里最容易犯的错是”先把 conv 输出存下来再 pool”,而正确的做法是每算完 2×2 区域就地取 max。
  6. 忽视 DNN 层的多样性:不同层的矩阵维度天差地别(MobileNet 的 1×1 逐点卷积是纯 GEMM、3×3 depthwise 卷积几乎没有通道复用、FC 层 1024×1000……),”一种调度打天下”不成立——这也是 cuDNN 提供多种算法(direct/implicit gemm/winograd 等)的原因。
  7. 数值精度想当然:低精度(FP16/INT8)会改变结果;张量核对低精度有巨大吞吐优势,但精度-性能权衡需要验证(且不同硬件支持不同)。

五、思考题(带答案)

Q1:把卷积映射为矩阵乘法时,为什么 implicit GEMM 比 explicit GEMM(im2col)好?”不物化矩阵”到底省了什么?

A1:explicit GEMM 为了把卷积喂给通用 GEMM 库,必须先把”卷积矩阵”((W×H) 行 × (R×S×C) 列)整体物化到内存——每个输入元素会被复制到 R×S 个不同的窗口行里,因此读激活张量的 DRAM 流量放大 R×S 倍(3×3 卷积 = 9 倍),还要占用巨额额外存储(尤其大 batch 时输出可达 1 GB 级)。implicit GEMM 不构造完整矩阵,而是每次只在 GPU 片上 shared memory 里物化一个子块(索引仍直接指向原始权重/激活张量),用调优的 shared-memory GEMM(CUTLASS)做子块乘法——子块来自 DRAM 的流量与直接卷积相同,不增加任何片外流量,也不需要额外片外存储。省下的是”9 倍读放大 + 大矩阵存储 + 构造矩阵的额外遍历”。

Q2:FlashAttention 风格的融合注意力,为什么”多做了计算”却仍然更快?它多做了哪些计算?

A2:分块 softmax 要求每并入一个新块 Sᵢⱼ 时,用新出现的行最大值 m_new 重新缩放已经累加进 O 的所有旧块(乘以 e^{m_prev−m_new}),同时 running l 也要同步重缩放——这就是”额外的计算”,朴素版没有这步(因为它一次看到整行)。但换来的是:① 空间从 O(N²) 降到 O(N)(从不物化 N×N 分数矩阵,长序列不再爆内存);② 算术强度大幅提升——每个 i 步只读 Q、K、V 三个块、做两次矩阵乘(QᵢKⱼᵀ 与 PᵢⱼVⱼ)+ 少量行归约,O 块常驻缓存,DRAM 往返从”每步读写整个 N×N 矩阵”降到”每步读写 3 个 N×d 块”。在 GPU 上(带宽是稀缺资源、矩阵乘有张量核)”多算一点、少搬很多”几乎总是净赢——这正是 FlashAttention 的原理。

Q3:为什么要对 DNN 做”层融合”(如 Conv+Scale/Bias+MaxPool)?请用 roofline 的视角解释,并说明 Assignment 4 里你会如何实现 maxpool 的融合。

A3:单独执行每一层时,每层都把自己的输出完整写 DRAM、下一层再完整读回来——中间张量(N×H×W×K,可到 1 GB)被反复搬运,而这些层的算术强度极低(scale/bias 每元素 1 次运算、pool 每 4 元素 1 次取 max),在 roofline 上必然落在带宽受限区,吞吐被带宽摁死。融合后中间结果不落 DRAM,只有最终输出(N×H/2×W/2×K)写内存——DRAM 流量减一个数量级,算术强度大幅提升,程序向 compute bound 区移动。maxpool 融合的实现:把输出像素循环按 2×2 分块,每个线程(或每组线程)算完 2×2 四个 conv 输出后,在寄存器里取 max 直接写 pool 输出——conv 的中间 2×2 块从不离开片上。这与示例 1 的 loop fusion 和示例 3 的分块是同一思想在不同层级的应用:让数据在最近的存储层级被复用