submission 741375
LunNova · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 291 lines, June 9 Researcher Reciprocity License v1.0.
sub_triton_ck_v2b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-741375?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:478d4c743d724c03afb2f44fe0e8f46c3d861f10540d612033d6f18e8b39d9d0
license declaredunknown
license concludedunknown
authorsLunNova
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
sub_triton_ck_v2b.py291 lines
"""
sub_triton_ck_v2b.py — Hybrid: Triton stage1 for shapes where it wins, CK stage1 fallback for rest.
Both paths use CK ASM stage2. Shape dispatch via sk = m | (e << 16) | (inter_dim << 32).
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import _moe_sorting_impl, get_block_size_M, use_nt, get_ksplit
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
# ═══════════════════════════════════════════════════════════════════
# Shapes where Triton stage1 wins (live benchmarks 2026-04-05)
# Key: m | (E << 16) | (inter_dim << 32)
# Default: CK stage1 (faster on 5/7 contest shapes)
# ═══════════════════════════════════════════════════════════════════
_TRITON_STAGE1_SHAPES = {
128 | (33 << 16) | (512 << 32), # de512_E32_bs128: triton 145 vs ck 148 (E=33 incl shared)
512 | (33 << 16) | (512 << 32), # de512_E32_bs512: triton 242 vs ck 252 (E=33 incl shared)
}
# ═══════════════════════════════════════════════════════════════════
# Fused MOE Stage1: Triton GEMM (gate+up) + SiLU + scatter
# Hardcoded config: BLOCK_N=64, BLOCK_K=512, warps=4, stages=2
# (autotuned winner for dep=512, Kp=3584 on gfx950)
# ═══════════════════════════════════════════════════════════════════
@triton.jit
def _moe_stage1_fused(
a_ptr, w_ptr, a2_ptr,
a_sc_ptr, w_sc_ptr,
sorted_ids_ptr, sorted_expert_ids_ptr, num_valid_ids_ptr,
M, d_expert_pad, K_packed, num_valid_padded,
stride_a_m, stride_a_k,
stride_w_e, stride_w_bn,
stride_as_blk,
stride_ws_blk, w_sc_expert_offset,
stride_a2_m, stride_a2_n,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
K_PACKED: tl.constexpr,
top_k: tl.constexpr,
):
SCALE_GRP: tl.constexpr = 32
K_PACK: tl.constexpr = BLOCK_K // 2
K_SC: tl.constexpr = BLOCK_K // SCALE_GRP
NUM_K_ITERS: tl.constexpr = (2 * K_PACKED) // BLOCK_K
pid = tl.program_id(0)
num_n_blks = tl.cdiv(d_expert_pad, BLOCK_N)
pid_m = pid // num_n_blks
pid_n = pid % num_n_blks
num_valid = tl.load(num_valid_ids_ptr)
if pid_m * BLOCK_M >= num_valid:
return
expert_id = tl.load(sorted_expert_ids_ptr + pid_m)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
valid_m = offs_m < num_valid
sid = tl.load(sorted_ids_ptr + offs_m, mask=valid_m, other=0)
token_ids = (sid & 0xFFFFFF)
topk_ids = sid >> 24
valid_tok = token_ids < M
safe_tids = tl.where(valid_tok, token_ids, 0)
a_mask = valid_m & valid_tok
offs_k_a = tl.arange(0, K_PACK)
a_ptrs_base = a_ptr + safe_tids[:, None] * stride_a_m + offs_k_a[None, :] * stride_a_k
offs_asm = pid_m * (BLOCK_M // 32) + tl.arange(0, BLOCK_M // 32)
offs_ks_flat = tl.arange(0, K_SC * 32)
a_sc_ptrs_base = a_sc_ptr + offs_asm[:, None] * stride_as_blk + offs_ks_flat[None, :]
w_base = w_ptr + expert_id * stride_w_e
offs_k_shuf = tl.arange(0, K_PACK * 16)
gate_bn_tile = pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)
b_gate_ptrs_base = w_base + gate_bn_tile[:, None] * stride_w_bn + offs_k_shuf[None, :]
ws_base_blk = expert_id * w_sc_expert_offset
gate_bsn = ws_base_blk + pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
b_gate_sc_ptrs_base = w_sc_ptr + gate_bsn[:, None] * stride_ws_blk + offs_ks_flat[None, :]
up_n_offset = d_expert_pad // 16
up_bn_tile = (up_n_offset + pid_n * (BLOCK_N // 16)) + tl.arange(0, BLOCK_N // 16)
b_up_ptrs_base = w_base + up_bn_tile[:, None] * stride_w_bn + offs_k_shuf[None, :]
up_sc_offset = d_expert_pad // 32
up_bsn = ws_base_blk + up_sc_offset + pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
b_up_sc_ptrs_base = w_sc_ptr + up_bsn[:, None] * stride_ws_blk + offs_ks_flat[None, :]
acc_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_up = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
a_ptrs = a_ptrs_base
a_sc_ptrs = a_sc_ptrs_base
b_gate_ptrs = b_gate_ptrs_base
b_gate_sc_ptrs = b_gate_sc_ptrs_base
b_up_ptrs = b_up_ptrs_base
b_up_sc_ptrs = b_up_sc_ptrs_base
for _ in range(NUM_K_ITERS):
a = tl.load(a_ptrs, mask=a_mask[:, None], other=0)
a_sc_raw = tl.load(a_sc_ptrs)
a_scales = (
a_sc_raw
.reshape(BLOCK_M // 32, K_SC // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_M, K_SC)
)
b_gate_raw = tl.load(b_gate_ptrs)
b_gate = (
b_gate_raw
.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_N, K_PACK)
.trans(1, 0)
)
b_gate_sc_raw = tl.load(b_gate_sc_ptrs)
b_gate_scales = (
b_gate_sc_raw
.reshape(BLOCK_N // 32, K_SC // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_N, K_SC)
)
acc_gate = tl.dot_scaled(a, a_scales, "e2m1", b_gate, b_gate_scales, "e2m1", acc_gate)
b_up_raw = tl.load(b_up_ptrs)
b_up = (
b_up_raw
.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_N, K_PACK)
.trans(1, 0)
)
b_up_sc_raw = tl.load(b_up_sc_ptrs)
b_up_scales = (
b_up_sc_raw
.reshape(BLOCK_N // 32, K_SC // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_N, K_SC)
)
acc_up = tl.dot_scaled(a, a_scales, "e2m1", b_up, b_up_scales, "e2m1", acc_up)
a_ptrs += K_PACK * stride_a_k
a_sc_ptrs += K_SC * 32
b_gate_ptrs += K_PACK * 16
b_gate_sc_ptrs += K_SC * 32
b_up_ptrs += K_PACK * 16
b_up_sc_ptrs += K_SC * 32
result = (acc_gate * tl.sigmoid(acc_gate)) * acc_up
flat_row = token_ids * top_k + topk_ids
safe_row = tl.where(a_mask, flat_row, 0)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
out_ptrs = a2_ptr + safe_row[:, None] * stride_a2_m + offs_n[None, :] * stride_a2_n
out_mask = a_mask[:, None] & (offs_n[None, :] < d_expert_pad)
tl.store(out_ptrs, result.to(tl.bfloat16), mask=out_mask)
# ═══════════════════════════════════════════════════════════════════
# Entry point
# ═══════════════════════════════════════════════════════════════════
def custom_kernel(data: input_t) -> output_t:
(
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
M = hidden_states.shape[0]
topk = topk_ids.shape[1]
E = gate_up_weight_shuffled.shape[0]
N_gateup = gate_up_weight_shuffled.shape[1]
inter_dim = N_gateup // 2
model_dim = down_weight_shuffled.shape[1]
d_hidden_pad = config["d_hidden_pad"]
d_expert_pad = config["d_expert_pad"]
K_packed = d_hidden_pad // 2
K_scale = d_hidden_pad // 32
dtype = torch.bfloat16
device = hidden_states.device
block_m = get_block_size_M(M, topk, E, inter_dim)
ksplit = get_ksplit(M, topk, E, inter_dim, model_dim)
non_temporal = use_nt(M, topk, E)
# ── 1. Token sorting (opus) ──
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = _moe_sorting_impl(
topk_ids, topk_weights, E, model_dim, dtype, block_m,
expert_mask=None, num_local_tokens=None, dispatch_policy=0, use_opus=True,
)
num_valid_padded = sorted_ids.shape[0]
# ── 2. Quantize activations (fused with sort → shuffled scales) ──
a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
hidden_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=M,
topk=1,
block_size=block_m,
)
# ── 3. Stage1: dispatch triton vs CK based on shape ──
sk = M | (E << 16) | (inter_dim << 32)
use_ck = sk not in _TRITON_STAGE1_SHAPES
a2 = torch.empty((M, topk, inter_dim), dtype=dtype, device=device)
if use_ck:
# CK stage1: gate_up GEMM + SiLU
w1_scale_e8m0 = gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0)
aiter.ck_moe_stage1_fwd(
a1, gate_up_weight_shuffled, down_weight_shuffled,
sorted_ids, sorted_expert_ids, num_valid_ids,
a2, topk, "",
w1_scale_e8m0, a1_scale, block_m,
None, QuantType.per_1x32, ActivationType.Silu,
ksplit, non_temporal, dtype,
)
else:
# Triton fused stage1: GEMM + SiLU + scatter
a2_flat = a2.view(-1, inter_dim)
w1_u8 = gate_up_weight_shuffled.view(torch.uint8)
a1_u8 = a1.view(torch.uint8)
a1_sc_u8 = a1_scale.view(torch.uint8)
w1_sc_u8 = gate_up_weight_scale_shuffled.view(torch.uint8)
stride_w_bn = 16 * K_packed
stride_as_blk = K_scale * 32
stride_ws_blk = K_scale * 32
w_sc_expert_offset = N_gateup // 32
BLOCK_N = 64
grid = (triton.cdiv(num_valid_padded, block_m) *
triton.cdiv(d_expert_pad, BLOCK_N),)
_moe_stage1_fused[grid](
a1_u8, w1_u8, a2_flat,
a1_sc_u8, w1_sc_u8,
sorted_ids, sorted_expert_ids, num_valid_ids,
M, d_expert_pad, K_packed, num_valid_padded,
a1_u8.stride(0), a1_u8.stride(1),
w1_u8.stride(0), stride_w_bn,
stride_as_blk,
stride_ws_blk, w_sc_expert_offset,
a2_flat.stride(0), a2_flat.stride(1),
BLOCK_M=block_m,
BLOCK_N=BLOCK_N,
BLOCK_K=512,
K_PACKED=K_packed,
top_k=topk,
num_warps=4,
num_stages=2,
)
# ── 4. Quantize intermediate (for CK stage2) ──
a2_flat = a2.view(-1, inter_dim)
a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
a2_flat,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=M,
topk=topk,
block_size=block_m,
)
a2_q = a2_q.view(M, topk, -1)
# ── 5. CK ASM Stage2: down GEMM ──
w2_scale_e8m0 = down_weight_scale_shuffled.view(dtypes.fp8_e8m0)
aiter.ck_moe_stage2_fwd(
a2_q, gate_up_weight_shuffled, down_weight_shuffled,
sorted_ids, sorted_expert_ids, num_valid_ids,
moe_buf, topk, "",
w2_scale_e8m0, a2_scale, block_m,
sorted_weights, QuantType.per_1x32, ActivationType.Silu,
non_temporal,
)
return moe_buf
scrolls · 291 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