submission 696235
dc1312 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 603 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-696235?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, mxfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:d704331211ecf01987023abce2c5684c5d846b1b5805b2ddc8d9962efdfdd3ff
license declaredunknown
license concludedunknown
authorsdc1312
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
fp4 = torch.empty((M, dep // 2), dtype=torch.uint8, device=gemm_out.device)fused-epilogue
- Inline fused_silu_mul_quant with cached buffersnum-warps = 1
num_warps=1,stages = 2
num_warps=nw1, num_stages=2,tile-m = 128
BLOCK_M=128, QB=32,tile-n = 128
SR_BN = 128Kernel source
submission.py603 lines
"""
v67: Grid-aware BN1 + inlined allocs + combined v60 optimizations.
- BN1=128 for dep=256 when total_sorted<2000 (small grid → need big CTAs)
- BN1=64 for dep=256 when total_sorted>=2000 (enough parallelism → more N-tiles)
- Inline fused_silu_mul_quant with cached buffers
- All v60 sort/scatter/buffer optimizations
"""
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
_BUF = {}
def _get_buf(key, numel, dtype, device):
if key not in _BUF or _BUF[key].numel() < numel or _BUF[key].dtype != dtype:
_BUF[key] = torch.empty(numel, dtype=dtype, device=device)
return _BUF[key][:numel]
@triton.jit
def _dynamic_mxfp4_quant_kernel(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m, stride_x_n,
stride_x_fp4_m, stride_x_fp4_n,
stride_bs_m, stride_bs_n,
M: tl.constexpr, N: tl.constexpr,
scaleN: tl.constexpr,
scaleM_pad: tl.constexpr,
scaleN_pad: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
SCALING_MODE: tl.constexpr,
SHUFFLE: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_x_m = tl.cast(stride_x_m, tl.int64)
stride_x_n = tl.cast(stride_x_n, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)
amax = tl.max(tl.abs(x), axis=1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
e = (qx >> 23) & 0xFF
m = qx & 0x7FFFFF
E8_BIAS: tl.constexpr = 127
E2_BIAS: tl.constexpr = 1
adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
e2m1_value = tl.reshape(e2m1_value, [BLOCK_SIZE, MXFP4_QUANT_BLOCK_SIZE // 2, 2])
evens, odds = tl.split(e2m1_value)
out_tensor = evens | (odds << 4)
out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 2)
out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
bs_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
bs_offs_n = pid_n
bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N)[None, :]
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)
def dynamic_mxfp4_quant(x):
M, N = x.shape
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
scaleN_valid = triton.cdiv(N, 32)
scaleN = triton.cdiv(scaleN_valid, 8) * 8
blockscale = torch.empty(
(triton.cdiv(M, 256) * 256, scaleN),
dtype=torch.uint8, device=x.device,
)
grid = (triton.cdiv(M, 128), scaleN_valid)
_dynamic_mxfp4_quant_kernel[grid](
x, x_fp4, blockscale,
*x.stride(), *x_fp4.stride(), *blockscale.stride(),
M=M, N=N, scaleN=scaleN_valid,
scaleM_pad=triton.cdiv(M, 32) * 32,
scaleN_pad=scaleN,
BLOCK_SIZE=128, MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0, SHUFFLE=False,
)
blockscale = blockscale[:M, :scaleN_valid].contiguous()
return (x_fp4, blockscale)
@triton.jit
def _fused_silu_mul_quant_kernel(
gemm_out_ptr, fp4_ptr, scale_ptr,
dep,
stride_gm, stride_gn,
stride_fp4_m, stride_fp4_n,
stride_sc_m, stride_sc_n,
M,
BLOCK_M: tl.constexpr,
QB: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_gm = tl.cast(stride_gm, tl.int64)
stride_gn = tl.cast(stride_gn, tl.int64)
stride_fp4_m = tl.cast(stride_fp4_m, tl.int64)
stride_fp4_n = tl.cast(stride_fp4_n, tl.int64)
m_offs = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
n_offs = pid_n * QB + tl.arange(0, QB)
m_mask = m_offs < M
n_mask = n_offs < dep
gate_offs = m_offs[:, None] * stride_gm + n_offs[None, :] * stride_gn
up_offs = m_offs[:, None] * stride_gm + (n_offs[None, :] + dep) * stride_gn
mask = m_mask[:, None] & n_mask[None, :]
gate = tl.load(gemm_out_ptr + gate_offs, mask=mask, other=0).to(tl.float32)
up = tl.load(gemm_out_ptr + up_offs, mask=mask, other=0).to(tl.float32)
x = (gate * tl.sigmoid(gate) * up).to(tl.bfloat16).to(tl.float32)
amax = tl.max(tl.abs(x), axis=1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_unb = tl.log2(amax).floor() - 2
scale_unb = tl.clamp(scale_unb, min=-127, max=127)
quant_scale = tl.exp2(-scale_unb)
qx = x * quant_scale
bs_e8m0 = scale_unb.to(tl.uint8) + 127
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
e = (qx >> 23) & 0xFF
m = qx & 0x7FFFFF
E8_BIAS: tl.constexpr = 127
E2_BIAS: tl.constexpr = 1
adj = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adj, m)
e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
e2m1_val = ((s >> 28) | e2m1_tmp).to(tl.uint8)
e2m1_val = tl.reshape(e2m1_val, [BLOCK_M, QB // 2, 2])
evens, odds = tl.split(e2m1_val)
packed = evens | (odds << 4)
out_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
out_n = pid_n * QB // 2 + tl.arange(0, QB // 2)
out_offs = out_m[:, None] * stride_fp4_m + out_n[None, :] * stride_fp4_n
out_mask = (out_m < M)[:, None] & (out_n < (dep // 2))[None, :]
tl.store(fp4_ptr + out_offs, packed, mask=out_mask)
sc_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
sc_offs = sc_m[:, None] * stride_sc_m + pid_n * stride_sc_n
tl.store(scale_ptr + sc_offs, bs_e8m0, mask=(sc_m < M)[:, None])
def fused_silu_mul_quant(gemm_out, dep):
M = gemm_out.shape[0]
fp4 = torch.empty((M, dep // 2), dtype=torch.uint8, device=gemm_out.device)
scaleN = triton.cdiv(dep, 32)
scale = torch.empty((M, scaleN), dtype=torch.uint8, device=gemm_out.device)
grid = (triton.cdiv(M, 128), scaleN)
_fused_silu_mul_quant_kernel[grid](
gemm_out, fp4, scale,
dep,
gemm_out.stride(0), gemm_out.stride(1),
fp4.stride(0), fp4.stride(1),
scale.stride(0), scale.stride(1),
M,
BLOCK_M=128, QB=32,
)
return fp4, scale
@triton.jit
def _remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
tall_xcds = GRID_MN % NUM_XCDS
tall_xcds = tl.where(tall_xcds == 0, NUM_XCDS, tall_xcds)
xcd = pid % NUM_XCDS
local_pid = pid // NUM_XCDS
new_pid = tl.where(
xcd < tall_xcds,
xcd * pids_per_xcd + local_pid,
tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid,
)
return new_pid
# ---- Triton counting sort kernels ----
@triton.jit
def _moe_count_kernel(topk_ids_ptr, counts_ptr, total, BLOCK: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < total
ids = tl.load(topk_ids_ptr + offs, mask=mask, other=0).to(tl.int32)
tl.atomic_add(counts_ptr + ids, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)
@triton.jit
def _moe_offsets_kernel(
counts_ptr, offsets_ptr, cum_blocks_ptr,
E,
BM: tl.constexpr,
BLOCK_E: tl.constexpr,
):
idx = tl.arange(0, BLOCK_E)
mask = idx < E
counts = tl.load(counts_ptr + idx, mask=mask, other=0).to(tl.int64)
token_offsets = tl.cumsum(counts, axis=0)
tl.store(offsets_ptr + 1 + idx, token_offsets, mask=mask)
n_blocks = tl.where(counts > 0, (counts + BM - 1) // BM, tl.zeros([BLOCK_E], dtype=tl.int64))
block_offsets = tl.cumsum(n_blocks, axis=0)
tl.store(cum_blocks_ptr + 1 + idx, block_offsets, mask=mask)
@triton.jit
def _moe_scatter_kernel(
topk_ids_ptr, topk_weights_ptr,
sorted_token_idx_ptr, sorted_weights_ptr,
offsets_ptr, write_counts_ptr,
reverse_idx_ptr, token_counts_ptr,
topk, total,
BLOCK: tl.constexpr,
):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < total
ids = tl.load(topk_ids_ptr + offs, mask=mask, other=0).to(tl.int32)
weights = tl.load(topk_weights_ptr + offs, mask=mask, other=0.0)
token_idx = (offs // topk).to(tl.int32)
old = tl.atomic_add(write_counts_ptr + ids, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)
base = tl.load(offsets_ptr + ids, mask=mask, other=0).to(tl.int32)
write_pos = (base + old).to(tl.int64)
tl.store(sorted_token_idx_ptr + write_pos, token_idx, mask=mask)
tl.store(sorted_weights_ptr + write_pos, weights, mask=mask)
tk_old = tl.atomic_add(token_counts_ptr + token_idx, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)
rev_pos = token_idx.to(tl.int64) * topk + tk_old.to(tl.int64)
tl.store(reverse_idx_ptr + rev_pos, write_pos.to(tl.int32), mask=mask)
@triton.jit
def _scatter_reduce_kernel(
src_ptr, reverse_idx_ptr, dst_ptr,
M, dh,
stride_sm, stride_sn,
topk,
BN: tl.constexpr, TOPK: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_sm = tl.cast(stride_sm, tl.int64)
stride_sn = tl.cast(stride_sn, tl.int64)
cols = pid_n * BN + tl.arange(0, BN)
col_mask = cols < dh
acc = tl.zeros([BN], dtype=tl.float32)
rev_base = pid_m * topk
for k in range(TOPK):
sorted_pos = tl.load(reverse_idx_ptr + rev_base + k).to(tl.int64)
vals = tl.load(src_ptr + sorted_pos * stride_sm + cols * stride_sn, mask=col_mask, other=0.0)
acc += vals
tl.store(dst_ptr + pid_m.to(tl.int64) * dh + cols, acc.to(tl.bfloat16), mask=col_mask)
# ---- Batched MoE GEMM kernel (supports variable BM) ----
@triton.jit
def _batched_moe_gemm_fp4(
a_ptr, b_ptr, c_ptr,
a_sc_ptr, b_sc_ptr,
token_idx_ptr, sorted_weights_ptr,
cum_blocks_ptr, expert_offsets_ptr,
E, N, K_half,
stride_am, stride_ak,
stride_be, stride_bn, stride_bk,
stride_cm, stride_cn,
stride_asm, stride_ask,
stride_bse, stride_bsn, stride_bsk,
max_m_blocks, total_tokens,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
EVEN_K: tl.constexpr, INDIRECT: tl.constexpr,
PRESHUFFLE: tl.constexpr, APPLY_WEIGHTS: tl.constexpr,
SEARCH_ITERS: tl.constexpr,
):
SG: tl.constexpr = 32
stride_am = tl.cast(stride_am, tl.int64)
stride_ak = tl.cast(stride_ak, tl.int64)
stride_be = tl.cast(stride_be, tl.int64)
stride_bn = tl.cast(stride_bn, tl.int64)
stride_bk = tl.cast(stride_bk, tl.int64)
stride_cm = tl.cast(stride_cm, tl.int64)
stride_cn = tl.cast(stride_cn, tl.int64)
stride_asm = tl.cast(stride_asm, tl.int64)
stride_ask = tl.cast(stride_ask, tl.int64)
stride_bse = tl.cast(stride_bse, tl.int64)
stride_bsn = tl.cast(stride_bsn, tl.int64)
stride_bsk = tl.cast(stride_bsk, tl.int64)
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_asm > 0)
tl.assume(stride_ask > 0)
tl.assume(stride_bsn > 0)
tl.assume(stride_bsk > 0)
pid_raw = tl.program_id(0)
nn = tl.cdiv(N, BN)
grid_mn = max_m_blocks * nn
pid = _remap_xcd(pid_raw, grid_mn)
pid_n = pid // max_m_blocks
pid_mb = pid % max_m_blocks
if pid_mb >= max_m_blocks:
return
lo = tl.cast(0, tl.int64)
hi = tl.cast(E, tl.int64)
pid_mb_64 = tl.cast(pid_mb, tl.int64)
for _ in range(SEARCH_ITERS):
mid = (lo + hi + 1) // 2
v = tl.load(cum_blocks_ptr + mid)
cond = v <= pid_mb_64
lo = tl.where(cond, mid, lo)
hi = tl.where(cond, hi, mid - 1)
expert_id = lo
if expert_id >= E:
return
block_start = tl.load(cum_blocks_ptr + expert_id)
local_block = pid_mb_64 - block_start
expert_token_start = tl.load(expert_offsets_ptr + expert_id)
expert_token_end = tl.load(expert_offsets_ptr + expert_id + 1)
row_start = expert_token_start + local_block * BM
if row_start >= expert_token_end:
return
sorted_rows = row_start + tl.arange(0, BM).to(tl.int64)
row_mask = (sorted_rows < total_tokens) & (sorted_rows < expert_token_end)
if INDIRECT:
a_rows = tl.load(token_idx_ptr + sorted_rows, mask=row_mask, other=0).to(tl.int64)
else:
a_rows = sorted_rows
cols = pid_n * BN + tl.arange(0, BN)
col_mask = cols < N
hk = tl.arange(0, BK // 2)
ap = a_ptr + a_rows[:, None] * stride_am + hk[None, :] * stride_ak
ks = tl.arange(0, BK // SG)
asp = a_sc_ptr + a_rows[:, None] * stride_asm + ks[None, :] * stride_ask
b_base = b_ptr + expert_id * stride_be
bsc_base = b_sc_ptr + expert_id * stride_bse
if PRESHUFFLE:
offs_bn_shuf = pid_n * (BN // 16) + tl.arange(0, BN // 16)
offs_k_shuf = tl.arange(0, (BK // 2) * 16)
bp = b_base + offs_bn_shuf[:, None].to(tl.int64) * stride_bn + offs_k_shuf[None, :].to(tl.int64) * stride_bk
bsp = bsc_base + cols[:, None].to(tl.int64) * stride_bsn + ks[None, :] * stride_bsk
else:
bp = b_base + hk[:, None] * stride_bk + cols[None, :].to(tl.int64) * stride_bn
bsp = bsc_base + cols[:, None].to(tl.int64) * stride_bsn + ks[None, :] * stride_bsk
acc = tl.zeros((BM, BN), dtype=tl.float32)
nk = tl.cdiv(K_half, BK // 2)
if PRESHUFFLE:
for _ in range(nk):
asc = tl.load(asp, mask=row_mask[:, None], other=0)
bsc = tl.load(bsp, mask=col_mask[:, None], other=0)
a = tl.load(ap, mask=row_mask[:, None], other=0)
b_raw = tl.load(bp)
b = (b_raw
.reshape(1, BN // 16, BK // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BN, BK // 2)
.trans(1, 0))
acc = tl.dot_scaled(a, asc, "e2m1", b, bsc, "e2m1", acc)
ap += (BK // 2) * stride_ak
asp += (BK // SG) * stride_ask
bp += (BK // 2) * 16 * stride_bk
bsp += (BK // SG) * stride_bsk
elif EVEN_K:
for _ in range(nk):
asc = tl.load(asp, mask=row_mask[:, None], other=0)
bsc = tl.load(bsp, mask=col_mask[:, None], other=0)
a = tl.load(ap, mask=row_mask[:, None], other=0)
b = tl.load(bp, mask=col_mask[None, :], other=0)
acc = tl.dot_scaled(a, asc, "e2m1", b, bsc, "e2m1", acc)
ap += (BK // 2) * stride_ak
bp += (BK // 2) * stride_bk
asp += (BK // SG) * stride_ask
bsp += (BK // SG) * stride_bsk
else:
K_rem = K_half
for _ in range(nk):
asc = tl.load(asp, mask=row_mask[:, None], other=0)
bsc = tl.load(bsp, mask=col_mask[:, None], other=0)
a = tl.load(ap, mask=row_mask[:, None] & (hk[None, :] < K_rem), other=0)
b = tl.load(bp, mask=(hk[:, None] < K_rem) & col_mask[None, :], other=0)
acc = tl.dot_scaled(a, asc, "e2m1", b, bsc, "e2m1", acc)
ap += (BK // 2) * stride_ak
bp += (BK // 2) * stride_bk
asp += (BK // SG) * stride_ask
bsp += (BK // SG) * stride_bsk
K_rem -= BK // 2
if APPLY_WEIGHTS:
w = tl.load(sorted_weights_ptr + sorted_rows, mask=row_mask, other=0)
acc = acc * w[:, None]
cp = c_ptr + sorted_rows[:, None] * stride_cm + cols[None, :].to(tl.int64) * stride_cn
tl.store(cp, acc, mask=row_mask[:, None] & col_mask[None, :])
def custom_kernel(data):
(
hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config,
) = data
dh = config["d_hidden"]
dep = config["d_expert_pad"]
dhp = config["d_hidden_pad"]
M = hidden_states.shape[0]
E = gate_up_weight.shape[0]
topk = topk_ids.shape[1]
dev = hidden_states.device
total_sorted = M * topk
if total_sorted == 0:
return torch.zeros((M, dh), dtype=torch.bfloat16, device=dev)
guw = gate_up_weight_shuffled.view(torch.uint8).view(E, (2 * dep) // 16, (dhp // 2) * 16)
dw = down_weight_shuffled.view(torch.uint8).view(E, dhp // 16, (dep // 2) * 16)
guw_sc = gate_up_weight_scale.view(torch.uint8).reshape(E, 2 * dep, dhp // 32)
dw_sc = down_weight_scale.view(torch.uint8).reshape(E, dhp, dep // 32)
hs = F.pad(hidden_states, (0, dhp - dh)) if dhp > dh else hidden_states
hs_fp4, hs_sc = dynamic_mxfp4_quant(hs)
estimated_per_expert = total_sorted / E
if estimated_per_expert >= 64 and dep <= 512:
BM_GEMM = 128
elif estimated_per_expert >= 64:
BM_GEMM = 64
elif estimated_per_expert >= 4:
BM_GEMM = 32
else:
BM_GEMM = 16
flat_ids = topk_ids.reshape(-1)
flat_weights = topk_weights.reshape(-1)
# Consolidated int32 sort buffers: [expert_counts(E) | write_counts(E) | token_counts(M)]
i32_sz = 2 * E + M
sort_i32 = _get_buf("si32", i32_sz, torch.int32, dev)
sort_i32.zero_()
expert_counts = sort_i32[:E]
write_counts = sort_i32[E:2*E]
token_counts = sort_i32[2*E:]
# Consolidated int64 sort buffers: [expert_offsets(E+1) | cum_blocks(E+1)]
i64_sz = 2 * (E + 1)
sort_i64 = _get_buf("si64", i64_sz, torch.int64, dev)
sort_i64.zero_()
expert_offsets = sort_i64[:E+1]
cum_blocks = sort_i64[E+1:]
sorted_token_idx = _get_buf("sti", total_sorted, torch.int32, dev)
sorted_weights = _get_buf("sw", total_sorted, flat_weights.dtype, dev)
reverse_idx = _get_buf("ri", M * topk, torch.int32, dev)
SORT_BLOCK = 256
_moe_count_kernel[(triton.cdiv(total_sorted, SORT_BLOCK),)](
flat_ids, expert_counts, total_sorted, BLOCK=SORT_BLOCK,
)
BLOCK_E = triton.next_power_of_2(E)
_moe_offsets_kernel[(1,)](
expert_counts, expert_offsets, cum_blocks,
E, BM=BM_GEMM, BLOCK_E=BLOCK_E,
num_warps=1,
)
_moe_scatter_kernel[(triton.cdiv(total_sorted, SORT_BLOCK),)](
flat_ids, flat_weights,
sorted_token_idx, sorted_weights,
expert_offsets, write_counts,
reverse_idx, token_counts,
topk, total_sorted, BLOCK=SORT_BLOCK,
)
max_m_blocks = min(total_sorted, (total_sorted + BM_GEMM - 1) // BM_GEMM + E)
si = 6 if E <= 64 else 9
nw1 = 4 if BM_GEMM >= 64 else 2
max_bk = 512 if BM_GEMM >= 64 else 1024
# BN1=128 for dep=256 when grid is small; BN1=64 when enough parallelism
BN1 = 128 if (dep <= 256 and BM_GEMM <= 32 and total_sorted < 2000) else 64
BK1 = 128
for bk in [1024, 512, 256, 128]:
if bk <= min(dhp, max_bk) and dhp % bk == 0:
BK1 = bk
break
even_k1 = (dhp % BK1) == 0
f32_sz = total_sorted * max(2 * dep, dhp)
gemm_f32 = _get_buf("gf32", f32_sz, torch.float32, dev)
gemm1_out = gemm_f32[:total_sorted * 2 * dep].view(total_sorted, 2 * dep)
nn1 = triton.cdiv(2 * dep, BN1)
_batched_moe_gemm_fp4[(max_m_blocks * nn1,)](
hs_fp4, guw, gemm1_out,
hs_sc, guw_sc,
sorted_token_idx, sorted_weights,
cum_blocks, expert_offsets,
E, 2 * dep, dhp // 2,
hs_fp4.stride(0), hs_fp4.stride(1),
guw.stride(0), guw.stride(1), guw.stride(2),
gemm1_out.stride(0), gemm1_out.stride(1),
hs_sc.stride(0), hs_sc.stride(1),
guw_sc.stride(0), guw_sc.stride(1), guw_sc.stride(2),
max_m_blocks, total_sorted,
BM=BM_GEMM, BN=BN1, BK=BK1, EVEN_K=even_k1, INDIRECT=True,
PRESHUFFLE=True, APPLY_WEIGHTS=False,
SEARCH_ITERS=si,
num_warps=nw1, num_stages=2,
)
# Inline fused_silu_mul_quant with cached buffers
inter_fp4 = _get_buf("ifp4", total_sorted * dep // 2, torch.uint8, dev).view(total_sorted, dep // 2)
scaleN_inter = triton.cdiv(dep, 32)
inter_sc = _get_buf("isc", total_sorted * scaleN_inter, torch.uint8, dev).view(total_sorted, scaleN_inter)
_fused_silu_mul_quant_kernel[(triton.cdiv(total_sorted, 128), scaleN_inter)](
gemm1_out, inter_fp4, inter_sc,
dep,
gemm1_out.stride(0), gemm1_out.stride(1),
inter_fp4.stride(0), inter_fp4.stride(1),
inter_sc.stride(0), inter_sc.stride(1),
total_sorted,
BLOCK_M=128, QB=32,
)
# GEMM2: BN2=256 for dep=256 (BK=256, fits LDS), BN2=128 otherwise
BN2 = 256 if dep <= 256 else 128
nw2 = 4
BK2 = 128
for bk in [512, 256, 128]:
if bk <= dep and dep % bk == 0:
BK2 = bk
break
even_k2 = (dep % BK2) == 0
gemm2_out = gemm_f32[:total_sorted * dhp].view(total_sorted, dhp)
nn2 = triton.cdiv(dhp, BN2)
_batched_moe_gemm_fp4[(max_m_blocks * nn2,)](
inter_fp4, dw, gemm2_out,
inter_sc, dw_sc,
sorted_token_idx, sorted_weights,
cum_blocks, expert_offsets,
E, dhp, dep // 2,
inter_fp4.stride(0), inter_fp4.stride(1),
dw.stride(0), dw.stride(1), dw.stride(2),
gemm2_out.stride(0), gemm2_out.stride(1),
inter_sc.stride(0), inter_sc.stride(1),
dw_sc.stride(0), dw_sc.stride(1), dw_sc.stride(2),
max_m_blocks, total_sorted,
BM=BM_GEMM, BN=BN2, BK=BK2, EVEN_K=even_k2, INDIRECT=False,
PRESHUFFLE=True, APPLY_WEIGHTS=True,
SEARCH_ITERS=si,
num_warps=nw2, num_stages=2,
)
out_bf16 = _get_buf("out_bf16", M * dh, torch.bfloat16, dev).view(M, dh)
SR_BN = 128
_scatter_reduce_kernel[(M, triton.cdiv(dh, SR_BN))](
gemm2_out, reverse_idx, out_bf16,
M, dh,
gemm2_out.stride(0), gemm2_out.stride(1),
topk,
BN=SR_BN, TOPK=topk,
num_warps=4,
)
return out_bf16
scrolls · 603 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON