From 6242b09066db994cf30177f543e8ccd2ee66af16 Mon Sep 17 00:00:00 2001 From: qazal <77887910+Qazalin@users.noreply.github.com> Date: Fri, 28 Aug 2026 14:59:35 +0800 Subject: [PATCH] cleaner mxfp4 gemm prelude (#17796) * cleaner mxfp4 prelude * rename sgprs * min diff --- extra/gemm/gemm_mxfp4.py | 183 +++++++++++++++------------------------ 1 file changed, 68 insertions(+), 115 deletions(-) diff --git a/extra/gemm/gemm_mxfp4.py b/extra/gemm/gemm_mxfp4.py index 678e1ffd5d..4edbf17799 100644 --- a/extra/gemm/gemm_mxfp4.py +++ b/extra/gemm/gemm_mxfp4.py @@ -20,34 +20,39 @@ def v_mfma_fp4(dst, a, b, opsel, opsel_hi, scale_a, scale_b): def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k = Kernel() scale_k = K // 32 + k.emit(s_and_b32(s[1], s[1], LIT, 65535)) if (tile_m, tile_n) == (128, 512): - k.emit(s_and_b32(s[1], s[1], LIT, 65535)) k.emit(s_mov_b32(s[47], s[2])) k.emit(s_mov_b32(s[48], s[3])) - k.emit(s_load_dwordx2(s[4:5], s[0:1], s[0], 0, 0, 0, 0, 1)) - k.emit(s_mov_b32(s[8], 0)) - k.emit(s_mov_b32(s[9], 0)) - k.emit(s_load_dwordx2(s[12:13], s[0:1], s[0], 8, 0, 0, 0, 1)) - k.emit(s_load_dwordx2(s[16:17], s[0:1], s[0], 16, 0, 0, 0, 1)) - k.emit(s_mov_b32(s[36], N)) - k.emit(s_mov_b32(s[37], K)) - k.emit(s_mov_b32(s[38], K)) - k.emit(s_mov_b32(s[43], M)) - k.emit(s_mov_b32(s[44], N)) - k.emit(s_mov_b32(s[45], K)) - k.emit(s_load_dwordx2(s[20:21], s[0:1], s[0], 24, 0, 0, 0, 1)) - k.emit(s_load_dwordx2(s[24:25], s[0:1], s[0], 32, 0, 0, 0, 1)) - k.emit(s_mov_b32(s[39], scale_k)) - k.emit(s_mov_b32(s[40], scale_k)) - k.emit(v_lshrrev_b32_e32(v[1], 10)) - k.emit(v_lshrrev_b32_e32(v[2], 10, v[1])) - k.emit(v_and_b32_e32(v[2], LIT, v[2], 1023)) - k.emit(v_and_b32_e32(v[1], LIT, v[1], 1023)) - k.emit(v_and_b32_e32(v[0], LIT, v[0], 1023)) - k.emit(v_lshrrev_b32_e32(v[3], 6)) - k.emit(v_and_b32_e32(v[0], 63)) - k.emit(v_readfirstlane_b32_e32(v[46], v[3])) - k.emit(s_waitcnt(49279)) + k.emit(s_load_dwordx2(s[4:5], s[0:1], s[0], 0, 0, 0, 0, 1)) + k.emit(s_mov_b32(s[8], 0)) + k.emit(s_mov_b32(s[9], 0)) + k.emit(s_load_dwordx2(s[12:13], s[0:1], s[0], 8, 0, 0, 0, 1)) + k.emit(s_load_dwordx2(s[16:17], s[0:1], s[0], 16, 0, 0, 0, 1)) + k.emit(s_mov_b32(s[36], N)) + k.emit(s_mov_b32(s[37], K)) + k.emit(s_mov_b32(s[38], K)) + k.emit(s_mov_b32(s[43], M)) + k.emit(s_mov_b32(s[44], N)) + k.emit(s_mov_b32(s[45], K)) + k.emit(s_load_dwordx2(s[20:21], s[0:1], s[0], 24, 0, 0, 0, 1)) + k.emit(s_load_dwordx2(s[24:25], s[0:1], s[0], 32, 0, 0, 0, 1)) + k.emit(s_mov_b32(s[39], scale_k)) + k.emit(s_mov_b32(s[40], scale_k)) + k.emit(v_lshrrev_b32_e32(v[1], 10)) + k.emit(v_lshrrev_b32_e32(v[2], 10, v[1])) + k.emit(v_and_b32_e32(v[2], LIT, v[2], 1023)) + k.emit(v_and_b32_e32(v[1], LIT, v[1], 1023)) + k.emit(v_and_b32_e32(v[0], LIT, v[0], 1023)) + k.emit(v_lshrrev_b32_e32(v[3], 6)) + k.emit(v_and_b32_e32(v[0], 63)) + if (tile_m, tile_n) == (256, 256): + k.emit(s_mov_b32(s[49], s[2])) + k.emit(s_mov_b32(s[47], s[3])) + k.emit(v_readfirstlane_b32_e32(v[46], v[3])) + k.emit(s_waitcnt(49279)) + + if (tile_m, tile_n) == (128, 512): for i in range(2): k.emit(s_mov_b32(s[6 + i * 8], -16)) k.emit(s_mov_b32(s[10 + i * 12], -16)) @@ -1213,31 +1218,6 @@ def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k.emit(s_waitcnt()) k.emit(s_endpgm()) elif (tile_m, tile_n) == (192, 256): - k.emit(s_and_b32(s[1], s[1], LIT, 65535)) - k.emit(s_load_dwordx2(s[4:5], s[0:1], s[0], 0, 0, 0, 0, 1)) - k.emit(s_mov_b32(s[8], 0)) - k.emit(s_mov_b32(s[9], 0)) - k.emit(s_load_dwordx2(s[12:13], s[0:1], s[0], 8, 0, 0, 0, 1)) - k.emit(s_load_dwordx2(s[16:17], s[0:1], s[0], 16, 0, 0, 0, 1)) - k.emit(s_mov_b32(s[36], N)) - k.emit(s_mov_b32(s[37], K)) - k.emit(s_mov_b32(s[38], K)) - k.emit(s_mov_b32(s[43], M)) - k.emit(s_mov_b32(s[44], N)) - k.emit(s_mov_b32(s[45], K)) - k.emit(s_load_dwordx2(s[20:21], s[0:1], s[0], 24, 0, 0, 0, 1)) - k.emit(s_load_dwordx2(s[24:25], s[0:1], s[0], 32, 0, 0, 0, 1)) - k.emit(s_mov_b32(s[39], scale_k)) - k.emit(s_mov_b32(s[40], scale_k)) - k.emit(v_lshrrev_b32_e32(v[1], 10)) - k.emit(v_lshrrev_b32_e32(v[2], 10, v[1])) - k.emit(v_and_b32_e32(v[2], LIT, v[2], 1023)) - k.emit(v_and_b32_e32(v[1], LIT, v[1], 1023)) - k.emit(v_and_b32_e32(v[0], LIT, v[0], 1023)) - k.emit(v_lshrrev_b32_e32(v[3], 6)) - k.emit(v_and_b32_e32(v[0], 63)) - k.emit(v_readfirstlane_b32_e32(v[46], v[3])) - k.emit(s_waitcnt(49279)) k.emit(s_mul_i32(s[63], LIT, 8, 192)) k.emit(v_cvt_f32_u32_e32(v[4], s[63])) k.emit(s_sub_i32(s[62], 0, s[63])) @@ -2234,49 +2214,22 @@ def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k.emit(s_waitcnt()) k.emit(s_endpgm()) elif (tile_m, tile_n) == (256, 256): - k.emit(s_and_b32(s[1], s[1], LIT, 65535)) - k.emit(s_load_dwordx2(s[4:5], s[0:1], s[0], 0, 0, 0, 0, 1)) - k.emit(s_mov_b32(s[8], 0)) - k.emit(s_mov_b32(s[9], 0)) - k.emit(s_load_dwordx2(s[12:13], s[0:1], s[0], 8, 0, 0, 0, 1)) - k.emit(s_load_dwordx2(s[16:17], s[0:1], s[0], 16, 0, 0, 0, 1)) - k.emit(s_mov_b32(s[40], N)) - k.emit(s_mov_b32(s[41], K)) - k.emit(s_mov_b32(s[42], K)) - k.emit(s_mov_b32(s[43], M)) - k.emit(s_mov_b32(s[44], N)) - k.emit(s_mov_b32(s[45], K)) - k.emit(s_load_dwordx2(s[20:21], s[0:1], s[0], 24, 0, 0, 0, 1)) - k.emit(s_load_dwordx2(s[24:25], s[0:1], s[0], 32, 0, 0, 0, 1)) - k.emit(s_mov_b32(s[36], scale_k)) - k.emit(s_mov_b32(s[37], scale_k)) - k.emit(v_lshrrev_b32_e32(v[1], 10)) - k.emit(v_lshrrev_b32_e32(v[2], 10, v[1])) - k.emit(v_and_b32_e32(v[2], LIT, v[2], 1023)) - k.emit(v_and_b32_e32(v[1], LIT, v[1], 1023)) - k.emit(v_and_b32_e32(v[0], LIT, v[0], 1023)) - k.emit(v_lshrrev_b32_e32(v[3], 6)) - k.emit(v_and_b32_e32(v[0], 63)) - k.emit(s_mov_b32(s[46], s[2])) - k.emit(s_mov_b32(s[47], s[3])) - k.emit(v_readfirstlane_b32_e32(v[49], v[3])) - k.emit(s_waitcnt(49279)) k.emit(s_add_u32(s[55], s[44], LIT, 255)) k.emit(s_lshr_b32(s[54], s[55], 8)) k.emit(s_mul_i32(s[48], s[54], s[47])) - k.emit(s_add_i32(s[48], s[48], s[46])) + k.emit(s_add_i32(s[48], s[48], s[49])) k.emit(s_add_u32(s[55], s[43], LIT, 255)) k.emit(s_lshr_b32(s[52], s[55], 8)) k.emit(s_lshl_b32(s[52], s[52], 5)) - k.emit(s_mov_b32(s[46], 0)) + k.emit(s_mov_b32(s[49], 0)) k.label('L2_00E8') k.emit(s_cmp_lt_i32(s[48], s[52])) k.emit(s_cbranch_scc1(3), target='L2_00FC') k.emit(s_sub_i32(s[48], s[48], s[52])) - k.emit(s_add_i32(s[46], s[46], 32)) + k.emit(s_add_i32(s[49], s[49], 32)) k.emit(s_branch(65531), target='L2_00E8') k.label('L2_00FC') - k.emit(s_sub_i32(s[54], s[54], s[46])) + k.emit(s_sub_i32(s[54], s[54], s[49])) k.emit(s_cmp_lt_i32(s[54], 32)) k.emit(s_cbranch_scc1(3), target='L2_0114') k.emit(s_lshr_b32(s[47], s[48], 5)) @@ -2311,7 +2264,7 @@ def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k.emit(s_mul_i32(s[52], s[54], s[47])) k.emit(s_sub_i32(s[52], s[48], s[52])) k.label('L2_0194') - k.emit(s_add_i32(s[46], s[52], s[46])) + k.emit(s_add_i32(s[49], s[52], s[49])) k.emit(s_mov_b32(s[6], -16)) k.emit(s_mov_b32(s[10], -16)) k.emit(s_mov_b32(s[18], -16)) @@ -2328,18 +2281,18 @@ def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k.emit(s_or_b32(s[9], s[9], LIT, 262144)) k.emit(s_or_b32(s[17], s[17], LIT, 262144)) k.emit(s_or_b32(s[13], s[13], LIT, 262144)) - k.emit(s_lshr_b32(s[41], s[41], 1)) - k.emit(s_mul_i32(s[52], s[41], s[43])) + k.emit(s_lshr_b32(s[37], s[37], 1)) + k.emit(s_mul_i32(s[52], s[37], s[43])) k.emit(s_mov_b32(s[14], s[52])) - k.emit(s_lshr_b32(s[42], s[42], 1)) - k.emit(s_mul_i32(s[52], s[42], s[44])) + k.emit(s_lshr_b32(s[38], s[38], 1)) + k.emit(s_mul_i32(s[52], s[38], s[44])) k.emit(s_mov_b32(s[18], s[52])) k.emit(s_add_u32(s[52], s[43], 31)) k.emit(s_lshr_b32(s[52], s[52], 5)) k.emit(s_lshl_b32(s[52], s[52], 5)) - k.emit(s_mul_i32(s[53], s[52], s[36])) + k.emit(s_mul_i32(s[53], s[52], s[39])) k.emit(s_mov_b32(s[22], s[53])) - k.emit(s_mul_i32(s[53], s[44], s[37])) + k.emit(s_mul_i32(s[53], s[44], s[40])) k.emit(s_mov_b32(s[26], s[53])) k.emit(s_mov_b32(s[23], LIT, 131072)) k.emit(s_mov_b32(s[27], LIT, 131072)) @@ -2356,23 +2309,23 @@ def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k.emit(v_add_u32_e32(v[5], v[5], v[6])) k.emit(v_and_b32_e32(v[4], 1, v[4])) k.emit(v_add_u32_e32(v[5], v[5], v[4])) - k.emit(v_mul_lo_u32(v[212], s[41], v[5])) + k.emit(v_mul_lo_u32(v[212], s[37], v[5])) k.emit(v_and_b32_e32(v[4], 7)) k.emit(v_lshlrev_b32_e32(v[4], 4, v[4])) k.emit(v_add_u32_e32(v[212], v[212], v[4])) - k.emit(s_lshr_b32(s[52], s[49], 1)) + k.emit(s_lshr_b32(s[52], s[46], 1)) k.emit(s_mul_i32(s[52], s[52], 8)) - k.emit(s_and_b32(s[53], s[49], 1)) + k.emit(s_and_b32(s[53], s[46], 1)) k.emit(s_mul_i32(s[53], s[53], 2)) k.emit(s_add_u32(s[52], s[52], s[53])) k.emit(s_mul_i32(s[53], s[47], LIT, 256)) k.emit(s_add_u32(s[52], s[52], s[53])) - k.emit(s_mul_i32(s[52], s[41], s[52])) + k.emit(s_mul_i32(s[52], s[37], s[52])) k.emit(v_add_u32_e32(v[212], s[52], v[212])) - k.emit(s_mul_i32(s[52], s[41], 32)) + k.emit(s_mul_i32(s[52], s[37], 32)) for i in range(7): k.emit(v_add_u32_e32(v[213 + i * 1], s[52], v[212 + i * 1])) - k.emit(s_mul_i32(s[59], LIT, s[49], 1056)) + k.emit(s_mul_i32(s[59], LIT, s[46], 1056)) k.emit(s_add_u32(s[59], LIT, s[59], 4096)) k.emit(v_and_b32_e32(v[4], 15)) k.emit(v_lshrrev_b32_e32(v[5], 3, v[4])) @@ -2396,35 +2349,35 @@ def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k.emit(v_add_u32_e32(v[221], LIT, v[220], 33792)) k.emit(v_lshlrev_b32_e32(v[222], 2)) k.emit(s_mul_i32(s[52], s[47], LIT, 256)) - k.emit(s_mul_i32(s[53], s[49], 32)) + k.emit(s_mul_i32(s[53], s[46], 32)) k.emit(s_add_i32(s[52], s[53], s[52])) - k.emit(s_mul_i32(s[53], s[52], s[36])) + k.emit(s_mul_i32(s[53], s[52], s[39])) k.emit(v_add_u32_e32(v[222], s[53], v[222])) - k.emit(s_mul_i32(s[53], LIT, s[36], 128)) + k.emit(s_mul_i32(s[53], LIT, s[39], 128)) k.emit(v_add_u32_e32(v[223], s[53], v[222])) - k.emit(s_mul_i32(s[60], s[49], LIT, 256)) + k.emit(s_mul_i32(s[60], s[46], LIT, 256)) k.emit(s_add_i32(s[60], s[60], 0)) k.emit(v_lshlrev_b32_e32(v[224], 2)) k.emit(v_add_u32_e32(v[224], 0, v[224])) k.emit(v_lshlrev_b32_e32(v[225], 4)) - k.emit(s_mul_i32(s[52], s[46], LIT, 256)) - k.emit(s_mul_i32(s[53], s[49], 64)) + k.emit(s_mul_i32(s[52], s[49], LIT, 256)) + k.emit(s_mul_i32(s[53], s[46], 64)) k.emit(s_add_u32(s[52], s[52], s[53])) - k.emit(s_mul_i32(s[52], s[52], s[42])) + k.emit(s_mul_i32(s[52], s[52], s[38])) k.emit(v_add_u32_e32(v[225], s[52], v[225])) - k.emit(s_mul_i32(s[52], 16, s[42])) + k.emit(s_mul_i32(s[52], 16, s[38])) k.emit(v_add_u32_e32(v[226], s[52], v[225])) k.emit(v_add_u32_e32(v[227], s[52], v[226])) k.emit(v_add_u32_e32(v[228], s[52], v[227])) for i in range(4): k.emit(v_add_u32_e32(v[229 + i * 1], LIT, v[225 + i * 1], 1024)) k.emit(v_lshlrev_b32_e32(v[233], 2)) - k.emit(s_mul_i32(s[52], s[46], LIT, 256)) - k.emit(s_mul_i32(s[53], s[49], 64)) + k.emit(s_mul_i32(s[52], s[49], LIT, 256)) + k.emit(s_mul_i32(s[53], s[46], 64)) k.emit(s_add_i32(s[52], s[53], s[52])) - k.emit(s_mul_i32(s[53], s[52], s[37])) + k.emit(s_mul_i32(s[53], s[52], s[40])) k.emit(v_add_u32_e32(v[233], s[53], v[233])) - k.emit(s_mul_i32(s[52], 32, s[37])) + k.emit(s_mul_i32(s[52], 32, s[40])) k.emit(v_add_u32_e32(v[234], s[52], v[233])) k.emit(s_mov_b32(s[61], LIT, 128)) k.emit(s_mov_b32(s[62], LIT, 2048)) @@ -2510,18 +2463,18 @@ def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k.emit(ds_read_b32(v[201], v[224], v[0], v[0], 0, 0, 1)) k.emit(ds_read_b32(v[202], v[224], v[0], v[0], 0, 0, 2)) k.emit(ds_read_b32(v[203], v[224], v[0], v[0], 0, 0, 3)) - k.emit(s_lshl_b32(s[40], s[40], 1)) + k.emit(s_lshl_b32(s[36], s[36], 1)) k.emit(s_mul_i32(s[52], s[47], LIT, 256)) - k.emit(s_mul_hi_u32(s[53], s[52], s[40])) + k.emit(s_mul_hi_u32(s[53], s[52], s[36])) k.emit(s_add_u32(s[5], s[5], s[53])) - k.emit(s_mul_i32(s[53], s[52], s[40])) + k.emit(s_mul_i32(s[53], s[52], s[36])) k.emit(s_add_u32(s[4], s[4], s[53])) k.emit(s_addc_u32(s[5], 0, s[5])) k.emit(s_sub_i32(s[52], s[43], s[52])) - k.emit(s_mul_i32(s[52], s[52], s[40])) + k.emit(s_mul_i32(s[52], s[52], s[36])) k.emit(s_mov_b32(s[6], s[52])) k.emit(v_and_b32_e64(v[235], v[0], 15)) - k.emit(v_mul_lo_u32(v[235], v[235], s[40])) + k.emit(v_mul_lo_u32(v[235], v[235], s[36])) k.emit(v_lshrrev_b32_e32(v[4], 5)) k.emit(v_mul_i32_i24_e32(v[4], 16, v[4])) k.emit(v_add_u32_e32(v[235], v[4], v[235])) @@ -2529,12 +2482,12 @@ def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k.emit(v_and_b32_e32(v[4], 1, v[4])) k.emit(v_mul_i32_i24_e32(v[4], 32, v[4])) k.emit(v_add_u32_e32(v[235], v[4], v[235])) - k.emit(s_mul_i32(s[52], s[46], LIT, 256)) - k.emit(s_mul_i32(s[53], s[49], 64)) + k.emit(s_mul_i32(s[52], s[49], LIT, 256)) + k.emit(s_mul_i32(s[53], s[46], 64)) k.emit(s_add_i32(s[52], s[52], s[53])) k.emit(s_lshl_b32(s[52], s[52], 1)) k.emit(v_add_u32_e32(v[235], s[52], v[235])) - k.emit(s_mul_i32(s[53], s[40], 16)) + k.emit(s_mul_i32(s[53], s[36], 16)) for i in range(15): k.emit(v_add_u32_e64(v[236 + i * 1], v[235 + i * 1], s[53])) k.emit(s_mov_b32(s[50], 0)) @@ -2543,7 +2496,7 @@ def build_kernel(M: int, N: int, K: int, tile_m: int, tile_n: int): k.emit(s_cmp_lt_u32(LIT, s[51], 512 + i * -256)) k.emit(s_cselect_b32(s[61 + i * 1], s[61 + i * 1], 0)) k.emit(s_cselect_b32(s[63 + i * 1], s[63 + i * 1], 0)) - k.emit(s_cmp_lt_i32(s[49], 2)) + k.emit(s_cmp_lt_i32(s[46], 2)) k.emit(s_cbranch_scc0(1367), target='L2_25B8') k.label('L2_105C') k.emit(s_waitcnt(122))