Skip to content
Gains Summary
Main Navigation 首页 / Home
C++ 编程 / C++ Programming
系统与高性能 / Systems & Performance
Web 开发 / Web Development
人工智能 / Artificial Intelligence
工业软件 / Industrial Software
其他内容 / Other Topics
C++ 编程 / C++系统与性能 / SystemsWeb 开发 / Web人工智能 / AI工业软件 / Industrial

外观

Sidebar Navigation

← 人工智能 / Artificial Intelligence

AI 编译器 / AI Compilers

1. AI 编译器全景——为什么模型需要编译器 / The AI Compiler Landscape and Why Models Need Compilers

2. 编译原理速通——面向 ML 工程师的核心概念 / Compiler Fundamentals for Machine Learning Engineers

3. 中间表示基础——理解 IR 层级与 lowering 链路 / Intermediate Representation Levels and Lowering Pipelines

4. 计算图的构建与表示 / Building and Representing Computational Graphs

5. MLIR 架构、方言与渐进式降级 / MLIR Architecture, Dialects, and Progressive Lowering

6. 算子语义、广播、归约与形状推导 / Operator Semantics, Broadcasting, Reduction, and Shape Inference

7. 模型前端格式:ONNX、TFLite、HLO 与 SavedModel / Model Frontend Formats: ONNX, TFLite, HLO, and SavedModel

8. 图优化 Pass——经典优化在 ML 中的应用 / Graph Optimization Passes for Machine Learning

9. 算子融合——编译器最重要的性能优化 / Operator Fusion as a Core Compiler Optimization

10. 内存规划——Buffer 分配与生命周期管理 / Memory Planning, Buffer Allocation, and Lifetime Management

11. Layout 优化——数据排布转换与内存效率 / Layout Optimization for Data Movement and Memory Efficiency

12. 动态 Shape——符号分析与形状处理 / Dynamic Shapes, Symbolic Analysis, and Shape Processing

13. 硬件约束下的操作调度 / Operation Scheduling Under Hardware Constraints

14. 从模板、DSL 到 IR 降级的代码生成架构 / Code Generation Architectures from Templates and DSLs to IR Lowering

15. CPU 后端:SIMD、分块与多线程 / CPU Backends with SIMD, Tiling, and Multithreading

16. CUDA 后端:合并访存与 Tensor Core / CUDA Backends, Memory Coalescing, and Tensor Cores

17. NPU 后端:脉动阵列与端侧 AI 生态 / NPU Backends, Systolic Arrays, and Edge AI Ecosystems

18. Kernel 性能基础:Roofline 与 Occupancy / Kernel Performance Fundamentals with Roofline and Occupancy

19. CUTLASS 与分层 GEMM 模板 / CUTLASS and Hierarchical GEMM Templates

20. TVM Tensor Expression 与计算调度分离 / TVM Tensor Expressions and Compute-Schedule Separation

21. 使用 Triton 编写高性能 GPU Kernel / Triton for High-Performance GPU Kernels in Python

22. 基于成本模型与实测搜索的自动调度 / Automatic Scheduling with Cost Models and Measurement-Based Search

23. XLA 内部机制:HLO、融合与 SPMD / XLA Internals, HLO, Fusion, and SPMD

24. Torch-MLIR:从 PyTorch 算子到 MLIR 方言 / Torch-MLIR from PyTorch Operators to MLIR Dialects

25. torch.compile:Dynamo、AOTAutograd、Inductor 与 Triton / Torch Compile with Dynamo, AOTAutograd, Inductor, and Triton

26. 从 MLIR 经 LLVM 降级到机器码 / Lowering from MLIR Through LLVM to Machine Code

27. 量化——低精度推理的工程实践 / Engineering Low-Precision Inference with Quantization

28. 分布式编译与训练——多设备编排的编译器支持 / Compiler Support for Distributed Training and Multi-Device Orchestration

29. 生产调试——真实问题的编译器视角排查 / Production Debugging from the Compiler Perspective

30. 未来方向——AI 编译器的新挑战与机遇 / Future Challenges and Opportunities for AI Compilers

本页目录

📅 创建时间:2026-06-03 🏷️ 标签:#Triton #GPU #Python #JIT #BlockedISA #FlashAttention #torch.compile #Pythonic 📚 前置知识:[[/04-ai/01-llm-engineering/07-llm-evolution]](LLM 发展脉络)[[17-kernel-primer]](Kernel 开发入门) 📚 相关知识:[[19-tvm-te]](TVM TE)[[24-torch-compile]](torch.compile)[[18-cutlass]](CUTLASS)


使用 Triton 编写高性能 GPU Kernel / Triton for High-Performance GPU Kernels in Python ​

┌──────────────────────────────────────────────────────────────────────────────┐ │ 📖 场景:Triton 生成的代码质量能打败手写 CUDA 吗? │ ├──────────────────────────────────────────────────────────────────────────────┤ │ 你的团队在用 Triton 写 Attention kernel。PyTorch 工程师写起来很顺手。 │ │ 但你怀疑它生成的代码质量:Python JIT 真的能打败手写 CUDA 吗? │ │ 实测 Flash Attention Triton 版本和 CUDA 官方实现速度接近, │ │ 但比手写的简单 kernel 快 30%。为什么 Triton 能做到? │ └──────────────────────────────────────────────────────────────────────────────┘

第1节 Triton 设计哲学——让 PyTorch 工程师写高性能 GPU Kernel ​

1.1 为什么需要 Triton ​

GPU Kernel 开发有两个极端:

极端 1: 手写 CUDA
  - 优点: 完全控制,可以达到最高性能
  - 缺点: 学习曲线陡峭,容易出错,代码难维护
  
极端 2: cuBLAS/cuDNN
  - 优点: 开箱即用,性能好
  - 缺点: 不支持自定义操作(如自定义激活函数)

Triton 的目标: 填补中间的空白
  - 让你用 Python 写 kernel
  - 自动生成接近手写 CUDA 的代码
  - 但不需要你手动管理 shared memory / coalesced access
1
2
3
4
5
6
7
8
9
10
11
12

Triton 的核心洞察:

GPU kernel 的优化空间是有限的、规则的:

1. Tiling 策略: 把矩阵分成 tile,每个 tile 放入 shared memory
2. 向量化: SIMD 指令一次处理多个元素
3. 内存访问合并: 连续线程访问连续内存
4. 双缓冲: 计算和内存传输重叠

这些规则是通用的,可以被编译器自动推断。
Triton 的 Blocked ISA 让用户只需指定 tile shape,
编译器自动处理剩下的优化。
1
2
3
4
5
6
7
8
9
10

1.2 Triton vs 其他方案的定位 ​

┌──────────────────────────────────────────────────────────────────────┐
│                           性能上限                                   │
│                                                                      │
│    手写 CUDA (cuBLAS 团队)  ──────────────────────────► 理论峰值    │
│              │                                                       │
│              │                                                       │
│    CUTLASS   │  接近 cuBLAS,但需要大量模板知识                      │
│              │                                                       │
│              │                                                       │
│    Triton    │  达到手写 CUDA 的 80-95%                             │
│              │  Pythonic 接口                                        │
│              │                                                       │
│              │                                                       │
│    TVM TE    │  可达到接近手写 CUDA                                 │
│              │  但 schedule 参数空间巨大                             │
│              │                                                       │
│              │                                                       │
│    PyTorch   │  简单但性能受限(einsum, torch.mm)                  │
│              │                                                       │
└──────────────────────────────────────────────────────────────────────┘

选择指南:
  - 需要极致性能 + 愿意投入 → 手写 CUDA / CUTLASS
  - 需要定制 + Python 优先 → Triton
  - 需要自动搜索最优参数 → TVM AutoScheduler
  - 标准 GEMM/Conv → cuBLAS/cuDNN
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

第2节 @triton.jit 装饰器——Python 函数变成 GPU Kernel ​

2.1 JIT 编译模型 ​

python
import triton
import triton.language as tl

# @triton.jit: 把 Python 函数标记为 Triton kernel
# Triton 会在运行时 (JIT) 编译这个函数到 PTX/CUDA 代码

@triton.jit
def add_kernel(x_ptr, y_ptr, output_ptr, n_elements, BLOCK_SIZE: tl.constexpr):
    """
    y = x + y, 每个元素一个 thread 处理
    但实际每个 thread 处理 BLOCK_SIZE 个连续元素 (向量化)
    
    参数说明:
      x_ptr, y_ptr, output_ptr: GPU 内存指针 (通过 torch.Tensor_PTR 传入)
      n_elements: 元素总数
      BLOCK_SIZE: compile-time 常量 (tl.constexpr)
    """
    
    # 程序 ID: 识别当前 thread block 处理的元素范围
    pid = tl.program_id(axis=0)  # axis=0 表示一维 grid
    
    # 计算当前 block 处理的元素范围
    # 每个 block 处理 BLOCK_SIZE 个连续元素
    block_start = pid * BLOCK_SIZE
    
    # 元素索引: 从 block 起始位置开始
    offsets = block_start + tl.arange(0, BLOCK_SIZE)
    
    # 创建 mask: 只处理有效范围内的元素 (n_elements 以内的)
    # 超出 n_elements 的位置被 mask 住,不会真正执行
    mask = offsets < n_elements
    
    # 从 GPU 内存加载数据
    # tl.load: 从 x_ptr[offsets] 读取,mask 住越界的访问
    x = tl.load(x_ptr + offsets, mask=mask, other=0.0)
    y = tl.load(y_ptr + offsets, mask=mask, other=0.0)
    
    # 执行计算
    output = x + y
    
    # 写回 GPU 内存
    tl.store(output_ptr + offsets, output, mask=mask)


# 调用 Triton kernel
def add(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
    output = torch.empty_like(x)
    
    n_elements = output.numel()
    
    # launch 配置
    BLOCK_SIZE = 1024
    
    # 计算 grid 大小: 需要多少个 block
    # 向上取整: ceil(n_elements / BLOCK_SIZE)
    grid = (triton.cdiv(n_elements, BLOCK_SIZE),)
    
    # 启动 kernel
    add_kernel[grid](
        x, y, output,  # torch.Tensor 会自动转换为 device pointer
        n_elements,
        BLOCK_SIZE=BLOCK_SIZE,
    )
    
    return output
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
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65

2.2 tl.constexpr——Compile-Time 常量 ​

python
# tl.constexpr: 标记在编译时求值的常量
# 作用: 允许 Triton 在编译时做常量折叠和优化

@triton.jit
def kernel_with_constexpr(
    x_ptr, y_ptr, output_ptr,
    BLOCK_SIZE: tl.constexpr,   # 编译时常量
    DOUBLE: tl.constexpr,       # 另一个编译时常量
):
    # 编译时可以直接计算
    # 比如: 1 << DOUBLE 在编译时就是 1 << 2 = 4
    stride = 1 << DOUBLE  # 如果 DOUBLE=2,stride=4
    
    offsets = ...
    
    # DOUBLE 作为 bool 用
    if DOUBLE:
        # 如果 DOUBLE 是 True,编译器知道只执行这个分支
        # 可以消除 dead code
        ...
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20

第3节 Triton Blocked ISA——BLOCK 概念 ​

3.1 什么是 Blocked ISA ​

传统 CUDA 编程模型:
  - 你需要手动管理:
    1. Thread/block 的映射到数据
    2. Shared memory 的加载/同步
    3. 内存访问的 coalescing
    4. Warp 级别的同步

Triton Blocked ISA:
  - 你只指定 BLOCK 的大小和形状
  - Triton 编译器自动推断:
    1. 如何把 BLOCK 映射到 thread/block
    2. 如何利用 shared memory
    3. 如何合并内存访问
  - BLOCK = Tile = 你要处理的最小数据单元
1
2
3
4
5
6
7
8
9
10
11
12
13
14

3.2 常见的 BLOCK 形状 ​

python
# BLOCK 形状的选择影响性能

# 1D BLOCK: 适合向量操作
BLOCK_M = 1
BLOCK_N = 1024  # 每个 thread 处理 1024 个元素

# 2D BLOCK: 适合矩阵操作
BLOCK_M = 128   # 行方向
BLOCK_N = 128   # 列方向

# BLOCK 形状的经验法则:
# 
# GEMM:
#   BLOCK_M = 64-128 (行方向 tile)
#   BLOCK_N = 64-128 (列方向 tile)
#   BLOCK_K = 32-64 (reduce 方向)
#
# Flash Attention:
#   BLOCK_M = 64-128 (query block)
#   BLOCK_N = 64-128 (key block)
#   小的 BLOCK 减少 SRAM 需求,大的 BLOCK 提高数据复用
#
# Softmax:
#   BLOCK_M = 1 (每个 thread 处理一行)
#   BLOCK_N = 1024-4096 (每行元素数)
#   需要 reduce 操作,但行间独立
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

3.3 自动向量化 ​

python
@triton.jit
def vectorized_load_kernel(x_ptr, output_ptr, n_elements, VECTOR_SIZE: tl.constexpr):
    """
    Triton 自动把 tl.load/store 向量化
    VECTOR_SIZE 指定每个 memory transaction 的大小
    """
    
    pid = tl.program_id(axis=0)
    
    # 注意: offsets 是 BLOCK_SIZE 个元素的索引
    # Triton 会识别这是连续访问,自动生成 vectorized load
    offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    
    # 如果 BLOCK_SIZE = 256 且 VECTOR_SIZE = 4
    # Triton 会生成 256/4 = 64 个 vectorized load
    # 每个 load 4 个 float = 16 bytes
    
    # 加载向量化的数据
    # x_ptr + offsets 指向连续内存
    # Triton 识别这是向量化访问,生成 64 个 4-element 的 vectorized load
    x = tl.load(x_ptr + offsets, mask=offsets < n_elements, other=0.0)
    
    output = x * 2.0
    
    tl.store(output_ptr + offsets, output, mask=offsets < n_elements)
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

第4节 Triton DSL——tl.load, tl.store, tl.dot, tl.constexpr ​

4.1 tl.load / tl.store ​

python
# tl.load: 从 GPU 内存加载数据到寄存器
# tl.store: 把寄存器数据写回 GPU 内存

@triton.jit
def load_store_examples(
    ptr,          # GPU 内存指针
    offsets,       # 索引 (可以是向量)
    mask,         # 掩码 (越界的位置)
    other,        # 越界时的默认值
    BLOCK_SIZE: tl.constexpr
):
    # 基础用法
    x = tl.load(ptr + offsets, mask=mask, other=0.0)
    
    # 展开指针 (pointer arithmetric)
    # ptr_tensor + offsets 等价于 ptr[offsets]
    # 你也可以手动计算偏移
    base_ptr = ptr
    x = tl.load(base_ptr + offsets, mask=mask, other=0.0)
    
    # 跨步加载 (strided load)
    # 用于非连续内存访问
    # 例如: 读取矩阵的一列
    # ptr: 矩阵起始地址
    # offsets: 行索引
    # stride: 列跨度 (矩阵的列数)
    row_offsets = offsets * stride
    x = tl.load(ptr + row_offsets, mask=mask, other=0.0)
    
    # 缓存策略提示
    # tl.RW0, tl.RW1 指定 cache modifier
    # tl.load(ptr, cache_modifier="ca")  # cache as streaming
    
    # store 用法相同
    tl.store(ptr + offsets, x, mask=mask)
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

4.2 tl.dot——矩阵乘 ​

python
# tl.dot: 执行矩阵乘累加
# 这是 Triton 最强大的功能之一

@triton.jit
def matmul_kernel(
    # Pointers for matrices
    a_ptr, b_ptr, c_ptr,
    # Matrix dimensions
    M, N, K,
    # Strides
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_cm, stride_cn,
    # Meta-parameters
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    GROUP_M: tl.constexpr,  # 用于 swizzle
):
    """
    实现: C = A @ B + C
    
    A: M×K (行主序)
    B: K×N (行主序)
    C: M×N (行主序)
    """
    
    # 程序 ID: 二维 grid
    # pid_m: 行方向 block index
    # pid_n: 列方向 block index
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    
    # 计算当前 block 的 C 矩阵范围
    # 每 block 处理 BLOCK_M × BLOCK_N 的 C
    rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)  # C 的行索引
    rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)  # C 的列索引
    
    # A 和 B 的索引
    # A: rm (行), k (列) - k 遍历 K
    # B: k (行), rn (列) - k 遍历 K
    
    # 初始化累加器
    # 用 fp32 累加,精度更高
    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    
    # K 方向循环
    for k in range(0, K, BLOCK_K):
        # A 的索引: rm (行), k+k_idx (列)
        ram = rm[:, None]  # shape: (BLOCK_M, 1)
        rak = k + tl.arange(0, BLOCK_K)  # shape: (BLOCK_K,)
        a_offsets = ram * stride_am + rak * stride_ak  # 全局索引
        
        # B 的索引: k+k_idx (行), rn (列)
        rak = k + tl.arange(0, BLOCK_K)  # shape: (BLOCK_K,)
        rbn = rn[None, :]  # shape: (1, BLOCK_N)
        b_offsets = rak * stride_bk + rbn * stride_bn
        
        # 加载 A 和 B 的 tile
        # shape: (BLOCK_M, BLOCK_K) 和 (BLOCK_K, BLOCK_N)
        # 自动处理 alignment 和 masking
        a_mask = (rm < M)[:, None] & (rak < K)[None, :]
        b_mask = (rak < K)[:, None] & (rn < N)[None, :]
        
        a = tl.load(a_ptr + a_offsets, mask=a_mask, other=0.0)
        b = tl.load(b_ptr + b_offsets, mask=b_mask, other=0.0)
        
        # 执行矩阵乘
        # accumulator += a @ b
        # a: (BLOCK_M, BLOCK_K)
        # b: (BLOCK_K, BLOCK_N)
        # result: (BLOCK_M, BLOCK_N)
        accumulator += tl.dot(a, b)
    
    # 把累加器 (fp32) 转换回 output dtype (通常是 fp16)
    c = accumulator.to(tl.float16)
    
    # 写回 C
    cm = rm[:, None]
    cn = rn[None, :]
    c_offsets = cm * stride_cm + cn * stride_cn
    c_mask = (rm < M)[:, None] & (rn < N)[None, :]
    
    tl.store(c_ptr + c_offsets, c, mask=c_mask)
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
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82

4.3 tl.constexpr 的高级用法 ​

python
@triton.jit
def advanced_constexpr(
    x_ptr, y_ptr,
    N: tl.constexpr,
    ACTIVATION: tl.constexpr,  # 编译时选择激活函数
    SCALE_BIAS: tl.constexpr,  # 编译时决定是否做 scale+bias
):
    x = tl.load(x_ptr)
    
    # 编译时分支消除
    if ACTIVATION == "relu":
        x = tl.where(x > 0, x, 0.0)
    elif ACTIVATION == "gelu":
        x = 0.5 * x * (1.0 + tl.tanh(0.797884 * (x + 0.044715 * x * x * x)))
    elif ACTIVATION == "silu":
        x = x * tl.sigmoid(x)
    
    # 编译时展开循环
    if SCALE_BIAS:
        scale = tl.load(y_ptr)
        bias = tl.load(y_ptr + 1)
        x = x * scale + bias
    
    return x
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24

第5节 Triton 的自动优化 ​

5.1 自动选择 Memory Access Pattern ​

python
@triton.jit
def autotuned_matmul(...):
    """
    Triton 自动处理:
    1. Coalesced memory access
    2. Bank conflict 避免 (通过 swizzle)
    3. L2 cache 友好的访问模式
    """
    ...
1
2
3
4
5
6
7
8
9
Triton 的自动 coalescing:

场景: 加载 A 矩阵的 tile (行 rm, 列 k 范围)
    rm = [0, 1, 2, ..., 127]
    k_range = [0, 1, 2, ..., 31]

问题: 如果用简单的索引 a[rm, k_range]
      每个 thread 访问的内存地址不连续

Triton 编译器识别:
    原始索引: a[0,0], a[1,0], a[2,0], ... (列优先)
    连续访问: a[0,0], a[0,1], a[0,2], ... (行优先)

编译器自动重排:
    1. 交换 axis 顺序
    2. 生成 coalesced 的 load 指令
    3. 或生成 transpose + coalesced load
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17

5.2 自动选择 Tiling Factor ​

python
# Triton 2.0 引入了 autotuning
# 可以自动搜索最优的 BLOCK_SIZE, num_warps 等参数

@triton.autotune(
    configs=[
        # configs 是一个配置列表,Triton 会自动尝试
        # 每个 config 是 {key: value} 的字典
        triton.Config({'BLOCK_M': 64, 'BLOCK_N': 64, 'BLOCK_K': 32}, num_warps=4),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'BLOCK_K': 32}, num_warps=8),
        triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128, 'BLOCK_K': 32}, num_warps=8),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64}, num_warps=8),
    ],
    # key: 用于分组的参数(相同 key 共享编译结果)
    key=['M', 'N', 'K'],
)
@triton.jit
def autotuned_gemm(...):
    ...
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18

5.3 自动管理 Shared Memory ​

python
# Triton 2.0 之前: 用户需要手动申请 shared memory
# Triton 2.0+ : 自动管理

@triton.jit
def auto_smem_matmul(...):
    # Triton 编译器自动决定:
    # 1. 需要多少 shared memory
    # 2. 如何安排 tile 的 layout
    # 3. 何时 flush shared memory
    
    # 加载 A tile 到 "自动分配的 shared memory"
    # triton 的 load/store 指令隐式使用 shared memory 缓存
    a = tl.load(a_ptr + a_offsets, mask=a_mask, other=0.0, eviction_policy="evict_last")
    # eviction_policy="evict_last": L2 cache 友好的策略
    # 先加载的数据最后被 evict,提高 cache 利用率
    
    # ... 计算 ...
    
    return
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19

第6节 Flash Attention 在 Triton 中的实现 ​

6.1 Flash Attention 的核心思想 ​

标准 Attention 实现的问题:
  1. 需要完整的 S×S attention matrix (S=序列长度)
  2. 内存复杂度 O(S²),无法处理长序列

Flash Attention 的核心思想:
  1. 不存储完整的 attention matrix
  2. 分块计算 online softmax
  3. 边算边归一化,只保留必要的统计量 (m, l)
  
  Online Softmax:
    m_i = max(m_{i-1}, x_i)      # max 值
    l_i = l_{i-1} * exp(m_{i-1} - m_i) + exp(x_i - m_i)  # sum
    softmax_i = exp(x_i - m_i) / l_i
    
    只需要 O(1) 额外内存!
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15

6.2 Triton Flash Attention 实现 ​

python
@triton.jit
def flash_attention_kernel(
    Q, K, V, Out,               # 指针
    stride_qb, stride_qh, stride_qm, stride_qk,  # Q 的 stride
    stride_kb, stride_kh, stride_kn, stride_kk,
    stride_vb, stride_vh, stride_vn, stride_vk,
    stride_ob, stride_oh, stride_om, stride_ok,
    # 维度
    B, H, N, D,                 # batch, heads, seq_len, head_dim
    # Meta-parameters
    BLOCK_M: tl.constexpr,      # Q block 大小
    BLOCK_N: tl.constexpr,      # K/V block 大小
    BLOCK_D: tl.constexpr,       # head_dim (必须 <= 128)
):
    """
    Flash Attention 的 Triton 实现
    
    公式:
      O_i = softmax(Q_i @ K^T / sqrt(D)) @ V
            = sum_j (exp(q_i · k_j - m_i) / l_i) * v_j
    """
    
    # 获取程序 ID
    bid_h = tl.program_id(0)  # head index
    bid_b = tl.program_id(1)  # batch index
    start_m = tl.program_id(2)  # Q block index
    
    # Q_i: 当前 query block
    # shape: (BLOCK_M, D)
    rm = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
    rn = tl.arange(0, BLOCK_N)  # 用于加载 K
    rd = tl.arange(0, BLOCK_D)
    
    # 计算 Q block 的内存偏移
    q_base = Q + bid_b * stride_qb + bid_h * stride_qh
    # 加载 Q
    q_offsets = rm[:, None] * stride_qm + rd[None, :] * stride_qk
    q_mask = (rm[:, None] < N) & (rd[None, :] < D)
    q = tl.load(q_base + q_offsets, mask=q_mask, other=0.0)
    
    # 初始化累加器
    # m_i: max 值 (用于 softmax)
    # l_i: sum 值 (用于归一化)
    m_i = tl.zeros((BLOCK_M,), dtype=tl.float32) + float("-inf")  # 初始化为 -inf
    l_i = tl.zeros((BLOCK_M,), dtype=tl.float32) + 1e-6  # 初始化很小的值
    
    # accumulator: 存放 O_i 的累加结果
    acc = tl.zeros((BLOCK_M, BLOCK_D), dtype=tl.float32)
    
    # 遍历 K, V 的 block
    for start_n in range(0, N, BLOCK_N):
        # 加载 K block: shape (BLOCK_N, D)
        kn = start_n + rn
        k_offsets = kn[:, None] * stride_kn + rd[None, :] * stride_kk
        k_mask = (kn[:, None] < N) & (rd[None, :] < D)
        k = tl.load(K + bid_b * stride_kb + bid_h * stride_kh + k_offsets, 
                    mask=k_mask, other=0.0)
        
        # 计算 Q @ K^T
        # q: (BLOCK_M, D), k: (BLOCK_N, D)
        # q @ k^T: (BLOCK_M, BLOCK_N)
        qk = tl.dot(q, tl.trans(k))  # transpose k
        
        # 加上 causal mask (如果需要)
        # mask: rm[:, None] >= start_n + rn[None, :]
        # 即: 当前位置 q_i 只能 attend 到 k_j,其中 j <= i
        
        # Scale by sqrt(D)
        scale = 1.0 / tl.sqrt(D.to(tl.float32))
        qk = qk * scale
        
        # 计算 block 内的 max 和 sum (online softmax)
        m_ij = tl.max(qk, axis=1)  # max across K
        m_i_new = tl.maximum(m_i, m_ij)  # 更新 global max
        
        # 计算 exp(qk - m_i_new)
        qk_minus_m = qk - m_i_new[:, None]
        p = tl.exp(qk_minus_m)
        
        # 计算 block 内的 l
        l_ij = tl.sum(p, axis=1)
        l_i_new = tl.exp(m_i - m_i_new) * l_i + l_ij
        
        # 计算 softmax 归一化后的 attention
        p_scaled = p / l_ij[:, None]
        
        # 加载 V block: shape (BLOCK_N, D)
        v_offsets = kn[:, None] * stride_vn + rd[None, :] * stride_vk
        v_mask = (kn[:, None] < N) & (rd[None, :] < D)
        v = tl.load(V + bid_b * stride_vb + bid_h * stride_vh + v_offsets,
                    mask=v_mask, other=0.0)
        
        # 计算 attention @ V
        # p_scaled: (BLOCK_M, BLOCK_N), v: (BLOCK_N, D)
        # result: (BLOCK_M, D)
        pv = tl.dot(p_scaled.to(tl.float16), v)
        
        # 累加到 acc
        # 需要 scale 旧的 acc
        acc_scale = tl.exp(m_i - m_i_new) * l_i / l_i_new
        acc = acc * acc_scale[:, None] + pv / l_i_new[:, None]
        
        # 更新 m 和 l
        m_i = m_i_new
        l_i = l_i_new
    
    # 写回输出
    out = acc.to(tl.float16)
    o_base = Out + bid_b * stride_ob + bid_h * stride_oh
    o_offsets = rm[:, None] * stride_om + rd[None, :] * stride_ok
    o_mask = (rm[:, None] < N) & (rd[None, :] < D)
    tl.store(o_base + o_offsets, out, mask=o_mask)
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
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112

第7节 Triton vs CUDA C++——什么时候选哪个 ​

7.1 选 Triton 的场景 ​

场景原因
PyTorch 项目和 PyTorch 无缝集成,不需要单独的编译流程
快速原型Python 开发速度快,JIT 编译即跑
中等复杂度算子GEMM、Attention、Softmax、LayerNorm 等
需要 fusionTriton 容易 fusion 多个操作
研究实验容易尝试不同的算法变体

7.2 选手写 CUDA 的场景 ​

场景原因
极致性能手写 CUDA 可以达到 100% 峰值
复杂控制流Triton 对复杂的分支/循环支持有限
动态 shapeTriton 的动态 shape 支持是实验性的
特殊硬件特性需要手动利用特殊的硬件指令
生产级库需要长期维护和优化的核心库

7.3 Triton vs 其他方案对比表 ​

维度TritonTVM TECUTLASS手写 CUDA
接口易用性⭐⭐⭐⭐⭐ Pythonic⭐⭐⭐ 需要写 schedule⭐ 需要 C++ 模板⭐ 需要手写底层
性能⭐⭐⭐⭐ 手写 CUDA 的 80-95%⭐⭐⭐⭐ 取决于 schedule⭐⭐⭐⭐⭐ 接近峰值⭐⭐⭐⭐⭐ 理论峰值
灵活性⭐⭐⭐⭐ 足够大多数场景⭐⭐⭐⭐⭐ 高度可定制⭐⭐⭐⭐ 模板化⭐⭐⭐⭐⭐ 完全控制
学习曲线低高很高很高
编译时间中等 (JIT)高中等无 (预编译)
动态 shape有限✅ 支持❌ 不支持❌ 不支持
适用场景PyTorch 集成、研究自动搜索优化生产级 GEMM极致优化库

第8节 Triton 的局限性 ​

8.1 动态 Shape 支持 ​

python
# Triton 对动态 shape 的支持有限

# 场景 1: 编译时知道 shape
@triton.jit
def static_shape_kernel(ptr, N: tl.constexpr):  # N 是编译时常量
    # 可以在编译时展开循环,优化代码
    for i in range(N):  # 编译时展开
        ...

# 场景 2: 运行时才知道 shape
@triton.jit  
def dynamic_shape_kernel(ptr, N):  # N 是运行时值
    # Triton 无法在编译时优化
    # 循环保持为循环,代码质量下降
    for i in range(N):  # 运行时循环,无法展开
        ...

# 解决方案: 用 autotune 覆盖常见 shape
@triton.autotune(configs=[
    triton.Config({'BLOCK_SIZE': 512}),
    triton.Config({'BLOCK_SIZE': 1024}),
    triton.Config({'BLOCK_SIZE': 2048}),
], key=['n_elements'])
@triton.jit
def tuned_kernel(ptr, n_elements, BLOCK_SIZE: tl.constexpr):
    ...
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

8.2 复杂 Control Flow ​

python
# Triton 不擅长复杂的控制流

# ❌ 不好的例子: 复杂分支
@triton.jit
def bad_control_flow(x, mode: tl.constexpr):
    if mode == "relu":
        x = tl.where(x > 0, x, 0.0)
    elif mode == "gelu":
        x = 0.5 * x * (1.0 + tl.tanh(...))
    elif mode == "silu":
        x = x * tl.sigmoid(x)
    elif mode == "identity":
        pass
    
    # 更好的做法: 用 lookup table 或分开实现
    # 或者接受编译出来的代码包含所有分支

# ✅ 更好的例子: 简化控制流
@triton.jit
def simple_activation(x, do_gelu: tl.constexpr, do_silu: tl.constexpr):
    # 编译时消除不需要的分支
    if do_gelu:
        x = gelu(x)
    if do_silu:
        x = x * tl.sigmoid(x)
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

升华 ​

┌────────────────────────────────────────────────────────────────────────────┐
│                          🚀 Triton 核心原则                                  │
├────────────────────────────────────────────────────────────────────────────┤
│                                                                            │
│  ① 选 Triton 而不是手写 CUDA 的条件:                                       │
│     你需要定制 + 性能要求在 80-95% 内 + Python 优先                         │
│     否则应该用 cuBLAS/cuDNN                                               │
│                                                                            │
│  ② @triton.jit 把 Python 变成 GPU 代码:                                   │
│     编译器自动处理 coalesced access、shared memory、vectorization           │
│     你只需要指定 BLOCK 大小                                                 │
│                                                                            │
│  ③ Flash Attention 是 Triton 的杀手级应用:                                │
│     分块 online softmax 只用 O(1) 额外内存                                  │
│     Triton 实现比手写 CUDA 简洁得多                                         │
│                                                                            │
│  ④ autotune 是性能关键:                                                   │
│     不要手动猜 BLOCK_SIZE,用 autotune 搜索最优配置                         │
│                                                                            │
│  ⑤ 动态 shape 是弱点:                                                     │
│     编译时 shape 能达到最佳性能                                             │
│     运行时 shape 需要用 autotune 覆盖常见情况                               │
│                                                                            │
└────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24

"AI 可查 vs 必须理解"清单 ​

必须理解(不理解就等于不会):

  • 🔴 Blocked ISA 的概念:为什么指定 BLOCK 大小就能自动优化
  • 🔴 @triton.jit 的 JIT 编译模型:Python 函数如何变成 GPU 代码
  • 🔴 tl.load/tl.store 的 mask 机制:为什么需要 mask,越界访问会发生什么
  • 🔴 Flash Attention 的 online softmax:m 和 l 的递推公式,为什么能省内存
  • 🔴 autotune 的机制:Triton 如何在运行时选择最优配置

AI 可查(知道去哪查就行):

  • ✅ Triton 具体的 API 签名(官方文档)
  • ✅ 特定 GPU 架构的 optimal BLOCK_SIZE(benchmark 数据)
  • ✅ Triton 的底层代码生成(PTX/SASS 输出)
  • ✅ triton.Config 的完整参数(官方文档)
  • ✅ Triton 和 torch.compile 的集成细节(PyTorch 文档)

学习状态:🟡 开始学习

最后更新于:

Pager
上一篇20. TVM Tensor Expression 与计算调度分离 / TVM Tensor Expressions and Compute-Schedule Separation
下一篇22. 基于成本模型与实测搜索的自动调度 / Automatic Scheduling with Cost Models and Measurement-Based Search

持续记录,持续成长

Copyright © Tidenflow