从零开始学习 TileLang:从 Python、PyTorch 到高性能 GPU Kernel 的完整开发教程
从零开始学习 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,也可能存在尚未稳定的变化。
对于初学者,建议:
- 优先使用稳定版本;
- 使用与安装版本匹配的官方文档;
- 示例报错时先核对版本;
- 不要直接混用不同版本的教程代码。
第四章: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=0∑N−1Xi,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=maxjXi,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,j−mi)
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)=2xlog2(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=0∑K−1Ai,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
)
逻辑上包含:
- 矩阵乘法;
- Bias Add;
- 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=N1j∑Xi,j
σi2=1N∑jXi,j2−μi2\sigma_i^2=\frac{1}{N}\sum_jX_{i,j}^2-\mu_i^2σi2=N1j∑Xi,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(...):
...
然后在 forward 和 backward 中分别调用 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=N1∑jXi,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(mold−mnew)
lnew=αlold+∑jexp(sj−mnew)l_{\mathrm{new}}=\alpha l_{\mathrm{old}}+\sum_j\exp(s_j-m_{\mathrm{new}})lnew=αlold+j∑exp(sj−mnew)
onew=αoold+∑jexp(sj−mnew)Vjo_{\mathrm{new}}=\alpha o_{\mathrm{old}}+\sum_j\exp(s_j-m_{\mathrm{new}})V_jonew=αoold+j∑exp(sj−mnew)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 |
| Softmax | Reduce、指数、内存 | Mask、Scale、Dropout 融合 |
| GEMM | Tensor Core 与数据复用 | 特殊形状、量化、Epilogue |
| LayerNorm | 内存带宽与 Reduce | Residual、Bias、量化融合 |
| RMSNorm | 内存带宽与 Reduce | LLM 特殊融合 |
| FlashAttention | GEMM、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 最重要的三个问题是:
- 一个线程块负责哪一块输出?
- 数据什么时候从 Global Memory 搬到 Shared Memory?
- 中间结果应该放在 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
每个阶段都会对比三种实现:
- 纯 Python:帮助理解算法。
- PyTorch:生产环境中最常见的高层写法。
- TileLang:自己控制 GPU Kernel 的分块、内存和流水线。
TileLang 的 API 仍在快速演进,因此不同版本中的少量函数签名可能发生变化。本文代码采用官方文档当前推荐的 @tilelang.jit、T.Kernel、T.Parallel、T.copy、T.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.Pipelined、T.copy、T.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 Intensity≈12 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=0∑N−1Xij
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-1N−1
次依赖加法,而且不能充分并行。
树形归约可在大约:
log2N\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=maxjxjm=\max_j x_jm=jmaxxj
yi=exi−m∑jexj−my_i=\frac{e^{x_i-m}}{\sum_j e^{x_j-m}}yi=∑jexj−mexi−m
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。
从逻辑上需要:
- 读取输入,求最大值;
- 计算指数;
- 求指数和;
- 除以指数和;
- 写出结果。
如果每一步都单独启动 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=0∑K−1AikBkj
矩阵形状:
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.Pipelined、T.copy、T.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=1∑Nxj
σ2=1N∑j=1N(xj−μ)2\sigma^2=\frac{1}{N}\sum_{j=1}^{N}(x_j-\mu)^2σ2=N1j=1∑N(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=1∑Nxj2+ϵ
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}}}α=emold−mnew
新块指数和:
lblock=∑jesj−mnewl_{\text{block}}=\sum_j e^{s_j-m_{\text{new}}}lblock=j∑esj−mnew
更新指数和:
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+j∑esj−mnewVj
全部 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_max、T.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;
exp2与exp转换;- 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 等关键控制能力。
更多推荐


所有评论(0)