"""Sweep tile configs for the GEMM / GEGLU / attention kernels at the shapes laya actually hits.""" import torch, sys, time, itertools, json sys.path.insert(0, "kernels"); import tl_kernels as K dev = "cuda" def bench(f, n=20): for _ in range(3): f() torch.cuda.synchronize(); t = time.perf_counter() for _ in range(n): f() torch.cuda.synchronize(); return (time.perf_counter() - t) / n * 1e3 res = {"gemm": {}, "geglu": {}, "attn": {}} Ms = [128, 1024, 8192, 28672] cfgs = [(bm, bn, th) for bm in (64, 128) for bn in (64, 128, 256) for th in (128, 256)] for (N, Kd) in [(2304, 768), (768, 768), (768, 1152)]: W = torch.randn(N, Kd, device=dev, dtype=torch.bfloat16) * 0.02; b = torch.zeros(N, device=dev) for M in Ms: A = torch.randn(M, Kd, device=dev, dtype=torch.bfloat16); C = torch.empty(M, N, device=dev, dtype=torch.bfloat16) tref = bench(lambda: torch.nn.functional.linear(A, W)) best = None for bm, bn, th in cfgs: try: k = K.gemm_kernel(N, Kd, bm=bm, bn=bn, bk=64, stages=3, threads=th); ms = bench(lambda: k(A, W, b, C)) if best is None or ms < best[0]: best = (ms, (bm, bn, th)) except Exception as e: print("fail", N, Kd, M, (bm, bn, th), str(e)[:60], flush=True) tf = 2 * M * N * Kd / best[0] / 1e9 print(f"gemm N={N} K={Kd} M={M}: best {best[1]} {best[0]:.3f}ms ({tf:.1f} TFLOPS) cublas {tref:.3f}ms (default 64x128x128thr)", flush=True) res["gemm"][f"{N},{Kd},{M}"] = best for M in Ms: A = torch.randn(M, 768, device=dev, dtype=torch.bfloat16); Wi = torch.randn(2304, 768, device=dev, dtype=torch.bfloat16) * 0.02; C = torch.empty(M, 1152, device=dev, dtype=torch.bfloat16) best = None for bm, bn, th in cfgs: if bn > 128: continue try: k = K.gemm_geglu_kernel(1152, 768, bm=bm, bn=bn, bk=64, stages=3, threads=th); ms = bench(lambda: k(A, Wi, C)) if best is None or ms < best[0]: best = (ms, (bm, bn, th)) except Exception as e: print("fail geglu", M, (bm, bn, th), str(e)[:60], flush=True) print(f"geglu M={M}: best {best[1]} {best[0]:.3f}ms ({2*M*2304*768/best[0]/1e9:.1f} TFLOPS)", flush=True) res["geglu"][str(M)] = best H, Dh = 12, 64 for (B, L) in [(32, 1024), (32, 128), (4, 128)]: qkv = torch.randn(B, L, 3, H, Dh, device=dev, dtype=torch.bfloat16); lens = torch.full((B,), L - 3, device=dev, dtype=torch.int32); O = torch.empty(B, L, H * Dh, device=dev, dtype=torch.bfloat16) for window in (0, 65): best = None for bm, bn, st, th in itertools.product((64, 128), (64, 128), (1, 2), (128, 256)): try: k = K.attn_kernel(B, L, H, Dh, window=window, bm=bm, bn=bn, stages=st, threads=th); ms = bench(lambda: k(qkv, lens, O)) if best is None or ms < best[0]: best = (ms, (bm, bn, st, th)) except Exception as e: print("fail attn", (B, L, window), (bm, bn, st, th), str(e)[:60], flush=True) fl = 4 * B * H * L * L * Dh if window == 0 else 4 * B * H * L * (2 * 65 + 1) * Dh print(f"attn B={B} L={L} window={window}: best {best[1]} {best[0]:.3f}ms ({fl/best[0]/1e9:.1f} TFLOPS)", flush=True) res["attn"][f"{B},{L},{window}"] = best json.dump(res, open("kernels/tune_results.json", "w"), indent=1) print("DONE")