import torch, math, time, sys sys.path.insert(0, "kernels") import tl_kernels as K torch.manual_seed(0) dev = "cuda" def err(a, b): return (a.float() - b.float()).abs().max().item() def bench(f, n=50): 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 M, Kd = 8 * 128, 768 A = torch.randn(M, Kd, device=dev, dtype=torch.bfloat16) # gemm variants for N, bias, act in [(2304, False, "none"), (768, True, "gelu"), (3072, True, "relu")]: W = torch.randn(N, Kd, device=dev, dtype=torch.bfloat16) * 0.02 b = torch.randn(N, device=dev) C = torch.empty(M, N, device=dev, dtype=torch.bfloat16) k = K.gemm_kernel(N, Kd, bias=bias, act=act) k(A, W, b, C) ref = A.float() @ W.float().T + (b if bias else 0) ref = {"none": ref, "gelu": torch.nn.functional.gelu(ref), "relu": torch.relu(ref)}[act] print(f"gemm N={N} bias={bias} act={act}: maxerr={err(C, ref):.4f} tl={bench(lambda: k(A, W, b, C)):.3f}ms torch={bench(lambda: torch.nn.functional.linear(A, W, b.bfloat16() if bias else None)):.3f}ms") # geglu F = 1152; Wi = torch.randn(2 * F, Kd, device=dev, dtype=torch.bfloat16) * 0.02 C = torch.empty(M, F, device=dev, dtype=torch.bfloat16); k = K.gemm_geglu_kernel(F, Kd); k(A, Wi, C) x = (A.float() @ Wi.float().T); ref = torch.nn.functional.gelu(x[:, :F]) * x[:, F:] def tref(): i, g = torch.nn.functional.linear(A, Wi).chunk(2, -1); return torch.nn.functional.gelu(i) * g print(f"geglu: maxerr={err(C, ref):.4f} (ref scale {ref.abs().max():.2f}) tl={bench(lambda: k(A, Wi, C)):.3f}ms torch={bench(tref):.3f}ms") # add_ln X = torch.randn(M, Kd, device=dev, dtype=torch.bfloat16); R = torch.randn(M, Kd, device=dev, dtype=torch.bfloat16) w = torch.rand(Kd, device=dev) + 0.5; bb = torch.randn(Kd, device=dev) for residual, bias in [(True, False), (False, False), (True, True)]: X2 = X.clone(); Y = torch.empty_like(X) k = K.add_ln_kernel(Kd, residual=residual, bias=bias); k(X2, R, w, bb, Y) xr = (X.float() + R.float()) if residual else X.float() xr_b = xr.bfloat16().float() if residual else xr # kernel stores the residual stream in bf16 ref = torch.nn.functional.layer_norm(xr_b, (Kd,), w, bb if bias else None, 1e-5) print(f"add_ln residual={residual} bias={bias}: maxerr Y={err(Y, ref):.4f} X={err(X2, xr):.4f} tl={bench(lambda: k(X2, R, w, bb, Y)):.3f}ms torch={bench(lambda: torch.nn.functional.layer_norm(X2 + R, (Kd,), w.bfloat16(), None, 1e-5)):.3f}ms") # rope H, Dh, L = 12, 64, 128; B = M // L qkv = torch.randn(M, 3 * H * Dh, device=dev, dtype=torch.bfloat16) inv = 1.0 / (10000 ** (torch.arange(0, Dh, 2, device=dev).float() / Dh)); pos = torch.arange(L, device=dev).float() fr = torch.outer(pos, inv); cos, sin = fr.cos().contiguous(), fr.sin().contiguous() q2 = qkv.clone(); k = K.rope_kernel(H, Dh, L); k(q2, cos, sin) def rot(x): # x [B,L,H,Dh] c = torch.cat([cos, cos], -1)[None, :, None]; s = torch.cat([sin, sin], -1)[None, :, None] x1, x2 = x[..., :Dh // 2], x[..., Dh // 2:] return x * c + torch.cat([-x2, x1], -1) * s v = qkv.float().view(B, L, 3, H, Dh); ref = v.clone(); ref[:, :, 0] = rot(v[:, :, 0]); ref[:, :, 1] = rot(v[:, :, 1]) print(f"rope: maxerr={err(q2.view(B, L, 3, H, Dh), ref):.4f} tl={bench(lambda: k(q2, cos, sin)):.3f}ms") # attention for L, window in [(128, 0), (128, 64), (1024, 0), (1024, 64)]: B = 4; qkv = torch.randn(B, L, 3, H, Dh, device=dev, dtype=torch.bfloat16) lens = torch.tensor([L, L - 5, L // 2 + 3, 7], device=dev, dtype=torch.int32) O = torch.empty(B, L, H * Dh, device=dev, dtype=torch.bfloat16) k = K.attn_kernel(B, L, H, Dh, window=window); k(qkv, lens, O) q, kk, vv = [qkv[:, :, i].transpose(1, 2).float() for i in range(3)] idx = torch.arange(L, device=dev) mask = (idx[None, :] < lens[:, None])[:, None, None, :].expand(B, 1, L, L) if window: mask = mask & ((idx[:, None] - idx[None, :]).abs() <= window)[None, None] ref = torch.nn.functional.scaled_dot_product_attention(q, kk, vv, attn_mask=mask).transpose(1, 2).reshape(B, L, -1) valid = (idx[None, :] < lens[:, None]) e = (O.float() - ref)[valid].abs().max().item() def sref(): return torch.nn.functional.scaled_dot_product_attention(qkv[:, :, 0].transpose(1, 2), qkv[:, :, 1].transpose(1, 2), qkv[:, :, 2].transpose(1, 2), attn_mask=mask) print(f"attn L={L} window={window}: maxerr(valid rows)={e:.4f} finite={torch.isfinite(O).all().item()} tl={bench(lambda: k(qkv, lens, O)):.3f}ms torch-sdpa={bench(sref):.3f}ms")