strix-halo-kernels / test_bench.py
axjns's picture
Fused RMSNorm + GEGLU/SwiGLU Triton kernels for gfx1151
e20beb1 verified
Raw
History Blame Contribute Delete
3.89 kB
"""Correctness + throughput checks for the gfx1151 Triton kernels."""
import sys, torch, triton
import rmsnorm, geglu
DEV = "cuda"
FAILS = []
def check(name, got, want, atol, rtol):
ok = torch.allclose(got.float(), want.float(), atol=atol, rtol=rtol)
md = (got.float() - want.float()).abs().max().item()
print(f" {'PASS' if ok else 'FAIL'} {name:44} max|diff|={md:.3e}")
if not ok:
FAILS.append(name)
def bench(fn, *a):
return triton.testing.do_bench(lambda: fn(*a), warmup=50, rep=200)
print(f"gfx: {torch.cuda.get_device_properties(0).gcnArchName} triton {triton.__version__}\n")
print("== RMSNorm correctness ==")
for dtype, atol in ((torch.float32, 1e-5), (torch.float16, 2e-2), (torch.bfloat16, 6e-2)):
for rows, cols in ((4, 512), (32, 2048), (7, 4096), (128, 5120), (3, 8192)):
x = torch.randn(rows, cols, dtype=dtype, device=DEV)
w = torch.randn(cols, dtype=dtype, device=DEV)
check(f"{str(dtype).split('.')[-1]:9} {rows}x{cols}",
rmsnorm.rms_norm(x, w), rmsnorm.rms_norm_ref(x, w), atol, 1e-2)
print("\n== RMSNorm 3-D / non-contiguous ==")
x = torch.randn(2, 17, 2048, dtype=torch.float16, device=DEV)
w = torch.randn(2048, dtype=torch.float16, device=DEV)
check("3-D (2,17,2048)", rmsnorm.rms_norm(x, w), rmsnorm.rms_norm_ref(x, w), 2e-2, 1e-2)
xt = torch.randn(2048, 64, dtype=torch.float16, device=DEV).t() # (64, 2048), non-contiguous
wt = torch.randn(2048, dtype=torch.float16, device=DEV)
assert not xt.is_contiguous()
check("non-contiguous (64,2048)", rmsnorm.rms_norm(xt, wt), rmsnorm.rms_norm_ref(xt, wt), 2e-2, 1e-2)
print("\n== Gated activations correctness ==")
for dtype, atol in ((torch.float32, 1e-5), (torch.float16, 3e-2), (torch.bfloat16, 8e-2)):
for shape in ((8, 4096), (2, 512, 3072), (1, 1024, 8192)):
g = torch.randn(shape, dtype=dtype, device=DEV)
u = torch.randn(shape, dtype=dtype, device=DEV)
t = str(dtype).split('.')[-1]
check(f"swiglu {t:9} {shape}", geglu.swiglu(g, u), geglu.swiglu_ref(g, u), atol, 1e-2)
check(f"geglu {t:9} {shape}", geglu.geglu(g, u), geglu.geglu_ref(g, u), atol, 1e-2)
print("\n== chunked form ==")
x = torch.randn(4, 512, 6144, dtype=torch.float16, device=DEV)
g, u = x.chunk(2, dim=-1)
check("geglu_chunked", geglu.geglu_chunked(x), geglu.geglu_ref(g, u), 3e-2, 1e-2)
print("\n== edge cases ==")
try:
rmsnorm.rms_norm(torch.randn(4, 128, device=DEV), torch.randn(64, device=DEV))
print(" FAIL mismatched weight not rejected"); FAILS.append("weight guard")
except ValueError:
print(" PASS mismatched weight rejected")
try:
geglu.geglu_chunked(torch.randn(4, 127, device=DEV))
print(" FAIL odd last dim not rejected"); FAILS.append("chunk guard")
except ValueError:
print(" PASS odd last dim rejected")
e = geglu.swiglu(torch.zeros(0, 64, device=DEV), torch.zeros(0, 64, device=DEV))
print(f" PASS empty input -> {tuple(e.shape)}")
print("\n== throughput vs eager (fp16) ==")
print(f" {'shape':>22} {'eager':>9} {'triton':>9} {'speedup':>8}")
for rows, cols in ((512, 4096), (2048, 4096), (4096, 5120), (8192, 2048)):
x = torch.randn(rows, cols, dtype=torch.float16, device=DEV)
w = torch.randn(cols, dtype=torch.float16, device=DEV)
a = bench(rmsnorm.rms_norm_ref, x, w)
b = bench(rmsnorm.rms_norm, x, w)
print(f" rmsnorm {rows:>5}x{cols:<5} {a:8.3f}ms {b:8.3f}ms {a/b:7.2f}x")
for rows, cols in ((512, 8192), (2048, 8192), (4096, 4096)):
g = torch.randn(rows, cols, dtype=torch.float16, device=DEV)
u = torch.randn(rows, cols, dtype=torch.float16, device=DEV)
a = bench(geglu.geglu_ref, g, u)
b = bench(geglu.geglu, g, u)
print(f" geglu {rows:>5}x{cols:<5} {a:8.3f}ms {b:8.3f}ms {a/b:7.2f}x")
print()
if FAILS:
print(f"{len(FAILS)} FAILURE(S): {FAILS}")
sys.exit(1)
print("all checks passed")