From c9b60caf8ccb599ccace4b4dc6b164ffa4395ed2 Mon Sep 17 00:00:00 2001 From: wozeparrot Date: Fri, 24 Jul 2026 07:25:15 -0700 Subject: [PATCH] gptoss: moe gemm kernels (#17178) --- extra/gemm/moe_gemm.py | 111 +++++++++++++++ extra/thunder/amd/grouped_mxfp8_gemm.cpp | 158 ++++++++++++++++++++++ extra/thunder/amd/grouped_mxfp8_wgrad.cpp | 141 +++++++++++++++++++ 3 files changed, 410 insertions(+) create mode 100644 extra/gemm/moe_gemm.py create mode 100644 extra/thunder/amd/grouped_mxfp8_gemm.cpp create mode 100644 extra/thunder/amd/grouped_mxfp8_wgrad.cpp diff --git a/extra/gemm/moe_gemm.py b/extra/gemm/moe_gemm.py new file mode 100644 index 0000000000..7efc41f4c6 --- /dev/null +++ b/extra/gemm/moe_gemm.py @@ -0,0 +1,111 @@ +import functools, pathlib +from tinygrad import Tensor, dtypes +from tinygrad.uop.ops import UOp, Ops, KernelInfo +from tinygrad.renderer import Estimates +from tinygrad.runtime.support.compiler_amd import HIPCCCompiler +from extra.gemm.cdna_asm_gemm import quantize_mxfp8, _mx_block_scale, _mx_block_scale_3d + +@functools.cache +def custom_hk_grouped_mxfp8_gemm(C:UOp, A:UOp, B:UOp, scale_A:UOp, scale_B:UOp, *extra:UOp, dname:str, n_experts:int) -> UOp: + M, K = A.shape + E, N, K2 = B.shape + assert K == K2, f"{A.shape} {B.shape}" + assert E == n_experts, f"{E} != {n_experts}" + threads = UOp.special(64 * 8, "lidx0") + workgroups = UOp.special((M // 256) * (N // 256), "gidx0") + sink_inputs = (C.base, A.base, B.base, scale_A.base, scale_B.base, extra[0].base, extra[1].base, extra[2].base, threads, workgroups) + sink = UOp.sink(*sink_inputs, + arg=KernelInfo(f"hk_grouped_mxfp8_gemm_{E}_{M}_{N}_{K}", + estimates=Estimates(ops=2*M*N*K, mem=(M*K+E*N*K)*A.dtype.itemsize+M*N*C.dtype.itemsize))) + kittens_path = pathlib.Path(__file__).parent.parent/"thunder"/"amd" + src = (kittens_path/"grouped_mxfp8_gemm.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}", + f"-DGEMM_E={E}"]).compile_cached(src) + return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.LINEAR, src=(*sink.src, sink)), UOp(Ops.SOURCE, arg=src), + UOp(Ops.BINARY, arg=lib))) + +@functools.cache +def custom_hk_grouped_mxfp8_wgrad(C:UOp, A:UOp, B:UOp, scale_A:UOp, scale_B:UOp, expert_off:UOp, *, dname:str, n_experts:int) -> UOp: + N, M = A.shape + K, M2 = B.shape + assert M == M2, f"{A.shape} {B.shape}" + E = n_experts + threads = UOp.special(64 * 8, "lidx0") + workgroups = UOp.special(E * (N // 256) * (K // 256), "gidx0") + sink = UOp.sink(C.base, A.base, B.base, scale_A.base, scale_B.base, expert_off.base, threads, workgroups, + arg=KernelInfo(f"hk_grouped_mxfp8_wgrad_{E}_{M}_{N}_{K}", + estimates=Estimates(ops=2*M*N*K, mem=(N*M+K*M)*A.dtype.itemsize+E*N*K*C.dtype.itemsize))) + kittens_path = pathlib.Path(__file__).parent.parent/"thunder"/"amd" + src = (kittens_path/"grouped_mxfp8_wgrad.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"-DWGRAD_M={M}", f"-DWGRAD_N={N}", f"-DWGRAD_K={K}", + f"-DWGRAD_E={E}"]).compile_cached(src) + return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.LINEAR, src=(*sink.src, sink)), UOp(Ops.SOURCE, arg=src), + UOp(Ops.BINARY, arg=lib))) + +def grouped_mx_wgrad(g:Tensor, xg:Tensor, expert_off:Tensor, n_experts:int) -> Tensor: + from extra.llama_kernels.transpose_quantize_mxfp8 import transpose_quantize_mxfp8 + M, N = g.shape + M2, K = xg.shape + assert M == M2, f"{g.shape} {xg.shape}" + assert M % 128 == 0 and N % 256 == 0 and K % 256 == 0, f"wgrad needs M%128,N%256,K%256, got {g.shape} {xg.shape}" + gT, _, g_si = transpose_quantize_mxfp8(g.contiguous()) + xT, _, x_si = transpose_quantize_mxfp8(xg.contiguous()) + dname = (g.device[0] if isinstance(g.device, tuple) else g.device).split(":")[0] + is_multi = isinstance(g.device, tuple) + inv = Tensor.invalids(1, n_experts * N, K, dtype=dtypes.bfloat16, device=g.device) + out = Tensor(inv.uop.multi(0), device=g.device) if is_multi else inv + out = Tensor.custom_kernel(out, gT, xT, g_si, x_si, expert_off, + fxn=functools.partial(custom_hk_grouped_mxfp8_wgrad, dname=dname, n_experts=n_experts))[0] + out = out.sum(0) if is_multi else out.squeeze(0) + return out.reshape(n_experts, N, K) + +def mx_pack_3d(e8:Tensor) -> Tensor: + E, rows, scale_K = e8.shape + return e8.reshape(E, rows, scale_K // 4, 4).bitcast(dtypes.uint32).reshape(E, rows, scale_K // 4).permute(0, 2, 1).contiguous() + +@functools.cache +def custom_grouped_mx_gemm_bw(gradient:UOp, kernel:UOp, w_stored:bool=False) -> tuple: + inputs = kernel.src[1:] + aq = Tensor(inputs[1], device=inputs[1].device) + bq = Tensor(inputs[2], device=inputs[2].device) + ae8 = Tensor(inputs[5], device=inputs[5].device) + be8 = Tensor(inputs[6], device=inputs[6].device) + E, N = bq.shape[0], bq.shape[1] + M, K = aq.shape + g = Tensor(gradient, device=aq.device).reshape(M, N).cast(dtypes.bfloat16) + x_phys = (aq.cast(dtypes.bfloat16) * _mx_block_scale(ae8).cast(dtypes.bfloat16)) + w_phys = (bq.cast(dtypes.bfloat16) * _mx_block_scale_3d(be8).cast(dtypes.bfloat16)) + expert_off = Tensor(inputs[7], device=inputs[7].device) + grad_x = grouped_mx_gemm(g, w_phys.transpose(1, 2), expert_off) + grad_w = grouped_mx_wgrad(g, x_phys, expert_off, E) + grad_xq = grad_x * _mx_block_scale(ae8).cast(dtypes.bfloat16) + grad_wq = grad_w.contiguous() if w_stored else (grad_w * _mx_block_scale_3d(be8).cast(dtypes.bfloat16)).contiguous() + return (None, grad_xq.uop, grad_wq.uop) + tuple(None for _ in inputs[3:]) + +_grouped_bw_stored = functools.partial(custom_grouped_mx_gemm_bw, w_stored=True) + +def grouped_mx_gemm(x:Tensor, w:Tensor|tuple[Tensor, Tensor], expert_off:Tensor) -> Tensor: + if (pre_quantized := isinstance(w, tuple)): + w_q, w_e8 = w + E, N, K2 = w_q.shape + else: + E, N, K2 = w.shape + M, K = x.shape + assert K == K2, f"shape mismatch {x.shape} {w.shape}" + assert M % 256 == 0 and N % 256 == 0 and K % 128 == 0, f"grouped mxfp8 needs M%256,N%256,K%128, got {x.shape} {w.shape}" + dname = (x.device[0] if isinstance(x.device, tuple) else x.device).split(":")[0] + x_q, x_e8, x_si = quantize_mxfp8(x) + if not pre_quantized: w_q, w_e8, _ = quantize_mxfp8(w) + w_si = mx_pack_3d(w_e8) + xe_in, out_shape = x_e8.reshape(M, K // 32), (M, N) + if isinstance(x.device, tuple) and (row_axis := x.uop.axis) is not None: + ndev = len(x.device) + out = Tensor(Tensor.invalids(*(s // ndev if i == row_axis else s for i, s in enumerate(out_shape)), + dtype=dtypes.bfloat16, device=x.device).uop.multi(row_axis), device=x.device) + else: + out = Tensor.invalids(*out_shape, dtype=dtypes.bfloat16, device=x.device) + return Tensor.custom_kernel(out, x_q, w_q, x_si, w_si, xe_in, w_e8, expert_off, + fxn=functools.partial(custom_hk_grouped_mxfp8_gemm, dname=dname, n_experts=E), + grad_fxn=(_grouped_bw_stored if pre_quantized else custom_grouped_mx_gemm_bw))[0] diff --git a/extra/thunder/amd/grouped_mxfp8_gemm.cpp b/extra/thunder/amd/grouped_mxfp8_gemm.cpp new file mode 100644 index 0000000000..2a41c41c6c --- /dev/null +++ b/extra/thunder/amd/grouped_mxfp8_gemm.cpp @@ -0,0 +1,158 @@ +#include "kittens.cuh" + +using namespace kittens; + +#ifndef GEMM_M +constexpr int GEMM_M = 8192; +#endif +#ifndef GEMM_N +constexpr int GEMM_N = 8192; +#endif +#ifndef GEMM_K +constexpr int GEMM_K = 8192; +#endif +#ifndef GEMM_E +constexpr int GEMM_E = 8; +#endif + +// Kernel +constexpr int NUM_WARPS = 8; +constexpr int WARPS_ROW = 2; +constexpr int WARPS_COL = 4; +constexpr int BLOCK_ROW = 256; +constexpr int BLOCK_COL = 256; +constexpr int BLOCK_K = 128; +constexpr int HALF_ROW = BLOCK_ROW / 2; +constexpr int HALF_COL = BLOCK_COL / 2; +constexpr int REG_M = BLOCK_ROW / WARPS_ROW / 2; +constexpr int REG_N = BLOCK_COL / WARPS_COL / 2; + +using G = kittens::group; + +__global__ __launch_bounds__(512, 2) void grouped_mxfp8_gemm_kernel(bf16 *C_ptr, fp8e4m3 *A_ptr, fp8e4m3 *B_ptr, fp8e8m0 *scale_A_ptr, fp8e8m0 *scale_B_ptr, + const uint8_t *__restrict__ a_e8_unused, + const uint8_t *__restrict__ b_e8_unused, + const int *__restrict__ expert_off) { + constexpr int M = GEMM_M, N = GEMM_N, K = GEMM_K, E = GEMM_E; + + kittens::gl A{A_ptr, nullptr, nullptr, nullptr, nullptr}; + kittens::gl B{B_ptr, nullptr, nullptr, nullptr, nullptr}; // all experts stacked on rows + kittens::gl C{C_ptr, nullptr, nullptr, nullptr, nullptr}; + + constexpr int k_iters = K / BLOCK_K; + constexpr int NUM_THREADS = NUM_WARPS * WARP_THREADS; + + kittens::gl scale_A_gl{scale_A_ptr, nullptr, nullptr, nullptr, nullptr}; + kittens::gl scale_B_gl{scale_B_ptr, nullptr, nullptr, nullptr, nullptr}; + + using ST_A = st_fp8e4m3; + using ST_B = st_fp8e4m3; + using ST_Scale = st; + using RT_A = rt_fp8e4m3; + using RT_B = rt_fp8e4m3; + using RT_C = rt_fl; + + __shared__ ST_A As[2][2]; + __shared__ ST_B Bs[2][2]; + __shared__ ST_Scale scale_A_smem[2], scale_B_smem[2]; + + RT_A a; + RT_B b0, b1; + RT_C cA, cB, cC, cD; + zero(cA); zero(cB); zero(cC); zero(cD); + + constexpr int tiles_M = M / BLOCK_ROW; + constexpr int tiles_N = N / BLOCK_COL; + const int NUM_XCDS = 8; + const int WGM = 8; + int wgid = chiplet_transform_chunked(blockIdx.x, gridDim.x, NUM_XCDS, WGM * WGM); + int num_wgid_in_group = WGM * tiles_N; + int group_id = wgid / num_wgid_in_group; + int first_pid_m = group_id * WGM; + int group_size_m = min(tiles_M - first_pid_m, WGM); + int block_row = first_pid_m + ((wgid % num_wgid_in_group) % group_size_m); + int block_col = (wgid % num_wgid_in_group) / group_size_m; + int block_m = block_row * BLOCK_ROW; + int block_n = block_col * BLOCK_COL; + + int e = 0; + #pragma unroll + for (int i = 1; i < E; i++) e += (expert_off[i] <= block_row * BLOCK_ROW); + e = __builtin_amdgcn_readfirstlane(e); + const int bcol_base = e * (N / HALF_COL); // expert base in B row-tile (128-row) units + const int sb_base = e * (k_iters * tiles_N); // expert base into scale_B batches + + int warp_m = warpid() / WARPS_COL; + int warp_n = warpid() % WARPS_COL; + + using T = fp8e4m3; + constexpr int bpt = ST_A::underlying_subtile_bytes_per_thread; + constexpr int bpm = bpt * NUM_THREADS; + constexpr int copies_A = HALF_ROW * BLOCK_K * sizeof(T) / bpm; + constexpr int copies_B = HALF_COL * BLOCK_K * sizeof(T) / bpm; + uint32_t sw_A[copies_A], sw_B[copies_B]; + G::prefill_swizzled_offsets(As[0][0], A, sw_A); + G::prefill_swizzled_offsets(Bs[0][0], B, sw_B); + + const T *a_base = (const T *)&A[{0, 0, 0, 0}]; + const T *b_base = (const T *)&B[{0, 0, 0, 0}]; + const int a_row_stride = A.template stride<2>() * sizeof(T); + const int b_row_stride = B.template stride<2>() * sizeof(T); + i32x4 a_srd = make_srsrc(a_base, (uint32_t)((uint64_t)M * a_row_stride), a_row_stride); + i32x4 b_srd = make_srsrc(b_base, (uint32_t)((uint64_t)E * N * b_row_stride), b_row_stride); + + const int wid = warpid() % NUM_WARPS; + constexpr int elem_per_warp = (16 / sizeof(T)) * kittens::WARP_THREADS; + uint32_t a_lds_00 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&As[0][0].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t a_lds_01 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&As[0][1].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t a_lds_10 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&As[1][0].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t a_lds_11 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&As[1][1].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t b_lds_00 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&Bs[0][0].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t b_lds_01 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&Bs[0][1].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t b_lds_10 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&Bs[1][0].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t b_lds_11 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&Bs[1][1].data[0]) + wid * elem_per_warp * sizeof(T))); + + int a_row_h0 = warp_m * REG_M; + int a_row_h1 = HALF_ROW + warp_m * REG_M; + int b_row_h0 = warp_n * REG_N; + int b_row_h1 = HALF_COL + warp_n * REG_N; + + + uint32_t a_lds[2][2] = {{a_lds_00, a_lds_01}, {a_lds_10, a_lds_11}}; + uint32_t b_lds[2][2] = {{b_lds_00, b_lds_01}, {b_lds_10, b_lds_11}}; + + #pragma unroll 1 + for (int kk = 0; kk < k_iters; kk++) { + G::load(As[0][0], A, {0, 0, block_row * 2, kk}, sw_A, a_srd, a_base, __builtin_amdgcn_readfirstlane(a_lds[0][0])); + G::load(As[0][1], A, {0, 0, block_row * 2 + 1, kk}, sw_A, a_srd, a_base, __builtin_amdgcn_readfirstlane(a_lds[0][1])); + G::load(Bs[0][0], B, {0, 0, bcol_base + block_col * 2, kk}, sw_B, b_srd, b_base, __builtin_amdgcn_readfirstlane(b_lds[0][0])); + G::load(Bs[0][1], B, {0, 0, bcol_base + block_col * 2 + 1, kk}, sw_B, b_srd, b_base, __builtin_amdgcn_readfirstlane(b_lds[0][1])); + G::load(scale_A_smem[0], scale_A_gl, {kk * tiles_M + block_row, 0, 0, 0}); + G::load(scale_B_smem[0], scale_B_gl, {sb_base + kk * tiles_N + block_col, 0, 0, 0}); + asm volatile("s_waitcnt vmcnt(0)"); + asm volatile("s_waitcnt lgkmcnt(0)"); + __builtin_amdgcn_s_barrier(); + + fp8e8m0_4 sa_h0 = pack_scales(scale_A_smem[0].data, a_row_h0); + fp8e8m0_4 sa_h1 = pack_scales(scale_A_smem[0].data, a_row_h1); + fp8e8m0_4 sb_h0 = pack_scales(scale_B_smem[0].data, b_row_h0); + fp8e8m0_4 sb_h1 = pack_scales(scale_B_smem[0].data, b_row_h1); + + auto bs0 = subtile_inplace(Bs[0][0], {warp_n, 0}); load(b0, bs0); + auto bs1 = subtile_inplace(Bs[0][1], {warp_n, 0}); load(b1, bs1); + auto as0 = subtile_inplace(As[0][0], {warp_m, 0}); load(a, as0); + asm volatile("s_waitcnt lgkmcnt(0)"); + mma_ABt_scaled(cA, a, b0, cA, &sa_h0, &sb_h0); + mma_ABt_scaled(cB, a, b1, cB, &sa_h0, &sb_h1); + auto as1 = subtile_inplace(As[0][1], {warp_m, 0}); load(a, as1); + asm volatile("s_waitcnt lgkmcnt(0)"); + mma_ABt_scaled(cC, a, b0, cC, &sa_h1, &sb_h0); + mma_ABt_scaled(cD, a, b1, cD, &sa_h1, &sb_h1); + __builtin_amdgcn_s_barrier(); + } + + 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}); + store(C, cC, {0, 0, block_row * WARPS_ROW * 2 + WARPS_ROW + warp_m, block_col * WARPS_COL * 2 + warp_n}); + store(C, cD, {0, 0, block_row * WARPS_ROW * 2 + WARPS_ROW + warp_m, block_col * WARPS_COL * 2 + WARPS_COL + warp_n}); +} diff --git a/extra/thunder/amd/grouped_mxfp8_wgrad.cpp b/extra/thunder/amd/grouped_mxfp8_wgrad.cpp new file mode 100644 index 0000000000..f74b97f316 --- /dev/null +++ b/extra/thunder/amd/grouped_mxfp8_wgrad.cpp @@ -0,0 +1,141 @@ +#include "kittens.cuh" + +using namespace kittens; + +#ifndef WGRAD_M +constexpr int WGRAD_M = 8192; +#endif +#ifndef WGRAD_N +constexpr int WGRAD_N = 8192; +#endif +#ifndef WGRAD_K +constexpr int WGRAD_K = 8192; +#endif +#ifndef WGRAD_E +constexpr int WGRAD_E = 8; +#endif + +constexpr int NUM_WARPS = 8; +constexpr int WARPS_ROW = 2; +constexpr int WARPS_COL = 4; +constexpr int BLOCK_ROW = 256; +constexpr int BLOCK_COL = 256; +constexpr int BLOCK_K = 128; +constexpr int HALF_ROW = BLOCK_ROW / 2; +constexpr int HALF_COL = BLOCK_COL / 2; +constexpr int REG_M = BLOCK_ROW / WARPS_ROW / 2; +constexpr int REG_N = BLOCK_COL / WARPS_COL / 2; + +using G = kittens::group; + +__global__ __launch_bounds__(512, 2) void grouped_mxfp8_wgrad_kernel(bf16 *C_ptr, fp8e4m3 *A_ptr, fp8e4m3 *B_ptr, + fp8e8m0 *scale_A_ptr, fp8e8m0 *scale_B_ptr, + const int *__restrict__ expert_off) { + constexpr int M = WGRAD_M, N = WGRAD_N, K = WGRAD_K, E = WGRAD_E; + + kittens::gl A{A_ptr, nullptr, nullptr, nullptr, nullptr}; // g^T + kittens::gl B{B_ptr, nullptr, nullptr, nullptr, nullptr}; // x^T + kittens::gl C{C_ptr, nullptr, nullptr, nullptr, nullptr}; // grad_w, experts stacked + + constexpr int m_blocks = M / BLOCK_K; // 128-wide blocks along the contraction (token) axis + constexpr int tiles_N = N / BLOCK_ROW; + constexpr int tiles_K = K / BLOCK_COL; + constexpr int NUM_THREADS = NUM_WARPS * WARP_THREADS; + + kittens::gl scale_A_gl{scale_A_ptr, nullptr, nullptr, nullptr, nullptr}; + kittens::gl scale_B_gl{scale_B_ptr, nullptr, nullptr, nullptr, nullptr}; + + using ST_A = st_fp8e4m3; + using ST_B = st_fp8e4m3; + using ST_Scale = st; + using RT_A = rt_fp8e4m3; + using RT_B = rt_fp8e4m3; + using RT_C = rt_fl; + + __shared__ ST_A As[2]; + __shared__ ST_B Bs[2]; + __shared__ ST_Scale scale_A_smem, scale_B_smem; + + RT_A a; + RT_B b0, b1; + RT_C cA, cB, cC, cD; + zero(cA); zero(cB); zero(cC); zero(cD); + + const int wg = blockIdx.x; + const int e = wg / (tiles_N * tiles_K); + const int rem = wg % (tiles_N * tiles_K); + const int block_row = rem / tiles_K; // over grad_w rows (N) + const int block_col = rem % tiles_K; // over grad_w cols (K) + + const int o0 = __builtin_amdgcn_readfirstlane(expert_off[e]); + const int kk0 = o0 / BLOCK_K; + const int nk = (__builtin_amdgcn_readfirstlane(expert_off[e + 1]) - o0) / BLOCK_K; + + const int warp_m = warpid() / WARPS_COL; + const int warp_n = warpid() % WARPS_COL; + + using T = fp8e4m3; + constexpr int bpt = ST_A::underlying_subtile_bytes_per_thread; + constexpr int bpm = bpt * NUM_THREADS; + constexpr int copies_A = HALF_ROW * BLOCK_K * sizeof(T) / bpm; + constexpr int copies_B = HALF_COL * BLOCK_K * sizeof(T) / bpm; + uint32_t sw_A[copies_A], sw_B[copies_B]; + G::prefill_swizzled_offsets(As[0], A, sw_A); + G::prefill_swizzled_offsets(Bs[0], B, sw_B); + + const T *a_base = (const T *)&A[{0, 0, 0, 0}]; + const T *b_base = (const T *)&B[{0, 0, 0, 0}]; + const int a_row_stride = A.template stride<2>() * sizeof(T); + const int b_row_stride = B.template stride<2>() * sizeof(T); + i32x4 a_srd = make_srsrc(a_base, (uint32_t)((uint64_t)N * a_row_stride), a_row_stride); + i32x4 b_srd = make_srsrc(b_base, (uint32_t)((uint64_t)K * b_row_stride), b_row_stride); + + const int wid = warpid() % NUM_WARPS; + constexpr int elem_per_warp = (16 / sizeof(T)) * kittens::WARP_THREADS; + uint32_t a_lds_0 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&As[0].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t a_lds_1 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&As[1].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t b_lds_0 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&Bs[0].data[0]) + wid * elem_per_warp * sizeof(T))); + uint32_t b_lds_1 = __builtin_amdgcn_readfirstlane(static_cast(reinterpret_cast(&Bs[1].data[0]) + wid * elem_per_warp * sizeof(T))); + + const int a_row_h0 = warp_m * REG_M; + const int a_row_h1 = HALF_ROW + warp_m * REG_M; + const int b_row_h0 = warp_n * REG_N; + const int b_row_h1 = HALF_COL + warp_n * REG_N; + + #pragma unroll 1 + for (int t = 0; t < nk; t++) { + const int kk = kk0 + t; + G::load(As[0], A, {0, 0, block_row * 2, kk}, sw_A, a_srd, a_base, __builtin_amdgcn_readfirstlane(a_lds_0)); + G::load(As[1], A, {0, 0, block_row * 2 + 1, kk}, sw_A, a_srd, a_base, __builtin_amdgcn_readfirstlane(a_lds_1)); + G::load(Bs[0], B, {0, 0, block_col * 2, kk}, sw_B, b_srd, b_base, __builtin_amdgcn_readfirstlane(b_lds_0)); + G::load(Bs[1], B, {0, 0, block_col * 2 + 1, kk}, sw_B, b_srd, b_base, __builtin_amdgcn_readfirstlane(b_lds_1)); + G::load(scale_A_smem, scale_A_gl, {kk * tiles_N + block_row, 0, 0, 0}); + G::load(scale_B_smem, scale_B_gl, {kk * tiles_K + block_col, 0, 0, 0}); + asm volatile("s_waitcnt vmcnt(0)"); + asm volatile("s_waitcnt lgkmcnt(0)"); + __builtin_amdgcn_s_barrier(); + + fp8e8m0_4 sa_h0 = pack_scales(scale_A_smem.data, a_row_h0); + fp8e8m0_4 sa_h1 = pack_scales(scale_A_smem.data, a_row_h1); + fp8e8m0_4 sb_h0 = pack_scales(scale_B_smem.data, b_row_h0); + fp8e8m0_4 sb_h1 = pack_scales(scale_B_smem.data, b_row_h1); + + auto bs0 = subtile_inplace(Bs[0], {warp_n, 0}); load(b0, bs0); + auto bs1 = subtile_inplace(Bs[1], {warp_n, 0}); load(b1, bs1); + auto as0 = subtile_inplace(As[0], {warp_m, 0}); load(a, as0); + asm volatile("s_waitcnt lgkmcnt(0)"); + mma_ABt_scaled(cA, a, b0, cA, &sa_h0, &sb_h0); + mma_ABt_scaled(cB, a, b1, cB, &sa_h0, &sb_h1); + auto as1 = subtile_inplace(As[1], {warp_m, 0}); load(a, as1); + asm volatile("s_waitcnt lgkmcnt(0)"); + mma_ABt_scaled(cC, a, b0, cC, &sa_h1, &sb_h0); + mma_ABt_scaled(cD, a, b1, cD, &sa_h1, &sb_h1); + __builtin_amdgcn_s_barrier(); + } + + const int crow_base = e * (N / REG_M); // grad_w rows are experts stacked; store coord is in REG_M units + store(C, cA, {0, 0, crow_base + block_row * WARPS_ROW * 2 + warp_m, block_col * WARPS_COL * 2 + warp_n}); + store(C, cB, {0, 0, crow_base + block_row * WARPS_ROW * 2 + warp_m, block_col * WARPS_COL * 2 + WARPS_COL + warp_n}); + store(C, cC, {0, 0, crow_base + block_row * WARPS_ROW * 2 + WARPS_ROW + warp_m, block_col * WARPS_COL * 2 + warp_n}); + store(C, cD, {0, 0, crow_base + block_row * WARPS_ROW * 2 + WARPS_ROW + warp_m, block_col * WARPS_COL * 2 + WARPS_COL + warp_n}); +}