From 528d35e30668577e68f825b4e71b11497be9ed17 Mon Sep 17 00:00:00 2001 From: wozeparrot Date: Fri, 1 May 2026 08:14:41 +0800 Subject: [PATCH] llama speed 4 (#15993) --- examples/mlperf/model_train.py | 4 ++++ examples/mlperf/models/flat_llama.py | 6 +++--- extra/gemm/cdna_asm_gemm.py | 27 +++++++++++++------------- extra/thunder/amd/fa.py | 4 +--- extra/thunder/amd/gemm_fp8.cpp | 29 +++++++++++++++++++++++++++- 5 files changed, 50 insertions(+), 20 deletions(-) diff --git a/examples/mlperf/model_train.py b/examples/mlperf/model_train.py index be9d129c9d..e8e1febb7b 100644 --- a/examples/mlperf/model_train.py +++ b/examples/mlperf/model_train.py @@ -1446,6 +1446,10 @@ def train_llama3(): idx = next(j for j, p in enumerate(optim.params) if p is w) optim.master_params[idx].assign((optim.master_params[idx] * w._inv_scale.reshape(-1, *([1]*(w.ndim-1)))).contiguous()) + # realize everything here + if optim.master_params: Tensor.realize(*optim.master_params) + Tensor.realize(*optim.params, *fp8_inv_scales, *fp8_amax, *fp8_grad_amax) + @TinyJit def minibatch(tokens:Tensor): if is_dp: tokens = tokens.to(None).shard(device, 0) diff --git a/examples/mlperf/models/flat_llama.py b/examples/mlperf/models/flat_llama.py index 926a76c70f..de0643c1b9 100644 --- a/examples/mlperf/models/flat_llama.py +++ b/examples/mlperf/models/flat_llama.py @@ -158,14 +158,14 @@ class FlatTransformer: xq, xk = apply_rotary_emb(xq, xk, freqs_cis) xq, xk, xv = xq.cast(dtypes.bfloat16), xk.cast(dtypes.bfloat16), xv.cast(dtypes.bfloat16) - xq, xk, xv = xq.transpose(1, 2), xk.transpose(1, 2), xv.transpose(1, 2) if getenv("HK_FLASH_ATTENTION"): from extra.thunder.amd.fa import flash_attention attn, *save = flash_attention(xq, xk, xv, is_causal=True) saves.extend(save) else: - attn = xq.scaled_dot_product_attention(xk, xv, is_causal=True, enable_gqa=True) - attn = attn.transpose(1, 2).reshape(bsz, seqlen, -1) + xq, xk, xv = xq.transpose(1, 2), xk.transpose(1, 2), xv.transpose(1, 2) + attn = xq.scaled_dot_product_attention(xk, xv, is_causal=True, enable_gqa=True).transpose(1, 2) + attn = attn.reshape(bsz, seqlen, -1) out, *ret = matmul(attn, wo, amax_x=amax_xo, w_inv_scale=s_o, grad_amax_state=grad_amax_xo) new_amaxs.extend(ret[:1]) diff --git a/extra/gemm/cdna_asm_gemm.py b/extra/gemm/cdna_asm_gemm.py index d0c65c491d..5dd3a54ad7 100644 --- a/extra/gemm/cdna_asm_gemm.py +++ b/extra/gemm/cdna_asm_gemm.py @@ -2628,21 +2628,24 @@ def custom_asm_gemm(C:UOp, A:UOp, B:UOp, dname:str) -> UOp: # ** FP8 GEMM custom kernel @functools.cache -def custom_hk_fp8_gemm(C:UOp, A:UOp, B:UOp, X_s:UOp, W_s:UOp, *extra:UOp, dname:str) -> UOp: - # A is (batch, M, K), B is (N, K) transposed, X_s is x_scale, W_s is w_scale — kernel multiplies by both. - # extra is unused fwd inputs (e.g. grad_amax_state) plumbed through so the bwd can read them via kernel.src. +def custom_hk_fp8_gemm(C:UOp, A:UOp, B:UOp, *args:UOp, dname:str, scale_mode:int=3) -> UOp: + # scale_mode: 0=no scale, 1=x only, 2=w only, 3=both + n_scales = (1 if scale_mode & 1 else 0) + (1 if scale_mode & 2 else 0) + scales, extra = args[:n_scales], args[n_scales:] M, K = A.shape[0]*A.shape[1], A.shape[2] N, K2 = B.shape[(1 if B.ndim == 3 else 0):] assert K == K2, f"{A.shape} {B.shape}" block_size = 256 threads = UOp.special(64 * 8, "lidx0") workgroups = UOp.special((M // block_size) * (N // block_size), "gidx0") - sink = UOp.sink(C.base, A.base, B.base, X_s.base, W_s.base, threads, workgroups, + sink_inputs = (C.base, A.base, B.base) + tuple(s.base for s in scales) + (threads, workgroups) + sink = UOp.sink(*sink_inputs, arg=KernelInfo(f"hk_fp8_gemm_{M}_{N}_{K}", estimates=Estimates(ops=2*M*N*K, mem=(M*K+N*K)*A.dtype.itemsize+M*N*C.dtype.itemsize))) kittens_path = pathlib.Path(__file__).parent.parent/"thunder"/"amd" src = (kittens_path/"gemm_fp8.cpp").read_text() lib = HIPCCCompiler("gfx950", [f"-I{(kittens_path/'include').as_posix()}", "-std=c++20", "-DKITTENS_CDNA4", "-ffast-math", - "-DHIP_ENABLE_WARP_SYNC_BUILTINS", f"-DGEMM_M={M}", f"-DGEMM_N={N}", f"-DGEMM_K={K}"]).compile_cached(src) + "-DHIP_ENABLE_WARP_SYNC_BUILTINS", f"-DGEMM_M={M}", f"-DGEMM_N={N}", f"-DGEMM_K={K}", + f"-DSCALE_MODE={scale_mode}"]).compile_cached(src) return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.DEVICE, arg=dname), UOp(Ops.LINEAR, src=(*sink.src, sink)), UOp(Ops.SOURCE, arg=src), UOp(Ops.BINARY, arg=lib))) @@ -2699,8 +2702,7 @@ def custom_uop_gemm(C:UOp, A:UOp, B:UOp) -> UOp: def custom_gemm_bw(gradient:UOp, kernel:UOp): inputs = kernel.src[1:] - # fp8 scaled gemm has 5 inputs (out, a, b, x_scale, w_scale) optionally plus grad_amax_state (6 total); plain gemm has 3 - if len(inputs) >= 5: + if inputs[1].dtype == FP8_DTYPE: grad_amax_state = inputs[5] if len(inputs) == 6 else None out, a, b, s_x, s_w = inputs[:5] a_t, b_t, g_t = Tensor(a, device=a.device), Tensor(b, device=a.device), Tensor(gradient, device=a.device) @@ -2720,8 +2722,7 @@ def custom_gemm_bw(gradient:UOp, kernel:UOp): # dgrad: uses g_scale * x_scale * w_scale grad_a = asm_gemm(g_fp8, b_t, x_scale=g_scale * s_x_t, w_scale=s_w_t) # wgrad: no w_scale - _one = Tensor(1.0, dtype=dtypes.float, device=a.device) - grad_b = asm_gemm(g_fp8.permute(2, 0, 1).reshape(g_t.shape[-1], -1), a_t.reshape(-1, a_t.shape[-1]), x_scale=g_scale * s_x_t, w_scale=_one) + grad_b = asm_gemm(g_fp8.permute(2, 0, 1).reshape(g_t.shape[-1], -1), a_t.reshape(-1, a_t.shape[-1]), x_scale=g_scale * s_x_t) # Attach the delayed-amax store effect (if any) to grad_a so realizing grads commits the amax update. ret = (None, grad_a.uop.after(store_effect), grad_b.uop, None, None) if len(inputs) == 6: ret = ret + (None,) @@ -2774,11 +2775,11 @@ def asm_gemm(a:Tensor, b:Tensor, x_scale:Tensor|None=None, w_scale:Tensor|None=N if arch.startswith("gfx950") and getenv("USE_ASM", 1): # fp8 gemm computes a@b.T, kernel multiplies output by x_scale * w_scale before bf16 store if a.dtype == FP8_DTYPE: - _one = lambda: Tensor(1.0, dtype=dtypes.float, device=a.device) - xs = x_scale if x_scale is not None else _one() - ws = w_scale if w_scale is not None else _one() + scales = tuple(s for s in (x_scale, w_scale) if s is not None) + scale_mode = (1 if x_scale is not None else 0) | (2 if w_scale is not None else 0) extra = [grad_amax_state] if grad_amax_state is not None else [] - out = Tensor.custom_kernel(out, a, b.T, xs, ws, *extra, fxn=functools.partial(custom_hk_fp8_gemm, dname=dname), grad_fxn=custom_gemm_bw)[0] + fxn = functools.partial(custom_hk_fp8_gemm, dname=dname, scale_mode=scale_mode) + out = Tensor.custom_kernel(out, a, b.T, *scales, *extra, fxn=fxn, grad_fxn=custom_gemm_bw)[0] else: out = Tensor.custom_kernel(out, a, b, fxn=functools.partial(custom_asm_gemm, dname=dname), grad_fxn=custom_gemm_bw)[0] else: diff --git a/extra/thunder/amd/fa.py b/extra/thunder/amd/fa.py index e22cb5f55a..87d5f058ac 100644 --- a/extra/thunder/amd/fa.py +++ b/extra/thunder/amd/fa.py @@ -55,8 +55,6 @@ def flash_attention(xq, xk, xv, attn_mask:Tensor|None=None, is_causal:bool=False assert attn_mask is None, "attn_mask not supported" assert is_causal, "only causal attention supported" - xq, xk, xv = xq.transpose(1, 2), xk.transpose(1, 2), xv.transpose(1, 2) - B, N, H, D = xq.shape H_KV = xk.shape[2] assert D == 128, "only D=128 supported" @@ -81,7 +79,7 @@ def flash_attention(xq, xk, xv, attn_mask:Tensor|None=None, is_causal:bool=False attn, l_vec = Tensor.custom_kernel(attn, l_vec, xq, xk, xv, fxn=functools.partial(custom_fa_forward, device=single_device, arch=arch, B=B_local, N=N, H=H_local, H_KV=H_KV_local, D=D), grad_fxn=grad)[:2] - return attn.transpose(1, 2), attn, l_vec + return attn, attn, l_vec @functools.cache def custom_fa_forward(o:UOp, l_vec:UOp, q:UOp, k:UOp, v:UOp, device:str, arch:str, B:int, N:int, H:int, H_KV:int, D:int): diff --git a/extra/thunder/amd/gemm_fp8.cpp b/extra/thunder/amd/gemm_fp8.cpp index 6bdff15252..ee3136125b 100644 --- a/extra/thunder/amd/gemm_fp8.cpp +++ b/extra/thunder/amd/gemm_fp8.cpp @@ -93,7 +93,20 @@ constexpr int NUM_WARPS = 8; using G = kittens::group; -__global__ __launch_bounds__(512, 2) void hk_fp8_gemm(bf16 *C_ptr, fp8e4m3 *A_ptr, fp8e4m3 *B_ptr, float *x_scale_ptr, float *w_scale_ptr) { +// scale_mode: 0=no scale, 1=x only, 2=w only, 3=both +#ifndef SCALE_MODE +#define SCALE_MODE 3 +#endif + +__global__ __launch_bounds__(512, 2) void hk_fp8_gemm(bf16 *C_ptr, fp8e4m3 *A_ptr, fp8e4m3 *B_ptr +#if SCALE_MODE == 1 + , float *x_scale_ptr +#elif SCALE_MODE == 2 + , float *w_scale_ptr +#elif SCALE_MODE == 3 + , float *x_scale_ptr, float *w_scale_ptr +#endif +) { constexpr int M = GEMM_M, N = GEMM_N, K = GEMM_K; kittens::gl A{A_ptr, nullptr, nullptr, nullptr, nullptr}; @@ -333,11 +346,25 @@ __global__ __launch_bounds__(512, 2) void hk_fp8_gemm(bf16 *C_ptr, fp8e4m3 *A_pt } // apply x_scale * w_scale before bf16 store to prevent overflow +#if SCALE_MODE == 1 + float scale = *x_scale_ptr; + mul(cA, cA, scale); + mul(cB, cB, scale); + mul(cC, cC, scale); + mul(cD, cD, scale); +#elif SCALE_MODE == 2 + float scale = *w_scale_ptr; + mul(cA, cA, scale); + mul(cB, cB, scale); + mul(cC, cC, scale); + mul(cD, cD, scale); +#elif SCALE_MODE == 3 float scale = *x_scale_ptr * *w_scale_ptr; mul(cA, cA, scale); mul(cB, cB, scale); mul(cC, cC, scale); mul(cD, cD, scale); +#endif store(C, cA, {0, 0, block_row * WARPS_ROW * 2 + warp_m, block_col * WARPS_COL * 2 + warp_n}); store(C, cB, {0, 0, block_row * WARPS_ROW * 2 + warp_m, block_col * WARPS_COL * 2 + WARPS_COL + warp_n});