Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
9ca026e
[jit][moe] prefill fp8 placeholder version
luocheng25 May 12, 2026
15ec3aa
[jit][moe] moe_2stage_down support fp8 ptpc
May 12, 2026
21521c1
[jit][moe] moe_2stage_gateup support fp8 ptpc
luocheng25 May 13, 2026
2ffbfdb
[jit][moe] add profiling info per kernel
luocheng25 May 13, 2026
7c18c22
[jit][moe] fix accuracy bug in down_kernel
May 13, 2026
1f57749
[jit][moe] fix crash bug in moe_2stage_down
May 13, 2026
5af6908
[jit][moe] add ref for moe_2stage_gateup
luocheng25 May 13, 2026
ee99a84
[jit][moe] support dynamic block-size-M to moe_2stage_down
May 14, 2026
ff4697d
[jit][moe] use store_dw4 + reduce instead of atomic_add
May 14, 2026
78ff984
[jit][moe] support xcd swizzle
luocheng25 May 14, 2026
ab0bba0
[jit][moe] allow TILE_M=256 in moe_2stage_down
May 15, 2026
85da1f8
[jit][moe] support dynamic block-size-M to moe_2stage_gateup
luocheng25 May 15, 2026
5cf98c1
[jit][moe] TILE_M=64 instead of 128, finetune down
May 15, 2026
f71a5d8
[jit][moe] fix incorrect dynamic block-size-M selection for gateup
luocheng25 May 18, 2026
2c03a58
[jit][moe] refactor test_moe, call ref when prec_fp8_t
tingqli May 20, 2026
c3006fd
[jit][moe] support prec_fp8_t in prefill
tingqli May 20, 2026
8dab15b
[jit][moe] support pref_fp_t in decoding
luocheng25 May 20, 2026
e4c8c13
[jit][moe] default use preshuffle on for aiter
luocheng25 May 20, 2026
7b3c6cb
[jit][moe] dynamic schedule gateup
luocheng25 May 20, 2026
45ed5b2
[jit][moe] dynamic schedule down
luocheng25 May 21, 2026
a23151c
[jit][moe] refactor to support dyn and static schedule
luocheng25 May 21, 2026
05939a8
[jit][moe] test case for qwen3.5 & hunyuan
luocheng25 May 25, 2026
65a0262
[jit][moe] fix using sgpr in v_mul_u32_u24 to compute prefetch offset
luocheng25 May 28, 2026
8a58dea
[jit][moe] fix out of bound access expert_id if dyn scheduler is true
luocheng25 Jun 2, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
91 changes: 62 additions & 29 deletions src/contrib/common/gemm.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,7 +134,8 @@ def __init__(self, J,
mfma_MN:int,
wave_size:list,
wave_cnt:list,
K, N):
K, N,
is_fp8=False):
self.K = K
self.N = N

Expand All @@ -144,7 +145,11 @@ def __init__(self, J,
assert wave_size_M % mfma_MN == 0
assert wave_size_N % mfma_MN == 0
# *2 for 8-bf16/fp16 so DWORDx4 lane-size can be used
self.mfma_K = (2*8 if mfma_MN == 32 else 2*16)
if is_fp8:
assert mfma_MN == 16, 'fp8 only support 16x16x16'
self.mfma_K = 2 * 32
else:
self.mfma_K = (2*8 if mfma_MN == 32 else 2*16)

# number of C/D regs per wave
wave_nCM = wave_size_M // mfma_MN
Expand All @@ -171,18 +176,25 @@ def __init__(self, J,
self.wave_nCM = wave_nCM
self.wave_nCN = wave_nCN
self.wave_nCK = wave_nCK
self.is_fp8 = is_fp8
if is_fp8:
self.sizeof_a = 1
self.sizeof_b = 1
else:
self.sizeof_a = J.sizeof_bf16
self.sizeof_b = J.sizeof_bf16

def run(self, loaderA, loaderB, buff_c, M, debug_warp, skip_load):
J = self.J

LDSA_size = self.wg_M * self.wg_K * J.sizeof_bf16
LDSB_size = self.wg_N * self.wg_K * J.sizeof_bf16
LDSA_size = self.wg_M * self.wg_K * self.sizeof_a
LDSB_size = self.wg_N * self.wg_K * self.sizeof_b
ldsA = J.alloc_lds(LDSA_size)
ldsB = J.alloc_lds(LDSB_size)

# prefetch in memory-coalescing way
# each lane prefetch DWORDx4 which is 8xhalf
num_lanes_per_row = J.div(J.sizeof_bf16 * self.wg_K, J.sizeof_DW4)
num_lanes_per_row = J.div(self.sizeof_a * self.wg_K, J.sizeof_DW4)
dw4_prefetch_MN = self.wave_cnt * J.div(64, num_lanes_per_row)
assert dw4_prefetch_MN >= 1
print(f"{self.wg_M=} {num_lanes_per_row=} {dw4_prefetch_MN=}")
Expand All @@ -194,38 +206,38 @@ def swizzle(row, col):

# each swizzle generates a new vaddr pattern, precompute all of them
# each wave reads its own part
ds_readA_vaddr = J.gpr(self.wave_nCM, self.wave_nCK, "vu32")
ds_readB_vaddr = J.gpr(self.wave_nCN, self.wave_nCK, "vu32")
ds_readA_vaddr = J.gpr(1, self.wave_nCK, "vu32")
ds_readB_vaddr = J.gpr(1, self.wave_nCK, "vu32")
# wave location
warp_id_m = J.warp_id // self.wave_cnt_N
warp_id_n = J.warp_id % self.wave_cnt_N
warp_offset_m = warp_id_m * self.wave_size_M
warp_offset_n = warp_id_n * self.wave_size_N
for m in range(self.wave_nCM):
for m in range(1):
for k in range(self.wave_nCK):
row = J.lane_id % self.mfma_MN + warp_offset_m + m*self.mfma_MN
col = J.lane_id // self.mfma_MN + (k * self.mfma_K * J.sizeof_bf16) // J.sizeof_DW4
col = J.lane_id // self.mfma_MN + (k * self.mfma_K * self.sizeof_a) // J.sizeof_DW4
swizzle_col = swizzle(row, col) % (num_lanes_per_row)
ds_readA_vaddr[m, k] = J.gpr((row * (self.wg_K * J.sizeof_bf16)) + swizzle_col*(J.sizeof_DW4))
ds_readA_vaddr[m, k] = J.gpr((row * (self.wg_K * self.sizeof_a)) + swizzle_col*(J.sizeof_DW4))

for n in range(self.wave_nCN):
for n in range(1):
for k in range(self.wave_nCK):
row = J.lane_id % self.mfma_MN + warp_offset_n + n*self.mfma_MN
col = J.lane_id // self.mfma_MN + (k * self.mfma_K * J.sizeof_bf16) // J.sizeof_DW4
col = J.lane_id // self.mfma_MN + (k * self.mfma_K * self.sizeof_b) // J.sizeof_DW4
swizzle_col = swizzle(row, col) % (num_lanes_per_row)
ds_readB_vaddr[n, k] = J.gpr((row * (self.wg_K * J.sizeof_bf16)) + swizzle_col*(J.sizeof_DW4))
ds_readB_vaddr[n, k] = J.gpr((row * (self.wg_K * self.sizeof_b)) + swizzle_col*(J.sizeof_DW4))

Creg_size = (self.mfma_MN * self.mfma_MN)//64
mfma_C = self.J.gpr(self.wave_nCM, self.wave_nCN, Creg_size, f"af32")
ABReg_size = (self.mfma_MN * self.mfma_K * 2//4)//64
ABReg_size = (self.mfma_MN * self.mfma_K * self.sizeof_a//4)//64
mfma_A = J.gpr(self.wave_nCM, self.wave_nCK, ABReg_size, "vbf16x2")
mfma_B = J.gpr(self.wave_nCN, self.wave_nCK, ABReg_size, "vbf16x2")

def ds_readA(m, k):
J.ds_read_b128(mfma_A[m,k], ds_readA_vaddr[m, k], mod=f"offset:{ldsA}") # vaddr, vdata offset gds
J.ds_read_b128(mfma_A[m,k], ds_readA_vaddr[0, k], mod=f"offset:{ldsA + m*self.mfma_MN*self.wg_K * self.sizeof_a}") # vaddr, vdata offset gds

def ds_readB(n, k):
J.ds_read_b128(mfma_B[n,k], ds_readB_vaddr[n,k], mod=f"offset:{ldsB}") # vaddr, vdata offset gds
J.ds_read_b128(mfma_B[n,k], ds_readB_vaddr[0,k], mod=f"offset:{ldsB+n*self.mfma_MN*self.wg_K * self.sizeof_b}") # vaddr, vdata offset gds

J.debug_setup((J.blockIdx.x[0] == 0) & (J.blockIdx.y[0] == 0) & (J.warp_id == debug_warp))

Expand Down Expand Up @@ -256,7 +268,7 @@ def ds_readB(n, k):
mfma_C[:] = 0

# prelog 1: ds_write + prefetch
k_offset[0] = 0 if skip_load else (k_offset[0] + self.wg_K * J.sizeof_bf16)
k_offset[0] = 0 if skip_load else (k_offset[0] + self.wg_K * self.sizeof_a)
loaderA.reset_offset(k_offset)
loaderB.reset_offset(k_offset)

Expand Down Expand Up @@ -297,6 +309,8 @@ def ds_readB(n, k):
16:("v_mfma_f32_16x16x32_bf16",16) if is_cdna4 else ("v_mfma_f32_16x16x16_bf16",16),
32:("v_mfma_f32_32x32x16_bf16",32) if is_cdna4 else ("v_mfma_f32_32x32x8_bf16",32)
}
if self.is_fp8:
mfma_info[16] = ("v_mfma_f32_16x16x32_fp8_fp8", 16)
mfma_name = mfma_info[self.mfma_MN][0]
mfma_cycles = mfma_info[self.mfma_MN][1]

Expand Down Expand Up @@ -327,11 +341,29 @@ def mfma_generator(k):
mfma_C[m,n])
cur_k = J.gpr("su32", 0)
k_loop_cnt = self.K//self.wg_K
mfma_cnt_dict = {
(256, 256) : {
'ds_read': 2,
'ds_write': 2,
'prefetch': 8
},
(128, 256) : {
'ds_read': 2,
'ds_write': 3,
'prefetch': 3
},
(64, 256) : {
'ds_read': 1,
'ds_write': 2,
'prefetch': 1
}
}
mfma_cnt = mfma_cnt_dict.get((self.wg_M, self.wg_N), mfma_cnt_dict[(256, 256)])

with J.While(cur_k[0] < k_loop_cnt):
#for unroll in range(k_loop_cnt):

k_offset[0] = 0 if skip_load else (k_offset[0] + self.wg_K * J.sizeof_bf16)
k_offset[0] = 0 if skip_load else (k_offset[0] + self.wg_K * self.sizeof_a)
loaderA.reset_offset(k_offset)
loaderB.reset_offset(k_offset)

Expand All @@ -344,12 +376,13 @@ def mfma_generator(k):
for k in range(self.wave_nCK//2, self.wave_nCK):
for m in range(self.wave_nCM):
ds_readA(m, k)
J.emit(mfma0, 16*2)
J.emit(mfma0, 16*mfma_cnt['ds_read'])
for n in range(self.wave_nCN):
ds_readB(n, k)
J.emit(mfma0, 16*2)
J.emit(mfma0, 16*mfma_cnt['ds_read'])

# ensure all waves has been finished reading LDS, so ds_write can overwrite it
J.emit(mfma0, 16*2)
J.s_waitcnt(mod=f"lgkmcnt(0)")
J.s_barrier()

Expand All @@ -358,17 +391,17 @@ def mfma_generator(k):
J.emit([mfma0, mfma1], 16)
J.s_waitcnt(mod=f"vmcnt({num_prefetch_N + num_prefetch_M - 1})")
loaderA.ds_write(r, ldsA)
J.emit([mfma0, mfma1], 16*2)
J.emit([mfma0, mfma1], 16*mfma_cnt['ds_write'])
loaderA.prefetch(r)
J.emit([mfma0, mfma1], 16*8)
J.emit([mfma0, mfma1], 16*mfma_cnt['prefetch'])

for r in range(num_prefetch_N):
J.emit([mfma0, mfma1], 16)
J.s_waitcnt(mod=f"vmcnt({num_prefetch_N + num_prefetch_M - 1})")
loaderB.ds_write(r, ldsB)
J.emit([mfma0, mfma1], 16*2)
J.emit([mfma0, mfma1], 16*mfma_cnt['ds_write'])
loaderB.prefetch(r)
J.emit([mfma0, mfma1], 16*8)
J.emit([mfma0, mfma1], 16*mfma_cnt['prefetch'])

# enure mfma0 finished using part0 of mfma_A/mfma_B, before ds_read0 overwrites them
# (most likely already empty and some part of mfma1 has been consumed)
Expand All @@ -379,10 +412,10 @@ def mfma_generator(k):
for k in range(0,self.wave_nCK//2):
for m in range(self.wave_nCM):
ds_readA(m, k)
J.emit([mfma1], 16*2)
J.emit([mfma1], 16*mfma_cnt['ds_read'])
for n in range(self.wave_nCN):
ds_readB(n, k)
J.emit([mfma1], 16*2)
J.emit([mfma1], 16*mfma_cnt['ds_read'])
J.emit(mfma1)
cur_k[0] += 1

Expand Down Expand Up @@ -456,7 +489,7 @@ def mfma_generator(k):
for n in range(self.wave_nCN):
row = J.lane_id % self.mfma_MN + warp_offset_m + m*self.mfma_MN
col = J.lane_id // self.mfma_MN + n * (self.mfma_MN * J.sizeof_fp32 // J.sizeof_DW4)
voffset = J.gpr(row * (N*J.sizeof_fp32) + warp_offset_n*J.sizeof_fp32 + col*J.sizeof_DW4)
voffset = J.gpr(row * (self.N*J.sizeof_fp32) + warp_offset_n*J.sizeof_fp32 + col*J.sizeof_DW4)
if self.mfma_MN == 16:
buff_c.store_dwordx4(mfma_C[m,n], voffset, 0)
elif self.mfma_MN == 32:
Expand Down Expand Up @@ -503,7 +536,7 @@ def gemm_kernel(J, K, N, M01, GroupNum,
gemm.wg_N, J.sizeof_bf16*gemm.wg_K, stride_bytes,
total_wave_cnt, swizzle_row_div, skip_load)
else:
loaderB = MFMA_DW4Loader_preshuffled(J, pB, actual_wg_M * K * J.sizeof_bf16, mfma_MN,
loaderB = MFMA_DW4Loader_preshuffled(J, pB, gemm.wg_N * K * J.sizeof_bf16, mfma_MN,
gemm.wg_N, J.sizeof_bf16*gemm.wg_K, stride_bytes,
total_wave_cnt, swizzle_row_div, skip_load)

Expand Down Expand Up @@ -665,7 +698,7 @@ def test_gemm(mfma_MN, wave_size, wave_cnt, A_preshuffled = False, B_preshuffled
#assert 0
#test_gemm(32, [128, 128], [2, 2], A_preshuffled = False, B_preshuffled = False)
#test_gemm(16, [128, 128], [2, 2], A_preshuffled = False, B_preshuffled = False)
test_gemm(16, [64, 64], [2, 2], A_preshuffled = False, B_preshuffled = True)
test_gemm(16, [32, 128], [2, 2], A_preshuffled = False, B_preshuffled = True)
#test_gemm(16, [128, 128], [2, 2], A_preshuffled = True, B_preshuffled = True)
#test_gemm(32, [128, 128], [2, 2], A_preshuffled = True, B_preshuffled = True)
assert 0
Expand Down
32 changes: 23 additions & 9 deletions src/contrib/common/gemm_splitk.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,10 +22,13 @@ def gemm_splitk(J:JIT,
BLOCK_TILE_SIZE_N = 32,
BLOCK_TILE_SIZE_M = 16,
USE_FP4_SHUFFLE_WEIGHT=False,
fp8_ptpc=True
quant_type_str='no',
):
assert BLOCK_TILE_SIZE_M % 16 == 0, f'BLOCK_TILE_SIZE_M must be multiple of 16, current {BLOCK_TILE_SIZE_M=}'
assert BLOCK_TILE_SIZE_N % 32 == 0, f'BLOCK_TILE_SIZE_N must be multiple of 32, current {BLOCK_TILE_SIZE_N=}'
fp8_ptpc = True if quant_type_str == 'per_Token' else False
fp8_per_tensor = True if quant_type_str == 'per_Tensor' else False
fp8_block128 = True if quant_type_str == 'per_1x128' else False
sizeof_f32 = 4
sizeof_bf16 = 2
sizeof_w = sizeof_bf16 if weight_dtype == torch.bfloat16 else 1
Expand Down Expand Up @@ -77,11 +80,12 @@ def gemm_splitk(J:JIT,
k_scale_n = div_up(div_up(K, 32), 8) // num_split_k
v_w_scale = J.gpr(B_horz // 2, k_scale_n, 'vf32', align=4)
k_scale_n_next_read_idx = 0
elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and fp8_ptpc:
v_w_scale = J.gpr(B_horz, 4, 'vf32')
# for n in range(B_horz):
# J.global_load_dwordx4(v_w_scale[n], voffset_scale[n], p_w_scale)
elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and not fp8_ptpc:
elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz):
if fp8_per_tensor:
v_w_scale = J.gpr(2, 'vf32')
elif fp8_ptpc:
v_w_scale = J.gpr(B_horz, 4, 'vf32')
else:
QUAN_BLK_SZ=128
#[ping-pong, ngroups]
v_w_scale = J.gpr(2, 2, 'vf32')
Expand All @@ -106,7 +110,7 @@ def load_gen(pp_reg_id, k=None):
J.global_load_dword(v_w_scale[n, k_scale_n_next_read_idx], voffset_scale[n], p_w_scale, mod=f'offset:{k_scale_n_next_read_idx * 64 * sizeof_f32}')
k_scale_n_next_read_idx += 1
k_scale_wip = B_horz // 2
elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and not fp8_ptpc:
elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and fp8_block128:
J.global_load_dword(v_w_scale[pp_reg_id, 0], voffset_scale[0], p_w_scale, mod=f'offset:{(k+1) * k_step_wg // QUAN_BLK_SZ * sizeof_f32}')
J.global_load_dword(v_w_scale[pp_reg_id, 1], voffset_scale[1], p_w_scale, mod=f'offset:{(k+1) * k_step_wg // QUAN_BLK_SZ * sizeof_f32}')
k_scale_wip = 2
Expand Down Expand Up @@ -165,7 +169,7 @@ def delayed_fma():
next(gen)
next(gen)

elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and fp8_ptpc:
elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and (fp8_ptpc or fp8_per_tensor):
v_w_f32 = J.gpr(2, 2, 2, 'vf32', align=4)
v_w_bf16 = J.gpr(B_horz, 2, 2, 'vf32', align=4)
# kl = 16 would be divided into 2 steps. Each accumulate 8 in K dimension.
Expand Down Expand Up @@ -198,7 +202,7 @@ def delayed_fma():
yield J.v_mfma_f32_16x16x16_bf16(C_reg[n, m], v_w_bf16[n, 0], A_reg[pp_reg_id, m, i, 0], C_reg[n, m])
yield J.v_mfma_f32_16x16x16_bf16(C_reg[n, m], v_w_bf16[n, 1], A_reg[pp_reg_id, m, i, 1], C_reg[n, m])

elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and not fp8_ptpc:
elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and fp8_block128:
# cdna3 path:
if is_cdna4 == False:
v_w_f32 = J.gpr(2, 2, 2, 'vf32', align=4)
Expand Down Expand Up @@ -289,6 +293,9 @@ def tail(pp_reg_id, k=None):
for n in range(B_horz):
J.global_load_dwordx4(v_w_scale[n], voffset_scale[n], p_w_scale)
J.s_waitcnt(mod=f"vmcnt({B_horz})")
elif (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and fp8_per_tensor:
J.global_load_dword(v_w_scale[0], voffset_scale[0], p_w_scale)
J.s_waitcnt(mod=f"vmcnt(1)")
else:
J.s_waitcnt(mod=f"vmcnt(0)")

Expand Down Expand Up @@ -344,3 +351,10 @@ def tail(pp_reg_id, k=None):
for m in range(A_vert):
J.v_pk_mul_f32(C_reg[n, m, 0:1], C_reg[n, m, 0:1], v_w_scale[n, 0:1])
J.v_pk_mul_f32(C_reg[n, m, 2:3], C_reg[n, m, 2:3], v_w_scale[n, 2:3])
if (weight_dtype == torch.float8_e4m3fn or weight_dtype == torch.float8_e4m3fnuz) and fp8_per_tensor:
J.s_waitcnt(mod=f"vmcnt(0)")
v_w_scale[1] = v_w_scale[0]
for n in range(B_horz):
for m in range(A_vert):
J.v_pk_mul_f32(C_reg[n, m, 0:1], C_reg[n, m, 0:1], v_w_scale)
J.v_pk_mul_f32(C_reg[n, m, 2:3], C_reg[n, m, 2:3], v_w_scale)
Loading