KernelBench hard · RTX PRO 6000
Sonic MoE Kimi K3 (256k)
manually audited: clean
Genuine MoE up-projection: a custom CuTeDSL (CUTLASS Python DSL) warp-specialized persistent grouped GEMM on tcgen05 as the primary Blackwell path (device-side expert/tile scheduler reading expert_offsets, TMA-streamed A slices via domain_offset, paired gate/up TMEM accumulators, fp32 u*silu(g) epilogue with TMA full-tile stores and predicated SIMT ragged-edge stores), plus a real Triton grouped-GEMM fallback (interleaved gate/up weights + TMA descriptors, and a fully-masked generic pointer kernel). No sonic_moe import, no torch.matmul/bmm/F.linear, no cached or constant output path, no grader interaction. The 0.0760 fraction is a real timing. Result was regraded after an infra-only original-check timeout; regrade ran with numeric stress on and passed.
Kernel source (redacted)
"""Top-K MoE FFN up-projection: variable-length grouped GEMM + fused SwiGLU.
Per expert e: h_e = silu(x_e @ W_gate[e]) * (x_e @ W_up[e])
with x_e = hidden_states[offsets[e]:offsets[e+1]], offsets an (E+1,) prefix sum.
Primary path (Blackwell): a custom CuTeDSL (CUTLASS Python DSL) warp-specialized
persistent grouped GEMM:
- device-side tile scheduler warp reads expert_offsets in GMEM and broadcasts
(expert, m_tile, n_tile) work items through an smem pipeline — fully dynamic
tile counts, zero host<->device sync;
- TMA warp streams A slices (single global descriptor + domain_offset per
expert) and per-expert gate/up B chunks (two descriptors, L-sliced);
- tcgen05 MMA warp issues gate & up GEMMs into paired TMEM accumulators
sharing the A smem tile;
- epilogue warps read paired accumulators, apply u * silu(g) in fp32, and
store bf16 via TMA (full tiles) or predicated SIMT stores (ragged edges).
Fallback path (any GPU / missing DSL / odd shapes): custom Triton grouped GEMM
(interleaved-weights TMA warp-specialized kernel when H/I allow; otherwise a
fully-masked generic pointer kernel).
"""
from __future__ import annotations
import os
_HERE = os.path.dirname(os.path.abspath(__file__))
os.environ.setdefault("CUTE_DSL_CACHE_DIR", os.path.join(_HERE, ".cute_dsl_cache"))
from typing import Optional, Tuple, Type, Union
import torch
import torch.nn as nn
# ============================================================================
# CuTeDSL (CUTLASS Python DSL) fast path
# ============================================================================
_CUTE_OK = False
_CUTE_ERR = None
try:
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
import cutlass.pipeline as pipeline
import cutlass.utils as utils
from cutlass import Boolean, Int32, const_expr
from cutlass.cute.nvgpu import cpasync, tcgen05
from cutlass.cutlass_dsl import (
dsl_user_op,
extract_mlir_values,
new_from_mlir_values,
)
from cutlass.cute.runtime import from_dlpack, make_fake_stream, make_ptr
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
from cutlass.utils.gemm.sm100 import (
transform_partitioned_tensor_layout,
epilogue_tmem_copy_and_partition,
epilogue_smem_copy_and_partition,
)
_CUTE_OK = True
except Exception as e: # pragma: no cover
import traceback
_CUTE_ERR = traceback.format_exc()
LOG2E = 1.4426950408889634
if _CUTE_OK:
class MoEWorkTileInfo:
"""Work tile: (expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt, token_offset, tokens_i)."""
FIELDS = 8 # padded to 16B multiples (6 used)
def __init__(self, expert_idx, tile_m_idx, tile_n_idx, k_tile_cnt, token_offset, tokens_i):
self.expert_idx = expert_idx
self.tile_m_idx = tile_m_idx
self.tile_n_idx = tile_n_idx
self.k_tile_cnt = k_tile_cnt
self.token_offset = token_offset
self.tokens_i = tokens_i
@property
def is_valid_tile(self):
return self.expert_idx >= Int32(0)
def __extract_mlir_values__(self):
vals = []
for f in (self.expert_idx, self.tile_m_idx, self.tile_n_idx,
self.k_tile_cnt, self.token_offset, self.tokens_i):
vals.extend(extract_mlir_values(f))
return vals
def __new_from_mlir_values__(self, values):
assert len(values) == 6
return MoEWorkTileInfo(
new_from_mlir_values(self.expert_idx, [values[0]]),
new_from_mlir_values(self.tile_m_idx, [values[1]]),
new_from_mlir_values(self.tile_n_idx, [values[2]]),
new_from_mlir_values(self.k_tile_cnt, [values[3]]),
new_from_mlir_values(self.token_offset, [values[4]]),
new_from_mlir_values(self.tokens_i, [values[5]]),
)
def to_rmem_tensor(self):
rmem = cute.make_rmem_tensor((self.FIELDS,), Int32)
rmem[0] = self.expert_idx
rmem[1] = self.tile_m_idx
rmem[2] = self.tile_n_idx
rmem[3] = self.k_tile_cnt
rmem[4] = self.token_offset
rmem[5] = self.tokens_i
rmem[6] = Int32(0)
rmem[7] = Int32(0)
return rmem
@staticmethod
def from_rmem_tensor(rmem):
return MoEWorkTileInfo(rmem[0], rmem[1], rmem[2], rmem[3], rmem[4], rmem[5])
class MoESchedParams:
def __init__(self, expert_cnt, intermediate, hidden, cta_tile_mnk, cluster_shape_mn, raster=0, group_m=8):
self.expert_cnt = expert_cnt if isinstance(expert_cnt, Int32) else Int32(expert_cnt)
self.intermediate = intermediate if isinstance(intermediate, Int32) else Int32(intermediate)
self.hidden = hidden if isinstance(hidden, Int32) else Int32(hidden)
self.cta_tile_mnk = cta_tile_mnk
self.cluster_shape_mn = cluster_shape_mn
self.raster = raster
self.group_m = group_m
@property
def cluster_tile_m(self):
return self.cta_tile_mnk[0] * self.cluster_shape_mn[0]
@property
def cluster_tile_n(self):
return self.cta_tile_mnk[1] * self.cluster_shape_mn[1]
def __extract_mlir_values__(self):
vals = []
vals.extend(extract_mlir_values(self.expert_cnt))
vals.extend(extract_mlir_values(self.intermediate))
vals.extend(extract_mlir_values(self.hidden))
return vals
def __new_from_mlir_values__(self, values):
return MoESchedParams(
new_from_mlir_values(self.expert_cnt, [values[0]]),
new_from_mlir_values(self.intermediate, [values[1]]),
new_from_mlir_values(self.hidden, [values[2]]),
self.cta_tile_mnk,
self.cluster_shape_mn,
)
class MoETileScheduler:
"""Persistent scheduler: enumerates (expert, cluster_tile_m, cluster_tile_n)
tiles; expert-major with short-side-first local rasterization."""
def __init__(self, params, offs, num_persistent_clusters, cur_idx, cta_id_in_cluster,
cur_expert, expert_tile_start, expert_tile_end):
self.params = params
self.offs = offs
self.num_persistent_clusters = num_persistent_clusters
self.cur_idx = cur_idx
self.cta_id_in_cluster = cta_id_in_cluster
self.cur_expert = cur_expert
self.expert_tile_start = expert_tile_start
self.expert_tile_end = expert_tile_end
@staticmethod
@cute.jit
def create(params, offs, block_idx, grid_dim):
num_persistent_clusters = cute.size(grid_dim) // cute.size(params.cluster_shape_mn)
bidx, bidy, bidz = block_idx
cur_idx = Int32(bidz)
cta_id_in_cluster = (
Int32(bidx % params.cluster_shape_mn[0]),
Int32(bidy % params.cluster_shape_mn[1]),
Int32(0),
)
return MoETileScheduler(
params, offs, num_persistent_clusters, cur_idx, cta_id_in_cluster,
Int32(0), Int32(0), Int32(0),
)
def __extract_mlir_values__(self):
vals = []
vals.extend(self.params.__extract_mlir_values__())
vals.extend(extract_mlir_values(self.offs))
vals.extend(extract_mlir_values(self.num_persistent_clusters))
vals.extend(extract_mlir_values(self.cur_idx))
vals.extend(extract_mlir_values(self.cta_id_in_cluster))
vals.extend(extract_mlir_values(self.cur_expert))
vals.extend(extract_mlir_values(self.expert_tile_start))
vals.extend(extract_mlir_values(self.expert_tile_end))
return vals
def __new_from_mlir_values__(self, values):
idx = 0
params = self.params.__new_from_mlir_values__(values[idx:idx + 3])
idx += 3
offs_len = len(extract_mlir_values(self.offs))
offs = new_from_mlir_values(self.offs, values[idx:idx + offs_len])
idx += offs_len
npc = new_from_mlir_values(self.num_persistent_clusters, [values[idx]]); idx += 1
cur = new_from_mlir_values(self.cur_idx, [values[idx]]); idx += 1
cta_id = new_from_mlir_values(self.cta_id_in_cluster, values[idx:idx + 3]); idx += 3
ce = new_from_mlir_values(self.cur_expert, [values[idx]]); idx += 1
ets = new_from_mlir_values(self.expert_tile_start, [values[idx]]); idx += 1
ete = new_from_mlir_values(self.expert_tile_end, [values[idx]]); idx += 1
return MoETileScheduler(params, offs, npc, cur, cta_id, ce, ets, ete)
@dsl_user_op
@cute.jit
def _tiles_for_expert(self, expert_idx, *, loc=None, ip=None):
tokens_i = self.offs[expert_idx]
if expert_idx > Int32(0):
tokens_i = tokens_i - self.offs[expert_idx - 1]
tile_m_cnt = (tokens_i + self.params.cluster_tile_m - 1) // self.params.cluster_tile_m
tile_n_cnt = (self.params.intermediate + self.params.cluster_tile_n - 1) // self.params.cluster_tile_n
return tile_m_cnt * tile_n_cnt
@dsl_user_op
@cute.jit
def _advance_expert(self, idx, *, loc=None, ip=None):
if self.expert_tile_end == Int32(0):
self.expert_tile_end = self._tiles_for_expert(Int32(0))
while idx >= self.expert_tile_end and self.cur_expert < self.params.expert_cnt:
self.cur_expert = self.cur_expert + 1
self.expert_tile_start = self.expert_tile_end
if self.cur_expert < self.params.expert_cnt:
self.expert_tile_end = self.expert_tile_end + self._tiles_for_expert(self.cur_expert)
@dsl_user_op
@cute.jit
def initial_work_tile_info(self, *, loc=None, ip=None):
return self._work_for_idx(self.cur_idx)
@dsl_user_op
@cute.jit
def advance_to_next_work(self, *, loc=None, ip=None):
self.cur_idx = self.cur_idx + self.num_persistent_clusters
return self._work_for_idx(self.cur_idx)
@dsl_user_op
@cute.jit
def _work_for_idx(self, idx, *, loc=None, ip=None):
self._advance_expert(idx)
is_valid = self.cur_expert < self.params.expert_cnt
info = MoEWorkTileInfo(Int32(-1), Int32(0), Int32(0), Int32(0), Int32(0), Int32(0))
if is_valid:
e = self.cur_expert
token_offset = Int32(0)
if e > Int32(0):
token_offset = self.offs[e - 1]
tokens_i = self.offs[e] - token_offset
tile_m_cnt = (tokens_i + self.params.cluster_tile_m - 1) // self.params.cluster_tile_m
tile_n_cnt = (self.params.intermediate + self.params.cluster_tile_n - 1) // self.params.cluster_tile_n
local = idx - self.expert_tile_start
cm = Int32(-1)
cn = Int32(-1)
if const_expr(self.params.raster == 0):
# n-fastest (short-side rule)
if tile_m_cnt <= tile_n_cnt:
cm = local % tile_m_cnt
cn = local // tile_m_cnt
else:
cn = local % tile_n_cnt
cm = local // tile_n_cnt
if const_expr(self.params.raster == 1):
# m-fastest: B chunk reused by consecutive items
cm = local % tile_m_cnt
cn = local // tile_m_cnt
if const_expr(self.params.raster == 2):
# supertiles of (group_m x all n), m-fast inside
gm = cutlass.Int32(self.params.group_m)
per_super = gm * tile_n_cnt
super_idx = local // per_super
r = local % per_super
mg = r % gm
cn = r // gm
cm = super_idx * gm + mg
if const_expr(self.params.raster == 3):
# supertiles of (group_m x all n), n-fast inside
gm = cutlass.Int32(self.params.group_m)
per_super = gm * tile_n_cnt
super_idx = local // per_super
r = local % per_super
mg = r // tile_n_cnt
cn = r % tile_n_cnt
cm = super_idx * gm + mg
cta_m = cm * self.params.cluster_shape_mn[0] + self.cta_id_in_cluster[0]
cta_n = cn * self.params.cluster_shape_mn[1] + self.cta_id_in_cluster[1]
k_tile_cnt = (self.params.hidden + self.params.cta_tile_mnk[2] - 1) // self.params.cta_tile_mnk[2]
info = MoEWorkTileInfo(e, cta_m, cta_n, k_tile_cnt, token_offset, tokens_i)
return info
# ----------------------------------------------------------------------------
# Main kernel
# ----------------------------------------------------------------------------
class MoeUpSwigluKernel:
def __init__(
self,
mma_tiler_mn: Tuple[int, int] = (128, 128),
cluster_shape_mn: Tuple[int, int] = (1, 1),
use_2cta_instrs: bool = False,
acc_dtype: Type[cutlass.Numeric] = cutlass.Float32,
num_acc_stage: int = 2,
epi_mode: int = 0,
k_mult: int = 4,
num_ab_stage_override: int = 0,
hot_loop: int = 0,
raster: int = 0,
group_m: int = 8,
):
self.raster = raster
self.group_m = group_m
self.hot_loop = hot_loop
self.epi_mode = epi_mode
self.k_mult = k_mult
self.num_ab_stage_override = num_ab_stage_override
self.acc_dtype = acc_dtype
self.use_2cta_instrs = use_2cta_instrs
self.cluster_shape_mn = cluster_shape_mn
self.mma_tiler_mn = mma_tiler_mn
self.mma_tiler = (*mma_tiler_mn, 1)
self.arch = "sm_100"
self.cta_group = tcgen05.CtaGroup.TWO if use_2cta_instrs else tcgen05.CtaGroup.ONE
self.occupancy = 1
self.epilogue_warp_id = (0, 1, 2, 3)
self.mma_warp_id = 4
self.tma_warp_id = 5
self.threads_per_cta = 32 * 6
self.epilog_sync_bar_id = 1
self.tmem_alloc_sync_bar_id = 2
self.tmem_dealloc_sync_bar_id = 3
self.num_acc_stage = acc if False else num_acc_stage # tiles; each tile has (gate, up) fragments
self.num_sched_stages = 4
def _create_tiled_mma(self):
return utils.sm100.make_trivial_tiled_mma(
self.a_dtype, self.a_major_mode, self.b_major_mode,
self.acc_dtype, self.cta_group, self.mma_tiler[:2],
)
def _setup_attributes(self):
tiled_mma = self._create_tiled_mma()
mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2])
self.mma_tiler = (self.mma_tiler[0], self.mma_tiler[1], mma_inst_shape_k * self.k_mult)
self.cta_tile_shape_mnk = (
self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape),
self.mma_tiler[1],
self.mma_tiler[2],
)
self.cluster_layout_vmnk = cute.tiled_divide(
cute.make_layout((*self.cluster_shape_mn, 1)),
(tiled_mma.thr_id.shape,),
)
self.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2])
self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1])
self.is_a_mcast = self.num_mcast_ctas_a > 1
self.is_b_mcast = self.num_mcast_ctas_b > 1
self.epi_tile = utils.sm100.compute_epilogue_tile_shape(
self.cta_tile_shape_mnk, self.use_2cta_instrs, self.c_layout, self.c_dtype,
)
c_smem_layout_1 = utils.sm100.make_smem_layout_epi(self.c_dtype, self.c_layout, self.epi_tile, 1)
self.smem_capacity = utils.get_smem_capacity_in_bytes()
a_stage_one = utils.sm100.make_smem_layout_a(tiled_mma, self.mma_tiler, self.a_dtype, 1)
b_stage_one = utils.sm100.make_smem_layout_b(tiled_mma, self.mma_tiler, self.b_dtype, 1)
a_bytes = cute.size_in_bytes(self.a_dtype, a_stage_one)
b_bytes = cute.size_in_bytes(self.b_dtype, b_stage_one)
ab_bytes_per_stage = a_bytes + 2 * b_bytes # A + gate B + up B
mbar_helpers_bytes = 1024
c_bytes_per_stage = cute.size_in_bytes(self.c_dtype, c_smem_layout_1)
num_c_stage = 2
c_bytes = c_bytes_per_stage * num_c_stage
self.num_ab_stage = (self.smem_capacity // self.occupancy - (mbar_helpers_bytes + c_bytes)) // ab_bytes_per_stage
if self.num_ab_stage_override > 0:
self.num_ab_stage = self.num_ab_stage_override
num_c_stage += (
self.smem_capacity
- self.occupancy * ab_bytes_per_stage * self.num_ab_stage
- self.occupancy * (mbar_helpers_bytes + c_bytes)
) // (self.occupancy * c_bytes_per_stage)
self.num_c_stage = num_c_stage
self.a_smem_layout_staged = utils.sm100.make_smem_layout_a(
tiled_mma, self.mma_tiler, self.a_dtype, self.num_ab_stage
)
self.b_smem_layout_staged = utils.sm100.make_smem_layout_b(
tiled_mma, self.mma_tiler, self.b_dtype, self.num_ab_stage
)
self.c_smem_layout_staged = utils.sm100.make_smem_layout_epi(
self.c_dtype, self.c_layout, self.epi_tile, self.num_c_stage
)
# TMEM columns: acc fragment with 2 tiles x (gate, up)
acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
tCtAcc_fake = tiled_mma.make_fragment_C(cute.append(acc_shape, self.num_acc_stage * 2))
self.num_tmem_alloc_cols = utils.get_num_tmem_alloc_cols(tCtAcc_fake, arch=self.arch)
@cute.jit
def __call__(
self,
a_ptr: cute.Pointer,
bg_ptr: cute.Pointer,
bu_ptr: cute.Pointer,
c_ptr: cute.Pointer,
offs_ptr: cute.Pointer,
m: cutlass.Int32,
n: cutlass.Int32,
k: cutlass.Int32,
expert_cnt: cutlass.Int32,
max_active_clusters: cutlass.Constexpr,
stream: cuda.CUstream,
):
c1 = cutlass.Int32(1)
c0 = cutlass.Int32(0)
a_gemm = cute.make_tensor(
a_ptr, cute.make_layout((m, k, c1), stride=(k, c1, c0))
)
mat_bg = cute.make_tensor(
bg_ptr, cute.make_layout((n, k, expert_cnt), stride=(c1, n, n * k))
)
mat_bu = cute.make_tensor(
bu_ptr, cute.make_layout((n, k, expert_cnt), stride=(c1, n, n * k))
)
c_gemm = cute.make_tensor(
c_ptr, cute.make_layout((m, n, c1), stride=(n, c1, c0))
)
offs = cute.make_tensor(offs_ptr, cute.make_layout((expert_cnt,), stride=(c1,)))
self.a_dtype = a_gemm.element_type
self.b_dtype = mat_bg.element_type
self.c_dtype = c_gemm.element_type
# A: K-major; B: N-major; C: N-major (row-major)
self.a_major_mode = utils.LayoutEnum.ROW_MAJOR.mma_major_mode()
self.b_major_mode = utils.LayoutEnum.COL_MAJOR.mma_major_mode()
self.c_layout = utils.LayoutEnum.ROW_MAJOR
self._setup_attributes()
tiled_mma = self._create_tiled_mma()
a_op = utils.sm100.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id)
a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0))
tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
a_op, a_gemm, a_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape,
)
b_op = utils.sm100.cluster_shape_to_tma_atom_B(self.cluster_shape_mn, tiled_mma.thr_id)
b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0))
tma_atom_bg, tma_tensor_bg = cute.nvgpu.make_tiled_tma_atom_B(
b_op, mat_bg, b_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape,
)
tma_atom_bu, tma_tensor_bu = cute.nvgpu.make_tiled_tma_atom_B(
b_op, mat_bu, b_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape,
)
a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout)
b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout)
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
self.num_tma_load_bytes = (a_copy_size + 2 * b_copy_size) * atom_thr_size
epi_smem_layout = cute.select(self.c_smem_layout_staged, mode=[0, 1])
tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(), c_gemm, epi_smem_layout, self.epi_tile
)
sched_params = MoESchedParams(
expert_cnt, n, k, self.cta_tile_shape_mnk, self.cluster_shape_mn,
raster=self.raster, group_m=self.group_m,
)
grid = (
self.cluster_shape_mn[0],
self.cluster_shape_mn[1],
max_active_clusters,
)
self.kernel(
tiled_mma,
tma_atom_a, tma_tensor_a,
tma_atom_bg, tma_tensor_bg,
tma_atom_bu, tma_tensor_bu,
tma_atom_c, tma_tensor_c,
c_gemm,
offs,
sched_params,
self.cluster_layout_vmnk,
self.a_smem_layout_staged,
self.b_smem_layout_staged,
self.c_smem_layout_staged,
self.epi_tile,
).launch(
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
stream=stream,
min_blocks_per_mp=self.occupancy,
)
@cute.kernel
def kernel(
self,
tiled_mma: cute.TiledMma,
tma_atom_a: cute.CopyAtom,
tma_tensor_a: cute.Tensor,
tma_atom_bg: cute.CopyAtom,
tma_tensor_bg: cute.Tensor,
tma_atom_bu: cute.CopyAtom,
tma_tensor_bu: cute.Tensor,
tma_atom_c: cute.CopyAtom,
tma_tensor_c: cute.Tensor,
c_gmem: cute.Tensor, # plain (M, N, 1) gmem tensor for C
offs: cute.Tensor,
sched_params: MoESchedParams,
cluster_layout_vmnk: cute.Layout,
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
epi_tile: cute.Tile,
):
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
if warp_idx == self.tma_warp_id:
cpasync.prefetch_descriptor(tma_atom_a)
cpasync.prefetch_descriptor(tma_atom_bg)
cpasync.prefetch_descriptor(tma_atom_bu)
cpasync.prefetch_descriptor(tma_atom_c)
use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2
bidx, bidy, bidz = cute.arch.block_idx()
mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
is_leader_cta = mma_tile_coord_v == 0
cta_rank_in_cluster = cute.arch.make_warp_uniform(cute.arch.block_idx_in_cluster())
block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(cta_rank_in_cluster)
tidx, _, _ = cute.arch.thread_idx()
@cute.struct
class SharedStorage:
ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage * 2]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage * 2]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_holding_buf: cutlass.Int32
smem = utils.SmemAllocator()
storage = smem.allocate(SharedStorage)
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
ab_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, num_tma_producer
)
ab_producer, ab_consumer = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
num_stages=self.num_ab_stage,
producer_group=ab_pipeline_producer_group,
consumer_group=ab_pipeline_consumer_group,
tx_count=self.num_tma_load_bytes,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
).make_participants()
acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_acc_consumer_threads = len(self.epilogue_warp_id) * (2 if use_2cta_instrs else 1)
acc_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, num_acc_consumer_threads
)
acc_pipeline = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
num_stages=self.num_acc_stage,
producer_group=acc_pipeline_producer_group,
consumer_group=acc_pipeline_consumer_group,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
)
tmem_alloc_barrier = pipeline.NamedBarrier(
barrier_id=self.tmem_alloc_sync_bar_id,
num_threads=32 * len((self.mma_warp_id, *self.epilogue_warp_id)),
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.epilogue_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr.ptr,
)
pipeline_init_arrive(cluster_shape_mn=cluster_layout_vmnk, is_relaxed=True)
sA = smem.allocate_tensor(
element_type=self.a_dtype,
layout=a_smem_layout_staged.outer,
byte_alignment=128,
swizzle=a_smem_layout_staged.inner,
)
sBg = smem.allocate_tensor(
element_type=self.b_dtype,
layout=b_smem_layout_staged.outer,
byte_alignment=128,
swizzle=b_smem_layout_staged.inner,
)
sBu = smem.allocate_tensor(
element_type=self.b_dtype,
layout=b_smem_layout_staged.outer,
byte_alignment=128,
swizzle=b_smem_layout_staged.inner,
)
sC = smem.allocate_tensor(
element_type=self.c_dtype,
layout=c_smem_layout_staged.outer,
byte_alignment=128,
swizzle=c_smem_layout_staged.inner,
)
a_full_mcast_mask = None
b_full_mcast_mask = None
if const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs):
a_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
)
b_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1
)
a_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape)
b_cta_layout = cute.make_layout(cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape)
tCrA = tiled_mma.make_fragment_A(sA)
tCrBg = tiled_mma.make_fragment_B(sBg)
tCrBu = tiled_mma.make_fragment_B(sBu)
acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
# (MMA, MMA_M, MMA_N, 2*num_acc_stage): stage index = slot*2 + gu
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, self.num_acc_stage * 2)
)
pipeline_init_wait(cluster_shape_mn=cluster_layout_vmnk)
# =================================================================
# TMA load warp
# =================================================================
if warp_idx == self.tma_warp_id:
scheduler = MoETileScheduler.create(
sched_params, offs, cute.arch.block_idx(), cute.arch.grid_dim()
)
work_tile_info = scheduler.initial_work_tile_info()
while work_tile_info.is_valid_tile:
e = work_tile_info.expert_idx
token_offset = work_tile_info.token_offset
tokens_i = work_tile_info.tokens_i
# ---- per-expert partitions ----
real_a = cute.domain_offset((token_offset, 0, 0), tma_tensor_a)
real_a = cute.make_tensor(
real_a.iterator,
cute.make_layout((tokens_i, real_a.shape[1], Int32(1)), stride=real_a.stride),
)
real_bg = cute.domain_offset((0, 0, e), tma_tensor_bg)
real_bg = cute.make_tensor(
real_bg.iterator,
cute.make_layout((real_bg.shape[0], real_bg.shape[1], Int32(1)), stride=real_bg.stride),
)
real_bu = cute.domain_offset((0, 0, e), tma_tensor_bu)
real_bu = cute.make_tensor(
real_bu.iterator,
cute.make_layout((real_bu.shape[0], real_bu.shape[1], Int32(1)), stride=real_bu.stride),
)
gA_mkl = cute.local_tile(real_a, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None))
gBg_nkl = cute.local_tile(real_bg, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None))
gBu_nkl = cute.local_tile(real_bu, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None))
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
tCgA = thr_mma.partition_A(gA_mkl)
tCgBg = thr_mma.partition_B(gBg_nkl)
tCgBu = thr_mma.partition_B(gBu_nkl)
tAsA, tAgA = cpasync.tma_partition(
tma_atom_a,
block_in_cluster_coord_vmnk[2],
a_cta_layout,
cute.group_modes(sA, 0, 3),
cute.group_modes(tCgA, 0, 3),
)
tBgsB, tBgBg = cpasync.tma_partition(
tma_atom_bg,
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sBg, 0, 3),
cute.group_modes(tCgBg, 0, 3),
)
tBusB, tBgBu = cpasync.tma_partition(
tma_atom_bu,
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sBu, 0, 3),
cute.group_modes(tCgBu, 0, 3),
)
# ---- inner loop over this expert's tiles ----
same_expert = Boolean(work_tile_info.expert_idx == e)
while (work_tile_info.is_valid_tile) & same_expert:
k_tile_cnt = work_tile_info.k_tile_cnt
mma_tile_m = work_tile_info.tile_m_idx // cute.size(tiled_mma.thr_id.shape)
if cutlass.const_expr((self.hot_loop & 1) != 0):
mma_tile_m = Int32(0)
tAgA_slice = tAgA[(None, mma_tile_m, None, 0)]
if cutlass.const_expr((self.hot_loop & 2) != 0):
tile_n_hot = Int32(0)
else:
tile_n_hot = work_tile_info.tile_n_idx
tBgBg_slice = tBgBg[(None, tile_n_hot, None, 0)]
tBgBu_slice = tBgBu[(None, tile_n_hot, None, 0)]
ab_producer.reset()
peek_ab_empty_status = ab_producer.try_acquire()
for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
handle = ab_producer.acquire_and_advance(peek_ab_empty_status)
cute.copy(
tma_atom_a, tAgA_slice[(None, handle.count)], tAsA[(None, handle.index)],
tma_bar_ptr=handle.barrier, mcast_mask=a_full_mcast_mask,
)
cute.copy(
tma_atom_bg, tBgBg_slice[(None, handle.count)], tBgsB[(None, handle.index)],
tma_bar_ptr=handle.barrier, mcast_mask=b_full_mcast_mask,
)
cute.copy(
tma_atom_bu, tBgBu_slice[(None, handle.count)], tBusB[(None, handle.index)],
tma_bar_ptr=handle.barrier, mcast_mask=b_full_mcast_mask,
)
peek_ab_empty_status = cutlass.Boolean(1)
if handle.count + 1 < k_tile_cnt:
peek_ab_empty_status = ab_producer.try_acquire()
work_tile_info = scheduler.advance_to_next_work()
same_expert = Boolean(work_tile_info.expert_idx == e)
ab_producer.tail()
# =================================================================
# MMA warp
# =================================================================
if warp_idx == self.mma_warp_id:
tmem.wait_for_alloc()
tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout)
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
scheduler = MoETileScheduler.create(
sched_params, offs, cute.arch.block_idx(), cute.arch.grid_dim()
)
work_tile_info = scheduler.initial_work_tile_info()
num_kblocks = cute.size(tCrA, mode=[2])
while work_tile_info.is_valid_tile:
k_tile_cnt = work_tile_info.k_tile_cnt
if is_leader_cta:
gi = acc_producer_state.index * 2
tCtAcc_g = tCtAcc_base[(None, None, None, gi)]
tCtAcc_u = tCtAcc_base[(None, None, None, gi + 1)]
ab_consumer.reset()
peek_ab_full_status = ab_consumer.try_wait()
acc_pipeline.producer_acquire(acc_producer_state)
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
handle = ab_consumer.wait_and_advance(peek_ab_full_status)
for kblk_idx in cutlass.range(num_kblocks, unroll_full=True):
kblk_crd = (None, None, kblk_idx, handle.index)
cute.gemm(tiled_mma, tCtAcc_g, tCrA[kblk_crd], tCrBg[kblk_crd], tCtAcc_g)
cute.gemm(tiled_mma, tCtAcc_u, tCrA[kblk_crd], tCrBu[kblk_crd], tCtAcc_u)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
handle.release()
peek_ab_full_status = cutlass.Boolean(1)
if handle.count + 1 < k_tile_cnt:
peek_ab_full_status = ab_consumer.try_wait()
acc_pipeline.producer_commit(acc_producer_state)
acc_producer_state.advance()
work_tile_info = scheduler.advance_to_next_work()
acc_pipeline.producer_tail(acc_producer_state)
# =================================================================
# Epilogue warps
# =================================================================
if warp_idx < self.mma_warp_id:
tmem.allocate(self.num_tmem_alloc_cols)
tmem.wait_for_alloc()
tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout)
tCtAcc_transformed = transform_partitioned_tensor_layout(tCtAcc_base)
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, 32 * len(self.epilogue_warp_id)
)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=self.num_c_stage, producer_group=c_producer_group
)
epilog_sync_barrier = pipeline.NamedBarrier(
barrier_id=self.epilog_sync_bar_id,
num_threads=32 * len(self.epilogue_warp_id),
)
num_tiles_executed = cutlass.Int32(0)
scheduler = MoETileScheduler.create(
sched_params, offs, cute.arch.block_idx(), cute.arch.grid_dim()
)
work_tile_info = scheduler.initial_work_tile_info()
while work_tile_info.is_valid_tile:
e = work_tile_info.expert_idx
token_offset = work_tile_info.token_offset
tokens_i = work_tile_info.tokens_i
# ---- per-expert C/epilogue partitions ----
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
real_c_tma = cute.domain_offset((token_offset, 0, 0), tma_tensor_c)
real_c_tma = cute.make_tensor(
real_c_tma.iterator,
cute.make_layout((tokens_i, real_c_tma.shape[1], Int32(1)), stride=real_c_tma.stride),
)
gC_tma_mnl = cute.local_tile(real_c_tma, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None))
tCgC_tma = thr_mma.partition_C(gC_tma_mnl)
tCgC_tma_t = transform_partitioned_tensor_layout(tCgC_tma)
real_c_g = cute.domain_offset((token_offset, 0, 0), c_gmem)
real_c_g = cute.make_tensor(
real_c_g.iterator,
cute.make_layout((tokens_i, real_c_g.shape[1], Int32(1)), stride=real_c_g.stride),
)
gC_g_mnl = cute.local_tile(real_c_g, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None))
tCgC_g = thr_mma.partition_C(gC_g_mnl)
tCgC_g_t = transform_partitioned_tensor_layout(tCgC_g)
cC_mnl = cute.make_identity_tensor(real_c_g.shape)
gCc_mnl = cute.local_tile(cC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None))
tCcC = thr_mma.partition_C(gCc_mnl)
tCcC_t = transform_partitioned_tensor_layout(tCcC)
(tiled_copy_t2r, tTR_tAcc_base, tTR_rAcc) = epilogue_tmem_copy_and_partition(
self, tidx, tCtAcc_transformed, tCgC_tma_t, epi_tile, use_2cta_instrs
)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)
tTR_rAcc_u = cute.make_rmem_tensor(tTR_rAcc.shape, self.acc_dtype)
tCgC_tma_epi = cute.flat_divide(tCgC_tma_t, epi_tile)
bSG_sC, bSG_gC_partitioned = cpasync.tma_partition(
tma_atom_c, 0, cute.make_layout(1),
cute.group_modes(sC, 0, 2),
cute.group_modes(tCgC_tma_epi, 0, 2),
)
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
tCgC_g_epi = cute.flat_divide(tCgC_g_t, epi_tile)
tTR_gC_partitioned = thr_copy_t2r.partition_D(tCgC_g_epi)
cC_epi = cute.flat_divide(tCcC_t, epi_tile)
tTR_cC_partitioned = thr_copy_t2r.partition_D(cC_epi)
tiled_copy_r2s, tRS_rC, tRS_sC = epilogue_smem_copy_and_partition(
self, tiled_copy_t2r, tTR_rC, tidx, sC
)
# ---- inner loop over this expert's tiles ----
same_expert = Boolean(work_tile_info.expert_idx == e)
while (work_tile_info.is_valid_tile) & same_expert:
k_tile_cnt = work_tile_info.k_tile_cnt
tile_m = work_tile_info.tile_m_idx
tile_n = work_tile_info.tile_n_idx
is_full = (tile_m * self.mma_tiler[0] + self.mma_tiler[0]) <= tokens_i
mma_tile_coord_mnl = (
tile_m // cute.size(tiled_mma.thr_id.shape),
tile_n,
Int32(0),
)
bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
tTR_gC = tTR_gC_partitioned[(None, None, None, None, None, *mma_tile_coord_mnl)]
tTR_gC = cute.group_modes(tTR_gC, 3, cute.rank(tTR_gC))
tTR_cC = tTR_cC_partitioned[(None, None, None, None, None, *mma_tile_coord_mnl)]
tTR_cC = cute.group_modes(tTR_cC, 3, cute.rank(tTR_cC))
gi = acc_consumer_state.index * 2
tTR_tAcc_g = tTR_tAcc_base[(None, None, None, None, None, gi)]
tTR_tAcc_u = tTR_tAcc_base[(None, None, None, None, None, gi + 1)]
tTR_tAcc_g = cute.group_modes(tTR_tAcc_g, 3, cute.rank(tTR_tAcc_g))
tTR_tAcc_u = cute.group_modes(tTR_tAcc_u, 3, cute.rank(tTR_tAcc_u))
acc_pipeline.consumer_wait(acc_consumer_state)
subtile_cnt = cute.size(tTR_tAcc_g.shape, mode=[3])
num_prev_subtiles = num_tiles_executed * subtile_cnt
for subtile_idx in cutlass.range(subtile_cnt, unroll_full=True):
if cutlass.const_expr(self.epi_mode <= 1):
cute.copy(tiled_copy_t2r, tTR_tAcc_g[(None, None, None, subtile_idx)], tTR_rAcc)
cute.copy(tiled_copy_t2r, tTR_tAcc_u[(None, None, None, subtile_idx)], tTR_rAcc_u)
if subtile_idx == subtile_cnt - 1:
cute.arch.fence_view_async_tmem_load()
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
if cutlass.const_expr(self.epi_mode == 0):
g_vec = tTR_rAcc.load()
u_vec = tTR_rAcc_u.load()
sig = 1.0 / (1.0 + cute.math.exp2(g_vec * (-LOG2E), fastmath=True))
out_vec = u_vec * (g_vec * sig)
tTR_rC.store(out_vec.to(self.c_dtype))
if is_full:
c_buffer = (num_prev_subtiles + subtile_idx) % self.num_c_stage
cute.copy(tiled_copy_r2s, tRS_rC, tRS_sC[(None, None, None, c_buffer)])
cute.arch.fence_proxy("async.shared", space="cta")
epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilogue_warp_id[0]:
cute.copy(tma_atom_c, bSG_sC[(None, c_buffer)], bSG_gC[(None, subtile_idx)])
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
epilog_sync_barrier.arrive_and_wait()
else:
tTR_gC_sub = tTR_gC[(None, None, None, subtile_idx)]
tTR_cC_sub = tTR_cC[(None, None, None, subtile_idx)]
mcl = cute.max_common_layout(tTR_rC.layout, tTR_gC_sub.layout)
num_bits = min(tTR_gC_sub.iterator.alignment * 8, cute.size(mcl) * self.c_dtype.width, 256)
simt_atom = cute.make_copy_atom(
cute.nvgpu.CopyR2GOp(), self.c_dtype,
num_bits_per_copy=num_bits,
l1c_evict_priority=cutlass.cute.nvgpu.common.CacheEvictionPriority.NO_ALLOCATE,
)
pred_C = cute.make_rmem_tensor((1, *tTR_cC_sub.shape[1:]), Boolean)
for m_idx in cutlass.range(cute.size(tTR_cC_sub.shape[1]), unroll_full=True):
for n_idx in cutlass.range(cute.size(tTR_cC_sub.shape[2]), unroll_full=True):
coord = tTR_cC_sub[(0, m_idx, n_idx)]
pred_C[(0, m_idx, n_idx)] = coord[0] < tokens_i
cute.copy(simt_atom, tTR_rC, tTR_gC_sub, pred=pred_C)
num_tiles_executed = num_tiles_executed + 1
work_tile_info = scheduler.advance_to_next_work()
same_expert = Boolean(work_tile_info.expert_idx == e)
c_pipeline.producer_tail()
tmem.relinquish_alloc_permit()
tmem.free(tmem_ptr)
@cute.jit
def _sched_read(self, sched_pipeline, state, copy_atom, buf_tensor):
sched_pipeline.consumer_wait(state)
rmem = cute.make_rmem_tensor((MoEWorkTileInfo.FIELDS,), Int32)
cute.copy(copy_atom, buf_tensor[(None, state.index)], rmem)
info = MoEWorkTileInfo.from_rmem_tensor(rmem)
cute.arch.fence_acq_rel_cta()
sched_pipeline.consumer_release(state)
return info
# ============================================================================
# Triton fallback kernels
# ============================================================================
_TRITON_OK = False
try:
import triton
import triton.language as tl
_TRITON_OK = True
except Exception: # pragma: no cover
pass
if _TRITON_OK:
@triton.jit
def _scan_tiles_kernel(offsets_ptr, tile_start_ptr, total_ptr, E, BLOCK_M: tl.constexpr, CHUNK: tl.constexpr):
total = tl.zeros((), dtype=tl.int32)
for start in range(0, E, CHUNK):
offs = start + tl.arange(0, CHUNK)
mask = offs < E
s = tl.load(offsets_ptr + offs, mask=mask, other=0)
e = tl.load(offsets_ptr + offs + 1, mask=mask, other=0)
n = tl.where(mask, e - s, 0)
nt = (n + BLOCK_M - 1) // BLOCK_M
c = tl.cumsum(nt, 0)
tl.store(tile_start_ptr + offs, total + c - nt, mask=mask)
total += tl.sum(nt)
tl.store(tile_start_ptr + E, total)
tl.store(total_ptr, total)
@triton.jit
def _fill_tiles_kernel(offsets_ptr, tile_start_ptr, tile_e_ptr, tile_m_ptr,
BLOCK_M: tl.constexpr, SPL: tl.constexpr):
e = tl.program_id(0)
s = tl.load(offsets_ptr + e)
en = tl.load(offsets_ptr + e + 1)
nt = (en - s + BLOCK_M - 1) // BLOCK_M
base = tl.load(tile_start_ptr + e)
for i in range(0, nt, SPL):
offs = i + tl.arange(0, SPL)
mask = offs < nt
tl.store(tile_e_ptr + base + offs, tl.full((SPL,), 0, tl.int32) + e, mask=mask)
tl.store(tile_m_ptr + base + offs, s + offs * BLOCK_M, mask=mask)
@triton.jit
def _moe_up_tma_kernel(
a_ptr, wb_ptr, c_ptr,
offsets_ptr, tile_e_ptr, tile_m_ptr, total_ptr,
H, I2, EH, T_perm, NB,
NUM_SMS: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
WARP_SPECIALIZE: tl.constexpr,
):
start_pid = tl.program_id(0)
total = tl.load(total_ptr) * NB
a_desc = tl.make_tensor_descriptor(a_ptr, shape=[T_perm, H], strides=[H, 1], block_shape=[BLOCK_M, BLOCK_K])
b_desc = tl.make_tensor_descriptor(wb_ptr, shape=[EH, I2], strides=[I2, 1], block_shape=[BLOCK_K, BLOCK_N])
k_tiles = tl.cdiv(H, BLOCK_K)
for work in tl.range(start_pid, total, NUM_SMS, flatten=True, warp_specialize=WARP_SPECIALIZE):
t = work // NB
pid_n = work % NB
e = tl.load(tile_e_ptr + t)
m0 = tl.load(tile_m_ptr + t)
row_end = tl.load(offsets_ptr + e + 1)
n0 = pid_n * BLOCK_N
b_row = e * H
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for ki in range(0, k_tiles):
offs_k = ki * BLOCK_K
a = a_desc.load([m0, offs_k])
b = b_desc.load([b_row + offs_k, n0])
acc = tl.dot(a, b, acc)
accr = tl.reshape(acc, (BLOCK_M, BLOCK_N // 2, 2))
g, u = tl.split(accr)
out = (u * (g * tl.sigmoid(g))).to(tl.bfloat16)
cn0 = pid_n * (BLOCK_N // 2)
offs_m = m0 + tl.arange(0, BLOCK_M)
offs_n = cn0 + tl.arange(0, BLOCK_N // 2)
c_ptrs = c_ptr + offs_m[:, None].to(tl.int64) * (I2 // 2) + offs_n[None, :]
tl.store(c_ptrs, out, mask=(offs_m < row_end)[:, None])
@triton.jit
def _moe_up_generic_kernel(
a_ptr, wg_ptr, wu_ptr, c_ptr,
offsets_ptr, tile_e_ptr, tile_m_ptr, total_ptr,
H, I, T_perm, NB,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
EVEN_K: tl.constexpr, EVEN_N: tl.constexpr,
):
pid = tl.program_id(0)
t = pid // NB
pid_n = pid % NB
total = tl.load(total_ptr)
if t >= total:
return
e = tl.load(tile_e_ptr + t).to(tl.int32)
m0 = tl.load(tile_m_ptr + t)
row_end = tl.load(offsets_ptr + e + 1)
offs_m = m0 + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
m_mask = offs_m < row_end
if EVEN_N:
n_mask = tl.full((BLOCK_N,), True, tl.int1)
else:
n_mask = offs_n < I
e64 = e.to(tl.int64)
b_row0 = e64 * H * I
a_ptrs = a_ptr + offs_m[:, None].to(tl.int64) * H + offs_k[None, :]
bg_ptrs = wg_ptr + b_row0 + offs_k[:, None].to(tl.int64) * I + offs_n[None, :]
bu_ptrs = wu_ptr + b_row0 + offs_k[:, None].to(tl.int64) * I + offs_n[None, :]
acc_g = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_u = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k_tiles = tl.cdiv(H, BLOCK_K)
for k in range(0, k_tiles):
if EVEN_K:
a = tl.load(a_ptrs, mask=m_mask[:, None], other=0.0)
bg = tl.load(bg_ptrs, mask=n_mask[None, :], other=0.0)
bu = tl.load(bu_ptrs, mask=n_mask[None, :], other=0.0)
else:
k_rem = H - k * BLOCK_K
a = tl.load(a_ptrs, mask=m_mask[:, None] & (offs_k[None, :] < k_rem), other=0.0)
bg = tl.load(bg_ptrs, mask=(offs_k[:, None] < k_rem) & n_mask[None, :], other=0.0)
bu = tl.load(bu_ptrs, mask=(offs_k[:, None] < k_rem) & n_mask[None, :], other=0.0)
acc_g = tl.dot(a, bg, acc_g)
acc_u = tl.dot(a, bu, acc_u)
a_ptrs += BLOCK_K
bg_ptrs += BLOCK_K * I
bu_ptrs += BLOCK_K * I
g = acc_g
out = acc_u * (g * tl.sigmoid(g))
c_ptrs = c_ptr + offs_m[:, None].to(tl.int64) * I + offs_n[None, :]
if EVEN_N:
store_mask = m_mask[:, None]
else:
store_mask = m_mask[:, None] & (offs_n[None, :] < I)
tl.store(c_ptrs, out.to(tl.bfloat16), mask=store_mask)
def _alloc_fn(size: int, alignment: int, stream):
return torch.empty(size, device="cuda", dtype=torch.int8)
triton.set_allocator(_alloc_fn)
_NUM_SMS_CACHE = [None]
def _num_sms():
if _NUM_SMS_CACHE[0] is None:
_NUM_SMS_CACHE[0] = torch.cuda.get_device_properties("cuda").multi_processor_count
return _NUM_SMS_CACHE[0]
# Per-(H, I) Triton fast-path configs
_TRITON_FAST_CFG = {
(4096, 1536): (128, 256, 64, 4, 4, True),
(2048, 1024): (128, 256, 64, 8, 4, True),
(2048, 4096): (128, 256, 64, 4, 3, True),
}
# Per-(H, I) CuTeDSL configs: (mma_tiler, cluster, use_2cta, num_acc_stage)
_CUTE_CFG = {
(4096, 1536): ((128, 128), (2, 1), False, 2),
(2048, 1024): ((128, 128), (2, 1), False, 2),
(2048, 4096): ((128, 128), (2, 1), False, 2),
}
_CUTE_CFG_DEFAULT = ((128, 128), (2, 1), False, 2)
class _CuteRunner:
"""Compiled-kernel wrapper. Caches the lightweight Python arg tuple per
pointer set; every launch still goes through the supported compiled() call."""
__slots__ = ("cfg", "compiled", "mac", "arg_cache")
def __init__(self, cfg):
self.cfg = cfg
self.compiled = None
self.mac = None
self.arg_cache = {}
def go(self, hidden, wg, wu, offs32, out, stream_i):
T_perm, H = hidden.shape
key = (
hidden.data_ptr(), wg.data_ptr(), wu.data_ptr(), offs32.data_ptr(),
out.data_ptr(), T_perm, stream_i,
)
args = self.arg_cache.get(key)
if args is None:
if len(self.arg_cache) > 64:
self.arg_cache.clear()
E = wg.shape[0]
I = wg.shape[2]
args = (
make_ptr(cutlass.BFloat16, hidden.data_ptr(), cutlass.AddressSpace.gmem, assumed_align=16),
make_ptr(cutlass.BFloat16, wg.data_ptr(), cutlass.AddressSpace.gmem, assumed_align=16),
make_ptr(cutlass.BFloat16, wu.data_ptr(), cutlass.AddressSpace.gmem, assumed_align=16),
make_ptr(cutlass.BFloat16, out.data_ptr(), cutlass.AddressSpace.gmem, assumed_align=16),
make_ptr(cutlass.Int32, offs32.data_ptr() + 4, cutlass.AddressSpace.gmem, assumed_align=4),
cutlass.Int32(T_perm), cutlass.Int32(I), cutlass.Int32(H), cutlass.Int32(E),
cuda.CUstream(stream_i),
)
if self.compiled is None:
mma_tiler, cluster, use2cta = self.cfg[:3]
kernel = MoeUpSwigluKernel(
mma_tiler_mn=mma_tiler, cluster_shape_mn=cluster, use_2cta_instrs=use2cta,
num_acc_stage=self.cfg[3],
)
self.mac = utils.HardwareInfo().get_max_active_clusters(cluster[0] * cluster[1])
self.compiled = cute.compile(kernel, *args[:-1], self.mac, make_fake_stream())
self.arg_cache[key] = args
self.compiled(*args)
return out
class Model(nn.Module):
"""MoE up-projection with grouped GEMM + fused SwiGLU (see module docstring)."""
def __init__(self, T_total: int, H: int, I: int, E: int, K: int): # noqa: E741
super().__init__()
self.T_total = T_total
self.H = H
self.I = I
self.E = E
self.K = K
self.W_gate = nn.Parameter(torch.empty(E, H, I, dtype=torch.bfloat16))
self.W_up = nn.Parameter(torch.empty(E, H, I, dtype=torch.bfloat16))
nn.init.normal_(self.W_gate, std=0.02)
nn.init.normal_(self.W_up, std=0.02)
self._cute_runner = None
self._cute_enabled = _CUTE_OK
self._wcat = None
self._wcat_key = None
# ---- CuTeDSL path -----------------------------------------------------
def _forward_cute(self, hidden_states, expert_offsets):
T_perm, H = hidden_states.shape
I, E = self.I, self.E
cfg = _CUTE_CFG.get((H, I), _CUTE_CFG_DEFAULT)
if self._cute_runner is None or self._cute_runner.cfg != cfg:
self._cute_runner = _CuteRunner(cfg)
if expert_offsets.dtype == torch.int32:
return self._cute_runner.go(
hidden_states, self.W_gate, self.W_up, expert_offsets,
torch.empty(T_perm, I, dtype=torch.bfloat16, device=hidden_states.device),
torch.cuda.current_stream().cuda_stream,
)
return self._cute_runner.go(
hidden_states, self.W_gate, self.W_up, expert_offsets.to(torch.int32),
torch.empty(T_perm, I, dtype=torch.bfloat16, device=hidden_states.device),
torch.cuda.current_stream().cuda_stream,
)
# ---- Triton paths ------------------------------------------------------
def _interleaved_weights(self) -> torch.Tensor:
key = (
self.W_gate.data_ptr(), self.W_up.data_ptr(),
self.W_gate._version, self.W_up._version,
)
if self._wcat_key != key or self._wcat is None:
E, H, I = self.E, self.H, self.I
cated = torch.cat([self.W_gate.unsqueeze(-1), self.W_up.unsqueeze(-1)], dim=-1)
self._wcat = cated.view(E * H, 2 * I).contiguous()
self._wcat_key = key
return self._wcat
def _forward_triton(self, hidden_states, expert_offsets):
T_perm, H = hidden_states.shape
I, E = self.I, self.E
dev = hidden_states.device
num_sms = _num_sms()
offs_i32 = expert_offsets if expert_offsets.dtype == torch.int32 else expert_offsets.to(torch.int32)
cfg = _TRITON_FAST_CFG.get((H, I))
use_fast = cfg is not None and (H % cfg[2] == 0) and ((2 * I) % cfg[1] == 0)
if use_fast:
BLOCK_M, BLOCK_N, BLOCK_K, warps, stages, ws = cfg
else:
BLOCK_M, BLOCK_N, BLOCK_K, warps, stages = 128, 128, 64, 8, 3
U = triton.cdiv(T_perm, BLOCK_M) + E
out = torch.empty(T_perm, I, dtype=torch.bfloat16, device=dev)
tile_start = torch.empty(E + 1, dtype=torch.int32, device=dev)
total = torch.empty(1, dtype=torch.int32, device=dev)
tile_e = torch.empty(U, dtype=torch.int32, device=dev)
tile_m = torch.empty(U, dtype=torch.int32, device=dev)
_scan_tiles_kernel[(1,)](offs_i32, tile_start, total, E, BLOCK_M=BLOCK_M, CHUNK=128, num_warps=4)
_fill_tiles_kernel[(E,)](offs_i32, tile_start, tile_e, tile_m, BLOCK_M=BLOCK_M, SPL=64, num_warps=1)
if use_fast:
wcat = self._interleaved_weights()
NB = (2 * I) // BLOCK_N
_moe_up_tma_kernel[(min(num_sms, U * NB),)](
hidden_states, wcat, out,
offs_i32, tile_e, tile_m, total,
H, 2 * I, E * H, T_perm, NB,
NUM_SMS=num_sms,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
WARP_SPECIALIZE=ws,
num_warps=warps, num_stages=stages,
)
else:
NBn = triton.cdiv(I, BLOCK_N)
_moe_up_generic_kernel[(U * NBn,)](
hidden_states, self.W_gate, self.W_up, out,
offs_i32, tile_e, tile_m, total,
H, I, T_perm, NBn,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
EVEN_K=(H % BLOCK_K == 0), EVEN_N=(I % BLOCK_N == 0),
num_warps=warps, num_stages=stages,
)
return out
def forward(self, hidden_states: torch.Tensor, expert_offsets: torch.Tensor) -> torch.Tensor:
H = hidden_states.shape[1]
I = self.I
if expert_offsets.dtype != torch.int32 or not expert_offsets.is_contiguous():
expert_offsets = expert_offsets.contiguous().to(torch.int32)
cute_ok = (
self._cute_enabled
and hidden_states.dtype == torch.bfloat16
and self.W_gate.is_contiguous()
and self.W_up.is_contiguous()
and hidden_states.is_contiguous()
and H % 64 == 0
and I % 128 == 0
and (H * 2) % 16 == 0
and (I * 2) % 16 == 0
and expert_offsets.shape[0] == self.E + 1
and expert_offsets.is_contiguous()
)
if cute_ok:
try:
return self._forward_cute(hidden_states, expert_offsets)
except Exception:
self._cute_enabled = False
if _TRITON_OK:
return self._forward_triton(hidden_states, expert_offsets)
raise RuntimeError(f"No usable backend (cute err: {_CUTE_ERR})")
# Module-level shape shims (mirroring reference.py)
T_total = 32768
H = 4096
I = 1536 # noqa: E741
E = 128
K = 8
def _build_routing(T_total: int, E: int, K: int, device: str = "cpu") -> torch.Tensor:
T_perm = T_total * K
base = T_perm // E
rem = T_perm - base * E
counts = torch.full((E,), base, dtype=torch.int32, device=device)
counts[:rem] += 1
offsets = torch.zeros(E + 1, dtype=torch.int32, device=device)
offsets[1:] = torch.cumsum(counts, dim=0)
return offsets
def get_inputs():
T_perm = T_total * K
hidden_states = torch.randn(T_perm, H, dtype=torch.bfloat16) * 0.1
expert_offsets = _build_routing(T_total, E, K)
return [hidden_states, expert_offsets]
def get_init_inputs():
return [T_total, H, I, E, K]
20260715_220809_kinetic-claude_kinetic-0715_06_sonic_moe_swiglu