跳转至

GPU Kernel 编程:Triton vs TileLang vs CUDA

更新日期:2026-04-15


一、Kernel 编程框架全景

框架 抽象层级 学习曲线 性能上限 典型场景
CUDA C++ 最底层 100% 库作者 / 极致优化
CUTLASS 中-底层 中-高 95-100% matmul / FlashAttention 内核
TileLang 中层(Tile DSL) 95% 自定义 Attention / MoE kernel
Triton 中-高层(Python) 90-95% 快速原型 / 科研
torch.compile 高层 极低 80-90% 业务代码自动加速
ThunderKittens 中层(C++ DSL) 95% H100 异构资源调度

实践笔记:如果你的目标是写自定义 Attention 变体(如 MLA、Sliding Window),建议从 Triton 入手验证正确性,再用 TileLang 优化性能。直接写 CUDA 的 ROI 很少值得,除非你在做基础设施级别的工作。


二、Triton 深入

Triton 是 OpenAI 开源的 Python-like DSL,专门用于写 GPU kernel。参考 Triton (Tillet et al., 2019)

2.1 Triton 编程模型

import triton
import triton.language as tl

@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
    # Triton 的核心概念: "program" = 一个 block
    # 你不直接管理 threads, Triton 编译器自动并行化

    pid = tl.program_id(0)  # 当前 block ID

    # 计算这个 block 负责的数据范围
    offsets = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offsets < n  # 边界处理

    # 加载 + 计算 + 存储
    x = tl.load(x_ptr + offsets, mask=mask)
    y = tl.load(y_ptr + offsets, mask=mask)
    tl.store(out_ptr + offsets, x + y, mask=mask)

# 启动:
def add(x, y):
    n = x.numel()
    out = torch.empty_like(x)
    grid = lambda meta: (triton.cdiv(n, meta['BLOCK']),)
    add_kernel[grid](x, y, out, n, BLOCK=1024)
    return out

2.2 Triton 的抽象层次

2.3 Triton Flash Attention (核心)

@triton.jit
def flash_attention_kernel(
    Q, K, V, Out,
    stride_qm, stride_kn, stride_vn, stride_om,
    sm_scale,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, D: tl.constexpr,
):
    # 每个 block 处理 Q 的一段
    block_idx = tl.program_id(0)

    # 加载 Q 的 block
    q_offs = block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
    q = tl.load(Q + q_offs[:, None] * stride_qm + tl.arange(0, D)[None, :])

    # 初始化 online softmax
    m = tl.full([BLOCK_M], -float('inf'), dtype=tl.float32)
    l = tl.zeros([BLOCK_M], dtype=tl.float32)
    o = tl.zeros([BLOCK_M, D], dtype=tl.float32)

    # 遍历所有 KV blocks
    for kv_start in range(0, N, BLOCK_N):
        # 加载 K, V
        k = tl.load(K + ...)
        v = tl.load(V + ...)

        # 计算注意力分数
        s = tl.dot(q, tl.trans(k)) * sm_scale

        # Online softmax 更新
        m_new = tl.maximum(m, tl.max(s, axis=1))
        alpha = tl.exp(m - m_new)
        p = tl.exp(s - m_new[:, None])
        l = alpha * l + tl.sum(p, axis=1)
        o = alpha[:, None] * o + tl.dot(p, v)
        m = m_new

    # 归一化并存储
    o = o / l[:, None]
    tl.store(Out + ..., o)

三、TileLang 深入

TileLang 是新兴的 tile-based DSL,抽象层次比 Triton 更高。DeepSeek 已在生产中使用。参考 TileLang (2024)

3.1 TileLang Flash Attention

from tilelang import Kernel, Tile

@Kernel
def flash_attention(
    Q: Tile[M, D], 
    K: Tile[N, D], 
    V: Tile[N, D]
) -> Tile[M, D]:
    # TileLang 以 Tile 为核心操作单元
    # 自动管理 shared memory / registers / tensor cores

    O = Tile.zeros(M, D)
    m = Tile.full(M, -float('inf'))
    l = Tile.zeros(M)

    for kv_tile in K.tiles(BLOCK_N):  # 自动分块迭代
        v_tile = V[kv_tile.index]

        # 矩阵乘法自动映射到 tensor core
        s = Q @ kv_tile.T * scale

        m_new = max(m, s.row_max())
        alpha = exp(m - m_new)
        p = exp(s - m_new.broadcast())
        l = alpha * l + p.row_sum()
        O = alpha.broadcast() * O + p @ v_tile
        m = m_new

    return O / l.broadcast()
    # 约 30 行 vs Triton 约 100 行
    # 性能: 与手写 CUDA 持平 (TileLang 论文数据)

3.2 TileLang 优势

3.3 TileLang FlashMLA

# MLA 的高性能 Kernel (DeepSeek-V3 使用)
@Kernel
def flash_mla(
    Q: Tile[M, D_q],
    C_kv: Tile[N, D_c],      # 压缩的 KV cache
    W_uk: Tile[D_c, D_k],    # 上投影矩阵
    W_uv: Tile[D_c, D_v]
) -> Tile[M, D_v]:

    O = Tile.zeros(M, D_v)
    m = Tile.full(M, -float('inf'))
    l = Tile.zeros(M)

    for c_tile in C_kv.tiles(BLOCK_N):
        # 关键: 在 shared memory 中即时恢复 K, V
        K_tile = c_tile @ W_uk
        V_tile = c_tile @ W_uv

        s = Q @ K_tile.T * scale
        # online softmax (同上)
        ...

    return O / l.broadcast()

四、何时手写 CUDA

场景 推荐 为什么
通用矩阵乘法 cuBLAS (不用写) NVIDIA 工程师针对每代 GPU 手动调优到极致,自己写不可能超过
标准 Attention Flash Attention library Tri Dao 团队持续优化(FA-⅔),直接调库最省事
自定义融合算子 Triton 或 TileLang 融合多个操作减少 HBM 读写是性能关键;库不支持的融合模式需要自己写
极致性能 + 熟悉 CUDA CUDA 当 Triton 编译器生成的代码离理论峰值差距 >20% 时,手写 CUDA 是最后手段
新硬件 / 新特性 CUDA (先行支持) Hopper 的 TMA、Warp Specialization 等新特性通常 CUDA 先支持,Triton 滞后 6-12 个月
快速原型 Triton 从想法到可运行 kernel 最快几小时,CUDA 可能要几天
生产部署 TileLang 或 CUDA TileLang 有跨平台优势;纯 NVIDIA 环境用 CUDA 更成熟

五、Warp Specialization (Hopper+)

H100 和 Blackwell 支持 Warp Specialization: 同一 block 内不同 warp 做不同工作。FlashAttention-3 的核心优化。

# 传统: 所有 warp 都做相同工作
# Warp Specialization: 
# - Producer warp: 从 HBM 加载数据 (用 TMA 指令)
# - Consumer warp: 计算 (用 Tensor Core)
# 两者 pipeline 并行执行

# Triton 3.x 支持:
@triton.jit
def flash_attn_hopper(Q, K, V, Out,
    num_consumer_groups: tl.constexpr = 2,
    num_buffers_warp_spec: tl.constexpr = 3,
):
    # Triton 编译器自动生成 producer-consumer warp 代码
    # 用户只需指定参数
    ...

参考 FlashAttention-3 (Shah et al., 2024)


六、工具链

工具 用途 何时用 / 备注
Triton 通用 GPU kernel DSL,Python 编写 自定义算子首选入口;PyTorch 2.x 的 torch.compile 后端之一
TileLang 高级 tile DSL,比 Triton 更抽象 需要跨平台或追求极致性能时用;DeepSeek 生产验证
CUTLASS NVIDIA C++ 高性能 GEMM/Attention 模板库 需要在 C++ 层面集成高性能 GEMM 时用;学习成本高但灵活
cuBLAS/cuDNN NVIDIA 官方黑盒库 GEMM/Conv 的默认选择,不用自己写;torch.mm 底层就是 cuBLAS
ROCm/HIP AMD GPU 编程栈(API 类似 CUDA) 用 AMD MI250/MI300 时必须;HIP 代码可用 hipify 从 CUDA 半自动转换
TVM 多后端编译器(CPU/GPU/NPU) 端侧部署和异构硬件编译优化;Apache 项目,社区活跃
XLA Google 编译器 (TPU 为主,也支持 GPU) 用 JAX/TPU 训练时的默认编译器;GPU 上不如 Triton 灵活
Nsight Compute NVIDIA Kernel profiling 工具 分析 kernel 的 occupancy、memory throughput、warp stall 原因的必备工具

参考文献


上级 · C. 分布式训练基础设施