submission 754196
divc13 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1141 lines, June 9 Researcher Reciprocity License v1.0.
submission_v329.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754196?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:5111877487356be8c8076cbab4ce0ac2801de6d89a302513cf2ddf14f4738fae
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15
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
def _fused_silu_mul_quant_kernel(num-warps = 4
num_warps=4,stages = 3
num_stages=3,tile-m = 128
BLOCK_M=128, QB=32,tile-n = 256
SR_BN = 256 if M >= 64 else 128Kernel source
submission_v329.py1141 lines
"""
v329: v268 + eliminate dhp padding waste.
dh=7168 is already divisible by 1024/512/256/128/64/32, so next_power_of_2
padding to 8192 wastes 12.5% of:
- GEMM1 K-iterations (8->7)
- GEMM2 N-tiles (64->56 for BN2=128)
- Quantization grid (256->224 scale groups)
- Buffer allocations
Weights keep their dhp-based shapes (pre-padded by harness), kernels just
access the first dh-relevant portion.
"""
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
@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, cache_modifier=".cg")
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, cache_modifier=".cg")
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, cache_modifier=".cg")
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], cache_modifier=".cg")
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_and_offsets_kernel(
topk_ids_ptr, counts_ptr, offsets_ptr, cum_blocks_ptr,
done_ptr, total, E, num_ctas,
BM: tl.constexpr, BLOCK: tl.constexpr, BLOCK_E: 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)
old_done = tl.atomic_add(done_ptr, 1)
if old_done == num_ctas - 1:
idx = tl.arange(0, BLOCK_E)
e_mask = idx < E
counts = tl.load(counts_ptr + idx, mask=e_mask, other=0).to(tl.int64)
token_offsets = tl.cumsum(counts, axis=0)
tl.store(offsets_ptr + 1 + idx, token_offsets, mask=e_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=e_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 _single_cta_sort_kernel(
topk_ids_ptr, topk_weights_ptr, sorted_token_idx_ptr, sorted_weights_ptr,
counts_ptr, offsets_ptr, cum_blocks_ptr, write_counts_ptr,
reverse_idx_ptr, token_counts_ptr,
total, E, M, topk,
BM: tl.constexpr, BLOCK: tl.constexpr, BLOCK_E: tl.constexpr,
):
idx = tl.arange(0, BLOCK_E)
e_mask = idx < E
tl.store(counts_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int32), mask=e_mask)
tl.store(write_counts_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int32), mask=e_mask)
e1_mask = idx < (E + 1)
tl.store(offsets_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int64), mask=e1_mask)
tl.store(cum_blocks_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int64), mask=e1_mask)
offs = tl.arange(0, BLOCK)
t_mask = offs < M
tl.store(token_counts_ptr + offs, tl.zeros([BLOCK], dtype=tl.int32), mask=t_mask)
tl.debug_barrier()
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)
tl.atomic_add(counts_ptr + ids, tl.full([BLOCK], 1, dtype=tl.int32), mask=mask)
tl.debug_barrier()
counts = tl.load(counts_ptr + idx, mask=e_mask, other=0).to(tl.int64)
token_offsets = tl.cumsum(counts, axis=0)
tl.store(offsets_ptr + 1 + idx, token_offsets, mask=e_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=e_mask)
tl.debug_barrier()
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 _fused_quant_sort_single_kernel(
x_ptr, x_fp4_ptr, x_sc_ptr,
stride_x_m, stride_x_n,
stride_fp4_m, stride_fp4_n,
stride_sc_m, stride_sc_n,
topk_ids_ptr, topk_weights_ptr,
sorted_token_idx_ptr, sorted_weights_ptr,
counts_ptr, offsets_ptr, cum_blocks_ptr, write_counts_ptr,
reverse_idx_ptr, token_counts_ptr,
total_sorted, E, M, topk,
scaleN, N_quant,
QUANT_GRID: tl.constexpr,
BM: tl.constexpr,
BLOCK_E: tl.constexpr,
):
pid = tl.program_id(0)
if pid < QUANT_GRID:
sxm = tl.cast(stride_x_m, tl.int64)
sxn = tl.cast(stride_x_n, tl.int64)
sfm = tl.cast(stride_fp4_m, tl.int64)
sfn = tl.cast(stride_fp4_n, tl.int64)
pid_m = pid // scaleN
pid_n = pid % scaleN
x_offs_m = pid_m * 128 + tl.arange(0, 128)
x_offs_n = pid_n * 32 + tl.arange(0, 32)
x_offs = x_offs_m[:, None] * sxm + x_offs_n[None, :] * sxn
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N_quant)[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 = 127
E2_BIAS = 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, [128, 16, 2])
evens, odds = tl.split(e2m1_value)
out_tensor = evens | (odds << 4)
out_offs_m = pid_m * 128 + tl.arange(0, 128)
out_offs_n = pid_n * 16 + tl.arange(0, 16)
out_offs = out_offs_m[:, None] * sfm + out_offs_n[None, :] * sfn
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N_quant // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask, cache_modifier=".cg")
bs_offs_m = pid_m * 128 + tl.arange(0, 128)
bs_offs_n = pid_n
bs_offs = bs_offs_m[:, None] * stride_sc_m + bs_offs_n[None, :] * stride_sc_n
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N_quant)[None, :]
tl.store(x_sc_ptr + bs_offs, bs_e8m0, mask=bs_mask, cache_modifier=".cg")
else:
idx = tl.arange(0, BLOCK_E)
e_mask = idx < E
tl.store(counts_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int32), mask=e_mask)
tl.store(write_counts_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int32), mask=e_mask)
e1_mask = idx < (E + 1)
tl.store(offsets_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int64), mask=e1_mask)
tl.store(cum_blocks_ptr + idx, tl.zeros([BLOCK_E], dtype=tl.int64), mask=e1_mask)
offs = tl.arange(0, 256)
t_mask = offs < M
tl.store(token_counts_ptr + offs, tl.zeros([256], dtype=tl.int32), mask=t_mask)
tl.debug_barrier()
mask = offs < total_sorted
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)
tl.atomic_add(counts_ptr + ids, tl.full([256], 1, dtype=tl.int32), mask=mask)
tl.debug_barrier()
counts = tl.load(counts_ptr + idx, mask=e_mask, other=0).to(tl.int64)
token_offsets = tl.cumsum(counts, axis=0)
tl.store(offsets_ptr + 1 + idx, token_offsets, mask=e_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=e_mask)
tl.debug_barrier()
token_idx = (offs // topk).to(tl.int32)
old = tl.atomic_add(write_counts_ptr + ids, tl.full([256], 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([256], 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 _fused_quant_count_kernel(
x_ptr, x_fp4_ptr, x_sc_ptr,
stride_x_m, stride_x_n,
stride_fp4_m, stride_fp4_n,
stride_sc_m, stride_sc_n,
topk_ids_ptr, counts_ptr, offsets_ptr, cum_blocks_ptr, done_ptr,
total_sorted, E, M, num_count_ctas,
scaleN, N_quant,
QUANT_GRID: tl.constexpr,
BM: tl.constexpr,
BLOCK_E: tl.constexpr,
):
pid = tl.program_id(0)
if pid < QUANT_GRID:
sxm = tl.cast(stride_x_m, tl.int64)
sxn = tl.cast(stride_x_n, tl.int64)
sfm = tl.cast(stride_fp4_m, tl.int64)
sfn = tl.cast(stride_fp4_n, tl.int64)
pid_m = pid // scaleN
pid_n = pid % scaleN
x_offs_m = pid_m * 128 + tl.arange(0, 128)
x_offs_n = pid_n * 32 + tl.arange(0, 32)
x_offs = x_offs_m[:, None] * sxm + x_offs_n[None, :] * sxn
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N_quant)[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 = 127
E2_BIAS = 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, [128, 16, 2])
evens, odds = tl.split(e2m1_value)
out_tensor = evens | (odds << 4)
out_offs_m = pid_m * 128 + tl.arange(0, 128)
out_offs_n = pid_n * 16 + tl.arange(0, 16)
out_offs = out_offs_m[:, None] * sfm + out_offs_n[None, :] * sfn
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N_quant // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask, cache_modifier=".cg")
bs_offs_m = pid_m * 128 + tl.arange(0, 128)
bs_offs_n = pid_n
bs_offs = bs_offs_m[:, None] * stride_sc_m + bs_offs_n[None, :] * stride_sc_n
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N_quant)[None, :]
tl.store(x_sc_ptr + bs_offs, bs_e8m0, mask=bs_mask, cache_modifier=".cg")
else:
count_pid = pid - QUANT_GRID
offs = count_pid * 256 + tl.arange(0, 256)
mask = offs < total_sorted
ids = tl.load(topk_ids_ptr + offs, mask=mask, other=0).to(tl.int32)
tl.atomic_add(counts_ptr + ids, tl.full([256], 1, dtype=tl.int32), mask=mask)
old_done = tl.atomic_add(done_ptr, 1)
if old_done == num_count_ctas - 1:
idx = tl.arange(0, BLOCK_E)
e_mask = idx < E
counts = tl.load(counts_ptr + idx, mask=e_mask, other=0).to(tl.int64)
token_offsets = tl.cumsum(counts, axis=0)
tl.store(offsets_ptr + 1 + idx, token_offsets, mask=e_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=e_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).to(tl.float32)
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, USE_CG: tl.constexpr = False,
OUTPUT_BF16: tl.constexpr = False,
NK: tl.constexpr = 0,
EVEN_N: tl.constexpr = False,
):
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 = tl.program_id(0)
nn = tl.cdiv(N, BN)
pid_mb = pid // nn
pid_n = pid % nn
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)
if PRESHUFFLE and USE_CG:
for _ in range(NK):
asc = tl.load(asp, mask=row_mask[:, None], other=0)
if EVEN_N:
bsc = tl.load(bsp, cache_modifier=".cg")
else:
bsc = tl.load(bsp, mask=col_mask[:, None], other=0, cache_modifier=".cg")
a = tl.load(ap, mask=row_mask[:, None], other=0)
b_raw = tl.load(bp, cache_modifier=".cg")
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 PRESHUFFLE:
for _ in range(NK):
asc = tl.load(asp, mask=row_mask[:, None], other=0)
if EVEN_N:
bsc = tl.load(bsp)
else:
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
store_mask = row_mask[:, None] if EVEN_N else (row_mask[:, None] & col_mask[None, :])
if OUTPUT_BF16:
tl.store(cp, acc.to(tl.bfloat16), mask=store_mask, cache_modifier=".cg")
else:
tl.store(cp, acc, mask=store_mask, cache_modifier=".cg")
@triton.jit
def _fused_gemm1_silu_quant_fp4(
a_ptr, b_ptr,
fp4_out_ptr, scale_out_ptr,
a_sc_ptr, b_sc_ptr,
token_idx_ptr,
cum_blocks_ptr, expert_offsets_ptr,
E, dep, K_half,
stride_am, stride_ak,
stride_be, stride_bn, stride_bk,
stride_fp4m, stride_fp4n,
stride_scm, stride_scn,
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,
SEARCH_ITERS: tl.constexpr, USE_CG: tl.constexpr = False,
NK: tl.constexpr = 1,
EVEN_N: tl.constexpr = False,
):
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_fp4m = tl.cast(stride_fp4m, tl.int64)
stride_fp4n = tl.cast(stride_fp4n, 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_asm > 0)
tl.assume(stride_ask > 0)
tl.assume(stride_bsn > 0)
tl.assume(stride_bsk > 0)
pid = tl.program_id(0)
nn = tl.cdiv(dep, BN)
pid_mb = pid // nn
pid_n = pid % nn
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)
a_rows = tl.load(token_idx_ptr + sorted_rows, mask=row_mask, other=0).to(tl.int64)
gate_cols = pid_n * BN + tl.arange(0, BN)
up_cols = dep + gate_cols
gate_col_mask = gate_cols < dep
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
offs_bn_gate = pid_n * (BN // 16) + tl.arange(0, BN // 16)
offs_bn_up = (dep // 16) + offs_bn_gate
offs_k_shuf = tl.arange(0, (BK // 2) * 16)
bp_gate = b_base + offs_bn_gate[:, None].to(tl.int64) * stride_bn + offs_k_shuf[None, :].to(tl.int64) * stride_bk
bp_up = b_base + offs_bn_up[:, None].to(tl.int64) * stride_bn + offs_k_shuf[None, :].to(tl.int64) * stride_bk
bsp_gate = bsc_base + gate_cols[:, None].to(tl.int64) * stride_bsn + ks[None, :] * stride_bsk
bsp_up = bsc_base + up_cols[:, None].to(tl.int64) * stride_bsn + ks[None, :] * stride_bsk
acc_gate = tl.zeros((BM, BN), dtype=tl.float32)
acc_up = tl.zeros((BM, BN), dtype=tl.float32)
if USE_CG and EVEN_N:
for _ in range(NK):
asc = tl.load(asp, mask=row_mask[:, None], other=0)
a = tl.load(ap, mask=row_mask[:, None], other=0)
bsc_g = tl.load(bsp_gate, cache_modifier=".cg")
b_raw_g = tl.load(bp_gate, cache_modifier=".cg")
b_g = (b_raw_g.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
acc_gate = tl.dot_scaled(a, asc, "e2m1", b_g, bsc_g, "e2m1", acc_gate)
bsc_u = tl.load(bsp_up, cache_modifier=".cg")
b_raw_u = tl.load(bp_up, cache_modifier=".cg")
b_u = (b_raw_u.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
acc_up = tl.dot_scaled(a, asc, "e2m1", b_u, bsc_u, "e2m1", acc_up)
ap += (BK // 2) * stride_ak
asp += (BK // SG) * stride_ask
bp_gate += (BK // 2) * 16 * stride_bk
bp_up += (BK // 2) * 16 * stride_bk
bsp_gate += (BK // SG) * stride_bsk
bsp_up += (BK // SG) * stride_bsk
elif USE_CG:
for _ in range(NK):
asc = tl.load(asp, mask=row_mask[:, None], other=0)
a = tl.load(ap, mask=row_mask[:, None], other=0)
bsc_g = tl.load(bsp_gate, mask=gate_col_mask[:, None], other=0, cache_modifier=".cg")
b_raw_g = tl.load(bp_gate, cache_modifier=".cg")
b_g = (b_raw_g.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
acc_gate = tl.dot_scaled(a, asc, "e2m1", b_g, bsc_g, "e2m1", acc_gate)
bsc_u = tl.load(bsp_up, mask=gate_col_mask[:, None], other=0, cache_modifier=".cg")
b_raw_u = tl.load(bp_up, cache_modifier=".cg")
b_u = (b_raw_u.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
acc_up = tl.dot_scaled(a, asc, "e2m1", b_u, bsc_u, "e2m1", acc_up)
ap += (BK // 2) * stride_ak
asp += (BK // SG) * stride_ask
bp_gate += (BK // 2) * 16 * stride_bk
bp_up += (BK // 2) * 16 * stride_bk
bsp_gate += (BK // SG) * stride_bsk
bsp_up += (BK // SG) * stride_bsk
elif EVEN_N:
for _ in range(NK):
asc = tl.load(asp, mask=row_mask[:, None], other=0)
a = tl.load(ap, mask=row_mask[:, None], other=0)
bsc_g = tl.load(bsp_gate)
b_raw_g = tl.load(bp_gate)
b_g = (b_raw_g.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
acc_gate = tl.dot_scaled(a, asc, "e2m1", b_g, bsc_g, "e2m1", acc_gate)
bsc_u = tl.load(bsp_up)
b_raw_u = tl.load(bp_up)
b_u = (b_raw_u.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
acc_up = tl.dot_scaled(a, asc, "e2m1", b_u, bsc_u, "e2m1", acc_up)
ap += (BK // 2) * stride_ak
asp += (BK // SG) * stride_ask
bp_gate += (BK // 2) * 16 * stride_bk
bp_up += (BK // 2) * 16 * stride_bk
bsp_gate += (BK // SG) * stride_bsk
bsp_up += (BK // SG) * stride_bsk
else:
for _ in range(NK):
asc = tl.load(asp, mask=row_mask[:, None], other=0)
a = tl.load(ap, mask=row_mask[:, None], other=0)
bsc_g = tl.load(bsp_gate, mask=gate_col_mask[:, None], other=0)
b_raw_g = tl.load(bp_gate)
b_g = (b_raw_g.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
acc_gate = tl.dot_scaled(a, asc, "e2m1", b_g, bsc_g, "e2m1", acc_gate)
bsc_u = tl.load(bsp_up, mask=gate_col_mask[:, None], other=0)
b_raw_u = tl.load(bp_up)
b_u = (b_raw_u.reshape(1, BN // 16, BK // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BN, BK // 2).trans(1, 0))
acc_up = tl.dot_scaled(a, asc, "e2m1", b_u, bsc_u, "e2m1", acc_up)
ap += (BK // 2) * stride_ak
asp += (BK // SG) * stride_ask
bp_gate += (BK // 2) * 16 * stride_bk
bp_up += (BK // 2) * 16 * stride_bk
bsp_gate += (BK // SG) * stride_bsk
bsp_up += (BK // SG) * stride_bsk
x = (acc_gate * tl.sigmoid(acc_gate) * acc_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, [BM, BN // 2, 2])
evens, odds = tl.split(e2m1_val)
packed = evens | (odds << 4)
out_m = sorted_rows
out_n = pid_n * BN // 2 + tl.arange(0, BN // 2)
out_offs = out_m[:, None] * stride_fp4m + out_n[None, :] * stride_fp4n
out_mask = row_mask[:, None] if EVEN_N else (row_mask[:, None] & (out_n < (dep // 2))[None, :])
tl.store(fp4_out_ptr + out_offs, packed, mask=out_mask, cache_modifier=".cg")
sc_n = pid_n
sc_offs = out_m[:, None] * stride_scm + sc_n * stride_scn
tl.store(scale_out_ptr + sc_offs, bs_e8m0, mask=row_mask[:, None], cache_modifier=".cg")
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)
guw_s0, guw_s1, guw_s2 = guw.stride()
dw_s0, dw_s1, dw_s2 = dw.stride()
guw_sc_s0, guw_sc_s1, guw_sc_s2 = guw_sc.stride()
dw_sc_s0, dw_sc_s1, dw_sc_s2 = dw_sc.stride()
hs = hidden_states
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 >= 8:
BM_GEMM = 32
else:
BM_GEMM = 16
flat_ids = topk_ids.reshape(-1)
flat_weights = topk_weights.reshape(-1)
scaleN_valid = triton.cdiv(dh, 32)
scaleN_pad = triton.cdiv(scaleN_valid, 8) * 8
pad_M_sc = triton.cdiv(M, 256) * 256
i32_sz = 2 * E + M + 1
i64_sz = 2 * (E + 1)
scaleN_inter = triton.cdiv(dep, 32)
use_fused = (BM_GEMM <= 16) or (BM_GEMM <= 32 and dep >= 512)
A = 256
def al(n):
return (n + A - 1) & ~(A - 1)
o = 0
o_hfp4 = o; s_hfp4 = M * (dh // 2); o += al(s_hfp4)
o_hsc = o; s_hsc = pad_M_sc * scaleN_pad; o += al(s_hsc)
o_sort = o
o_si32 = o; s_si32 = i32_sz * 4; o += al(s_si32)
o_si64 = o; s_si64 = i64_sz * 8; o += al(s_si64)
o_sort_end = o
o_stix = o; s_stix = total_sorted * 4; o += al(s_stix)
o_sw = o; s_sw = total_sorted * 2; o += al(s_sw)
o_rev = o; s_rev = total_sorted * 4; o += al(s_rev)
o_ifp4 = o; s_ifp4 = total_sorted * (dep // 2); o += al(max(s_ifp4, 1))
o_isc = o; s_isc = total_sorted * scaleN_inter; o += al(max(s_isc, 1))
if not use_fused:
o_g1 = o; s_g1 = total_sorted * 2 * dep * 4; o += al(s_g1)
o_g2 = o; s_g2 = total_sorted * dh * 2; o += al(s_g2)
o_out = o; s_out = M * dh * 2; o += al(s_out)
_buf = torch.empty(o, dtype=torch.uint8, device=dev)
hs_fp4 = _buf[o_hfp4:o_hfp4 + s_hfp4].view(M, dh // 2)
hs_sc_buf = _buf[o_hsc:o_hsc + s_hsc].view(pad_M_sc, scaleN_pad)
sort_i32 = _buf[o_si32:o_si32 + s_si32].view(torch.int32)
sort_i64 = _buf[o_si64:o_si64 + s_si64].view(torch.int64)
sorted_token_idx = _buf[o_stix:o_stix + s_stix].view(torch.int32)
sorted_weights = _buf[o_sw:o_sw + s_sw].view(torch.bfloat16)
reverse_idx = _buf[o_rev:o_rev + s_rev].view(torch.int32)
inter_fp4 = _buf[o_ifp4:o_ifp4 + s_ifp4].view(total_sorted, dep // 2) if s_ifp4 > 0 else _buf[o_ifp4:o_ifp4 + 1].view(1, 1)
inter_sc = _buf[o_isc:o_isc + s_isc].view(total_sorted, scaleN_inter) if s_isc > 0 else _buf[o_isc:o_isc + 1].view(1, 1)
if not use_fused:
gemm1_out = _buf[o_g1:o_g1 + s_g1].view(torch.float32).view(total_sorted, 2 * dep)
gemm2_out = _buf[o_g2:o_g2 + s_g2].view(torch.bfloat16).view(total_sorted, dh)
out_bf16 = _buf[o_out:o_out + s_out].view(torch.bfloat16).view(M, dh)
expert_counts = sort_i32[:E]
write_counts = sort_i32[E:2*E]
token_counts = sort_i32[2*E:2*E+M]
done_counter = sort_i32[2*E+M:]
expert_offsets = sort_i64[:E+1]
cum_blocks = sort_i64[E+1:]
SORT_BLOCK = 256
BLOCK_E = triton.next_power_of_2(E)
sort_grid = triton.cdiv(total_sorted, SORT_BLOCK)
quant_grid_m = triton.cdiv(M, 128)
quant_grid = quant_grid_m * scaleN_valid
if sort_grid == 1:
_fused_quant_sort_single_kernel[(quant_grid + 1,)](
hs, hs_fp4, hs_sc_buf,
*hs.stride(), *hs_fp4.stride(), *hs_sc_buf.stride(),
flat_ids, flat_weights, sorted_token_idx, sorted_weights,
expert_counts, expert_offsets, cum_blocks, write_counts,
reverse_idx, token_counts,
total_sorted, E, M, topk,
scaleN_valid, dh,
QUANT_GRID=quant_grid, BM=BM_GEMM, BLOCK_E=BLOCK_E,
num_warps=4,
)
else:
_buf[o_sort:o_sort_end].zero_()
_fused_quant_count_kernel[(quant_grid + sort_grid,)](
hs, hs_fp4, hs_sc_buf,
*hs.stride(), *hs_fp4.stride(), *hs_sc_buf.stride(),
flat_ids, expert_counts, expert_offsets, cum_blocks,
done_counter, total_sorted, E, M, sort_grid,
scaleN_valid, dh,
QUANT_GRID=quant_grid, BM=BM_GEMM, BLOCK_E=BLOCK_E,
num_warps=4,
)
_moe_scatter_kernel[(sort_grid,)](
flat_ids, flat_weights,
sorted_token_idx, sorted_weights,
expert_offsets, write_counts,
reverse_idx, token_counts,
topk, total_sorted, BLOCK=SORT_BLOCK,
)
hs_sc = hs_sc_buf[:M]
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
use_mfma16 = BM_GEMM <= 32
BK1 = 128
for bk in [1024, 512, 256, 128]:
if bk <= min(dh, max_bk) and dh % bk == 0:
BK1 = bk
break
even_k1 = (dh % BK1) == 0
NK1 = triton.cdiv(dh // 2, BK1 // 2)
if use_fused:
BN_fused = 32
nn_fused = triton.cdiv(dep, BN_fused)
_fused_gemm1_silu_quant_fp4[(max_m_blocks * nn_fused,)](
hs_fp4, guw,
inter_fp4, inter_sc,
hs_sc, guw_sc,
sorted_token_idx,
cum_blocks, expert_offsets,
E, dep, dh // 2,
hs_fp4.stride(0), hs_fp4.stride(1),
guw_s0, guw_s1, guw_s2,
inter_fp4.stride(0), inter_fp4.stride(1),
inter_sc.stride(0), inter_sc.stride(1),
hs_sc.stride(0), hs_sc.stride(1),
guw_sc_s0, guw_sc_s1, guw_sc_s2,
max_m_blocks, total_sorted,
BM=BM_GEMM, BN=BN_fused, BK=BK1,
EVEN_K=even_k1,
SEARCH_ITERS=si, USE_CG=(dep <= 512),
NK=NK1,
EVEN_N=(dep % BN_fused == 0),
num_warps=nw1, num_stages=(1 if BM_GEMM >= 32 else 2),
matrix_instr_nonkdim=16, schedule_hint="attention",
)
else:
if dep > 512:
BN1 = 256
elif dep <= 256 and BM_GEMM <= 32 and total_sorted < 2000:
BN1 = 128
elif dep == 512 and BM_GEMM >= 64:
BN1 = 128
else:
BN1 = 64
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, dh // 2,
hs_fp4.stride(0), hs_fp4.stride(1),
guw_s0, guw_s1, guw_s2,
gemm1_out.stride(0), gemm1_out.stride(1),
hs_sc.stride(0), hs_sc.stride(1),
guw_sc_s0, guw_sc_s1, guw_sc_s2,
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, USE_CG=(dep <= 512),
OUTPUT_BF16=False, NK=NK1,
EVEN_N=((2 * dep) % BN1 == 0),
num_warps=(2 if BM_GEMM >= 128 else nw1), num_stages=(3 if dep == 512 and BM_GEMM >= 64 else 2),
**({"matrix_instr_nonkdim": 16} if use_mfma16 else {}),
schedule_hint="attention",
)
_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,
num_stages=3,
)
BN2 = 256 if dep <= 256 else 128
nw2 = 4
BK2 = 128
bk2_max = 256 if dep == 512 else 512
for bk in [512, 256, 128]:
if bk <= min(dep, bk2_max) and dep % bk == 0:
BK2 = bk
break
even_k2 = (dep % BK2) == 0
NK2 = triton.cdiv(dep // 2, BK2 // 2)
nn2 = triton.cdiv(dh, 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, dh, dep // 2,
inter_fp4.stride(0), inter_fp4.stride(1),
dw_s0, dw_s1, dw_s2,
gemm2_out.stride(0), gemm2_out.stride(1),
inter_sc.stride(0), inter_sc.stride(1),
dw_sc_s0, dw_sc_s1, dw_sc_s2,
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, USE_CG=(dep <= 512),
OUTPUT_BF16=True, NK=NK2,
EVEN_N=(dh % BN2 == 0),
num_warps=nw2, num_stages=2,
**({"matrix_instr_nonkdim": 16} if use_mfma16 else {}),
schedule_hint="attention",
)
SR_BN = 256 if M >= 64 else 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, num_stages=3,
)
return out_bf16
scrolls · 1141 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