跳到正文

GPU 与 Triton 入门:从存储层级到第一个 kernel

GPU 的存储层级、SM 结构、occupancy 和 CUDA 执行层级;再用 GELU、softmax、row sum、matmul 四个例子入门 Triton,附录对照 TPU。GPU training optimization 系列的第一篇。

这个系列讲怎么让 Transformer 训练在 GPU 上跑得更快。这一篇先打基础:第 1 节介绍 GPU 的硬件,包括数据存在哪几层、一个 SM 里有什么、CUDA 的执行层级怎么对应到硬件;第 2 节用四个例子介绍 Triton,后面两篇的 kernel 都用它写。TPU 和 GPU 的对照放在附录 A。

下一篇谁偷走了 5090 的算力和显存用这些知识在 RTX 5090 上拆解一步训练的时间和显存,第三篇 FlashAttention 1–4 讲 attention 的 IO-aware kernel。

环境:RTX 5090 32 GB,torch 2.11.0+cu130,Triton 3.6.0。


1 GPU:越靠近计算单元越快
#

做性能分析时最常用的两个数是峰值算力和显存带宽。一个 kernel 受哪个限制,取决于它要搬多少数据、数据在哪一层被读写。这一节按这个顺序看 5090 的硬件:先看数据存在哪几层(1.1),再看负责计算的 SM 里有什么(1.2)。这两节只讲硬件。最后看软件:CUDA 的 thread、warp、block 怎么落到这些硬件上,一个 SM 上同时驻留多少 warp,怎么在等数据时不让计算单元闲着(1.3)。

1.1 存储层级:越近越快,也越小
#

处理器的频率早就不怎么涨了,算力的增长主要来自并行:更多的 SM、更宽的 Tensor core。显存带宽的增长比算力慢得多:过去 20 年,峰值算力涨了约 6 万倍(每两年 3.0 倍),DRAM 带宽只涨了约 100 倍(每两年 1.6 倍),芯片之间的互联带宽约 30 倍(图 1-1)。所以越来越多的 op 落在 roofline 的斜坡上,这是 memory-bound 越来越常见的原因。带宽跟不上,就只能少搬:让读进来的数据留在离计算单元近的地方,多用几次。

1997 到 2023 年峰值算力、DRAM 带宽和互联带宽的增长:算力每两年 3.0 倍,DRAM 带宽每两年 1.6 倍,互联带宽每两年 1.4 倍
图 1-1 峰值算力与 DRAM、互联带宽 20 年来的增长,纵轴是相对 1997 年的倍数(对数轴)。图来自 Gholami 等人的 AI and Memory Wall(IEEE Micro,2024)。

数据离计算单元越近,读写越快,容量也越小(图 1-2)。每个 SM 里有 register 和 L1 / shared memory,所有 SM 共享芯片上的 L2,显存在芯片外面。写 kernel 时,L1 和 L2 由硬件当作 cache 自动管理;能自己安排的只有 shared memory(以及 register)。所以 kernel 优化的套路都是一样的:把一块数据从显存读进 shared memory 或 register,在片上尽量多算几次,再写回去。第 2 节的分块矩阵乘、下一篇里的算子融合和之后的 FlashAttention 都是这个思路。

GPU 芯片(die),共 170 个 SMSMregister256 KBL1 / shared128 KBSMregister256 KBL1 / shared128 KBSMregister256 KBL1 / shared128 KB⋯SMregister256 KBL1 / shared128 KBL2 cache 96 MB,所有 SM 共享片上GDDR7 显存32 GB带宽 β1.79e12 B/s片外越靠近 SM 越快、越小:register > L1 / shared > L2 > 显存
图 1-2 RTX 5090 的存储层级示意。规格来自 NVIDIA RTX Blackwell 白皮书,官方的整芯片和 SM 结构图也在白皮书里。

5090 是消费级显卡,显存用的是 GDDR7(32 GB,512-bit,28 Gbps,带宽 1,792 GB/s),不是数据中心卡上的 HBM,也没有 ECC 和 NVLink。

register 和 L1 / shared memory 都在 SM 里面,下面看 SM 本身的结构。

1.2 一个 SM 里有什么
#

每个 SM 分成 4 个 SMSP(SM sub-partition),每个 SMSP 有自己的 warp 调度器、register、32 个 CUDA core(FP32 运算单元)和一个 Tensor core,4 个 SMSP 共享一块 L1 / shared memory(图 1-3)。几代 GPU 的规格对比见表 1-1。

SM(每颗 5090 有 170 个)SMSP 0warp 调度器register 64 KBCUDA core ×32Tensor core ×1SMSP 1warp 调度器register 64 KBCUDA core ×32Tensor core ×1SMSP 2warp 调度器register 64 KBCUDA core ×32Tensor core ×1SMSP 3warp 调度器register 64 KBCUDA core ×32Tensor core ×1L1 cache / shared memory 128 KB(4 个 SMSP 共享)
图 1-3 RTX 5090 一个 SM 的结构示意(SMSP 即 SM sub-partition)。每个 SMSP 有自己的 warp 调度器、register、32 个 CUDA core 和 1 个 Tensor core,4 个 SMSP 共享 L1 / shared memory。
表 1-1 几代 NVIDIA GPU 的规格
A100H100B200RTX 5090
SM 数108132148170
L240 MB50 MB—96 MB
显存80 GB HBM2e80 GB HBM3192 GB HBM3e32 GB GDDR7
显存带宽2.0e12 B/s3.35e12 B/s8e12 B/s1.79e12 B/s
每 SM 的 CUDA core(FP32)64128128128
每 SM 的 Tensor core4444
每 SM 的 L1 + shared192 KB256 KB256 KB128 KB
每 SM 的 register256 KB256 KB256 KB256 KB
每 SM 的 SMSP(warp 调度器)4444

来自 NVIDIA 各代架构白皮书和产品规格(A100 80GB SXM,H100 SXM)。B200 的 L2 没有找到可靠的公开数字,暂缺。

这些是硬件提供的资源。register 按 thread 分配,shared memory 按 block 分配,要知道一个 kernel 怎么用这些资源,得先看 CUDA 怎么把线程组织起来。

1.3 编程模型:thread、warp、block、grid
#

CUDA 的执行模型叫 SIMT(single instruction, multiple threads):同一个 warp 里的 32 个线程在同一时刻执行同一条指令。线程分四层:thread;warp,32 个 thread;block(也叫 CTA),若干个 warp;grid,一次 kernel 启动的全部 block。图 1-4 把四层画在一起,左边逐层放大,右边是每层跑在哪个硬件上。

软件:一次 kernel 启动硬件:跑在哪里(RTX 5090)整张 GPUblock 由硬件分到 170 个 SM一个 SMblock 整个驻留,不跨 SM一个 SMSP调度器每拍发一条指令一条 laneCUDA core,自己的 registergrid:所有 block,互相独立B0B1B2B3B4B5B6⋯block(Triton 的 program)shared memory:块内共享,可以同步warp 0warp 1warp 2warp 3warp:32 个 thread,同一时刻执行同一条指令thread:自己的 register 和下标
图 1-4 CUDA 的执行层级。左边从上到下逐层放大:grid 里的一个 block、block 里的一个 warp、warp 里的一个 thread;右边是每一层跑在哪个硬件上,对应图 1-3。

block 放到 SM 上以后就一直驻留,直到它的线程全部跑完。一个 SM 能同时驻留多少,有几条硬件上限(表 1-2)。warp 的上限分摊到每个 SMSP 的调度器;block 的上限记在整个 SM 上,因为一个 block 的 shared memory 和同步由 4 个 SMSP 共用。

表 1-2 一个 SM 的驻留上限
A100H100B200RTX 5090
每 SMSP 最多驻留 warp16161612
每 SM 最多驻留 warp64646448
每 SM 最多驻留 block32323224
每 thread 最多 register255255255255

来自 CUDA Programming Guide 的 compute capability 表(8.0、9.0、10.0、12.0)。

一个 block 的 warp 按编号轮流分到 4 个 SMSP:warp 0 到 SMSP 0,warp 1 到 SMSP 1,依此类推。驻留不等于在跑。一个 warp 等显存数据要几百个周期,这段时间里调度器改发另一个就绪的 warp;所有驻留 warp 的 register 一直留在 register file 里,换 warp 不用保存或恢复状态。所以一个 SM 上同时驻留的 warp 越多,访存延迟越容易藏住。图 1-5 画的是驻留 3 个 block、每个 4 个 warp 时的某一拍:每个 SMSP 有 3 个 warp,分别来自 3 个 block,调度器每个周期只发射其中一个;SMSP 2 的 3 个都在等数据,这一拍就空转。

一个 SM 的驻留情况整个 SM 最多驻留 24 个 block、48 个 warpSMSP 0warp 调度器最多驻留 12 个 warpB0B1B2SMSP 1warp 调度器最多驻留 12 个 warpB0B1B2SMSP 2空转warp 调度器最多驻留 12 个 warpB0B1B2SMSP 3warp 调度器最多驻留 12 个 warpB0B1B2B0、B1、B2 各有 4 个 warp:warp 0 到 3 依次分到 SMSP 0 到 3正在发射就绪,没被选中等数据空槽
图 1-5 一个 SM 上驻留 3 个 block(B0、B1、B2)时的某一拍。格子里写的是这个 warp 属于哪个 block,每个 SMSP 有 12 个格子,整个 SM 共 48 个。每个调度器每拍最多发射一个 warp;SMSP 2 的 3 个 warp 都在等数据,这一拍空转。

驻留几个 warp 才够、在 Triton 里怎么调,放到 2.5 节 写 kernel 时再看。

block 之间不能直接共享数据,只能通过显存交换,而这正是最慢的一层。所以要尽量把会重复读的数据放进同一个 block 里处理,这就是分块(tiling)。在 CUDA 里分块要自己管 shared memory 和线程同步,第 2 节的 Triton 把这部分交给编译器。


2 Triton:按 block 写 kernel
#

Triton 按 block 写 kernel,正好对应 1.3 节说的分块。CUDA 要写清楚每个线程做什么,控制最细,但 shared memory、线程同步这些都要自己管。Triton 只要写清楚每个线程块做什么:把一块数据读进来,在片上算完,再写回显存,块内怎么分给线程、要不要经过 shared memory 由编译器决定。Triton 把一个线程块叫作一个 program,概念对应见表 2-1。

表 2-1 CUDA 概念在 Triton 里的对应
CUDATriton写法能否直接控制
thread不暴露写不到 threadIdx不能
warp只给数量启动时传 num_warps=N,默认 4只能调数量
block(CTA)program@triton.jit 函数体就是一个 program 的代码;tl.program_id(axis) ≈ blockIdx,tl.num_programs(axis) ≈ gridDim主要的编程层
gridgridkernel[grid](...),例如 grid = (triton.cdiv(n, BLOCK),)自己定
shared memory由编译器分配没有对应语句不能

下面四个例子由浅入深:逐元素的 GELU,一行一个 program 的 softmax,行太长时分块累加的 row sum,以及分块乘加的矩阵乘;最后一节看每个 program 该用几个 warp。代码都在 RTX 5090 上和 PyTorch 对照过。下一篇手写的融合 RMSNorm 和之后的 FlashAttention 也都用 Triton 写。

2.1 GELU:逐元素 kernel
#

grid一个 program 做的事段 0段 1 ← pid=1段 2⋯grid = (cdiv(n, BLOCK),)一段一个 program一个 program(pid = 1):处理第 1 段的 1024 个元素算下标offsets读tl.load(mask)GELUtanh 近似写回tl.store(mask)
图 2-1 GELU kernel。左边是 grid:n 个元素按 BLOCK 切段,一段一个 program;右边虚线框是其中一个 program 做的事。

逐元素 op 最简单,每个元素各算各的,program 之间不需要配合。把 n 个元素切成每段 BLOCK 个,一段交给一个 program(图 2-1)。kernel 里先用 tl.program_id 算出自己负责的下标,读进来,算完写回:

@triton.jit
def gelu_kernel(x_ptr, y_ptr, n, BLOCK: tl.constexpr):
    pid = tl.program_id(axis=0)               # 第几个 program(≈ CUDA 的 blockIdx.x)
    offsets = pid * BLOCK + tl.arange(0, BLOCK)  # 这个 program 负责的 BLOCK 个下标
    mask = offsets < n                        # 最后一个 program 可能越界
    x = tl.load(x_ptr + offsets, mask=mask)   # 从显存读到 register
    # tanh 近似:0.5·x·(1 + tanh(√(2/π)·(x + 0.044715·x³))),tanh(a) = (e^{2a} − 1)/(e^{2a} + 1)
    a = 0.79788456 * (x + 0.044715 * x * x * x)
    e = tl.exp(2 * a)
    y = 0.5 * x * (1 + (e - 1) / (e + 1))
    tl.store(y_ptr + offsets, y, mask=mask)   # 写回显存

启动时只要定 grid,也就是 program 的个数。program 个数和 SM 个数无关:5090 有 170 个 SM,program 多于 SM 时,GPU 会一批一批地调度。

def gelu(x):
    y = torch.empty_like(x)
    n, BLOCK = x.numel(), 1024                # 每个 program 处理 1024 个元素
    grid = (triton.cdiv(n, BLOCK),)           # program 个数,不是 SM 个数
    gelu_kernel[grid](x, y, n, BLOCK=BLOCK)
    return y

Triton 会把 kernel 编译成 PTX,PTX 描述的是单个线程的指令。下面是 GELU 的 PTX 节选:

mov.u32        %r17, %ctaid.x;                // program 编号,即 blockIdx.x
mov.u32        %r20, %tid.x;                  // 线程编号
@%p1 ld.global.b32 { %r1 }, [ %rd1 + 0 ];     // 读显存,%p1 是 mask 算出的谓词
...                                           // 一共 8 条 ld.global
ex2.approx.f32 %r72, %r71;                    // tl.exp 编成以 2 为底的指数
@%p1 st.global.b32 [ %rd9 + 0 ], { %r9 };     // 写回显存

一个 program 处理 1024 个元素,默认 4 个 warp 共 128 个线程,所以每个线程分到 8 个元素,PTX 里正好是 8 条 ld.global 和 8 条 st.global。

GELU 的每个元素互不依赖,怎么切都可以。softmax 的每个元素都要用到整行的最大值和总和,切分就得按行来。

2.2 Softmax:一行一个 program
#

softmax 在 eager 下要 5 个 kernel,和下一篇里 attention 的 softmax 一样。对 [M, N] 的输入,一共读 5MN + 2M 个数,写 3MN + 2M 个:

def softmax_naive(x):                         # x: [M, N],按行
    m = x.max(dim=1, keepdim=True).values     # 读 MN,写 M
    z = x - m                                 # 读 MN + M,写 MN
    e = torch.exp(z)                          # 读 MN,写 MN
    s = e.sum(dim=1, keepdim=True)            # 读 MN,写 M
    return e / s                              # 读 MN + M,写 MN
grid一个 program 做的事行 0行 1 ← pid=1行 2⋯grid = (M,)一行一个 program一个 program(pid = 1):处理第 1 行,中间结果都在 register 里读一行tl.load减最大值x − tl.max(x)exp、求和tl.exp, tl.sum归一化写回tl.store
图 2-2 softmax kernel,画法同图 2-1。一行一个 program,整行读进来以后,中间结果不回显存。

如果一个 program 处理一整行(图 2-2),max、exp、求和、除都可以在 register 里做完,只读 MN、写 MN,读写量是原来的 1/4:

@triton.jit
def softmax_kernel(x_ptr, y_ptr, stride, N, BLOCK: tl.constexpr):
    row = tl.program_id(0)                    # 一个 program 处理一整行
    cols = tl.arange(0, BLOCK)                # BLOCK ≥ N,一次装下整行
    x = tl.load(x_ptr + row * stride + cols, mask=cols < N, other=float("-inf"))
    x = x - tl.max(x, axis=0)                 # 以下全在 register 里
    e = tl.exp(x)
    y = e / tl.sum(e, axis=0)
    tl.store(y_ptr + row * stride + cols, y, mask=cols < N)

def softmax(x):
    M, N = x.shape
    y = torch.empty_like(x)
    softmax_kernel[(M,)](x, y, x.stride(0), N, BLOCK=triton.next_power_of_2(N))
    return y

在 5090 上对 4096 × 4096 的 fp32 输入实测,eager 版 249 µs,Triton 版 80 µs,快了 3.1 倍,和 torch.softmax(79 µs)相当。

这个写法要求 BLOCK ≥ N,也就是一行能一次装进一个 program。行再长,一个 program 就只能一块一块地读。

2.3 Row sum:行太长时分块累加
#

grid一个 program 做的事行 0行 1 ← pid=1行 2⋯grid = (M,)一行一个 program一个 program(pid = 1):沿第 1 行一块一块地读初始化acc = tl.zeros读一块tl.load累加acc += x归约写回tl.sum, tl.store沿行循环,每次 TILE 个
图 2-3 row sum kernel,画法同图 2-1。一个 program 装不下整行,沿行循环读(红线),每次 TILE 个。

先看最简单的按行归约:求和。program 沿着行一块一块读(图 2-3),每块 TILE 个数加到一个长度为 TILE 的部分和向量上,最后再把这个向量加成一个数。循环里每个位置各加各的,不需要线程之间通信;只有最后的 tl.sum 做一次跨线程归约。

@triton.jit
def row_sum_kernel(x_ptr, out_ptr, N, TILE: tl.constexpr):
    row = tl.program_id(0)
    acc = tl.zeros([TILE], dtype=tl.float32)  # TILE 路部分和,不是一个标量
    for start in range(0, N, TILE):           # 行比 TILE 长:一块一块串行读
        cols = start + tl.arange(0, TILE)
        acc += tl.load(x_ptr + row * N + cols, mask=cols < N, other=0.0)
    tl.store(out_ptr + row, tl.sum(acc, axis=0))  # 最后只做一次跨线程归约

def row_sum(x, TILE=1024):
    M, N = x.shape
    out = torch.empty(M, device=x.device, dtype=x.dtype)
    row_sum_kernel[(M,)](x, out, N, TILE=TILE)
    return out

softmax 不能直接这样分块:要先知道整行的最大值才能算 exp,而最大值要读完整行才知道。边读边更新最大值、同时修正已经算好的部分和,就是 online softmax,FlashAttention 那篇会讲。

前面三个例子里每个数只读一次,分块只是为了把工作分给不同的 program。矩阵乘不同:A 的每个元素要和 B 的一整列相乘,分块是为了让读进片上的数据被多次使用。

2.4 Matmul:分块乘加,顺手融合 ReLU
#

grid一个 program 做的事(1,2)C 的分块,每块 BM × BNgrid = (cdiv(M, BM),cdiv(N, BN))一个 program(pid_m = 1,pid_n = 2):算出 C 的 (1, 2) 块初始化acc [BM, BN] = 0读 A、B 的块两次 tl.load乘加acc += tl.dotReLU 写回tl.store沿 K 循环,每次 BK
图 2-4 matmul kernel,画法同图 2-1。grid 是二维的,一个 program 负责 C 的一块,沿 K 循环(红线)读 A、B 的块累加。

矩阵乘按输出 C 分块,一个 program 负责 C 的一个 BM × BN 块,grid 是二维的(图 2-4)。program 沿 K 方向每次读 A 的一个 BM × BK 块和 B 的一个 BK × BN 块,用 tl.dot 乘加到累加器上。写回之前在 register 里做 ReLU,激活函数就不需要单独一个 kernel 再读写一遍 C:

@triton.jit
def matmul_relu_kernel(a_ptr, b_ptr, c_ptr, M, N, K,
                       sam, sak, sbk, sbn, scm, scn,
                       BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr):
    pid_m, pid_n = tl.program_id(0), tl.program_id(1)   # 负责 C 的第 (pid_m, pid_n) 块
    rm = pid_m * BM + tl.arange(0, BM)
    rn = pid_n * BN + tl.arange(0, BN)
    rk = tl.arange(0, BK)
    a_ptrs = a_ptr + rm[:, None] * sam + rk[None, :] * sak  # [BM, BK] 的指针
    b_ptrs = b_ptr + rk[:, None] * sbk + rn[None, :] * sbn  # [BK, BN] 的指针
    acc = tl.zeros([BM, BN], dtype=tl.float32)
    for k in range(0, K, BK):                 # 沿 K 一块一块乘加
        a = tl.load(a_ptrs, mask=(rm[:, None] < M) & (rk[None, :] + k < K), other=0.0)
        b = tl.load(b_ptrs, mask=(rk[:, None] + k < K) & (rn[None, :] < N), other=0.0)
        acc += tl.dot(a, b, input_precision="ieee")  # fp32 默认会走 tf32,这里要求精确的 fp32
        a_ptrs += BK * sak
        b_ptrs += BK * sbk
    acc = tl.maximum(acc, 0.0)                # ReLU 在 register 里做完,不多一次读写
    c_ptrs = c_ptr + rm[:, None] * scm + rn[None, :] * scn
    tl.store(c_ptrs, acc, mask=(rm[:, None] < M) & (rn[None, :] < N))

def matmul_relu(a, b, BM=64, BN=64, BK=32):
    M, K = a.shape
    _, N = b.shape
    c = torch.empty((M, N), device=a.device, dtype=torch.float32)
    grid = (triton.cdiv(M, BM), triton.cdiv(N, BN))
    matmul_relu_kernel[grid](a, b, c, M, N, K, *a.stride(), *b.stride(), *c.stride(), BM=BM, BN=BN, BK=BK)
    return c

fp32 输入时 tl.dot 默认走 tf32 的 Tensor core,结果和 fp32 的矩阵乘差 0.1 左右;要和 PyTorch(关掉 tf32)对上,需要 input_precision="ieee"。

2.5 num_warps:结合 grid 和 op 一起看
#

四个例子都用了默认的 num_warps=4,即每个 program 4 个 warp。每个 SMSP 实际驻留几个 warp,还要看一个 SM 上同时放了几个 program:每个 SMSP 的 warp 数 = 每个 SM 驻留的 program 数 × num_warps ÷ 4。一个 SM 能驻留几个 program,取决于 grid 够不够分(program 总数 ÷ 170),以及表 1-2 的几条上限。需要驻留几个 warp,则由 op 决定。下面先看 op 需要几个 warp,再看 grid 和 num_warps 怎么凑出这个数。

需要几个 warp,取决于每个 warp 要等多久。以 FMA(fused multiply-add,一条指令算 a × b + c,记 2 次浮点运算)为例:一条 FMA 大约 4 拍后才出结果,如果下一条要用这个结果,就得等它出来才能发。一个 SMSP 上只有 1 个 warp(A)时,4 拍里只有 1 拍在发指令;有 4 个 warp(A、B、C、D)时,A 等待的 3 拍由 B、C、D 填上,每一拍都有指令可发:

拍:          1  2  3  4  5  6  7  8
只有 A:      A  .  .  .  A  .  .  .     用了 1/4 的拍
A、B、C、D:  A  B  C  D  A  B  C  D     每拍都在发

所以要几个 warp,就看等待有几拍,要让别的 warp 把这几拍填满。图 2-5 在 5090 上验证了这一点。kernel 把 256 MiB 数组的每个元素读进来,连做 K 次 FMA(每次都用上一次的结果),再写回。K 越大 intensity 越高,四个 K 让 intensity 从 0.25 到 128,最后一个超过了 ridge point 58,是 compute-bound。grid 固定为 170 × m 个 program,每个 SMSP 就正好驻留 m 个 warp。

0%20%40%60%80%100%1428312416520624728832936104011441248每 SMSP每 SM速度,占这条线最快值的比例;横轴是驻留的 warp 数I = 0.25,显存瓶颈,每 SMSP 1 个 warp:1129.0 µs,4.76e+11 B/s,最快值的 33%I = 0.25,显存瓶颈,每 SMSP 2 个 warp:636.8 µs,8.43e+11 B/s,最快值的 59%I = 0.25,显存瓶颈,每 SMSP 3 个 warp:468.9 µs,1.14e+12 B/s,最快值的 80%I = 0.25,显存瓶颈,每 SMSP 4 个 warp:409.0 µs,1.31e+12 B/s,最快值的 91%I = 0.25,显存瓶颈,每 SMSP 5 个 warp:375.6 µs,1.43e+12 B/s,最快值的 99%I = 0.25,显存瓶颈,每 SMSP 6 个 warp:373.6 µs,1.44e+12 B/s,最快值的 100%I = 0.25,显存瓶颈,每 SMSP 7 个 warp:374.9 µs,1.43e+12 B/s,最快值的 100%I = 0.25,显存瓶颈,每 SMSP 8 个 warp:396.4 µs,1.35e+12 B/s,最快值的 94%I = 0.25,显存瓶颈,每 SMSP 9 个 warp:390.2 µs,1.38e+12 B/s,最快值的 96%I = 0.25,显存瓶颈,每 SMSP 10 个 warp:412.0 µs,1.30e+12 B/s,最快值的 91%I = 0.25,显存瓶颈,每 SMSP 11 个 warp:395.4 µs,1.36e+12 B/s,最快值的 94%I = 0.25,显存瓶颈,每 SMSP 12 个 warp:394.4 µs,1.36e+12 B/s,最快值的 95%I = 0.25,显存瓶颈I = 8,显存瓶颈,每 SMSP 1 个 warp:1245.6 µs,4.31e+11 B/s,最快值的 30%I = 8,显存瓶颈,每 SMSP 2 个 warp:714.3 µs,7.52e+11 B/s,最快值的 52%I = 8,显存瓶颈,每 SMSP 3 个 warp:503.2 µs,1.07e+12 B/s,最快值的 73%I = 8,显存瓶颈,每 SMSP 4 个 warp:410.6 µs,1.31e+12 B/s,最快值的 90%I = 8,显存瓶颈,每 SMSP 5 个 warp:381.0 µs,1.41e+12 B/s,最快值的 97%I = 8,显存瓶颈,每 SMSP 6 个 warp:368.9 µs,1.46e+12 B/s,最快值的 100%I = 8,显存瓶颈,每 SMSP 7 个 warp:373.9 µs,1.44e+12 B/s,最快值的 99%I = 8,显存瓶颈,每 SMSP 8 个 warp:377.8 µs,1.42e+12 B/s,最快值的 98%I = 8,显存瓶颈,每 SMSP 9 个 warp:375.9 µs,1.43e+12 B/s,最快值的 98%I = 8,显存瓶颈,每 SMSP 10 个 warp:377.3 µs,1.42e+12 B/s,最快值的 98%I = 8,显存瓶颈,每 SMSP 11 个 warp:376.0 µs,1.43e+12 B/s,最快值的 98%I = 8,显存瓶颈,每 SMSP 12 个 warp:373.6 µs,1.44e+12 B/s,最快值的 99%I = 8,显存瓶颈I = 32,显存瓶颈,每 SMSP 1 个 warp:1670.6 µs,3.21e+11 B/s,最快值的 22%I = 32,显存瓶颈,每 SMSP 2 个 warp:922.2 µs,5.82e+11 B/s,最快值的 40%I = 32,显存瓶颈,每 SMSP 3 个 warp:636.1 µs,8.44e+11 B/s,最快值的 58%I = 32,显存瓶颈,每 SMSP 4 个 warp:512.7 µs,1.05e+12 B/s,最快值的 73%I = 32,显存瓶颈,每 SMSP 5 个 warp:415.3 µs,1.29e+12 B/s,最快值的 90%I = 32,显存瓶颈,每 SMSP 6 个 warp:383.8 µs,1.40e+12 B/s,最快值的 97%I = 32,显存瓶颈,每 SMSP 7 个 warp:371.8 µs,1.44e+12 B/s,最快值的 100%I = 32,显存瓶颈,每 SMSP 8 个 warp:374.9 µs,1.43e+12 B/s,最快值的 99%I = 32,显存瓶颈,每 SMSP 9 个 warp:375.0 µs,1.43e+12 B/s,最快值的 99%I = 32,显存瓶颈,每 SMSP 10 个 warp:383.6 µs,1.40e+12 B/s,最快值的 97%I = 32,显存瓶颈,每 SMSP 11 个 warp:383.8 µs,1.40e+12 B/s,最快值的 97%I = 32,显存瓶颈,每 SMSP 12 个 warp:383.4 µs,1.40e+12 B/s,最快值的 97%I = 32,显存瓶颈I = 128,算力瓶颈,每 SMSP 1 个 warp:3255.4 µs,2.11e+13 FLOPS,最快值的 25%I = 128,算力瓶颈,每 SMSP 2 个 warp:1658.7 µs,4.14e+13 FLOPS,最快值的 49%I = 128,算力瓶颈,每 SMSP 3 个 warp:1147.0 µs,5.99e+13 FLOPS,最快值的 70%I = 128,算力瓶颈,每 SMSP 4 个 warp:877.2 µs,7.83e+13 FLOPS,最快值的 92%I = 128,算力瓶颈,每 SMSP 5 个 warp:914.5 µs,7.51e+13 FLOPS,最快值的 88%I = 128,算力瓶颈,每 SMSP 6 个 warp:877.4 µs,7.83e+13 FLOPS,最快值的 92%I = 128,算力瓶颈,每 SMSP 7 个 warp:877.1 µs,7.83e+13 FLOPS,最快值的 92%I = 128,算力瓶颈,每 SMSP 8 个 warp:865.1 µs,7.94e+13 FLOPS,最快值的 93%I = 128,算力瓶颈,每 SMSP 9 个 warp:814.5 µs,8.44e+13 FLOPS,最快值的 99%I = 128,算力瓶颈,每 SMSP 10 个 warp:808.5 µs,8.50e+13 FLOPS,最快值的 100%I = 128,算力瓶颈,每 SMSP 11 个 warp:865.7 µs,7.94e+13 FLOPS,最快值的 93%I = 128,算力瓶颈,每 SMSP 12 个 warp:828.0 µs,8.30e+13 FLOPS,最快值的 98%I = 128,算力瓶颈
图 2-5 驻留的 warp 数和速度(RTX 5090,256 MiB 的 fp32 数组,每个元素读 4 B、写 4 B,中间串 K 次相互依赖的 FMA,I = 2K / 8 FLOPs/B)。grid 正好是 170 × m 个 program、每个 4 个 warp,全部同时驻留,每个 SMSP 就是 m 个 warp。最快值:前三条是显存带宽 1.4e12 B/s,I = 128 是 8.5e13 FLOPS(fp32 峰值的 81%)。悬停可看数值。

compute-bound 的那条线(I = 128)和上面的推算一致:1 个 warp 时是 25%,正好 1/4;4 个 warp 时到 92%。memory-bound 的三条线等的是显存,一次读要几百拍,按说要更多 warp 才能填满;但显存每秒能送的数据有上限,warp 多到把显存喂满以后,再加也没用,实测在 5、6 个 warp 时持平。

反例是 2.4 节的矩阵乘。它每个 SM 只放得下 3 个 block(shared memory 先到上限),每个 SMSP 只有 3 个 warp,照样跑到 5.2e13 FLOPS,是 fp32 峰值的一半。因为一个线程负责很多个输出,这些 FMA 互不依赖,第 2 条不用等第 1 条的结果,一个 warp 自己就能每拍发一条:

拍:          1  2  3  4  5  6  7  8
矩阵乘的 A:  A  A  A  A  A  A  A  A     每条 FMA 互不依赖

知道了要几个 warp,再看 grid 和 num_warps 能不能凑出来。拿 2.3 节的 row sum 试两种形状,元素总数一样(都是 340 MiB),每个线程每轮都读 4 个数,只改 num_warps:

表 2-2 row sum 在不同 num_warps 下的耗时
形状每个 SM 分到的 program12481632
170 × 5242881439 µs232 µs214 µs213 µs212 µs211 µs
10880 × 819264212 µs212 µs211 µs211 µs211 µs217 µs

表头的 1 到 32 是 num_warps。TILE = 128 × num_warps,每个线程每轮读 4 个 fp32。

170 行时 grid 只有 170 个 program,每个 SM 只分到 1 个,每个 SMSP 的 warp 数就是 num_warps ÷ 4:取 1 时 4 个 SMSP 里只有 1 个有活,439 µs;取 4 以上每个 SMSP 都有 warp,降到 214 µs,带宽 1.67e12 B/s。10880 行时 program 远多于 SM,每个 SM 都塞到上限,num_warps 只改变同样多的 warp 怎么分组,6 种取值都是 211 到 217 µs。

num_warps 还决定每个线程要存多少数据:program 一次要存的数据平分给它的线程。2.2 节的 softmax 把一整行放在 register 里,行长 32768 时,num_warps=1 每个线程要存 1024 个数,远超每个线程 255 个 register 的上限,溢出到显存,耗时 705 µs;num_warps=8 每个线程存 128 个,不溢出,183 µs。

驻留 warp 数与上限之比叫 occupancy,Nsight Compute 会直接报出来。它不必占满,够把等待的拍数填满就行:图 2-5 的 kernel 4 到 6 个 warp 就够,矩阵乘 25% 的 occupancy 也够。调 num_warps 的顺序是:先看 program 比 SM 多多少,program 少时把 num_warps 调大;再看 register 会不会溢出;最后用 triton.autotune 微调。

第 2 节的四个例子是同一个模式:program 把一块数据读进 register,在片上做完能做的计算,再写回显存。下一篇在 5090 上拆解一步训练的时间和显存,用的就是这个模式:把 RMSNorm 的几个 kernel 融合成一个,少读写一遍中间结果。


附录 A TPU:更大的矩阵单元,更简单的控制
#

正文的 GPU 靠 kernel 自己分块来复用数据。TPU 面对同样的带宽问题,思路是把控制逻辑做轻,把矩阵乘单元做大。图 A-1 是 TPU v5p 的一个 TensorCore:Scalar unit 取指令、发给 VPU 和 MXU,整个核只有这一路控制,没有 warp 和调度器之间的切换;VPU 做逐元素运算和归约;矩阵乘交给 4 个 MXU,每个是 128 × 128 的 systolic array。片上存储只有 VMEM 和 VREG 两层,VMEM 是编译器管理的 scratchpad,不是硬件 cache,数据什么时候从 HBM 搬进来、搬多少,都由编译器(XLA)提前排好。

TensorCore(TPU v5p 每个芯片 2 个)Scalar unit:取指令,发给 VPU 和 MXU(只有这一路控制)VMEM:片上 scratchpad,由编译器显式搬入搬出VREG 256 KB(vector register)VPU8 × 128 lanesMXU128 × 128MXU128 × 128MXU128 × 128MXU128 × 128HBM95 GB带宽2.8e12 B/s没有 warp,也没有硬件管理的 L1 / L2:HBM → VMEM → VREG → VPU / MXU 每一步都由编译器排好
图 A-1 TPU v5p 一个 TensorCore 的结构示意,画法对应图 1-2。规格来自 How to Scale Your Model,原书也有一张 TensorCore 的结构图。

MXU 和 GPU 最大的不同在乘加单元怎么拿到数据(图 A-2)。GPU 上每个 CUDA core 做一次乘加,要从 register file 读两个数、写回一个;一个数被用很多次,得靠 kernel 自己分块,先搬进 shared memory 和 register 再反复读,2.4 节的 matmul 就是这样写的。systolic array 把复用做进了硬件:权重先装进 128 × 128 个格子,固定不动;输入 x 从左边进入,每拍向右传一格,经过的每一格把它乘上自己的权重,加到从上面传下来的部分和上,再把部分和往下传。一个 x 从 VMEM 读一次,在一行里被 128 个格子用到;稳定以后每拍从左边进 128 个数、从底部出 128 个结果,阵列里同时做 16384 次乘加,中间结果都不经过 register file。

GPU:CUDA core 从 register file 取数register fileFMAFMAFMAFMA每次乘加读 2 个数、写回 1 个一个数要被用很多次,得靠 kernel 自己分块:先搬进 shared memory / register,再反复读TPU:MXU 是 systolic array(示意 4 × 4,实际 128 × 128)W 先装进格子固定不动;x 每拍右移一格,部分和每拍下移一格w00w01w02w03x0w10w11w12w13x1w20w21w22w23x2w30w31w32w33x3y0y1y2y3
图 A-2 乘加单元怎么拿到数据。左边是 GPU 的 CUDA core,右边是 TPU 的 MXU;x 按行错开一拍进入,y 是每列的输出。动画演示见 How to Scale Your Model 附录 B。

代价是形状要规整。矩阵的维度要是 128 的倍数才能填满阵列,小矩阵和不规则的计算会浪费格子;数据流由编译器提前排好,也不能像 GPU 那样靠切换 warp 应付不确定的访存延迟。多卡之间怎么互联是另一个大区别,超出了本文的范围。两边术语的对应见表 A-1。

表 A-1 GPU 与 TPU 的术语对照
GPUTPU作用RTX 5090TPU v5p
SMTensorCore包含其他运算单元的核心1702
warp 调度器VPUSIMD 向量运算的调度6808
CUDA coreVPU ALU普通的 SIMD 运算单元21760—
L1 / shared memoryVMEM片上的快速存储21.3 MB128 MB
registerVREGvector register42.5 MB512 KB
Tensor coreMXU矩阵乘单元6808
显存HBM大容量主存32 GB GDDR795 GB

RTX 5090 的片上存储是 170 个 SM 加起来的总量。TPU 一列参考 Google 的 How to Scale Your Model。


References
#

  1. Gholami, Yao, Kim, Hooper, Mahoney, Keutzer. AI and Memory Wall. IEEE Micro, 2024. 图 1-1。
  2. NVIDIA. NVIDIA RTX Blackwell GPU Architecture(白皮书)。RTX 5090 的存储层级和 SM 结构,图 1-2、图 1-3。
  3. NVIDIA. NVIDIA A100 Tensor Core GPU Architecture、NVIDIA H100 Tensor Core GPU Architecture、NVIDIA Blackwell Architecture。表 1-1 的 A100、H100、B200 一列。
  4. NVIDIA. CUDA Programming Guide。各 compute capability 每 SM 的 warp、block、register 和 shared memory 上限,表 1-2,CUDA 执行层级。
  5. NVIDIA. Nsight Compute Kernel Profiling Guide。warp 调度状态和 occupancy。
  6. Austin et al. How to Scale Your Model, Google DeepMind, 2025. 附录 A 的 TPU v5p 规格、systolic array 的数据流和表 A-1。
  7. Jouppi et al. In-Datacenter Performance Analysis of a Tensor Processing Unit. ISCA 2017. 第一代 TPU 和它的 systolic array。
  8. Kung. Why Systolic Architectures?. IEEE Computer, 1982.
  9. Tillet, Kung, Cox. Triton: An Intermediate Language and Compiler for Tiled Neural Network Computations. MAPL 2019.
  10. Triton tutorials。第 2 节 softmax 和 matmul 例子的写法。
  11. Hendrycks, Gimpel. Gaussian Error Linear Units (GELUs). arXiv:1606.08415, 2016. 第 2.1 节的 tanh 近似。
  12. Milakov, Gimelshein. Online normalizer calculation for softmax. arXiv:1805.02867, 2018. 第 2.3 节提到的 online softmax。
GPU training optimization - 这篇文章属于一个选集。
→ GPU 与 Triton 入门:从存储层级到第一个 kernel (本文)

Noviorlu喵

多伦多大学计算机工程。热爱计算机图形学、渲染和 AI/ML。

留言

用 GitHub 账号登录后可以留言。