从零开始学习 TileLang:从 Python、PyTorch 到高性能 GPU Kernel 的完整开发教程

前言

在大模型训练、推理和科学计算中,真正决定性能的往往不是上层 Python 代码,而是底层 GPU Kernel,例如:

  • 向量加法;
  • Reduce;
  • Softmax;
  • GEMM;
  • LayerNorm;
  • RMSNorm;
  • FlashAttention;
  • MoE;
  • 量化与反量化;
  • 自定义融合算子。

对于普通开发者来说,最常见的选择是 PyTorch。PyTorch 使用简单、生态完善,并且内置了大量高性能算子。

但当现有算子不能满足需求,或者需要把多个操作融合成一个 GPU Kernel 时,通常需要进入 CUDA、Triton、CUTLASS 或 TileLang 这一层。

TileLang 的目标,就是在开发效率和底层性能控制之间取得平衡:

使用接近 Python 的语法描述 GPU Tile、内存搬运、并行计算和流水线,并最终编译成高性能 GPU Kernel。

本文面向没有 GPU Kernel 开发经验的初学者,从最基础的向量加法开始,逐步讲解:

  • TileLang 是什么;
  • TileLang 的项目背景;
  • Python、PyTorch 和 TileLang 的差异;
  • GPU Kernel 的基本运行方式;
  • TileLang 的安装与开发流程;
  • 向量加法;
  • Reduce 与 Softmax;
  • GEMM;
  • LayerNorm 与 RMSNorm;
  • FlashAttention;
  • 正确性验证;
  • 性能测试;
  • 调试方法;
  • 自动调优;
  • 常见性能问题;
  • 完整学习路线。

第一章:TileLang 是什么

1.1 TileLang 的定义

TileLang,也称 Tile Language,是一种面向高性能 AI 算子的领域特定语言。

它主要用于编写:

  • GPU Kernel;
  • 深度学习基础算子;
  • 大模型推理算子;
  • 矩阵乘法;
  • Attention;
  • 归一化;
  • 量化算子;
  • 融合算子。

TileLang 使用 Python 风格的语法,但它不是普通的 Python 数值计算库。

开发者编写的 TileLang 代码会经过编译器转换,最终生成 CUDA、HIP 或其他硬件后端可以执行的 GPU Kernel。

一个简化的编译流程如下:

Python
   │
   ▼
TileLang DSL
   │
   ▼
TVM TIR
   │
   ▼
Compiler Pass
   │
   ▼
CUDA / HIP / Metal / LLVM
   │
   ▼
GPU Kernel

TileLang 建立在 TVM 和 TIR 等编译技术之上。


1.2 TileLang 的核心思想:Tile

GPU 不会把一个大型矩阵一次性全部放入高速片上内存。

通常会先把矩阵切分成多个小块:

完整矩阵
┌───────────────────────────┐
│ Tile │ Tile │ Tile │ Tile │
├──────┼──────┼──────┼──────┤
│ Tile │ Tile │ Tile │ Tile │
├──────┼──────┼──────┼──────┤
│ Tile │ Tile │ Tile │ Tile │
└───────────────────────────┘

每个 Tile 会经历以下过程:

Global Memory
      │
      │ 数据搬运
      ▼
Shared Memory
      │
      │ 并行计算
      ▼
Registers / Fragment
      │
      │ 写回
      ▼
Global Memory

TileLang 允许开发者直接表达:

  • 一个线程块处理哪个 Tile;
  • 一个 Tile 的大小;
  • 数据如何进入 Shared Memory;
  • 中间结果放在哪里;
  • 如何调用矩阵乘法指令;
  • 如何组织软件流水线;
  • 如何在不同线程之间归约。

因此,TileLang 的重点不是“逐元素写程序”,而是“描述 Tile 级的数据流”。


1.3 TileLang 的项目背景

严格来说,TileLang 不是某一家商业公司的闭源产品。

根据公开项目资料,TileLang 具有明显的学术界与工业界合作背景:

  • 项目最初由北京大学相关研究团队发起;
  • 早期开发过程中有微软研究院参与和指导;
  • 一部分工作完成于研究人员在微软研究院实习期间;
  • 后续主要由 TileAI 社区和开源贡献者持续维护;
  • 官方代码仓库位于 tile-ai 组织下。

因此,更准确的描述是:

TileLang 是由高校研究团队发起、产业研究机构参与、目前由 TileAI 社区维护的开源高性能 Kernel 编程项目。

它不是 NVIDIA、微软或 OpenAI 某一家公司的内部专属语言。


1.4 TileLang 与其他方案的区别

Python

优点:

  • 最容易编写;
  • 适合教学;
  • 适合快速验证算法。

缺点:

  • 循环由解释器执行;
  • Python 对象开销很大;
  • 通常无法直接利用 GPU;
  • 不适合大规模高性能计算。

PyTorch

优点:

  • 开发效率高;
  • 支持自动求导;
  • 内置大量高性能算子;
  • 支持 CPU、CUDA、ROCm 等平台;
  • 生态成熟。

缺点:

  • 单个算子的底层实现不容易修改;
  • 多个算子之间可能产生中间 Tensor;
  • 特殊形状和特殊数据布局不一定最优;
  • 自定义复杂融合 Kernel 较困难。

CUDA

优点:

  • 控制能力最强;
  • 可以直接操作线程、Warp、Shared Memory 和硬件指令;
  • 能实现极致性能。

缺点:

  • 学习成本高;
  • 代码复杂;
  • 调试困难;
  • 需要处理大量硬件细节。

CUTLASS 和 CuTe

优点:

  • 能构建高性能矩阵乘法和 Tensor Core Kernel;
  • 模板化程度高;
  • 性能强。

缺点:

  • C++ 模板复杂;
  • 学习曲线陡峭;
  • 对初学者不友好。

Triton

优点:

  • Python 风格;
  • 开发效率高;
  • 比 CUDA 更容易入门;
  • 已被大量深度学习框架采用。

缺点:

  • 对某些复杂数据流和精细硬件控制的表达能力有限;
  • 复杂 Attention、特殊 Pipeline 或特殊 Tensor Core 调度可能需要更多技巧。

TileLang

优点:

  • Python 风格;
  • 直接表达 Tile 级数据流;
  • 可以控制 Shared Memory、Fragment、流水线和 GEMM;
  • 适合复杂 AI Kernel;
  • 能够查看生成的底层代码;
  • 支持自动调优。

缺点:

  • 比普通 PyTorch 和 Triton 更接近硬件;
  • 需要理解 GPU 内存层级;
  • 需要理解线程块、Warp 和 Tile;
  • 项目仍在快速发展,API 可能变化。

可以粗略理解为:

易用性高
   ↑
PyTorch
Triton
TileLang
CUTLASS / CuTe
CUDA
   ↓
控制能力强

这只是帮助理解的近似关系,并不代表严格的功能高低排序。


第二章:Python、PyTorch 与 TileLang 的底层区别

2.1 纯 Python 如何执行

例如:

result = []

for i in range(len(a)):
    result.append(a[i] + b[i])

底层过程大致为:

Python 解释器
      ↓
读取 Python 对象
      ↓
执行一次 Python 加法
      ↓
创建结果对象
      ↓
写入 Python list
      ↓
进入下一次循环

即使真正的浮点加法非常快,Python 循环、对象管理和解释器调度仍会占据大量时间。


2.2 PyTorch 如何执行

PyTorch 代码可能只有一行:

result = a + b

但底层不会由 Python 逐元素相加。

实际过程更接近:

Python API
    ↓
PyTorch Dispatcher
    ↓
选择设备和数据类型对应的实现
    ↓
启动 CUDA Kernel
    ↓
大量 GPU 线程并行计算

所以 PyTorch 快的原因不是 Python 本身快,而是 Python 只负责调用已经编译好的底层实现。


2.3 TileLang 如何执行

TileLang 允许开发者编写自己的 Kernel:

@tilelang.jit
def add(A, B):
    ...

底层过程是:

Python 构建 Kernel
      ↓
TileLang DSL
      ↓
TIR 中间表示
      ↓
编译器优化
      ↓
生成 CUDA 或其他后端代码
      ↓
编译并加载 GPU Kernel
      ↓
GPU 执行

TileLang 的优势不是让 Python 循环变快,而是允许开发者定义更合适的 GPU 执行方式。


第三章:安装 TileLang 开发环境

3.1 推荐环境

以本文所依据的官方文档状态为参考,推荐环境包括:

  • Python 3.10 或更高版本;
  • Linux 开发环境;
  • NVIDIA CUDA GPU;
  • 支持 CUDA 的 PyTorch;
  • 较新的 CUDA Toolkit;
  • 较新的 NVIDIA 驱动。

官方也支持或正在扩展:

  • AMD ROCm;
  • CPU;
  • Metal;
  • WebGPU 或实验性后端;
  • 其他加速器。

NVIDIA CUDA 平台通常拥有最成熟的文档、示例和工具支持。

TileLang 项目仍在快速迭代。此前核对文档时,官方文档页面显示版本为 0.1.12,而 PyPI 稳定包搜索结果仍可能显示 0.1.11。实际使用时应以当前安装版本的官方文档和示例为准。


3.2 创建 Python 虚拟环境

python -m venv tilelang-env
source tilelang-env/bin/activate

升级安装工具:

python -m pip install --upgrade pip

安装 PyTorch 和 TileLang:

pip install torch tilelang

验证:

python -c "import torch; import tilelang; print(torch.__version__); print(tilelang.__version__)"

3.3 检查 CUDA

import torch

print("CUDA available:", torch.cuda.is_available())

if torch.cuda.is_available():
    print("GPU:", torch.cuda.get_device_name())
    print(
        "Compute capability:",
        torch.cuda.get_device_capability(),
    )

如果 torch.cuda.is_available() 返回 False,需要优先检查:

  • 是否安装了支持 CUDA 的 PyTorch;
  • NVIDIA 驱动是否正常;
  • CUDA 版本是否兼容;
  • 当前 Python 环境是否正确;
  • 是否在没有 GPU 的环境中运行。

3.4 安装主分支版本

需要最新功能时可以使用:

pip install git+https://github.com/tile-ai/tilelang.git

主分支可能包含更新的 API,也可能存在尚未稳定的变化。

对于初学者,建议:

  1. 优先使用稳定版本;
  2. 使用与安装版本匹配的官方文档;
  3. 示例报错时先核对版本;
  4. 不要直接混用不同版本的教程代码。

第四章:GPU 开发必须掌握的基础

4.1 GPU 线程层级

一个典型 CUDA Kernel 包含:

Grid
 ├── Block 0
 │    ├── Thread 0
 │    ├── Thread 1
 │    └── ...
 ├── Block 1
 └── ...

在 TileLang 中,常见写法是:

with T.Kernel(
    number_of_blocks,
    threads=256,
) as block_id:
    ...

其中:

  • number_of_blocks 表示线程块数量;
  • threads=256 表示每个线程块使用 256 个线程;
  • block_id 类似 CUDA 的 blockIdx.x

线程内并行循环通常使用:

for thread_id in T.Parallel(256):
    ...

4.2 Warp

NVIDIA GPU 通常以 Warp 为基本执行单位。

一个 Warp 通常包含 32 个线程。

Warp 内线程会以锁步方式执行相同指令。若同一个 Warp 内不同线程走不同分支,就可能产生分支发散。

这也是为什么 GPU Kernel 中会关注:

  • Warp 利用率;
  • Warp Reduce;
  • Warp 分工;
  • 分支发散;
  • Warp 级矩阵乘法。

4.3 GPU 内存层级

Global Memory

特点:

  • 容量大;
  • 所有线程可以访问;
  • 延迟较高;
  • 对应输入和输出 Tensor;
  • 数据通常位于显存 HBM 或 GDDR 中。

Shared Memory

TileLang 中常用:

T.alloc_shared(...)

特点:

  • 一个线程块内部共享;
  • 速度远快于 Global Memory;
  • 容量有限;
  • 适合保存重复使用的 Tile。

Registers 与 Fragment

TileLang 中常用:

T.alloc_fragment(...)

特点:

  • 通常位于寄存器;
  • 访问速度最快;
  • 每个线程拥有自己的寄存器;
  • 容量非常有限;
  • 太多寄存器会降低 Occupancy;
  • 寄存器不足时可能溢出到 Local Memory。

4.4 为什么要做分块

假设矩阵乘法直接让每个线程从 Global Memory 读取数据:

线程 0 读取 A 和 B
线程 1 再次读取相同的 A 和 B
线程 2 再次读取相同的 A 和 B

相同数据会被重复读取很多次。

分块后:

一个线程块
    ↓
读取 A 的一个 Tile 到 Shared Memory
    ↓
读取 B 的一个 Tile 到 Shared Memory
    ↓
多个线程重复使用这些 Tile
    ↓
完成一部分矩阵乘法

这样可以显著提高数据复用率。


第五章:TileLang 的基本语法

5.1 @tilelang.jit

@tilelang.jit
def operation(...):
    ...

它表示该函数需要被 TileLang JIT 编译。

第一次运行时可能经历:

  • Python 函数执行;
  • TIR 构建;
  • 编译器优化;
  • GPU 代码生成;
  • CUDA 编译;
  • Module 加载。

因此,第一次调用通常比后续调用慢。


5.2 T.Tensor

A: T.Tensor((M, N), T.float16)

它声明:

  • Tensor 名称为 A
  • 形状为 M × N
  • 数据类型为 FP16。

5.3 T.empty

C = T.empty((M, N), T.float16)

它创建输出 Tensor。


5.4 T.Kernel

with T.Kernel(
    grid_x,
    grid_y,
    threads=128,
) as (block_x, block_y):
    ...

它定义 GPU Kernel 的线程网格。


5.5 T.Parallel

for i, j in T.Parallel(M, N):
    ...

表示两个维度上的并行循环。

它适合:

  • 逐元素计算;
  • 数据格式转换;
  • Scale;
  • Bias;
  • 激活函数;
  • 简单布局变换。

5.6 T.alloc_shared

tile = T.alloc_shared(
    (block_m, block_k),
    T.float16,
)

分配 Shared Memory。


5.7 T.alloc_fragment

accumulator = T.alloc_fragment(
    (block_m, block_n),
    T.float32,
)

分配寄存器级计算片段。


5.8 T.copy

T.copy(
    A[source_row, source_column],
    a_shared,
)

用于在不同内存层级之间复制 Tile。

常见方向包括:

Global Memory → Shared Memory
Global Memory → Fragment
Shared Memory → Fragment
Fragment → Shared Memory
Fragment → Global Memory

5.9 T.Pipelined

for k_block in T.Pipelined(
    number_of_blocks,
    num_stages=3,
):
    ...

用于构建软件流水线。

理想状态是让:

加载下一块数据

和:

计算当前块数据

发生重叠。

流水级越多不一定越快,因为每个流水阶段都可能需要额外的 Shared Memory 和寄存器。


5.10 T.gemm

T.gemm(
    a_shared,
    b_shared,
    accumulator,
)

执行 Tile 级矩阵乘法。

编译器会尝试将其映射到:

  • Tensor Core;
  • Warp-level MMA;
  • AMD Matrix Core;
  • 其他硬件矩阵乘法指令。

5.11 Reduce 原语

常见原语包括:

T.reduce_sum(...)
T.reduce_max(...)
T.reduce_min(...)

它们可以被降低为:

  • 线程内 Reduce;
  • Warp Reduce;
  • 线程块 Reduce;
  • 多阶段 Reduce。

5.12 两种常见编程风格

风格一:返回 T.prim_func

@tilelang.jit
def create_kernel(N: int):
    @T.prim_func
    def kernel(
        A: T.Tensor((N,), T.float32),
        C: T.Tensor((N,), T.float32),
    ):
        ...

    return kernel

调用:

kernel = create_kernel(1024)
kernel(a, c)

这种写法适合理解 Kernel 本身。

风格二:高级 @tilelang.jit

@tilelang.jit
def operation(A, B):
    M, N = T.const("M, N")

    A: T.Tensor((M, N), T.float32)
    B: T.Tensor((M, N), T.float32)

    C = T.empty((M, N), T.float32)

    with T.Kernel(...):
        ...

    return C

调用:

c = operation(a, b)

这种写法更接近普通 PyTorch 函数,也是当前官方示例中常见的形式。


第六章:正确性验证与性能测试

6.1 为什么先验证正确性

GPU Kernel 很容易出现:

  • 越界;
  • 数据类型错误;
  • Reduce 未清零;
  • 矩阵转置方向错误;
  • Shared Memory 布局错误;
  • 数值精度问题;
  • 尾部元素未处理;
  • 不同线程间同步错误。

因此开发顺序必须是:

先正确
  ↓
再稳定
  ↓
最后优化性能

6.2 使用 PyTorch 作为参考实现

推荐:

torch.testing.assert_close(
    actual,
    expected,
    rtol=1e-2,
    atol=1e-2,
)

对于 FP16、BF16、并行 Reduce 和不同累加顺序,通常不应该要求逐位完全相同。


6.3 不要直接用 time.time() 测 CUDA

CUDA Kernel 默认是异步执行的。

错误示例:

start = time.time()
result = operation(x)
end = time.time()

operation(x) 可能只是把任务提交给 GPU,CPU 计时结束时 GPU 还没有执行完成。


6.4 使用 TileLang Profiler

from tilelang.profiler import do_bench


def benchmark(fn, warmup=25, rep=100):
    return do_bench(
        fn,
        warmup=warmup,
        rep=rep,
        backend="event",
    )

调用:

latency_ms = benchmark(
    lambda: operation(x)
)

print(latency_ms)

do_bench 可以处理:

  • 预热;
  • 重复运行;
  • CUDA Event 计时;
  • 统计延迟;
  • 某些缓存控制;
  • CUPTI 或 CUDA Graph 后端。

6.5 排除 JIT 编译时间

推荐先编译:

compiled = operation.compile(...)

再运行:

compiled(...)

最后 benchmark:

do_bench(
    lambda: compiled(...),
    warmup=25,
    rep=100,
)

6.6 查看生成的 Kernel

print(compiled.get_kernel_source())

需要检查:

  • 是否生成了预期的 CUDA 代码;
  • 是否使用 Tensor Core;
  • 是否出现 Local Memory;
  • 是否有大量寄存器;
  • 是否正确向量化;
  • 是否正确生成异步拷贝;
  • 是否存在低效循环。

第七章:向量加法

目标:

Ci=Ai+BiC_i=A_i+B_iCi=Ai+Bi


7.1 纯 Python 实现

def python_vector_add(a, b):
    if len(a) != len(b):
        raise ValueError(
            "a and b must have the same length"
        )

    output = [0.0] * len(a)

    for i in range(len(a)):
        output[i] = a[i] + b[i]

    return output

底层特点:

  • Python 解释器逐次循环;
  • 每个数通常是 Python 对象;
  • 无法利用 GPU 大规模并行;
  • 适合小数据验证。

7.2 PyTorch 实现

import torch


def torch_vector_add(a, b):
    return a + b


n = 1 << 20

a = torch.randn(
    n,
    device="cuda",
    dtype=torch.float32,
)
b = torch.randn_like(a)

c = torch_vector_add(a, b)

PyTorch 会调用底层 Elementwise CUDA Kernel。


7.3 TileLang 实现

import torch
import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_vector_add(
    A,
    B,
    block_size: int = 256,
):
    N = T.const("N")

    A: T.Tensor((N,), T.float32)
    B: T.Tensor((N,), T.float32)

    C = T.empty((N,), T.float32)

    with T.Kernel(
        T.ceildiv(N, block_size),
        threads=block_size,
    ) as block_id:
        for thread_id in T.Parallel(block_size):
            index = (
                block_id * block_size
                + thread_id
            )

            if index < N:
                C[index] = A[index] + B[index]

    return C

测试:

n = 1 << 20

a = torch.randn(
    n,
    device="cuda",
    dtype=torch.float32,
)
b = torch.randn_like(a)

c_tilelang = tilelang_vector_add(a, b)
c_torch = a + b

torch.testing.assert_close(
    c_tilelang,
    c_torch,
)

7.4 底层执行方式

一个线程处理一个或少量元素:

Block 0
 ├── Thread 0 → C[0]
 ├── Thread 1 → C[1]
 ├── Thread 2 → C[2]
 └── ...

Block 1
 ├── Thread 0 → C[256]
 ├── Thread 1 → C[257]
 └── ...

7.5 性能差异

向量加法只做一次加法,却需要:

  • 读取一个 A
  • 读取一个 B
  • 写入一个 C

因此向量加法通常是内存带宽受限。

常见关系:

纯 Python 远慢于 PyTorch
PyTorch 通常接近 TileLang

TileLang 不一定能在单独的向量加法上击败 PyTorch。

TileLang 的机会主要来自融合。

例如 PyTorch:

y = a + b
z = torch.relu(y)
output = z * scale

可能对应多个 Kernel 和多个中间 Tensor。

TileLang 可以把它们融合成一个 Kernel:

output[index] = T.max(
    A[index] + B[index],
    0,
) * scale

这样可以减少:

  • Kernel 启动;
  • 中间 Tensor;
  • HBM 读写。

第八章:Reduce

8.1 Row Sum

目标:

Yi=∑j=0N−1Xi,jY_i=\sum_{j=0}^{N-1}X_{i,j}Yi=j=0N1Xi,j


8.2 纯 Python

def python_row_sum(matrix):
    output = []

    for row in matrix:
        total = 0.0

        for value in row:
            total += value

        output.append(total)

    return output

8.3 PyTorch

def torch_row_sum(x):
    return x.sum(dim=-1)

8.4 TileLang

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_row_sum(X):
    M, N = T.const("M, N")

    X: T.Tensor((M, N), T.float32)
    Y = T.empty((M,), T.float32)

    with T.Kernel(
        M,
        threads=128,
    ) as row:
        values = T.alloc_fragment(
            (1, N),
            T.float32,
        )
        total = T.alloc_fragment(
            (1,),
            T.float32,
        )

        T.copy(
            X[row, 0],
            values,
        )

        T.reduce_sum(
            values,
            total,
            dim=1,
            clear=True,
        )

        Y[row] = total[0]

    return Y

8.5 Reduce 的底层难点

Reduce 不是简单逐元素操作。

多个线程可能分别计算局部和:

Thread 0 → 局部和 0
Thread 1 → 局部和 1
Thread 2 → 局部和 2
...

然后需要继续合并:

局部和
  ↓
Warp Reduce
  ↓
Block Reduce
  ↓
最终结果

性能受以下因素影响:

  • Warp Shuffle;
  • Shared Memory;
  • 同步次数;
  • 行宽;
  • 线程数;
  • 数据类型;
  • 是否存在跨 Block Reduce。

第九章:Softmax

9.1 数学定义

稳定 Softmax 通常写为:

mi=max⁡jXi,jm_i=\max_jX_{i,j}mi=jmaxXi,j

Ei,j=exp⁡(Xi,j−mi)E_{i,j}=\exp(X_{i,j}-m_i)Ei,j=exp(Xi,jmi)

Yi,j=Ei,j∑kEi,kY_{i,j}=\frac{E_{i,j}}{\sum_kE_{i,k}}Yi,j=kEi,kEi,j

减去最大值是为了避免指数溢出。


9.2 纯 Python

import math


def python_softmax(row):
    maximum = max(row)

    exponentials = [
        math.exp(value - maximum)
        for value in row
    ]

    denominator = sum(exponentials)

    return [
        value / denominator
        for value in exponentials
    ]


def python_row_softmax(matrix):
    return [
        python_softmax(row)
        for row in matrix
    ]

9.3 PyTorch

import torch


def torch_softmax(x):
    return torch.softmax(
        x,
        dim=-1,
    )

torch.softmax 是专门优化的底层算子。

不应该简单认为它等价于多个独立的 PyTorch 操作:

maximum = x.max(
    dim=-1,
    keepdim=True,
).values

exponentials = torch.exp(
    x - maximum
)

output = exponentials / exponentials.sum(
    dim=-1,
    keepdim=True,
)

显式组合写法可能产生:

  • 多个 Kernel;
  • 多个中间 Tensor;
  • 额外显存读写。

不过,torch.compile 有可能捕获这些操作并进行融合。


9.4 TileLang 教学版

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_softmax(X):
    M, N = T.const("M, N")

    input_dtype = T.float16
    accumulation_dtype = T.float32

    X: T.Tensor((M, N), input_dtype)
    Y = T.empty((M, N), input_dtype)

    with T.Kernel(
        M,
        threads=128,
    ) as row:
        x = T.alloc_fragment(
            (1, N),
            input_dtype,
        )
        exponentials = T.alloc_fragment(
            (1, N),
            accumulation_dtype,
        )
        output = T.alloc_fragment(
            (1, N),
            input_dtype,
        )

        row_max = T.alloc_fragment(
            (1,),
            input_dtype,
        )
        row_sum = T.alloc_fragment(
            (1,),
            accumulation_dtype,
        )

        T.copy(
            X[row, 0],
            x,
        )

        T.reduce_max(
            x,
            row_max,
            dim=1,
            clear=True,
        )

        for _, column in T.Parallel(1, N):
            value = T.Cast(
                accumulation_dtype,
                x[0, column],
            )
            maximum = T.Cast(
                accumulation_dtype,
                row_max[0],
            )

            exponentials[0, column] = T.exp(
                value - maximum
            )

        T.reduce_sum(
            exponentials,
            row_sum,
            dim=1,
            clear=True,
        )

        for _, column in T.Parallel(1, N):
            output[0, column] = T.Cast(
                input_dtype,
                exponentials[0, column]
                / row_sum[0],
            )

        T.copy(
            output,
            Y[row, 0],
        )

    return Y

9.5 为什么这是教学版

该实现把整行放入 Fragment。

N 很大时,会出现:

  • 寄存器需求过高;
  • 寄存器溢出;
  • Occupancy 降低;
  • 编译失败;
  • 性能变差。

因此真实高性能 Softmax 通常采用:

  • 分块;
  • Warp Reduce;
  • 多次遍历;
  • Online Softmax;
  • Split;
  • Shared Memory;
  • 特殊线程映射。

9.6 Online Softmax

对于很长的一行,可以分 Tile 读取。

每次读取一个 Tile,并更新:

  • 当前最大值;
  • 当前指数和;
  • 当前归一化状态。

官方示例还可能使用:

exp2(x × log2(e))

代替直接计算自然指数。

关系为:

exp⁡(x)=2xlog⁡2(e)\exp(x)=2^{x\log_2(e)}exp(x)=2xlog2(e)

这是因为某些 GPU 上的二进制指数指令更适合底层实现。


9.7 Softmax 性能分析

Softmax 包含:

  • 最大值 Reduce;
  • 指数计算;
  • 求和 Reduce;
  • 除法;
  • 同步;
  • 多次内存访问。

影响性能的因素包括:

  • 行宽;
  • Warp 数量;
  • Reduce 策略;
  • 指数函数实现;
  • 是否产生中间 Tensor;
  • 是否与 Mask、Scale 或 Dropout 融合。

单独比较时:

纯 Python 远慢于 PyTorch
PyTorch Softmax 通常已经很强
调优后的 TileLang 可以接近或超过特定形状

TileLang 更有价值的场景通常是:

  • Softmax 与 Mask 融合;
  • Softmax 与 Scale 融合;
  • Softmax 与 Dropout 融合;
  • 特殊行宽;
  • Online Softmax;
  • Attention 内部融合。

第十章:GEMM 矩阵乘法

10.1 数学定义

Ci,j=∑k=0K−1Ai,kBk,jC_{i,j}=\sum_{k=0}^{K-1}A_{i,k}B_{k,j}Ci,j=k=0K1Ai,kBk,j

理论计算量约为:

2MNK2MNK2MNK

次浮点操作。


10.2 纯 Python

def python_matmul(a, b):
    m = len(a)
    k = len(a[0])
    n = len(b[0])

    output = [
        [0.0 for _ in range(n)]
        for _ in range(m)
    ]

    for i in range(m):
        for j in range(n):
            total = 0.0

            for reduction_index in range(k):
                total += (
                    a[i][reduction_index]
                    * b[reduction_index][j]
                )

            output[i][j] = total

    return output

该实现适合帮助理解矩阵乘法,但性能非常低。


10.3 PyTorch

def torch_matmul(a, b):
    return a @ b

在 GPU 上,PyTorch 通常会调用高度优化的矩阵乘法后端。

对于标准 GEMM,PyTorch 是非常强的性能基线。


10.4 TileLang

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_matmul(
    A,
    B,
    block_m: int = 128,
    block_n: int = 128,
    block_k: int = 32,
):
    M, N, K = T.const("M, N, K")

    input_dtype = T.float16
    accumulation_dtype = T.float32

    A: T.Tensor(
        (M, K),
        input_dtype,
    )
    B: T.Tensor(
        (K, N),
        input_dtype,
    )

    C = T.empty(
        (M, N),
        input_dtype,
    )

    with T.Kernel(
        T.ceildiv(N, block_n),
        T.ceildiv(M, block_m),
        threads=128,
    ) as (block_x, block_y):
        a_shared = T.alloc_shared(
            (block_m, block_k),
            input_dtype,
        )
        b_shared = T.alloc_shared(
            (block_k, block_n),
            input_dtype,
        )
        c_fragment = T.alloc_fragment(
            (block_m, block_n),
            accumulation_dtype,
        )

        T.clear(c_fragment)

        for reduction_block in T.Pipelined(
            T.ceildiv(K, block_k),
            num_stages=3,
        ):
            T.copy(
                A[
                    block_y * block_m,
                    reduction_block * block_k,
                ],
                a_shared,
            )

            T.copy(
                B[
                    reduction_block * block_k,
                    block_x * block_n,
                ],
                b_shared,
            )

            T.gemm(
                a_shared,
                b_shared,
                c_fragment,
            )

        T.copy(
            c_fragment,
            C[
                block_y * block_m,
                block_x * block_n,
            ],
        )

    return C

10.5 调用与验证

import torch


m = 1024
n = 1024
k = 1024

a = torch.randn(
    m,
    k,
    device="cuda",
    dtype=torch.float16,
)

b = torch.randn(
    k,
    n,
    device="cuda",
    dtype=torch.float16,
)

compiled = tilelang_matmul.compile(
    M=m,
    N=n,
    K=k,
    block_m=128,
    block_n=128,
    block_k=32,
)

c_tilelang = compiled(a, b)
c_torch = a @ b

torch.testing.assert_close(
    c_tilelang,
    c_torch,
    rtol=1e-2,
    atol=1e-2,
)

print(
    compiled.get_kernel_source()
)

10.6 GEMM 的底层数据流

A 的 Tile
    │
    ▼
Shared Memory
    │
    ├─────────┐
    │         │
    ▼         ▼
Tensor Core / MMA
    ▲         ▲
    │         │
    ├─────────┘
    │
B 的 Tile

计算结果累积在 Fragment 中:

FP16 输入
   ↓
Tensor Core 乘法
   ↓
FP32 累加
   ↓
转换为 FP16 或 BF16 输出

常见策略为:

输入:FP16 或 BF16
乘法:Tensor Core
累加:FP32
输出:FP16 或 BF16

10.7 PyTorch 与 TileLang 谁更快

对于标准 GEMM:

  • PyTorch 后端通常已经高度优化;
  • 未调优 TileLang 可能更慢;
  • Tile 大小不合适会明显降低性能;
  • 线程数和流水级不合适也会降低性能;
  • 不同 GPU 的最佳配置不同。

TileLang 更可能具有优势的场景包括:

  • 特殊矩阵形状;
  • 小矩阵;
  • 非常长或非常窄的矩阵;
  • 量化 GEMM;
  • 稀疏 GEMM;
  • 特殊数据布局;
  • GEMM 与 Bias 融合;
  • GEMM 与 Activation 融合;
  • GEMM 与其他 Epilogue 融合。

10.8 GEMM Epilogue 融合

PyTorch:

c = torch.relu(
    a @ b + bias
)

逻辑上包含:

  1. 矩阵乘法;
  2. Bias Add;
  3. ReLU。

TileLang 可以在写回之前处理:

for i, j in T.Parallel(
    block_m,
    block_n,
):
    c_fragment[i, j] = T.max(
        c_fragment[i, j]
        + bias[
            block_x * block_n + j
        ],
        0,
    )

融合后可能减少:

  • 输出中间 Tensor;
  • 再次读取矩阵结果;
  • 第二次或第三次 Kernel 启动;
  • HBM 流量。

10.9 GEMM 吞吐量

矩阵乘法的理论浮点操作数为:

FLOPs⁡=2MNK\operatorname{FLOPs}=2MNKFLOPs=2MNK

如果延迟单位为毫秒:

def gemm_tflops(
    m,
    n,
    k,
    latency_ms,
):
    operations = (
        2.0 * m * n * k
    )

    return (
        operations
        / latency_ms
        * 1e-9
    )

返回值近似为 TFLOP/s。


第十一章:LayerNorm

11.1 数学定义

对每一行:

μi=1N∑jXi,j\mu_i=\frac{1}{N}\sum_jX_{i,j}μi=N1jXi,j

σi2=1N∑jXi,j2−μi2\sigma_i^2=\frac{1}{N}\sum_jX_{i,j}^2-\mu_i^2σi2=N1jXi,j2μi2

Yi,j=γj(Xi,j−μi)1σi2+ϵ+βjY_{i,j}=\gamma_j(X_{i,j}-\mu_i)\frac{1}{\sqrt{\sigma_i^2+\epsilon}}+\beta_jYi,j=γj(Xi,jμi)σi2+ϵ1+βj


11.2 纯 Python

import math


def python_layer_norm(
    row,
    gamma,
    beta,
    eps=1e-5,
):
    width = len(row)

    mean = sum(row) / width

    variance = (
        sum(
            (value - mean) ** 2
            for value in row
        )
        / width
    )

    inverse_std = (
        1.0
        / math.sqrt(
            variance + eps
        )
    )

    return [
        gamma[i]
        * (row[i] - mean)
        * inverse_std
        + beta[i]
        for i in range(width)
    ]

11.3 PyTorch

import torch.nn.functional as F


def torch_layer_norm(
    x,
    gamma,
    beta,
    eps=1e-5,
):
    return F.layer_norm(
        x,
        normalized_shape=(
            x.shape[-1],
        ),
        weight=gamma,
        bias=beta,
        eps=eps,
    )

生产代码应优先使用 F.layer_norm,而不是手动拆成多个 Tensor 操作。


11.4 TileLang

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_layer_norm(
    X,
    Gamma,
    Beta,
    eps: float = 1e-5,
    block_m: int = 1,
):
    M, N = T.const("M, N")

    input_dtype = T.float16
    accumulation_dtype = T.float32

    X: T.Tensor(
        (M, N),
        input_dtype,
    )
    Gamma: T.Tensor(
        (N,),
        input_dtype,
    )
    Beta: T.Tensor(
        (N,),
        input_dtype,
    )

    Y = T.empty(
        (M, N),
        input_dtype,
    )

    with T.Kernel(
        T.ceildiv(M, block_m),
        threads=256,
    ) as block:
        x_shared = T.alloc_shared(
            (block_m, N),
            input_dtype,
        )
        gamma_shared = T.alloc_shared(
            (N,),
            input_dtype,
        )
        beta_shared = T.alloc_shared(
            (N,),
            input_dtype,
        )

        x = T.alloc_fragment(
            (block_m, N),
            accumulation_dtype,
        )
        x_squared = T.alloc_fragment(
            (block_m, N),
            accumulation_dtype,
        )

        row_sum = T.alloc_fragment(
            (block_m,),
            accumulation_dtype,
        )
        row_square_sum = T.alloc_fragment(
            (block_m,),
            accumulation_dtype,
        )
        row_mean = T.alloc_fragment(
            (block_m,),
            accumulation_dtype,
        )
        row_inverse_std = T.alloc_fragment(
            (block_m,),
            accumulation_dtype,
        )

        T.copy(
            X[
                block * block_m,
                0,
            ],
            x_shared,
        )

        T.copy(
            Gamma,
            gamma_shared,
        )

        T.copy(
            Beta,
            beta_shared,
        )

        for i, j in T.Parallel(
            block_m,
            N,
        ):
            value = T.Cast(
                accumulation_dtype,
                x_shared[i, j],
            )

            x[i, j] = value
            x_squared[i, j] = (
                value * value
            )

        T.reduce_sum(
            x,
            row_sum,
            dim=1,
            clear=True,
        )

        T.reduce_sum(
            x_squared,
            row_square_sum,
            dim=1,
            clear=True,
        )

        inverse_width = (
            T.float32(1.0)
            / T.Cast(
                accumulation_dtype,
                N,
            )
        )

        for i in T.Parallel(block_m):
            mean = (
                row_sum[i]
                * inverse_width
            )

            variance = (
                row_square_sum[i]
                * inverse_width
                - mean * mean
            )

            row_mean[i] = mean

            row_inverse_std[i] = T.rsqrt(
                variance + eps
            )

        for i, j in T.Parallel(
            block_m,
            N,
        ):
            normalized = (
                x[i, j]
                - row_mean[i]
            ) * row_inverse_std[i]

            result = (
                normalized
                * T.Cast(
                    accumulation_dtype,
                    gamma_shared[j],
                )
                + T.Cast(
                    accumulation_dtype,
                    beta_shared[j],
                )
            )

            x_shared[i, j] = T.Cast(
                input_dtype,
                result,
            )

        T.copy(
            x_shared,
            Y[
                block * block_m,
                0,
            ],
        )

    return Y

11.5 为什么 LayerNorm 适合融合

手动拆开的 PyTorch 版本:

mean = x.mean(
    dim=-1,
    keepdim=True,
)

centered = x - mean

variance = centered.square().mean(
    dim=-1,
    keepdim=True,
)

normalized = (
    centered
    * torch.rsqrt(
        variance + eps
    )
)

output = (
    normalized * gamma
    + beta
)

可能产生多个中间 Tensor。

专用 Kernel 可以让:

  • 均值;
  • 方差;
  • 归一化;
  • Scale;
  • Bias;

在一个或少量 Kernel 中完成。

不过,F.layer_norm 本身也已经是优化算子。因此正确的性能比较对象应当是:

F.layer_norm(...)

而不是教学用的拆分实现。


11.6 自动求导

官方 LayerNorm 示例不仅包含前向,还可能实现反向传播。

常见接入方式是:

class CustomLayerNormFunction(
    torch.autograd.Function
):
    @staticmethod
    def forward(...):
        ...

    @staticmethod
    def backward(...):
        ...

然后在 forwardbackward 中分别调用 TileLang Kernel。

这意味着 TileLang 不只可以写推理算子,也可以写训练需要的前向和反向 Kernel。


第十二章:RMSNorm

12.1 数学定义

RMSNorm 不计算均值,也不减均值:

ri=11N∑jXi,j2+ϵr_i=\frac{1}{\sqrt{\frac{1}{N}\sum_jX_{i,j}^2+\epsilon}}ri=N1jXi,j2+ϵ1

Yi,j=Xi,jriγjY_{i,j}=X_{i,j}r_i\gamma_jYi,j=Xi,jriγj


12.2 纯 Python

import math


def python_rms_norm(
    row,
    weight,
    eps=1e-6,
):
    mean_square = (
        sum(
            value * value
            for value in row
        )
        / len(row)
    )

    inverse_rms = (
        1.0
        / math.sqrt(
            mean_square + eps
        )
    )

    return [
        row[i]
        * inverse_rms
        * weight[i]
        for i in range(len(row))
    ]

12.3 PyTorch

import torch


def torch_rms_norm(
    x,
    weight,
    eps=1e-6,
):
    inverse_rms = torch.rsqrt(
        x.float()
        .square()
        .mean(
            dim=-1,
            keepdim=True,
        )
        + eps
    )

    return (
        x.float()
        * inverse_rms
        * weight.float()
    ).to(x.dtype)

如果当前 PyTorch 版本提供专用 RMSNorm API,应优先将专用 API 作为性能基线。


12.4 TileLang

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_rms_norm(
    X,
    Weight,
    eps: float = 1e-6,
    block_m: int = 1,
):
    M, N = T.const("M, N")

    input_dtype = T.float16
    accumulation_dtype = T.float32

    X: T.Tensor(
        (M, N),
        input_dtype,
    )

    Weight: T.Tensor(
        (N,),
        input_dtype,
    )

    Y = T.empty(
        (M, N),
        input_dtype,
    )

    with T.Kernel(
        T.ceildiv(M, block_m),
        threads=128,
    ) as block:
        values = T.alloc_fragment(
            (block_m, N),
            accumulation_dtype,
        )
        squares = T.alloc_fragment(
            (block_m, N),
            accumulation_dtype,
        )
        square_sum = T.alloc_fragment(
            (block_m,),
            accumulation_dtype,
        )
        inverse_rms = T.alloc_fragment(
            (block_m,),
            accumulation_dtype,
        )
        output = T.alloc_fragment(
            (block_m, N),
            input_dtype,
        )

        T.copy(
            X[
                block * block_m,
                0,
            ],
            values,
        )

        for i, j in T.Parallel(
            block_m,
            N,
        ):
            squares[i, j] = (
                values[i, j]
                * values[i, j]
            )

        T.reduce_sum(
            squares,
            square_sum,
            dim=1,
            clear=True,
        )

        for i in T.Parallel(block_m):
            inverse_rms[i] = T.rsqrt(
                square_sum[i] / N
                + eps
            )

        for i, j in T.Parallel(
            block_m,
            N,
        ):
            output[i, j] = T.Cast(
                input_dtype,
                values[i, j]
                * inverse_rms[i]
                * T.Cast(
                    accumulation_dtype,
                    Weight[j],
                ),
            )

        T.copy(
            output,
            Y[
                block * block_m,
                0,
            ],
        )

    return Y

12.5 Split-K 或分块版本

当归一化维度很大时,不适合把整行放入 Fragment。

可以采用:

第一次遍历
    ↓
分块累计平方和
    ↓
计算 inverse RMS
    ↓
第二次遍历
    ↓
归一化并写出

这种方式可以降低:

  • 寄存器压力;
  • Fragment 大小;
  • 编译复杂度;
  • Local Memory 溢出风险。

12.6 LayerNorm 与 RMSNorm 的性能区别

RMSNorm 不需要:

  • 求均值;
  • 减去均值;
  • 计算 Bias,具体取决于定义。

所以其计算过程通常比 LayerNorm 简单。

但二者在 GPU 上经常都是内存带宽受限:

  • 每个元素计算量不高;
  • 需要读取输入;
  • 需要读取权重;
  • 需要写出结果;
  • Reduce 需要线程协作。

TileLang 的优势主要来自融合,例如:

Residual Add
    ↓
RMSNorm
    ↓
Quantization

可以合并到一个 Kernel 中。


第十三章:FlashAttention

13.1 普通 Attention

Scaled Dot-Product Attention:

S=QKTDS=\frac{QK^T}{\sqrt{D}}S=DQKT

P=softmax⁡(S)P=\operatorname{softmax}(S)P=softmax(S)

O=PVO=PVO=PV

其中:

  • Q 是 Query;
  • K 是 Key;
  • V 是 Value;
  • D 是每个 Attention Head 的维度。

13.2 纯 Python

下面是单头二维教学版:

import math


def dot(a, b):
    return sum(
        x * y
        for x, y in zip(a, b)
    )


def python_attention(q, k, v):
    sequence_length = len(q)
    dimension = len(q[0])

    scale = (
        1.0
        / math.sqrt(dimension)
    )

    output = []

    for query_row in q:
        scores = [
            dot(
                query_row,
                key_row,
            ) * scale
            for key_row in k
        ]

        maximum = max(scores)

        probabilities = [
            math.exp(
                score - maximum
            )
            for score in scores
        ]

        denominator = sum(
            probabilities
        )

        probabilities = [
            probability
            / denominator
            for probability in probabilities
        ]

        output_row = [
            0.0
            for _ in range(dimension)
        ]

        for key_index in range(
            sequence_length
        ):
            for column in range(
                dimension
            ):
                output_row[column] += (
                    probabilities[key_index]
                    * v[key_index][column]
                )

        output.append(output_row)

    return output

13.3 PyTorch 直观实现

import math
import torch


def torch_attention_naive(
    q,
    k,
    v,
):
    scores = torch.matmul(
        q,
        k.transpose(-2, -1),
    ) / math.sqrt(
        q.shape[-1]
    )

    probabilities = torch.softmax(
        scores,
        dim=-1,
    )

    return torch.matmul(
        probabilities,
        v,
    )

这段实现会显式生成:

  • Score Tensor;
  • Probability Tensor。

当序列很长时,二者可能占用大量显存。


13.4 PyTorch SDPA

生产环境应优先比较:

import torch.nn.functional as F


def torch_attention_sdpa(
    q,
    k,
    v,
    causal=False,
):
    return (
        F.scaled_dot_product_attention(
            q,
            k,
            v,
            is_causal=causal,
        )
    )

PyTorch 会根据:

  • GPU;
  • 数据类型;
  • 输入形状;
  • Mask;
  • Causal 设置;
  • 后端可用性;

选择合适的 Attention 实现。

这可能包括:

  • 融合 Attention;
  • 内存高效 Attention;
  • FlashAttention 类后端;
  • 数学参考实现。

所以 TileLang FlashAttention 不应该只与朴素三步 PyTorch 实现比较。


13.5 普通 Attention 的内存问题

普通实现:

Q @ Kᵀ
   ↓
完整 Scores
   ↓
完整 Softmax
   ↓
完整 Probabilities
   ↓
Probabilities @ V

Score 矩阵大小与序列长度平方相关。

当序列长度增加时,显存和 HBM 流量快速上升。


13.6 FlashAttention 的核心思想

FlashAttention 不长期保存完整 Score 矩阵。

它按 Tile 处理:

读取 Q Tile
    ↓
读取 K Tile 0
    ↓
计算局部 Score
    ↓
更新在线 Softmax
    ↓
读取 V Tile 0
    ↓
更新局部输出
    ↓
读取 K Tile 1
    ↓
继续更新
    ↓
...

最终只写出 Attention 输出。

它减少了:

  • 完整 Score 写入 HBM;
  • 完整 Probability 写入 HBM;
  • 反复读取大型中间 Tensor;
  • 显存占用。

FlashAttention 是一种 IO-aware 算法。

重点不只是减少浮点运算,而是减少 HBM 与片上 SRAM 之间的数据移动。


13.7 Online Softmax 状态

对每个 Query 行维护:

  • 当前最大值;
  • 当前指数和;
  • 当前未归一化输出。

加入新的 Score Tile 后:

mnew=max⁡(mold,mtile)m_{\mathrm{new}}=\max(m_{\mathrm{old}},m_{\mathrm{tile}})mnew=max(mold,mtile)

α=exp⁡(mold−mnew)\alpha=\exp(m_{\mathrm{old}}-m_{\mathrm{new}})α=exp(moldmnew)

lnew=αlold+∑jexp⁡(sj−mnew)l_{\mathrm{new}}=\alpha l_{\mathrm{old}}+\sum_j\exp(s_j-m_{\mathrm{new}})lnew=αlold+jexp(sjmnew)

onew=αoold+∑jexp⁡(sj−mnew)Vjo_{\mathrm{new}}=\alpha o_{\mathrm{old}}+\sum_j\exp(s_j-m_{\mathrm{new}})V_jonew=αoold+jexp(sjmnew)Vj

最终:

O=olO=\frac{o}{l}O=lo

这保证分块计算仍然得到数学上等价的 Softmax 结果。


13.8 TileLang FlashAttention 教学实现

下面是非因果、固定长度、仅前向的简化版本。

它省略了:

  • Causal Mask;
  • Padding Mask;
  • Dropout;
  • 变长序列;
  • GQA;
  • MQA;
  • 反向传播;
  • 尾部非整除处理;
  • 架构专用调度;
  • 更复杂的 Warp 划分。
import tilelang
import tilelang.language as T


@tilelang.jit(
    out_idx=[3],
    pass_configs={
        tilelang.PassConfigKey
        .TL_ENABLE_FAST_MATH: True,
    },
)
def tilelang_flash_attention(
    batch: int,
    heads: int,
    sequence: int,
    dimension: int,
    block_m: int = 64,
    block_n: int = 64,
    num_stages: int = 2,
    threads: int = 128,
):
    input_dtype = T.float16
    accumulation_dtype = T.float32

    qkv_shape = (
        batch,
        heads,
        sequence,
        dimension,
    )

    softmax_scale = (
        (1.0 / dimension) ** 0.5
        * 1.4426950408889634
    )

    @T.prim_func
    def kernel(
        Q: T.Tensor(
            qkv_shape,
            input_dtype,
        ),
        K: T.Tensor(
            qkv_shape,
            input_dtype,
        ),
        V: T.Tensor(
            qkv_shape,
            input_dtype,
        ),
        O: T.Tensor(
            qkv_shape,
            input_dtype,
        ),
    ):
        with T.Kernel(
            T.ceildiv(
                sequence,
                block_m,
            ),
            heads,
            batch,
            threads=threads,
        ) as (
            query_block,
            head,
            batch_id,
        ):
            q_shared = T.alloc_shared(
                (
                    block_m,
                    dimension,
                ),
                input_dtype,
            )
            k_shared = T.alloc_shared(
                (
                    block_n,
                    dimension,
                ),
                input_dtype,
            )
            v_shared = T.alloc_shared(
                (
                    block_n,
                    dimension,
                ),
                input_dtype,
            )
            o_shared = T.alloc_shared(
                (
                    block_m,
                    dimension,
                ),
                input_dtype,
            )

            scores = T.alloc_fragment(
                (
                    block_m,
                    block_n,
                ),
                accumulation_dtype,
            )
            probabilities = (
                T.alloc_fragment(
                    (
                        block_m,
                        block_n,
                    ),
                    input_dtype,
                )
            )
            output = T.alloc_fragment(
                (
                    block_m,
                    dimension,
                ),
                accumulation_dtype,
            )

            running_max = (
                T.alloc_fragment(
                    (block_m,),
                    accumulation_dtype,
                )
            )
            previous_max = (
                T.alloc_fragment(
                    (block_m,),
                    accumulation_dtype,
                )
            )
            rescale = T.alloc_fragment(
                (block_m,),
                accumulation_dtype,
            )
            tile_sum = T.alloc_fragment(
                (block_m,),
                accumulation_dtype,
            )
            running_sum = (
                T.alloc_fragment(
                    (block_m,),
                    accumulation_dtype,
                )
            )

            T.copy(
                Q[
                    batch_id,
                    head,
                    query_block
                    * block_m,
                    0,
                ],
                q_shared,
            )

            T.fill(
                output,
                0,
            )
            T.fill(
                running_sum,
                0,
            )
            T.fill(
                running_max,
                -T.infinity(
                    accumulation_dtype
                ),
            )

            for key_block in T.Pipelined(
                T.ceildiv(
                    sequence,
                    block_n,
                ),
                num_stages=num_stages,
            ):
                T.copy(
                    K[
                        batch_id,
                        head,
                        key_block
                        * block_n,
                        0,
                    ],
                    k_shared,
                )

                T.clear(scores)

                T.gemm(
                    q_shared,
                    k_shared,
                    scores,
                    transpose_B=True,
                    policy=(
                        T.GemmWarpPolicy
                        .FullRow
                    ),
                )

                T.copy(
                    running_max,
                    previous_max,
                )

                T.fill(
                    running_max,
                    -T.infinity(
                        accumulation_dtype
                    ),
                )

                T.reduce_max(
                    scores,
                    running_max,
                    dim=1,
                    clear=False,
                )

                for i in T.Parallel(
                    block_m
                ):
                    running_max[i] = T.max(
                        running_max[i],
                        previous_max[i],
                    )

                    rescale[i] = T.exp2(
                        previous_max[i]
                        * softmax_scale
                        - running_max[i]
                        * softmax_scale
                    )

                for i, j in T.Parallel(
                    block_m,
                    dimension,
                ):
                    output[i, j] *= (
                        rescale[i]
                    )

                for i, j in T.Parallel(
                    block_m,
                    block_n,
                ):
                    scores[i, j] = T.exp2(
                        scores[i, j]
                        * softmax_scale
                        - running_max[i]
                        * softmax_scale
                    )

                T.reduce_sum(
                    scores,
                    tile_sum,
                    dim=1,
                    clear=True,
                )

                for i in T.Parallel(
                    block_m
                ):
                    running_sum[i] = (
                        running_sum[i]
                        * rescale[i]
                        + tile_sum[i]
                    )

                T.copy(
                    scores,
                    probabilities,
                )

                T.copy(
                    V[
                        batch_id,
                        head,
                        key_block
                        * block_n,
                        0,
                    ],
                    v_shared,
                )

                T.gemm(
                    probabilities,
                    v_shared,
                    output,
                    policy=(
                        T.GemmWarpPolicy
                        .FullRow
                    ),
                )

            for i, j in T.Parallel(
                block_m,
                dimension,
            ):
                output[i, j] /= (
                    running_sum[i]
                )

            T.copy(
                output,
                o_shared,
            )

            T.copy(
                o_shared,
                O[
                    batch_id,
                    head,
                    query_block
                    * block_m,
                    0,
                ],
            )

    return kernel

13.9 调用与验证

import torch
import torch.nn.functional as F


batch = 1
heads = 8
sequence = 1024
dimension = 64

q = torch.randn(
    batch,
    heads,
    sequence,
    dimension,
    device="cuda",
    dtype=torch.float16,
)

k = torch.randn_like(q)
v = torch.randn_like(q)

kernel = tilelang_flash_attention(
    batch=batch,
    heads=heads,
    sequence=sequence,
    dimension=dimension,
    block_m=64,
    block_n=64,
    num_stages=2,
    threads=128,
)

output_tilelang = kernel(
    q,
    k,
    v,
)

output_torch = (
    F.scaled_dot_product_attention(
        q,
        k,
        v,
    )
)

torch.testing.assert_close(
    output_tilelang,
    output_torch,
    rtol=2e-2,
    atol=2e-2,
)

官方 FlashAttention 示例目录通常还包含:

  • MHA;
  • GQA;
  • 前向;
  • 反向;
  • Causal;
  • 变长序列;
  • 不同布局;
  • 不同硬件架构优化。

第十四章:统一性能比较

14.1 性能不能只看语言

性能差异不是:

Python 语法
对比
TileLang 语法

真正比较的是:

解释器循环
对比
预编译底层 Kernel
对比
定制 GPU Kernel

14.2 各类算子的主要瓶颈

算子主要瓶颈TileLang 的主要机会
向量加法内存带宽多个逐元素操作融合
Reduce同步与线程协作特殊行宽和专用 Reduce
SoftmaxReduce、指数、内存Mask、Scale、Dropout 融合
GEMMTensor Core 与数据复用特殊形状、量化、Epilogue
LayerNorm内存带宽与 ReduceResidual、Bias、量化融合
RMSNorm内存带宽与 ReduceLLM 特殊融合
FlashAttentionGEMM、Softmax、HBM IO定制 Attention 数据流

14.3 纯 Python 为什么慢

主要原因:

  • 解释器循环;
  • Python 对象;
  • 无法自动使用 GPU;
  • 无法使用 Tensor Core;
  • 缺少向量化;
  • 缺少内存访问优化;
  • 无法进行 Warp Reduce。

14.4 PyTorch 为什么快

PyTorch 的优势包括:

  • 调用预编译 Kernel;
  • 使用高度优化的供应商库;
  • 自动选择后端;
  • 支持自动求导;
  • 支持 torch.compile
  • 常用算子已经做了大量优化。

因此,PyTorch 应该是默认选择。


14.5 TileLang 为什么可能更快

TileLang 的加速主要来自重新组织底层执行:

  • 跨算子融合;
  • 减少中间 Tensor;
  • 减少 HBM 读写;
  • 使用 Shared Memory;
  • 使用寄存器保存中间值;
  • 针对特定形状选择 Tile;
  • 使用软件流水线;
  • 映射到 Tensor Core;
  • 使用 Online Softmax;
  • 优化 Warp 分工;
  • 定制数据布局;
  • 定制量化与反量化。

14.6 不应该如何比较

不公平比较:

TileLang 融合 Kernel
对比
多个拆开的 PyTorch 操作

更公平的比较方式是:

  • Softmax 对比 torch.softmax
  • LayerNorm 对比 F.layer_norm
  • Attention 对比 scaled_dot_product_attention
  • GEMM 对比 torch.matmul
  • 编译 PyTorch 对比 torch.compile 后的实现;
  • 使用相同形状和相同数据类型;
  • 排除编译时间;
  • 使用同样的精度容差。

第十五章:完整 Benchmark 模板

import torch
from tilelang.profiler import do_bench


def check_close(
    actual,
    expected,
    rtol=1e-2,
    atol=1e-2,
):
    torch.testing.assert_close(
        actual,
        expected,
        rtol=rtol,
        atol=atol,
    )


def benchmark(
    fn,
    warmup=25,
    rep=100,
):
    return do_bench(
        fn,
        warmup=warmup,
        rep=rep,
        backend="event",
    )


def print_benchmark(
    name,
    fn,
):
    latency_ms = benchmark(fn)

    print(
        f"{name:<24}: "
        f"{latency_ms:.4f} ms"
    )

    return latency_ms

Softmax 示例:

def benchmark_softmax():
    x = torch.randn(
        4096,
        4096,
        device="cuda",
        dtype=torch.float16,
    )

    tilelang_softmax(x)

    torch_latency = (
        print_benchmark(
            "PyTorch softmax",
            lambda: torch.softmax(
                x,
                dim=-1,
            ),
        )
    )

    tilelang_latency = (
        print_benchmark(
            "TileLang softmax",
            lambda: tilelang_softmax(x),
        )
    )

    speedup = (
        torch_latency
        / tilelang_latency
    )

    print(
        f"Speedup: {speedup:.3f}x"
    )

第十六章:调试 TileLang

16.1 从最小输入开始

不要一开始使用:

Batch = 64
Sequence = 8192
Hidden = 8192

应先使用:

M = 2
N = 8
K = 4

这样更容易检查:

  • 每个元素;
  • 每个线程;
  • 每个 Tile;
  • 边界;
  • Reduce;
  • 转置。

16.2 分阶段验证

推荐顺序:

输入复制是否正确
      ↓
Shared Memory 是否正确
      ↓
Fragment 是否正确
      ↓
局部计算是否正确
      ↓
Reduce 是否正确
      ↓
写回是否正确
      ↓
流水线打开后是否正确

16.3 Kernel 内打印

可以使用:

T.print(value)

但 GPU 线程数量很多,打印会非常混乱。

只建议:

  • 使用非常小的输入;
  • 限制特定线程打印;
  • 临时调试;
  • 不要在性能测试中保留打印。

16.4 查看生成代码

compiled = operation.compile(...)
print(compiled.get_kernel_source())

这是判断编译器实际做了什么的重要方法。


16.5 常见错误排查顺序

是否越界
   ↓
形状是否正确
   ↓
数据类型是否正确
   ↓
Shared Memory 是否正确
   ↓
Fragment 是否过大
   ↓
Reduce 是否 clear
   ↓
T.gemm 的 transpose 是否正确
   ↓
流水线关闭后是否正确
   ↓
尾部非整除是否处理

16.6 调试工具类别

官方调试工具通常围绕三类问题:

生成问题

  • 编译失败;
  • TIR Lowering 失败;
  • 后端代码生成失败;
  • CUDA 编译失败。

正确性问题

  • 数值错误;
  • 越界;
  • Race Condition;
  • 未初始化;
  • Reduce 错误。

性能问题

  • 寄存器溢出;
  • Shared Memory 过大;
  • Occupancy 低;
  • Tensor Core 未使用;
  • 内存访问不合并;
  • Pipeline 无效。

可能使用的工具包括:

  • Pass Diff;
  • IR 可视化;
  • Kernel Source;
  • 程序缩减;
  • NVIDIA Nsight Compute;
  • NVIDIA Nsight Systems;
  • ROCm Profiler。

第十七章:常见性能问题

17.1 Fragment 太大

危险示例:

buffer = T.alloc_fragment(
    (128, 8192),
    T.float32,
)

可能导致:

  • 寄存器需求过高;
  • 寄存器溢出到 Local Memory;
  • Occupancy 降低;
  • Kernel 变慢;
  • 编译时间增加;
  • 编译失败。

解决方法:

  • 缩小 Tile;
  • 分块处理;
  • Split-K;
  • 两次遍历;
  • 使用 Shared Memory;
  • 减少同时存活的中间值。

17.2 Shared Memory 太大

Shared Memory 使用量与以下参数有关:

  • block_m
  • block_n
  • block_k
  • num_stages
  • 同时保留的输入 Tile 数量;
  • 是否双缓冲或多缓冲。

过大会导致:

  • 一个 SM 能同时驻留的线程块减少;
  • Occupancy 降低;
  • 超出硬件上限;
  • Kernel 无法启动。

17.3 Tile 太小

可能导致:

  • Tensor Core 利用率低;
  • 数据复用不足;
  • Kernel 数量多;
  • 调度开销比例大;
  • 指令效率低。

17.4 Tile 太大

可能导致:

  • 寄存器溢出;
  • Shared Memory 超限;
  • Occupancy 降低;
  • Block 数量不足;
  • 负载不均衡;
  • 尾部浪费严重。

17.5 流水级越多不一定越快

增加 num_stages 可能提高数据搬运与计算的重叠程度。

但也会增加:

  • Shared Memory;
  • 寄存器;
  • 调度复杂度;
  • 编译复杂度。

因此需要实测。


17.6 数据类型选择不合理

常见推荐:

输入:FP16 或 BF16
归约:FP32
GEMM 累加:FP32
输出:FP16 或 BF16

Softmax、LayerNorm 和 RMSNorm 如果完全使用 FP16 累加,可能出现:

  • 精度下降;
  • 溢出;
  • 下溢;
  • 误差增大。

17.7 没有处理尾部

如果形状不能整除 Tile 大小,需要处理:

N % block_size != 0
M % block_m != 0
K % block_k != 0

常见方式:

  • 边界判断;
  • Padding;
  • Masked Copy;
  • 特殊尾部 Kernel;
  • 要求输入形状对齐。

教学代码经常省略尾部处理,但生产代码不能忽略。


17.8 Kernel 启动开销

对于非常小的 Tensor,计算时间可能低于 Kernel 启动开销。

这时:

  • 自定义复杂 Kernel 不一定有优势;
  • 融合多个小操作更重要;
  • CUDA Graph 可能有帮助;
  • CPU 甚至可能更合适。

第十八章:自动调优

18.1 为什么需要自动调优

不同 GPU 和不同形状的最佳配置可能不同。

需要搜索:

  • block_m
  • block_n
  • block_k
  • threads
  • num_stages
  • Warp Policy;
  • Swizzle;
  • Shared Memory 布局;
  • Reduce 策略。

18.2 概念示例

import tilelang


configs = [
    {
        "block_m": 64,
        "block_n": 64,
        "block_k": 32,
        "num_stages": 2,
    },
    {
        "block_m": 128,
        "block_n": 128,
        "block_k": 32,
        "num_stages": 3,
    },
]


@tilelang.autotune(
    configs=configs,
    warmup=10,
    rep=20,
)
@tilelang.jit
def tuned_matmul(
    A,
    B,
    block_m=64,
    block_n=64,
    block_k=32,
    num_stages=2,
):
    ...

自动调优会尝试不同配置,并根据性能选出较优方案。


18.3 自动调优的代价

配置越多:

  • 编译时间越长;
  • 首次运行越慢;
  • 磁盘缓存越大;
  • 搜索成本越高;
  • 验证正确性的成本越高。

生产环境通常采用:

离线搜索最佳配置
      ↓
保存配置
      ↓
线上直接使用

第十九章:推荐学习路线

阶段一:理解线程并行

建议实现:

  • 向量加法;
  • ReLU;
  • Scale;
  • Clamp;
  • Add + ReLU;
  • Add + Scale。

重点掌握:

  • T.Kernel
  • T.Parallel
  • Block;
  • Thread;
  • 边界判断;
  • Global Memory。

阶段二:理解 Reduce

建议实现:

  • Row Sum;
  • Row Max;
  • Row Mean;
  • Softmax;
  • Online Softmax。

重点掌握:

  • Warp;
  • Fragment;
  • T.reduce_sum
  • T.reduce_max
  • 数值稳定性;
  • FP32 累加。

阶段三:理解内存层级

建议实现:

  • Tiled Copy;
  • Matrix Transpose;
  • GEMM;
  • GEMM + Bias;
  • GEMM + ReLU;
  • GEMM + Bias + Activation。

重点掌握:

  • Global Memory;
  • Shared Memory;
  • Fragment;
  • T.copy
  • 数据复用;
  • 内存合并访问;
  • Shared Memory Bank Conflict。

阶段四:理解 Tensor Core 和 Pipeline

重点实现:

  • FP16 GEMM;
  • BF16 GEMM;
  • 多 Stage GEMM;
  • 不同 Tile 配置。

重点掌握:

  • T.gemm
  • T.Pipelined
  • Tensor Core;
  • Warp Policy;
  • Tile 大小;
  • Pipeline Stage。

阶段五:理解归一化和融合

建议实现:

  • LayerNorm;
  • RMSNorm;
  • Residual + LayerNorm;
  • Residual + RMSNorm;
  • RMSNorm + Quantization。

重点掌握:

  • Reduce 与 Elementwise 融合;
  • 中间值生命周期;
  • Register Pressure;
  • Shared Memory 容量;
  • 两次遍历策略。

阶段六:复杂 Attention

建议实现:

  • 普通 Attention;
  • Online Softmax;
  • FlashAttention 前向;
  • Causal FlashAttention;
  • GQA;
  • Varlen Attention;
  • FlashAttention 反向。

重点掌握:

  • 两次 GEMM 融合;
  • Online Softmax;
  • Score Tile 不写 HBM;
  • Warp 分工;
  • Pipeline;
  • Mask;
  • 数值稳定性;
  • 不同 Head Dimension。

第二十章:实际项目中的最佳工作流

最合理的工作方式不是一开始就使用 TileLang 重写所有算子。

推荐流程:

先写 PyTorch 参考实现
        ↓
确认算法和模型收益
        ↓
使用 Profiler 找出真正热点
        ↓
判断瓶颈属于计算、内存还是启动开销
        ↓
只重写最耗时的少数 Kernel
        ↓
验证 TileLang 正确性
        ↓
查看生成的底层代码
        ↓
调 Tile、线程和流水线
        ↓
与 PyTorch 最优实现比较
        ↓
接入模型并测试端到端性能

必须区分:

单 Kernel 更快

和:

整个模型更快

一个 Kernel 加速并不一定能显著改善端到端性能。

原因可能包括:

  • 该 Kernel 占总时间比例很小;
  • 产生额外数据转换;
  • 编译开销;
  • 输入布局转换;
  • 调用频率不足;
  • CPU 调度成为新瓶颈;
  • 其他算子仍然占主导。

第二十一章:什么时候应该使用 TileLang

适合使用 TileLang:

  • PyTorch 中多个操作无法很好融合;
  • 存在大型中间 Tensor;
  • 需要特殊 Attention;
  • 需要量化 Kernel;
  • 需要特殊数据布局;
  • 需要特殊矩阵形状优化;
  • 需要控制 Shared Memory;
  • 需要使用 Tensor Core;
  • 需要定制 Pipeline;
  • 需要编写前向和反向;
  • 需要针对特定 GPU 调优。

不一定需要 TileLang:

  • 已有 PyTorch 算子已经足够快;
  • 标准 GEMM 已被供应商库很好覆盖;
  • 算子不是性能瓶颈;
  • 数据规模很小;
  • 项目更重视可维护性;
  • 团队缺少 GPU Kernel 调试经验;
  • 无法承担不同硬件的调优成本。

第二十二章:最终总结

纯 Python 适合:

  • 学习算法;
  • 编写参考实现;
  • 验证数学逻辑;
  • 处理小规模数据。

PyTorch 适合:

  • 绝大多数模型开发;
  • 自动求导;
  • 快速迭代;
  • 使用成熟算子;
  • 跨平台部署;
  • 生产训练和推理。

TileLang 适合:

  • 自定义高性能 GPU Kernel;
  • 跨算子融合;
  • 特殊形状优化;
  • 定制 Attention;
  • 定制矩阵乘法;
  • 量化和稀疏计算;
  • 精确控制 Tile、内存和流水线。

理解 TileLang 最重要的三个问题是:

  1. 一个线程块负责哪一块输出?
  2. 数据什么时候从 Global Memory 搬到 Shared Memory?
  3. 中间结果应该放在 Shared Memory 还是 Fragment?

进一步还需要思考:

  • 每个 Tile 是否被充分复用;
  • 是否减少了 HBM 访问;
  • 是否使用了 Tensor Core;
  • 是否产生寄存器溢出;
  • 是否降低了 Occupancy;
  • 是否值得进行跨算子融合;
  • 是否真正改善了端到端性能。

TileLang 的价值并不是替代 PyTorch,而是补充 PyTorch。

最合理的实践是:

使用 PyTorch 描述模型,使用 Profiler 找到热点,再使用 TileLang 重写少数真正值得优化的高成本 Kernel。

从学习角度,推荐严格按照以下顺序:

向量加法
   ↓
Reduce
   ↓
Softmax
   ↓
GEMM
   ↓
LayerNorm
   ↓
RMSNorm
   ↓
FlashAttention

只要真正理解了:

  • GPU 线程模型;
  • Tile;
  • 内存层级;
  • Reduce;
  • GEMM;
  • Online Softmax;
  • Kernel Fusion;
  • Pipeline;

就已经建立了从普通 Python 开发者走向高性能 AI Kernel 开发者所需的核心知识体系。
:::

后续可以在这篇教程基础上继续整理成“可运行项目版”,将每个算子拆成独立 Python 文件、测试文件和统一 benchmark 脚本。

我会把它整理成一份真正可照着练的入门教程:先搭环境和建立性能直觉,再逐个实现 6 类算子,并对照 Python、PyTorch、TileLang 的执行路径与瓶颈。示例会尽量保持同一输入、同一数学定义和同一验证方式。

TileLang 零基础完整开发教程

这份教程以 TileLang 0.1.12 官方文档和官方仓库当前示例为基准,带你沿着下面的路线学习:

向量加法 → Reduce → Softmax → GEMM → LayerNorm/RMSNorm → FlashAttention

每个阶段都会对比三种实现:

  1. 纯 Python:帮助理解算法。
  2. PyTorch:生产环境中最常见的高层写法。
  3. TileLang:自己控制 GPU Kernel 的分块、内存和流水线。

TileLang 的 API 仍在快速演进,因此不同版本中的少量函数签名可能发生变化。本文代码采用官方文档当前推荐的 @tilelang.jitT.KernelT.ParallelT.copyT.reduce_*T.gemm 风格。TileLang 官方将其定义为建立在 TVM 编译基础设施之上的 Python 风格高性能 Kernel DSL。


一、先理解三者的根本区别

1. 纯 Python

例如:

c = [a[i] + b[i] for i in range(len(a))]

实际执行路径大致是:

Python 解释器
    ↓
逐次执行 Python 字节码
    ↓
每次循环读取 Python 对象
    ↓
调用对象加法
    ↓
创建新的 Python 对象

问题不在加法本身,而在于:

  • Python 循环解释开销很大;
  • Python 浮点数是对象,不只是一个裸 float32
  • 每次索引和加法都涉及动态类型处理;
  • 很难自动利用 GPU;
  • 很难利用 SIMD、Tensor Core 和共享内存。

纯 Python 适合表达算法,不适合实现大规模数值 Kernel。


2. PyTorch

c = a + b

表面只有一行,底层通常是:

Python
    ↓
PyTorch Dispatcher
    ↓
选择 CUDA / CPU 实现
    ↓
调用预先编译好的 ATen Kernel
    ↓
CUDA Kernel 在 GPU 上并行执行

PyTorch 已经替你实现了大量高性能算子。

优点:

  • 简单;
  • 稳定;
  • 自动求导;
  • 算子库成熟;
  • 矩阵乘法通常直接调用 cuBLAS;
  • Attention 可调用优化后的 SDPA 或 FlashAttention 后端。

不足是:

  • 一个复杂表达式可能拆成多个 Kernel;
  • 每个 Kernel 都可能读写显存;
  • 自定义融合逻辑受现有算子边界限制;
  • 特殊模型结构可能没有最优现成实现;
  • 很难精确控制 Shared Memory、寄存器和流水线。

3. TileLang

TileLang 让你自己描述:

一个线程块处理哪块数据
数据怎样从显存搬到共享内存
哪些数据保存在寄存器中
哪些循环并行执行
哪些操作形成软件流水线
矩阵乘法怎样映射到 Tensor Core

大致编译路径是:

Python 风格 TileLang
    ↓
TIR / TileLang IR
    ↓
编译器优化 Pass
    ↓
CUDA C / HIP / LLVM 等后端代码
    ↓
GPU 或 CPU 可执行 Kernel

TileLang 显式提供 Global、Shared 和 Fragment 等内存层级,以及 T.PipelinedT.copyT.gemm 等高层原语。

一句话概括:

PyTorch 是“调用已有 Kernel”,TileLang 是“开发新的 Kernel”。


二、安装开发环境

1. 建议环境

入门阶段建议:

  • Ubuntu 20.04 或更新版本;
  • Python 3.10 或更新版本;
  • NVIDIA GPU;
  • 可用的 CUDA 版 PyTorch;
  • 独立虚拟环境。

TileLang 0.1.12 官方安装文档要求 Python 3.10 及以上,PyPI 包要求 glibc 2.28 及以上;官方同时提供源码构建和 pip CUDA toolchain 方案。

2. 创建环境

python3 -m venv tilelang-env
source tilelang-env/bin/activate

python -m pip install --upgrade pip
pip install torch
pip install tilelang

验证:

python -c "import torch; print(torch.__version__)"
python -c "import tilelang; print(tilelang.__version__)"

检查 GPU:

import torch

print("CUDA available:", torch.cuda.is_available())

if torch.cuda.is_available():
    print("GPU:", torch.cuda.get_device_name())
    print("CUDA version:", torch.version.cuda)

安装 GitHub 最新版:

pip install git+https://github.com/tile-ai/tilelang.git

不过学习阶段更推荐先使用稳定发布版,避免官方仓库主分支的接口更新影响示例。官方同时支持 PyPI 和 GitHub 安装。


三、创建统一的开发框架

建议先创建目录:

tilelang_tutorial/
├── common.py
├── 01_vector_add.py
├── 02_reduce.py
├── 03_softmax.py
├── 04_gemm.py
├── 05_norm.py
└── 06_flash_attention.py

common.py

from __future__ import annotations

import time
from collections.abc import Callable

import torch


def check_cuda() -> None:
    if not torch.cuda.is_available():
        raise RuntimeError("本教程的 GPU 示例需要可用的 CUDA GPU。")


def benchmark_cuda(
    fn: Callable[[], object],
    warmup: int = 20,
    repeat: int = 100,
) -> float:
    """
    返回平均耗时,单位为毫秒。
    """
    check_cuda()

    for _ in range(warmup):
        fn()

    torch.cuda.synchronize()

    start = torch.cuda.Event(enable_timing=True)
    end = torch.cuda.Event(enable_timing=True)

    start.record()

    for _ in range(repeat):
        fn()

    end.record()
    torch.cuda.synchronize()

    total_ms = start.elapsed_time(end)
    return total_ms / repeat


def benchmark_cpu(
    fn: Callable[[], object],
    repeat: int = 5,
) -> float:
    """
    返回平均耗时,单位为毫秒。
    """
    begin = time.perf_counter()

    for _ in range(repeat):
        fn()

    elapsed = time.perf_counter() - begin
    return elapsed * 1000 / repeat

为什么必须同步 GPU?

CUDA 默认异步执行。

下面的测试是错误的:

start = time.perf_counter()
c = a + b
elapsed = time.perf_counter() - start

Python 可能只测到了“提交 Kernel”的时间,而不是 GPU 完成计算的时间。

正确方式是:

torch.cuda.synchronize()
start = time.perf_counter()

c = a + b

torch.cuda.synchronize()
elapsed = time.perf_counter() - start

或者使用 CUDA Event。


四、TileLang 最重要的编程概念

1. @tilelang.jit

@tilelang.jit
def operation(...):
    ...

它把 TileLang 函数包装成可 JIT 编译的 Kernel。

典型使用:

kernel = operation.compile(...)
result = kernel(...)

TileLang 官方 JIT 模块负责将程序编译成可调用的 Kernel adapter。


2. T.Tensor

A: T.Tensor((M, N), T.float32)

表示:

  • 形状为 (M, N)
  • 数据类型为 float32
  • 通常位于设备 Global Memory。

3. T.Kernel

with T.Kernel(grid_size, threads=256) as block_id:
    ...

类似 CUDA:

kernel<<<grid_size, 256>>>(...)

其中 block_id 类似:

blockIdx.x

二维网格:

with T.Kernel(grid_x, grid_y, threads=128) as (bx, by):
    ...

4. T.Parallel

for i in T.Parallel(256):
    ...

表示并行迭代。

二维:

for i, j in T.Parallel(BM, BN):
    ...

编译器负责把循环映射到线程或向量化执行。


5. 内存层级

Global Memory

A: T.Tensor((M, N), dtype)

特点:

  • 容量大;
  • 延迟高;
  • GPU Kernel 间共享;
  • 通常对应显存。

Shared Memory

A_shared = T.alloc_shared((BM, BK), dtype)

特点:

  • 位于 GPU 芯片上;
  • 一个线程块共享;
  • 比显存快;
  • 容量有限;
  • 适合缓存重复使用的数据 Tile。

Fragment

C_local = T.alloc_fragment((BM, BN), T.float32)

通常用于:

  • 寄存器级中间值;
  • Tensor Core 累加器;
  • 线程或 warp 私有的计算片段。

官方语言基础文档把 Global、Shared 和 Fragment 作为核心软件管理内存层级。


6. T.copy

T.copy(
    A[by * BM, k * BK],
    A_shared,
)

表示把一个 Tile 从 Global Memory 搬到 Shared Memory。

TileLang 会基于源、目标形状和布局进行 Lowering。


7. T.Pipelined

for k in T.Pipelined(num_tiles, num_stages=3):
    T.copy(...)
    T.gemm(...)

其目标是让不同阶段重叠:

阶段 1:加载第 k+1 块
阶段 2:计算第 k 块
阶段 3:处理第 k-1 块

从而隐藏显存访问延迟。


五、示例一:向量加法

数学定义:

Ci=Ai+BiC_i=A_i+B_iCi=Ai+Bi


1. 纯 Python 实现

def python_vector_add(a: list[float], b: list[float]) -> list[float]:
    if len(a) != len(b):
        raise ValueError("a 和 b 长度必须相同")

    result = [0.0] * len(a)

    for i in range(len(a)):
        result[i] = a[i] + b[i]

    return result

底层特点:

  • 单线程;
  • Python 解释器执行循环;
  • 每个元素都要进行对象访问;
  • 无 GPU 并行。

2. PyTorch 实现

def torch_vector_add(
    a: torch.Tensor,
    b: torch.Tensor,
) -> torch.Tensor:
    return a + b

底层通常是:

一个 CUDA Elementwise Kernel
每个 GPU 线程处理一个或多个元素
从显存读取 A 和 B
执行一次加法
写回 C

每个元素大约需要:

  • 读取 A:4 字节;
  • 读取 B:4 字节;
  • 写入 C:4 字节;
  • 执行一次浮点加法。

因此 FP32 向量加法的算术强度约为:

Arithmetic Intensity≈1 FLOP12 bytes\text{Arithmetic Intensity}\approx\frac{1\ \text{FLOP}}{12\ \text{bytes}}Arithmetic Intensity12 bytes1 FLOP

这是典型的显存带宽受限算子,而不是算力受限算子。


3. TileLang 实现

from __future__ import annotations

import torch
import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_vector_add(
    A,
    B,
    block_size: int,
):
    N = T.const("N")

    A: T.Tensor((N,), T.float32)
    B: T.Tensor((N,), T.float32)

    C = T.empty((N,), T.float32)

    with T.Kernel(
        T.ceildiv(N, block_size),
        threads=block_size,
    ) as bx:
        for tx in T.Parallel(block_size):
            index = bx * block_size + tx

            if index < N:
                C[index] = A[index] + B[index]

    return C

调用:

def test_vector_add() -> None:
    n = 1 << 20

    a = torch.randn(n, device="cuda", dtype=torch.float32)
    b = torch.randn(n, device="cuda", dtype=torch.float32)

    kernel = tilelang_vector_add.compile(
        a,
        b,
        block_size=256,
    )

    actual = kernel(a, b)
    expected = a + b

    torch.testing.assert_close(actual, expected)
    print("Vector add passed.")

    print(kernel.get_kernel_source())


if __name__ == "__main__":
    test_vector_add()

TileLang 官方 ElementWise 教程也以逐元素加法解释动态形状、边界处理及其与 CUDA 的对应关系。


4. 三者性能差异

实现主要开销预期表现
Python解释器、对象访问、单线程极慢
PyTorch一个成熟 CUDA Kernel很快
TileLang自己生成 CUDA Kernel正确调度时接近 PyTorch
不合理 TileLang非合并访存、线程过少、边界分支可能慢于 PyTorch

向量加法中 TileLang 通常不容易显著超过 PyTorch,因为:

  • PyTorch 已经只启动一个简单 Kernel;
  • 计算本身只有一次加法;
  • 性能主要受显存带宽限制;
  • 没有多少可融合或复用的数据。

TileLang 的优势更常出现在“多个操作可以融合”的场景。


六、示例二:Reduce Sum

输入矩阵形状:

M × N

按行求和:

Yi=∑j=0N−1XijY_i=\sum_{j=0}^{N-1}X_{ij}Yi=j=0N1Xij


1. 纯 Python 实现

def python_row_sum(x: list[list[float]]) -> list[float]:
    output = []

    for row in x:
        total = 0.0

        for value in row:
            total += value

        output.append(total)

    return output

2. PyTorch 实现

def torch_row_sum(x: torch.Tensor) -> torch.Tensor:
    return torch.sum(x, dim=-1)

PyTorch 会调用归约 Kernel。

归约与向量加法不同,因为多个元素需要合并到一个结果:

线程 0:负责部分元素
线程 1:负责部分元素
线程 2:负责部分元素
...
线程块内进行树形归约
最后写入一个结果

最朴素的串行方式需要:

N−1N-1N1

次依赖加法,而且不能充分并行。

树形归约可在大约:

log⁡2N\log_2 Nlog2N

个同步阶段中完成。


3. TileLang 实现

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_row_sum(
    X,
    block_rows: int,
):
    M, N = T.const("M, N")

    X: T.Tensor((M, N), T.float32)
    Y = T.empty((M,), T.float32)

    with T.Kernel(
        T.ceildiv(M, block_rows),
        threads=128,
    ) as bx:
        x_local = T.alloc_fragment(
            (block_rows, N),
            T.float32,
        )
        row_sum = T.alloc_fragment(
            (block_rows,),
            T.float32,
        )

        T.copy(
            X[bx * block_rows : (bx + 1) * block_rows, :],
            x_local,
        )

        T.reduce_sum(
            x_local,
            row_sum,
            dim=1,
        )

        T.copy(
            row_sum,
            Y[bx * block_rows : (bx + 1) * block_rows],
        )

    return Y

调用与验证:

def test_row_sum() -> None:
    m = 1024
    n = 1024

    x = torch.randn(
        m,
        n,
        device="cuda",
        dtype=torch.float32,
    )

    kernel = tilelang_row_sum.compile(
        x,
        block_rows=1,
    )

    actual = kernel(x)
    expected = x.sum(dim=-1)

    torch.testing.assert_close(
        actual,
        expected,
        rtol=1e-4,
        atol=1e-4,
    )

    print("Reduce sum passed.")

4. Reduce 的性能关键

连续访存

对于 C-contiguous 的二维矩阵:

x[i, :]

是一段连续内存,因此按最后一维归约通常更友好。

而:

x[:, j]

是跨行访问,可能导致线程读取地址不连续。

累加精度

FP16 输入通常应使用 FP32 累加:

input_dtype = T.float16
accum_dtype = T.float32

否则长序列求和误差可能很大。

每行长度

  • 很短:一个 warp 处理多行;
  • 中等:一个 warp 或一个 block 处理一行;
  • 很长:分块归约,再进行第二级归约。

原子操作

多个线程块更新同一结果时,可以使用 atomic add,但原子竞争可能成为瓶颈。TileLang 提供原子加等操作。


七、示例三:Softmax

Softmax 的数学定义:

yi=exi∑jexjy_i=\frac{e^{x_i}}{\sum_j e^{x_j}}yi=jexjexi

直接计算可能溢出。例如 exp(1000) 超出常见浮点范围。

稳定版写法:

m=max⁡jxjm=\max_j x_jm=jmaxxj

yi=exi−m∑jexj−my_i=\frac{e^{x_i-m}}{\sum_j e^{x_j-m}}yi=jexjmexim


1. 纯 Python 实现

import math


def python_softmax(row: list[float]) -> list[float]:
    if not row:
        return []

    maximum = max(row)

    exponentials = [
        math.exp(value - maximum)
        for value in row
    ]

    denominator = sum(exponentials)

    return [
        value / denominator
        for value in exponentials
    ]

矩阵版本:

def python_row_softmax(
    x: list[list[float]],
) -> list[list[float]]:
    return [python_softmax(row) for row in x]

2. PyTorch 实现

import torch.nn.functional as F


def torch_softmax(x: torch.Tensor) -> torch.Tensor:
    return F.softmax(x, dim=-1)

PyTorch 会选择优化后的 Softmax Kernel。

从逻辑上需要:

  1. 读取输入,求最大值;
  2. 计算指数;
  3. 求指数和;
  4. 除以指数和;
  5. 写出结果。

如果每一步都单独启动 Kernel,会反复读写显存。成熟的 Softmax Kernel 会尽量在一个 Kernel 内完成,并把中间结果保存在寄存器或 Shared Memory。


3. TileLang 实现

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_softmax(
    X,
    block_rows: int,
):
    M, N = T.const("M, N")

    X: T.Tensor((M, N), T.float32)
    Y = T.empty((M, N), T.float32)

    with T.Kernel(
        T.ceildiv(M, block_rows),
        threads=128,
    ) as bx:
        x_local = T.alloc_fragment(
            (block_rows, N),
            T.float32,
        )
        row_max = T.alloc_fragment(
            (block_rows,),
            T.float32,
        )
        row_sum = T.alloc_fragment(
            (block_rows,),
            T.float32,
        )

        T.copy(
            X[bx * block_rows : (bx + 1) * block_rows, :],
            x_local,
        )

        T.reduce_max(
            x_local,
            row_max,
            dim=1,
        )

        for i, j in T.Parallel(block_rows, N):
            x_local[i, j] = T.exp(
                x_local[i, j] - row_max[i]
            )

        T.reduce_sum(
            x_local,
            row_sum,
            dim=1,
        )

        for i, j in T.Parallel(block_rows, N):
            x_local[i, j] = (
                x_local[i, j] / row_sum[i]
            )

        T.copy(
            x_local,
            Y[bx * block_rows : (bx + 1) * block_rows, :],
        )

    return Y

验证:

def test_softmax() -> None:
    m = 1024
    n = 1024

    x = torch.randn(
        m,
        n,
        device="cuda",
        dtype=torch.float32,
    )

    kernel = tilelang_softmax.compile(
        x,
        block_rows=1,
    )

    actual = kernel(x)
    expected = torch.softmax(x, dim=-1)

    torch.testing.assert_close(
        actual,
        expected,
        rtol=1e-4,
        atol=1e-4,
    )

    print("Softmax passed.")

官方 Attention 和 MLA 示例也使用 T.reduce_max、指数变换、T.reduce_sum 以及归一化构建在线 Softmax。


4. Softmax 为什么适合手写 Kernel?

因为它包含多个阶段:

max → subtract → exp → sum → divide

不融合时可能变成:

Kernel 1:max
Kernel 2:subtract
Kernel 3:exp
Kernel 4:sum
Kernel 5:divide

每个 Kernel 都会产生:

  • 启动开销;
  • 读显存;
  • 写显存;
  • 中间张量分配。

融合后:

Kernel 1:
读取 X
→ 寄存器中求 max
→ 寄存器中 exp
→ 寄存器中求 sum
→ 归一化
→ 写 Y

TileLang 的主要收益来自:

  • 减少 Kernel 数量;
  • 减少中间张量;
  • 减少显存流量;
  • 控制每行由 warp 还是 block 处理。

八、示例四:GEMM 矩阵乘法

定义:

C=ABC=ABC=AB

逐元素写成:

Cij=∑k=0K−1AikBkjC_{ij}=\sum_{k=0}^{K-1}A_{ik}B_{kj}Cij=k=0K1AikBkj

矩阵形状:

A: M × K
B: K × N
C: M × N

1. 纯 Python 实现

def python_matmul(
    a: list[list[float]],
    b: list[list[float]],
) -> list[list[float]]:
    m = len(a)
    k = len(a[0])
    n = len(b[0])

    if len(b) != k:
        raise ValueError("矩阵形状不兼容")

    c = [
        [0.0 for _ in range(n)]
        for _ in range(m)
    ]

    for i in range(m):
        for j in range(n):
            total = 0.0

            for p in range(k):
                total += a[i][p] * b[p][j]

            c[i][j] = total

    return c

复杂度:

O(MNK)O(MNK)O(MNK)

纯 Python 三重循环通常非常慢。


2. PyTorch 实现

def torch_matmul(
    a: torch.Tensor,
    b: torch.Tensor,
) -> torch.Tensor:
    return a @ b

GPU 上通常进入 cuBLAS 或相关优化实现。

成熟 GEMM 会使用:

  • 多级 Tile;
  • Shared Memory;
  • 寄存器分块;
  • Tensor Core;
  • 双缓冲或多阶段流水线;
  • 针对具体 GPU 架构的指令。

因此,普通规则 GEMM 中,自定义 TileLang 并不一定比 cuBLAS 快。TileLang 的价值通常在于:

  • 融合激活;
  • 特殊量化;
  • 特殊数据布局;
  • 非标准矩阵形状;
  • Grouped GEMM;
  • 与其他算子融合。

3. 为什么要分块?

假设直接计算一个输出:

C[i, j]

需要读取:

  • A 的第 i 行;
  • B 的第 j 列。

但相邻的输出会重复使用很多相同数据。

分块后:

A 的 BM × BK 块
B 的 BK × BN 块

被加载一次后,可以计算:

BM × BN

个输出的部分和。

这样减少了显存重复读取。


4. TileLang GEMM

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_matmul(
    A,
    B,
    block_m: int,
    block_n: int,
    block_k: int,
):
    M, N, K = T.const("M, N, K")

    A: T.Tensor((M, K), T.float16)
    B: T.Tensor((K, N), T.float16)

    C = T.empty((M, N), T.float16)

    with T.Kernel(
        T.ceildiv(N, block_n),
        T.ceildiv(M, block_m),
        threads=128,
    ) as (bx, by):
        A_shared = T.alloc_shared(
            (block_m, block_k),
            T.float16,
        )
        B_shared = T.alloc_shared(
            (block_k, block_n),
            T.float16,
        )
        C_local = T.alloc_fragment(
            (block_m, block_n),
            T.float32,
        )

        T.clear(C_local)

        for ko in T.Pipelined(
            T.ceildiv(K, block_k),
            num_stages=3,
        ):
            T.copy(
                A[
                    by * block_m,
                    ko * block_k,
                ],
                A_shared,
            )

            T.copy(
                B[
                    ko * block_k,
                    bx * block_n,
                ],
                B_shared,
            )

            T.gemm(
                A_shared,
                B_shared,
                C_local,
            )

        T.copy(
            C_local,
            C[
                by * block_m,
                bx * block_n,
            ],
        )

    return C

编译和运行:

def test_gemm() -> None:
    m = 1024
    n = 1024
    k = 1024

    a = torch.randn(
        m,
        k,
        device="cuda",
        dtype=torch.float16,
    )
    b = torch.randn(
        k,
        n,
        device="cuda",
        dtype=torch.float16,
    )

    kernel = tilelang_matmul.compile(
        a,
        b,
        block_m=128,
        block_n=128,
        block_k=32,
    )

    actual = kernel(a, b)
    expected = a @ b

    torch.testing.assert_close(
        actual,
        expected,
        rtol=1e-2,
        atol=1e-2,
    )

    print("GEMM passed.")
    print(kernel.get_kernel_source())

    profiler = kernel.get_profiler()
    latency_ms = profiler.do_bench()

    print(f"TileLang latency: {latency_ms:.4f} ms")

官方 GEMM 教程和仓库示例采用同样的核心结构:Shared Memory Tile、T.PipelinedT.copyT.gemm 和 Fragment 累加。


5. GEMM 的理论计算量

矩阵乘法大约执行:

2MNK2MNK2MNK

次浮点运算。

其中乘法和加法分别计为一次 FLOP。

吞吐率可以估算为:

TFLOPS=2MNKt×1012\text{TFLOPS}=\frac{2MNK}{t\times10^{12}}TFLOPS=t×10122MNK

其中 ttt 以秒为单位。

Python 计算 TFLOPS:

def gemm_tflops(
    m: int,
    n: int,
    k: int,
    latency_ms: float,
) -> float:
    operations = 2 * m * n * k
    seconds = latency_ms / 1000
    return operations / seconds / 1e12

6. GEMM 调优参数

block_m

一个线程块计算多少行输出。

block_n

一个线程块计算多少列输出。

block_k

一次加载多少个归约维元素。

threads

线程块线程数。

num_stages

软件流水线深度。

并不是越大越好:

  • Tile 太小:复用不足、Kernel 数量多;
  • Tile 太大:Shared Memory 占用高;
  • Fragment 太大:寄存器压力高;
  • stage 太多:占用更多 Shared Memory;
  • 线程太多:可能降低 occupancy。

TileLang 官方还提供自动调优能力,可编译、验证和测试多个配置并缓存最佳结果。


九、示例五:LayerNorm 与 RMSNorm

1. LayerNorm

对于一行输入:

μ=1N∑j=1Nxj\mu=\frac{1}{N}\sum_{j=1}^{N}x_jμ=N1j=1Nxj

σ2=1N∑j=1N(xj−μ)2\sigma^2=\frac{1}{N}\sum_{j=1}^{N}(x_j-\mu)^2σ2=N1j=1N(xjμ)2

yj=xj−μσ2+ϵγj+βjy_j=\frac{x_j-\mu}{\sqrt{\sigma^2+\epsilon}}\gamma_j+\beta_jyj=σ2+ϵxjμγj+βj

2. RMSNorm

RMSNorm 不减均值:

r=1N∑j=1Nxj2+ϵr=\sqrt{\frac{1}{N}\sum_{j=1}^{N}x_j^2+\epsilon}r=N1j=1Nxj2+ϵ

yj=xjrγjy_j=\frac{x_j}{r}\gamma_jyj=rxjγj

因此 RMSNorm 通常比完整 LayerNorm 少一部分计算。


3. 纯 Python LayerNorm

import math


def python_layer_norm(
    row: list[float],
    weight: list[float],
    bias: list[float],
    eps: float = 1e-5,
) -> list[float]:
    n = len(row)

    mean = sum(row) / n

    variance = sum(
        (value - mean) ** 2
        for value in row
    ) / n

    inv_std = 1.0 / math.sqrt(variance + eps)

    return [
        (row[i] - mean) * inv_std * weight[i] + bias[i]
        for i in range(n)
    ]

4. PyTorch LayerNorm

import torch.nn.functional as F


def torch_layer_norm(
    x: torch.Tensor,
    weight: torch.Tensor,
    bias: torch.Tensor,
    eps: float = 1e-5,
) -> torch.Tensor:
    return F.layer_norm(
        x,
        normalized_shape=(x.shape[-1],),
        weight=weight,
        bias=bias,
        eps=eps,
    )

PyTorch 的 LayerNorm 会使用优化实现,并尽量融合统计量计算和归一化。


5. PyTorch RMSNorm

def torch_rms_norm(
    x: torch.Tensor,
    weight: torch.Tensor,
    eps: float = 1e-6,
) -> torch.Tensor:
    variance = x.float().pow(2).mean(
        dim=-1,
        keepdim=True,
    )

    normalized = x * torch.rsqrt(
        variance + eps
    ).to(x.dtype)

    return normalized * weight

较新的 PyTorch 版本也可能提供对应的 RMSNorm 模块或函数,但手写参考实现更适合验证自定义 Kernel。


6. TileLang RMSNorm

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_rms_norm(
    X,
    Weight,
    block_rows: int,
    eps: float,
):
    M, N = T.const("M, N")

    X: T.Tensor((M, N), T.float32)
    Weight: T.Tensor((N,), T.float32)

    Y = T.empty((M, N), T.float32)

    with T.Kernel(
        T.ceildiv(M, block_rows),
        threads=128,
    ) as bx:
        x_local = T.alloc_fragment(
            (block_rows, N),
            T.float32,
        )
        square_local = T.alloc_fragment(
            (block_rows, N),
            T.float32,
        )
        square_sum = T.alloc_fragment(
            (block_rows,),
            T.float32,
        )

        T.copy(
            X[bx * block_rows : (bx + 1) * block_rows, :],
            x_local,
        )

        for i, j in T.Parallel(block_rows, N):
            square_local[i, j] = (
                x_local[i, j] * x_local[i, j]
            )

        T.reduce_sum(
            square_local,
            square_sum,
            dim=1,
        )

        for i in T.Parallel(block_rows):
            square_sum[i] = T.rsqrt(
                square_sum[i] / N + eps
            )

        for i, j in T.Parallel(block_rows, N):
            x_local[i, j] = (
                x_local[i, j]
                * square_sum[i]
                * Weight[j]
            )

        T.copy(
            x_local,
            Y[bx * block_rows : (bx + 1) * block_rows, :],
        )

    return Y

验证:

def test_rms_norm() -> None:
    m = 4096
    n = 1024
    eps = 1e-6

    x = torch.randn(
        m,
        n,
        device="cuda",
        dtype=torch.float32,
    )
    weight = torch.randn(
        n,
        device="cuda",
        dtype=torch.float32,
    )

    kernel = tilelang_rms_norm.compile(
        x,
        weight,
        block_rows=1,
        eps=eps,
    )

    actual = kernel(x, weight)
    expected = torch_rms_norm(
        x,
        weight,
        eps,
    )

    torch.testing.assert_close(
        actual,
        expected,
        rtol=1e-4,
        atol=1e-4,
    )

    print("RMSNorm passed.")

官方 RMSNorm 示例同样先计算平方、调用 T.reduce_sum、使用 T.rsqrt,再将归一化系数应用到原输入。


7. TileLang LayerNorm

@tilelang.jit
def tilelang_layer_norm(
    X,
    Weight,
    Bias,
    block_rows: int,
    eps: float,
):
    M, N = T.const("M, N")

    X: T.Tensor((M, N), T.float32)
    Weight: T.Tensor((N,), T.float32)
    Bias: T.Tensor((N,), T.float32)

    Y = T.empty((M, N), T.float32)

    with T.Kernel(
        T.ceildiv(M, block_rows),
        threads=128,
    ) as bx:
        x_local = T.alloc_fragment(
            (block_rows, N),
            T.float32,
        )
        centered_square = T.alloc_fragment(
            (block_rows, N),
            T.float32,
        )
        mean = T.alloc_fragment(
            (block_rows,),
            T.float32,
        )
        variance = T.alloc_fragment(
            (block_rows,),
            T.float32,
        )

        T.copy(
            X[bx * block_rows : (bx + 1) * block_rows, :],
            x_local,
        )

        T.reduce_sum(
            x_local,
            mean,
            dim=1,
        )

        for i in T.Parallel(block_rows):
            mean[i] = mean[i] / N

        for i, j in T.Parallel(block_rows, N):
            centered = x_local[i, j] - mean[i]
            centered_square[i, j] = centered * centered

        T.reduce_sum(
            centered_square,
            variance,
            dim=1,
        )

        for i in T.Parallel(block_rows):
            variance[i] = T.rsqrt(
                variance[i] / N + eps
            )

        for i, j in T.Parallel(block_rows, N):
            normalized = (
                x_local[i, j] - mean[i]
            ) * variance[i]

            x_local[i, j] = (
                normalized * Weight[j] + Bias[j]
            )

        T.copy(
            x_local,
            Y[bx * block_rows : (bx + 1) * block_rows, :],
        )

    return Y

8. Norm 性能差异在哪里?

若使用普通 PyTorch 组合:

mean = x.mean(...)
variance = ...
normalized = ...
output = normalized * weight + bias

理论上可能产生多个 Kernel 和中间张量。

成熟的 F.layer_norm 已经会调用融合实现,因此手写 TileLang 不一定更快。

TileLang 主要适合以下场景:

Residual Add
    ↓
Bias
    ↓
LayerNorm / RMSNorm
    ↓
Activation 或 Quantization

如果把这些步骤融合为一个 Kernel,可以减少:

  • 中间张量;
  • Kernel 启动;
  • 显存读写;
  • 类型转换。

十、示例六:FlashAttention

这是整条学习路线里最复杂的部分。

1. 普通 Attention

输入:

Q: B × H × S_q × D
K: B × H × S_k × D
V: B × H × S_k × D_v

计算:

S=QKTDS=\frac{QK^T}{\sqrt{D}}S=DQKT

P=softmax⁡(S)P=\operatorname{softmax}(S)P=softmax(S)

O=PVO=PVO=PV


2. 纯 Python Attention

这里只适合极小尺寸教学:

import math


def dot(a: list[float], b: list[float]) -> float:
    return sum(x * y for x, y in zip(a, b))


def python_attention(
    q: list[list[float]],
    k: list[list[float]],
    v: list[list[float]],
) -> list[list[float]]:
    sequence_q = len(q)
    sequence_k = len(k)
    head_dim = len(q[0])
    value_dim = len(v[0])

    scale = 1.0 / math.sqrt(head_dim)

    output = [
        [0.0] * value_dim
        for _ in range(sequence_q)
    ]

    for i in range(sequence_q):
        scores = [
            dot(q[i], k[j]) * scale
            for j in range(sequence_k)
        ]

        maximum = max(scores)

        probabilities = [
            math.exp(score - maximum)
            for score in scores
        ]

        denominator = sum(probabilities)
        probabilities = [
            value / denominator
            for value in probabilities
        ]

        for j in range(sequence_k):
            probability = probabilities[j]

            for d in range(value_dim):
                output[i][d] += (
                    probability * v[j][d]
                )

    return output

3. PyTorch 普通实现

import math
import torch


def torch_attention_naive(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
) -> torch.Tensor:
    scale = 1.0 / math.sqrt(q.shape[-1])

    scores = torch.matmul(
        q,
        k.transpose(-2, -1),
    ) * scale

    probabilities = torch.softmax(
        scores,
        dim=-1,
    )

    return torch.matmul(
        probabilities,
        v,
    )

问题在于会显式创建注意力矩阵:

B × H × S_q × S_k

当序列很长时,这个矩阵非常大。

例如:

B = 1
H = 32
S_q = S_k = 8192
FP16

注意力分数张量仅自身占用约:

1×32×8192×8192×2 bytes1\times32\times8192\times8192\times2\ \text{bytes}1×32×8192×8192×2 bytes

约为 4 GiB,尚未包含其他中间张量。


4. PyTorch SDPA

生产环境应优先考虑:

import torch.nn.functional as F


def torch_attention_sdpa(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
) -> torch.Tensor:
    return F.scaled_dot_product_attention(
        q,
        k,
        v,
        dropout_p=0.0,
        is_causal=False,
    )

PyTorch 会根据硬件、dtype、形状和设置选择可用后端。作为基准时,应明确区分:

  • 手工 matmul + softmax + matmul
  • PyTorch SDPA;
  • TileLang 自定义 FlashAttention。

否则对比可能不公平。


十一、FlashAttention 的核心思想

普通 Attention:

QK^T
  ↓
把整个 S 写入显存
  ↓
读取 S 做 Softmax
  ↓
把 P 写入显存
  ↓
读取 P 与 V 相乘

FlashAttention:

读取 Q Tile
    ↓
循环读取 K/V Tile
    ↓
局部计算 QK^T
    ↓
在线更新 Softmax
    ↓
立即乘 V
    ↓
不保存完整注意力矩阵

关键收益是:

  • 不把完整分数矩阵写回显存;
  • 不把完整概率矩阵写回显存;
  • 将分数 Tile 保存在寄存器或 Shared Memory;
  • QKᵀ、Softmax 和 PV 融合。

十二、在线 Softmax

假设已经处理了一部分 Key,已有:

  • 旧最大值 m_old
  • 旧指数和 l_old
  • 旧输出累积 O_old

处理新分块时,得到新块最大值 m_block

更新最大值:

mnew=max⁡(mold,mblock)m_{\text{new}}=\max(m_{\text{old}},m_{\text{block}})mnew=max(mold,mblock)

旧归一化量缩放:

α=emold−mnew\alpha=e^{m_{\text{old}}-m_{\text{new}}}α=emoldmnew

新块指数和:

lblock=∑jesj−mnewl_{\text{block}}=\sum_j e^{s_j-m_{\text{new}}}lblock=jesjmnew

更新指数和:

lnew=αlold+lblockl_{\text{new}}=\alpha l_{\text{old}}+l_{\text{block}}lnew=αlold+lblock

更新输出累积:

Onew=αOold+∑jesj−mnewVjO_{\text{new}}=\alpha O_{\text{old}}+\sum_j e^{s_j-m_{\text{new}}}V_jOnew=αOold+jesjmnewVj

全部 Key Tile 处理完后:

O=OnewlnewO=\frac{O_{\text{new}}}{l_{\text{new}}}O=lnewOnew

这样每次只需要保存少量统计量,不需要保存整个注意力矩阵。


十三、TileLang FlashAttention 教学版结构

完整的高性能 FlashAttention Kernel 很长,而且需要针对:

  • GPU 架构;
  • head dimension;
  • causal mask;
  • GQA/MQA;
  • FP16/BF16;
  • Hopper WGMMA;
  • TMA;
  • warp specialization;

分别调优。

下面给出核心骨架,重点是理解数据流。

import tilelang
import tilelang.language as T


@tilelang.jit
def tilelang_flash_attention(
    Q,
    K,
    V,
    block_m: int,
    block_n: int,
    block_d: int,
):
    B, H, SQ, SK, D = T.const(
        "B, H, SQ, SK, D"
    )

    Q: T.Tensor((B, H, SQ, D), T.float16)
    K: T.Tensor((B, H, SK, D), T.float16)
    V: T.Tensor((B, H, SK, D), T.float16)

    O = T.empty(
        (B, H, SQ, D),
        T.float16,
    )

    scale = 1.0 / T.sqrt(D)

    with T.Kernel(
        T.ceildiv(SQ, block_m),
        H,
        B,
        threads=256,
    ) as (q_block, head, batch):
        q_shared = T.alloc_shared(
            (block_m, block_d),
            T.float16,
        )
        k_shared = T.alloc_shared(
            (block_n, block_d),
            T.float16,
        )
        v_shared = T.alloc_shared(
            (block_n, block_d),
            T.float16,
        )

        scores = T.alloc_fragment(
            (block_m, block_n),
            T.float32,
        )
        output_acc = T.alloc_fragment(
            (block_m, block_d),
            T.float32,
        )

        row_max = T.alloc_fragment(
            (block_m,),
            T.float32,
        )
        row_max_old = T.alloc_fragment(
            (block_m,),
            T.float32,
        )
        row_sum = T.alloc_fragment(
            (block_m,),
            T.float32,
        )
        row_sum_block = T.alloc_fragment(
            (block_m,),
            T.float32,
        )
        correction = T.alloc_fragment(
            (block_m,),
            T.float32,
        )

        T.copy(
            Q[
                batch,
                head,
                q_block * block_m,
                0,
            ],
            q_shared,
        )

        T.clear(output_acc)
        T.fill(row_max, -T.infinity(T.float32))
        T.clear(row_sum)

        for kv_block in T.Pipelined(
            T.ceildiv(SK, block_n),
            num_stages=2,
        ):
            T.copy(row_max, row_max_old)

            T.copy(
                K[
                    batch,
                    head,
                    kv_block * block_n,
                    0,
                ],
                k_shared,
            )
            T.copy(
                V[
                    batch,
                    head,
                    kv_block * block_n,
                    0,
                ],
                v_shared,
            )

            T.clear(scores)

            T.gemm(
                q_shared,
                k_shared,
                scores,
                transpose_B=True,
            )

            for i, j in T.Parallel(
                block_m,
                block_n,
            ):
                scores[i, j] = scores[i, j] * scale

            T.fill(
                row_max,
                -T.infinity(T.float32),
            )

            T.reduce_max(
                scores,
                row_max,
                dim=1,
                clear=False,
            )

            for i in T.Parallel(block_m):
                row_max[i] = T.max(
                    row_max[i],
                    row_max_old[i],
                )

                correction[i] = T.exp(
                    row_max_old[i] - row_max[i]
                )

            for i, j in T.Parallel(
                block_m,
                block_n,
            ):
                scores[i, j] = T.exp(
                    scores[i, j] - row_max[i]
                )

            T.reduce_sum(
                scores,
                row_sum_block,
                dim=1,
            )

            for i in T.Parallel(block_m):
                row_sum[i] = (
                    row_sum[i] * correction[i]
                    + row_sum_block[i]
                )

            for i, j in T.Parallel(
                block_m,
                block_d,
            ):
                output_acc[i, j] = (
                    output_acc[i, j]
                    * correction[i]
                )

            T.gemm(
                scores,
                v_shared,
                output_acc,
            )

        for i, j in T.Parallel(
            block_m,
            block_d,
        ):
            output_acc[i, j] = (
                output_acc[i, j]
                / row_sum[i]
            )

        T.copy(
            output_acc,
            O[
                batch,
                head,
                q_block * block_m,
                0,
            ],
        )

    return O

这个代码表达了 FlashAttention 的核心,但在实际运行前还需要根据当前 TileLang 版本和 GPU 架构处理:

  • 尾块边界;
  • K/V 的实际 Tile 布局;
  • scores 到 GEMM 输入的 dtype 转换;
  • causal mask;
  • Shared Memory 限制;
  • Tensor Core 支持的 Tile 形状;
  • warp policy;
  • layout inference;
  • FP32 累加;
  • Hopper 或 Ampere 特定调度。

官方 MLA/Attention 示例使用两个 T.gemm 连接 QKᵀPV,中间通过 T.reduce_maxT.reduce_sum 和在线缩放维护 Softmax 状态。


十四、完整性能对比应该怎么做?

不要直接假设 TileLang 一定更快。应该针对自己的 GPU 和输入形状测量。

1. 向量加法基准

def benchmark_vector_add() -> None:
    n = 1 << 24

    a = torch.randn(
        n,
        device="cuda",
        dtype=torch.float32,
    )
    b = torch.randn_like(a)

    kernel = tilelang_vector_add.compile(
        a,
        b,
        block_size=256,
    )

    torch_ms = benchmark_cuda(
        lambda: a + b
    )

    tilelang_ms = benchmark_cuda(
        lambda: kernel(a, b)
    )

    print(f"PyTorch:  {torch_ms:.4f} ms")
    print(f"TileLang: {tilelang_ms:.4f} ms")

有效带宽估算:

def vector_add_bandwidth_gbs(
    n: int,
    latency_ms: float,
    element_size: int = 4,
) -> float:
    # A 读取 + B 读取 + C 写入
    bytes_moved = n * element_size * 3
    seconds = latency_ms / 1000
    return bytes_moved / seconds / 1e9

2. GEMM 基准

def benchmark_gemm() -> None:
    m = n = k = 4096

    a = torch.randn(
        m,
        k,
        device="cuda",
        dtype=torch.float16,
    )
    b = torch.randn(
        k,
        n,
        device="cuda",
        dtype=torch.float16,
    )

    kernel = tilelang_matmul.compile(
        a,
        b,
        block_m=128,
        block_n=128,
        block_k=32,
    )

    torch_ms = benchmark_cuda(
        lambda: a @ b,
        warmup=20,
        repeat=50,
    )

    tilelang_ms = benchmark_cuda(
        lambda: kernel(a, b),
        warmup=20,
        repeat=50,
    )

    print(f"PyTorch GEMM:  {torch_ms:.4f} ms")
    print(f"TileLang GEMM: {tilelang_ms:.4f} ms")

    print(
        "PyTorch TFLOPS:",
        gemm_tflops(m, n, k, torch_ms),
    )
    print(
        "TileLang TFLOPS:",
        gemm_tflops(m, n, k, tilelang_ms),
    )

3. Attention 基准

需要至少比较三个版本:

naive = lambda: torch_attention_naive(q, k, v)
sdpa = lambda: torch_attention_sdpa(q, k, v)
tilelang_fn = lambda: tilelang_kernel(q, k, v)

测试形状要明确:

Batch
Heads
Query length
Key length
Head dimension
dtype
causal / non-causal
GQA / MHA

只给出“Attention 快了多少”而不提供这些参数,几乎没有意义。


十五、通常会看到怎样的性能关系?

以下是常见趋势,不是固定数值。

1. 向量加法

纯 Python  <<  PyTorch ≈ 优化正确的 TileLang

原因:

  • Python 有解释器开销;
  • PyTorch 和 TileLang 都可生成单个并行 GPU Kernel;
  • 主要受显存带宽限制;
  • TileLang 没有很大的融合空间。

2. Reduce / Softmax

纯 Python  <<  普通拆分 PyTorch
                 ≤ 融合 PyTorch / TileLang

TileLang 可能受益于:

  • 单 Kernel 完成多阶段操作;
  • 中间值停留在寄存器;
  • 避免中间张量;
  • 针对固定行长度定制 warp 布局。

但 PyTorch 内置 softmax 已经是优化实现,因此真正合理的比较是:

TileLang Softmax
vs
torch.softmax

而不是把 PyTorch 手工拆成五个算子后宣称 TileLang 大幅领先。


3. GEMM

纯 Python  <<  普通自定义 GPU Kernel
                 ≤ TileLang 优化实现
                 ≈ 或 < cuBLAS / PyTorch

规则大型 GEMM 中,cuBLAS 非常难超越。

TileLang 更容易在以下场景产生优势:

  • GEMM + Bias + Activation;
  • GEMM + Quantization;
  • 不规则形状;
  • Grouped GEMM;
  • 特殊数据类型;
  • 特定模型结构。

4. LayerNorm/RMSNorm

纯 Python << 拆分 PyTorch < 融合 PyTorch ≈ 优化 TileLang

若 TileLang 把:

Residual + Norm + Weight + Quantization

融合成一个 Kernel,则可能明显减少显存流量。


5. FlashAttention

纯 Python << naive PyTorch Attention < 优化 SDPA / FlashAttention

TileLang 的目标不是简单击败错误的 naive 实现,而是用较高层的 DSL 写出接近专用 FlashAttention/CUTLASS Kernel 的实现。TileLang 官方仓库展示了 Flash Attention、MLA 等场景的性能结果,但实际性能依赖 GPU、形状、dtype 和调度配置,应使用官方基准脚本或自己的真实负载复测。


十六、如何判断算子是“算力受限”还是“带宽受限”?

带宽受限

典型算子:

  • 向量加法;
  • ReLU;
  • RMSNorm;
  • LayerNorm;
  • Softmax;
  • 很多逐元素融合算子。

特点:

每读取很多字节,只做少量计算

优化重点:

  • 减少显存读写;
  • 合并访存;
  • 算子融合;
  • 向量化加载;
  • 避免中间张量。

算力受限

典型算子:

  • 大型 GEMM;
  • 大型卷积;
  • Attention 中的矩阵乘法阶段。

特点:

同一份数据被重复用于大量乘加

优化重点:

  • Tensor Core;
  • Tile 复用;
  • Shared Memory;
  • 寄存器分块;
  • 软件流水线;
  • 合理 occupancy。

十七、正确性验证规范

1. 永远先写参考实现

def reference(x):
    return ...

再写 TileLang。

2. 使用随机输入

x = torch.randn(...)

还要额外测试:

  • 全零;
  • 全一;
  • 极大值;
  • 极小值;
  • 非整 Tile 尺寸;
  • 长度 1;
  • 奇数尺寸;
  • NaN 和 Inf。

3. 混合精度允许误差

FP32:

rtol=1e-5
atol=1e-5

FP16:

rtol=1e-2
atol=1e-2

注意这只是常见起点,不是所有算子的统一标准。

4. 查看生成代码

print(kernel.get_kernel_source())

官方 GEMM 示例直接提供 get_kernel_source() 和 profiler 使用方式。


十八、常见错误与排查

1. 越界访问

错误:

index = bx * block_size + tx
C[index] = A[index] + B[index]

当长度不是 block size 的整数倍时可能越界。

修复:

if index < N:
    C[index] = A[index] + B[index]

TileLang 当前控制流文档也说明了边界保护和安全内存访问 Lowering。


2. 结果有小误差

可能原因:

  • FP16 累加;
  • Reduction 顺序不同;
  • fast math;
  • exp2exp 转换;
  • Tensor Core 精度规则。

解决:

  • 使用 FP32 accumulator;
  • 合理设置容差;
  • 与 FP32 PyTorch 参考结果比较;
  • 不要用逐元素完全相等判断浮点结果。

3. TileLang 比 PyTorch 慢

常见原因:

  • Tile 尺寸不合适;
  • 线程太少;
  • Shared Memory 太多;
  • 寄存器溢出;
  • 非合并访存;
  • 没有使用 Tensor Core;
  • 小输入被 JIT 或启动开销主导;
  • 与高度优化的 cuBLAS 比较;
  • benchmark 没有 warm-up;
  • 没有同步 GPU;
  • 每轮都重新编译。

4. 编译失败

先检查:

print(kernel.get_kernel_source())

以及:

  • 输入 dtype;
  • Shape 是否匹配;
  • Tile 是否被硬件指令支持;
  • Shared Memory 是否超限;
  • CUDA 架构是否匹配;
  • 当前 TileLang 示例是否与安装版本一致。

TileLang 官方调试指南把问题分成生成错误、正确性错误和性能错误,并建议检查逐步 Lowering 的 IR;硬件性能问题则可进一步使用 Nsight Compute 或 rocProf。


十九、建议的练习顺序

第 1 周:建立 GPU 基础

完成:

向量加法
ReLU
Scale
Add + ReLU 融合

目标:

  • 理解 T.Kernel
  • 理解 block 和 thread;
  • 理解连续访存;
  • 学会正确 benchmark。

第 2 周:归约

完成:

Row Sum
Row Max
Softmax
LogSoftmax

目标:

  • 理解 Reduction;
  • 理解数值稳定性;
  • 理解 warp/block 内归约。

第 3 周:矩阵乘法

完成:

Naive GEMM
Shared Memory GEMM
Pipelined GEMM
GEMM + ReLU

目标:

  • 理解 Tile;
  • 理解 Shared Memory 复用;
  • 理解 Fragment;
  • 理解 Tensor Core。

第 4 周:归一化

完成:

RMSNorm
LayerNorm
Residual + RMSNorm

目标:

  • 理解融合;
  • 理解带宽瓶颈;
  • 理解 FP32 累加。

第 5 周以后:Attention

完成:

Naive Attention
Tiled QKᵀ
Online Softmax
Tiled PV
Causal Mask
完整 FlashAttention

目标:

  • 理解 IO-aware 算法;
  • 理解在线 Softmax;
  • 理解复杂 Kernel 的布局和流水线。

二十、最重要的性能结论

不要把 TileLang 理解为:

把 PyTorch 代码换一种语法写一遍,就自动变快。

真正决定性能的是:

算法
+ 数据布局
+ Tile 大小
+ 内存层级
+ 访存次数
+ Kernel 融合
+ 并行度
+ Tensor Core 利用率
+ 流水线

最值得牢记的是:

向量加法

优化目标是减少显存时间,几乎没有计算复用。

Softmax

优化目标是把 max、exp、sum 和 normalize 融合,并让中间值留在片上。

GEMM

优化目标是让每次从显存读入的数据参与尽可能多的乘加。

RMSNorm/LayerNorm

优化目标是融合统计量计算、归一化、缩放和相邻逐元素操作。

FlashAttention

优化目标是改变算法的数据流,不生成完整注意力矩阵,而不是仅仅把普通 Attention 翻译成 GPU 代码。

TileLang 的价值就在于:它让你用 Python 风格代码表达这些 GPU 级优化,同时仍然暴露 Shared Memory、Fragment、流水线和 Tile GEMM 等关键控制能力。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐