Designing High-Performance GPU Kernels with TileLang: Tensor-Core GEMM, Fused Softmax, FlashAttention, and Autotuning
@tilelang.jit(out_idx=[-1])
def make_matmul(M: int, N: int, K: int,
block_M: int = 128, block_N: int = 128, block_K: int = 32,
num_stages: int = 3, threads: int = 128,
use_swizzle: bool = False,
dtype: str = “float16”, accum_dtype: str = “float”):
@T.prim_func
def main(A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), dtype)):
with T.Kernel(T.ceildiv(N, block_N),
T.ceildiv(M, block_M),
threads=threads) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), dtype)
B_shared = T.alloc_shared((block_K, block_N), dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
if use_swizzle:
T.use_swizzle(panel_size=10, enable=True)
T.clear(C_local)
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=num_stages):
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 main
def smem_bytes(block_M, block_N, block_K, stages, itemsize=2):
return (block_M * block_K + block_K * block_N) * itemsize * stages
def section_2():
banner(“2. MEMORY HIERARCHY — tiled tensor-core GEMM”)
M = N = K = 2048
bm, bn, bk, st = 128, 128, 32, DEFAULT_STAGES
while smem_bytes(bm, bn, bk, st) > SMEM_CAP and st > 1:
st -= 1
print(f” config: {bm}x{bn}x{bk}, num_stages={st}, ”
f”smem={smem_bytes(bm,bn,bk,st)/1024:.0f} KB”)
kernel = make_matmul(M, N, K, bm, bn, bk, num_stages=st, threads=128)
a = torch.randn(M, K, device=DEV, dtype=torch.float16)
b = torch.randn(K, N, device=DEV, dtype=torch.float16)
c = kernel(a, b)
check(c, a @ b, “matmul 2048^3″)
flops = 2 * M * N * K
ms = bench(lambda: kernel(a, b))
ms_ref = bench(lambda: a @ b)
print(f” tilelang : {ms:7.3f} ms -> {flops/(ms*1e-3)/1e12:6.2f} TFLOP/s”)
print(f” cuBLAS : {ms_ref:7.3f} ms -> {flops/(ms_ref*1e-3)/1e12:6.2f} TFLOP/s”)
print(f” ratio : {ms_ref/ms*100:5.1f}% of cuBLAS from ~20 lines of Python”)
src = kernel.get_kernel_source()
for needle in (“mma.sync”, “wgmma”, “ldmatrix”, “cp.async”, “tl::gemm”):
if needle in src:
print(f” emitted: {needle}”)
return kernel
def section_3():
banner(“3. KNOBS — sweeping the schedule by hand”)
M = N = K = 2048
a = torch.randn(M, K, device=DEV, dtype=torch.float16)
b = torch.randn(K, N, device=DEV, dtype=torch.float16)
flops = 2 * M * N * K
candidates = [
(64, 64, 32, 2, 128, False),
(128, 128, 32, 2, 128, False),
(128, 128, 32, 2, 128, True),
(128, 128, 32, 3, 128, False),
(128, 128, 64, 2, 256, False),
(128, 256, 32, 2, 256, False),
]
print(f” {’tile’:>16} {‘stg’:>4} {‘thr’:>4} {‘swz’:>4} {‘smem’:>7} ”
f”{‘ms’:>8} {‘TFLOP/s’:>9}”)
results = []
for bm, bn, bk, st, thr, swz in candidates:
need = smem_bytes(bm, bn, bk, st)
tag = f”{bm}x{bn}x{bk}”
if need > SMEM_CAP:
print(f” {tag:>16} {st:>4} {thr:>4} {str(swz):>4} {need//1024:>5}KB ”
f” SKIPPED (over smem budget)”)
continue
try:
k = make_matmul(M, N, K, bm, bn, bk, num_stages=st,
threads=thr, use_swizzle=swz)
c = k(a, b)
ok = (c – (a @ b)).float().norm() / (a @ b).float().norm() < 2e-2
ms = bench(lambda: k(a, b), warmup=5, rep=20)
results.append((ms, tag, st, thr, swz))
print(f” {tag:>16} {st:>4} {thr:>4} {str(swz):>4} {need//1024:>5}KB ”
f”{ms:>8.3f} {flops/(ms*1e-3)/1e12:>9.2f}”
f”{” if ok else ‘ <- NUMERICALLY WRONG’}”)
except Exception as e:
print(f” {tag:>16} {st:>4} {thr:>4} {str(swz):>4} – ”
f”failed: {type(e).__name__}”)
if results:
best = min(results)
print(f”\n winner: {best[1]}, stages={best[2]}, threads={best[3]}, ”
f”swizzle={best[4]} ({best[0]:.3f} ms)”)
print(” Takeaway: the best schedule is arch- and shape-dependent, which is”)
print(” exactly why section 7 exists.”)
def make_matmul(M: int, N: int, K: int,
block_M: int = 128, block_N: int = 128, block_K: int = 32,
num_stages: int = 3, threads: int = 128,
use_swizzle: bool = False,
dtype: str = “float16”, accum_dtype: str = “float”):
@T.prim_func
def main(A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), dtype)):
with T.Kernel(T.ceildiv(N, block_N),
T.ceildiv(M, block_M),
threads=threads) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), dtype)
B_shared = T.alloc_shared((block_K, block_N), dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
if use_swizzle:
T.use_swizzle(panel_size=10, enable=True)
T.clear(C_local)
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=num_stages):
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 main
def smem_bytes(block_M, block_N, block_K, stages, itemsize=2):
return (block_M * block_K + block_K * block_N) * itemsize * stages
def section_2():
banner(“2. MEMORY HIERARCHY — tiled tensor-core GEMM”)
M = N = K = 2048
bm, bn, bk, st = 128, 128, 32, DEFAULT_STAGES
while smem_bytes(bm, bn, bk, st) > SMEM_CAP and st > 1:
st -= 1
print(f” config: {bm}x{bn}x{bk}, num_stages={st}, ”
f”smem={smem_bytes(bm,bn,bk,st)/1024:.0f} KB”)
kernel = make_matmul(M, N, K, bm, bn, bk, num_stages=st, threads=128)
a = torch.randn(M, K, device=DEV, dtype=torch.float16)
b = torch.randn(K, N, device=DEV, dtype=torch.float16)
c = kernel(a, b)
check(c, a @ b, “matmul 2048^3″)
flops = 2 * M * N * K
ms = bench(lambda: kernel(a, b))
ms_ref = bench(lambda: a @ b)
print(f” tilelang : {ms:7.3f} ms -> {flops/(ms*1e-3)/1e12:6.2f} TFLOP/s”)
print(f” cuBLAS : {ms_ref:7.3f} ms -> {flops/(ms_ref*1e-3)/1e12:6.2f} TFLOP/s”)
print(f” ratio : {ms_ref/ms*100:5.1f}% of cuBLAS from ~20 lines of Python”)
src = kernel.get_kernel_source()
for needle in (“mma.sync”, “wgmma”, “ldmatrix”, “cp.async”, “tl::gemm”):
if needle in src:
print(f” emitted: {needle}”)
return kernel
def section_3():
banner(“3. KNOBS — sweeping the schedule by hand”)
M = N = K = 2048
a = torch.randn(M, K, device=DEV, dtype=torch.float16)
b = torch.randn(K, N, device=DEV, dtype=torch.float16)
flops = 2 * M * N * K
candidates = [
(64, 64, 32, 2, 128, False),
(128, 128, 32, 2, 128, False),
(128, 128, 32, 2, 128, True),
(128, 128, 32, 3, 128, False),
(128, 128, 64, 2, 256, False),
(128, 256, 32, 2, 256, False),
]
print(f” {’tile’:>16} {‘stg’:>4} {‘thr’:>4} {‘swz’:>4} {‘smem’:>7} ”
f”{‘ms’:>8} {‘TFLOP/s’:>9}”)
results = []
for bm, bn, bk, st, thr, swz in candidates:
need = smem_bytes(bm, bn, bk, st)
tag = f”{bm}x{bn}x{bk}”
if need > SMEM_CAP:
print(f” {tag:>16} {st:>4} {thr:>4} {str(swz):>4} {need//1024:>5}KB ”
f” SKIPPED (over smem budget)”)
continue
try:
k = make_matmul(M, N, K, bm, bn, bk, num_stages=st,
threads=thr, use_swizzle=swz)
c = k(a, b)
ok = (c – (a @ b)).float().norm() / (a @ b).float().norm() < 2e-2
ms = bench(lambda: k(a, b), warmup=5, rep=20)
results.append((ms, tag, st, thr, swz))
print(f” {tag:>16} {st:>4} {thr:>4} {str(swz):>4} {need//1024:>5}KB ”
f”{ms:>8.3f} {flops/(ms*1e-3)/1e12:>9.2f}”
f”{” if ok else ‘ <- NUMERICALLY WRONG’}”)
except Exception as e:
print(f” {tag:>16} {st:>4} {thr:>4} {str(swz):>4} – ”
f”failed: {type(e).__name__}”)
if results:
best = min(results)
print(f”\n winner: {best[1]}, stages={best[2]}, threads={best[3]}, ”
f”swizzle={best[4]} ({best[0]:.3f} ms)”)
print(” Takeaway: the best schedule is arch- and shape-dependent, which is”)
print(” exactly why section 7 exists.”)
