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.")



Source link

Leave a Reply

Your email address will not be published. Required fields are marked *