laya-browser / code /kernels /test_kernels.py
cklxx's picture
v19s: WebChain real-site trajectories, format v5, webgym x7 + DAgger, harness fixes; replaces v17s
454b3e6 verified
Raw History Blame Contribute Delete
4.59 kB
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")