Task2 TileLang 零基础完成指南
这份指南为你提供完成 task2 TileLang 任务 需要的前置知识。
task2 任务目标:完成 solution.py 中的一维向量加法,通过测试,并得到 benchmark 结果。
如果你已经完成作业,想理解 JIT、GPU block、并行映射和性能问题,请阅读《从 a + b 到 GPU Kernel:TileLang 原理详解》。
0. 先看任务最终要做什么
输入是两个长度相同的一维张量 A 和 B,输出 C:
A = [1, 2, 3]
B = [4, 5, 6]
C = [5, 7, 9]每个位置的计算规则是:
C[i] = A[i] + B[i]你需要:
- 在
solution.py中完成tl_add_1d; - 使用
T.Kernel和T.Parallel; - 正确处理任意长度,包括最后一块放不满的情况;
- 通过
test_add.py; - 运行
benchmark_add.py并保存终端结果。
这个任务你要实践的是如何安排 GPU 上的并行工作。
1. 确认环境可以运行
TileLang kernel 需要在配置好的 GPU 环境中运行。先从仓库根目录加载环境:
cd /gollamago
source ./setup_env.sh
cd assignment/task2然后检查 Python 能否导入 TileLang 和 PyTorch:
python -c "import tilelang, torch; print('TileLang: OK'); print('GPU:', torch.cuda.is_available())"预期能看到:
TileLang: OK
GPU: True如果导入失败,先解决环境问题,不要修改 kernel 代码碰运气:
Could not find a built TileLang development tree:没有找到已编译的 TileLang,需要按提示设置TILELANG_ROOT;No module named tilelang:通常是没有在当前终端执行source ./setup_env.sh;GPU: False:当前 Python 环境没有识别到任务所需的 GPU。
source 只对当前终端会话生效。新开终端后需要重新执行。
2. 认识你要修改的代码
本指南以尚未完成的 solution.py 为起点。核心区域应当类似下面的骨架:
import tilelang
import tilelang.language as T
@tilelang.jit
def tl_add_1d(A, B, BLOCK_N: int):
N = T.const("N")
A: T.Tensor((N,), T.float16)
B: T.Tensor((N,), T.float16)
C = T.empty((N,), T.float16)
# 在这里完成 kernel
raise NotImplementedError("请根据步骤实现 Add")T 是什么?
文件开头的:
import tilelang.language as T把 tilelang.language 简称为 T,和 import numpy as np 一样。因此 T.Kernel、T.Parallel 和 T.ceildiv 都是 TileLang 提供的工具。另一行 import tilelang 主要用于 @tilelang.jit。
暂时只需要知道:
A、B是输入;C是等待写入的输出;N是向量长度;BLOCK_N是每个工作块计划处理的元素数量;float16是输入和输出的数据类型;@tilelang.jit让 TileLang 编译并运行这个函数。
如果你的文件已经有部分实现,不要机械地覆盖它。保留文件开头的声明,对照后面五步检查并补全逻辑即可。
3. 最少前置知识:为什么要“分块”
普通 Python 可以逐个元素计算:
for i in range(N):
C[i] = A[i] + B[i]但这表示按顺序完成 N 次循环,没有表达 GPU 并行。
GPU 更适合把工作拆成许多份,同时处理。假设:
N = 10
BLOCK_N = 4我们把 10 个元素按每组 4 个切分:
第 0 块:下标 0、1、2、3
第 1 块:下标 4、5、6、7
第 2 块:下标 8、9,以及两个超出范围的位置最后一块没有放满,称为 尾块。Task2 最关键的地方就是:
- 必须创建第 2 块,否则 8、9 会被漏掉;
- 又不能访问不存在的 10、11,否则会越界。
下面的五个步骤正好解决这两个问题。
4. 第一步:计算需要多少个工作块
我们第一步计算需要多少个工作块:
num_blocks = T.ceildiv(N, BLOCK_N)T.ceildiv 是向上取整除法。
ceildiv(10, 4) = 3
ceildiv(8, 4) = 2
ceildiv(1, 4) = 1不能写成:
num_blocks = N // BLOCK_N因为 10 // 4 只得到 2,会漏掉最后两个元素;1 // 4 甚至会得到 0。
此时代码的含义是:无论最后一块是否填满,都为它保留一个工作块。
5. 第二步:声明工作块网格并确定每块的起点
继续添加:
with T.Kernel(num_blocks, threads=BLOCK_N) as pid:
base_idx = pid * BLOCK_N注意第二行需要缩进。
这两行同时回答了四个问题:
总共启动多少个工作块? num_blocks
每个工作块使用多少线程? threads=BLOCK_N
当前是第几个工作块? pid
当前块从哪个元素开始? base_idx下面逐项解释。
T.Kernel(num_blocks, ...):启动多少个工作块
上一步已经计算出:
num_blocks = T.ceildiv(N, BLOCK_N)现在把 num_blocks 交给 T.Kernel,表示创建一维的工作块网格。例如 N=10、BLOCK_N=4 时,num_blocks=3,因此需要编号为 0、1、2 的三个工作块。
这里的“工作块”是 GPU block。每个 block 负责输入向量中的一个数据块,也就是一个 tile。
threads=BLOCK_N:每个工作块内部的线程配置
一个 GPU block 内部还包含多个可以协作并行的线程。threads=BLOCK_N 指定每个 block 使用的线程数量。
在这份入门练习中,我们让线程配置与 tile 长度使用同一个参数:
一个 tile 计划处理 BLOCK_N 个元素
一个 GPU block 配置 BLOCK_N 个线程这样便于表达“一个 block 并行处理一个 tile”。但现阶段不需要推导线程在硬件上如何调度,也不要由此得出“所有 kernel 都必须让线程数等于 tile 大小”。更复杂的 kernel 可能让一个线程处理多个元素,或者让 tile 大小与线程数不同;这些属于原理和性能调优内容。
as pid:当前工作块的编号
T.Kernel 描述的是所有工作块共同执行的代码。每个工作块运行到这里时,都通过 pid 得到自己的编号:
第一个工作块:pid = 0
第二个工作块:pid = 1
第三个工作块:pid = 2如果没有这个编号,每个工作块就不知道自己应该处理向量的哪一段。
base_idx:当前块负责的数据从哪里开始
base_idx = pid * BLOCK_N这是本步骤最需要理解的计算。每个工作块负责 BLOCK_N 个连续位置,所以:
当 BLOCK_N = 4 时:
pid = 0 → base_idx = 0
pid = 1 → base_idx = 4
pid = 2 → base_idx = 8base_idx 是当前 tile 的起始下标,不是最终要访问的元素下标。下一步还要加上 tile 内部的偏移 i:
最终下标 index = base_idx + i因此,不同工作块先通过 base_idx 找到各自的数据区域,再在区域内部并行处理不同位置。
这一节需要掌握到什么程度?
继续完成任务前,你只需能够回答:
num_blocks决定总共有多少个 GPU block;threads决定每个 block 的线程配置;pid区分当前是哪一个 block;base_idx = pid * BLOCK_N找到当前 tile 的起点。
至于线程如何被 GPU 调度、BLOCK_N 为什么影响性能,以及 tile 与线程数为什么可以不同,都不影响下一步填写代码,留到原理详解中再学习。
6. 第三步:在当前块内并行遍历
在 with 内继续添加:
for i in T.Parallel(BLOCK_N):
index = base_idx + i这里仍然要注意缩进关系:
with T.Kernel(...)
base_idx = ...
for i in T.Parallel(...)
index = ...T.Parallel(BLOCK_N) 表示这 BLOCK_N 次迭代可以并行执行;它不像普通 Python 的 range 那样表达顺序循环。现阶段只需把 i 理解为当前块内部的偏移,范围是 0 到 BLOCK_N - 1,不需要了解它如何对应到底层线程。
全局下标的计算公式是:
index = 当前块起点 + 块内偏移
= pid × BLOCK_N + i用 N=10、BLOCK_N=4 手算一次:
pid=0:index 为 0、1、2、3
pid=1:index 为 4、5、6、7
pid=2:index 为 8、9、10、11现在 0 到 9 都被覆盖了,但 10 和 11 越界,因此还不能直接读写张量。
7. 第四步:保护尾块
在循环内添加边界判断:
if index < N:这意味着只有真实存在的下标才能继续计算。
对于最后一个工作块:
8 < 10 → 执行
9 < 10 → 执行
10 < 10 → 跳过
11 < 10 → 跳过这个判断必须位于读取 A[index]、B[index] 和写入 C[index] 之前,因为三者都不能越界。
8. 第五步:计算并返回结果
在边界判断内写入输出:
C[index] = A[index] + B[index]离开 with 作用域后返回输出:
return C最终缩进结构应当是:
函数
计算工作块数量
Kernel
计算当前块起点
并行循环
计算全局下标
边界判断
执行加法并写回
返回输出同时删除原来的:
raise NotImplementedError("请根据步骤实现 Add")不要把它留在 return C 前面,否则函数会在返回结果前主动报错。
如果仍然无法组合代码,可以核对核心实现:
num_blocks = T.ceildiv(N, BLOCK_N)
with T.Kernel(num_blocks, threads=BLOCK_N) as pid:
base_idx = pid * BLOCK_N
for i in T.Parallel(BLOCK_N):
index = base_idx + i
if index < N:
C[index] = A[index] + B[index]
return C9. 先做一次人工检查
运行代码前,用下面四个问题检查自己的实现:
N=1、BLOCK_N=1024时,num_blocks是不是 1?- 不同的
pid是否会得到不同的base_idx? pid * BLOCK_N + i是否覆盖了0到N-1?- 所有张量访问是否都被
index < N保护?
如果四个答案都是“是”,核心逻辑通常就是完整的。
10. 运行测试
在 assignment/task2 目录运行:
python -m pytest -q test_add.py预期结果类似:
5 passed测试会覆盖长度:
1、127、1024、100003、1048576如果失败,按以下顺序排查:
出现 NotImplementedError
确认已经删除原来的 raise NotImplementedError(...)。
只有不能整除的长度失败
检查是否使用了 T.ceildiv,以及是否存在 if index < N。
长度 1 失败
检查块数是不是误用了 N // BLOCK_N。
报缩进或语法错误
对照第 8 节的缩进结构。Python 使用缩进表示代码的包含关系。
报 CPU、CUDA 或设备相关错误
确认环境检查中的 GPU: True,并确认当前终端执行过 source ./setup_env.sh。
11. 运行 benchmark
测试通过后运行:
python benchmark_add.py输出包含:
correct:结果是否正确;TileLang(us):TileLang kernel 的平均耗时;PyTorch(us):PyTorch 加法的平均耗时;max_abs_error:与参考结果的最大绝对误差。
本任务不要求 TileLang 比 PyTorch 更快。你需要确认每一行都是 PASS,然后保存终端截图。
任务描述还要求比较至少两种 BLOCK_N。当前 benchmark_add.py 的编译语句固定使用 BLOCK_N=1024:
k = tl_add_1d.compile(N=n, BLOCK_N=1024)要比较第二种取值,可以将它改为例如 256 后再次运行并记录结果,或者扩展脚本,让它依次测试 256 和 1024。不同设备上的快慢结果可能不同,这是正常现象。
12. 完成检查清单
提交前逐项确认:
-
solution.py不再包含会被执行的NotImplementedError; - 使用了
T.ceildiv; - 使用了
T.Kernel; - 使用了
T.Parallel; - 使用
index < N保护所有输入输出访问; -
return C位于函数末尾; -
python -m pytest -q test_add.py全部通过; - benchmark 每个长度均显示
PASS; - 已比较至少两种
BLOCK_N; - 已保存要求的终端截图。
到这里,Task2 的 TileLang 部分就完成了。
接下来可阅读《从 a + b 到 GPU Kernel:TileLang 原理详解》,进一步理解 tile 与 GPU block 的区别、T.Parallel 的并行语义、JIT 编译过程以及 BLOCK_N 对性能的影响。