submission 754773
flower2123 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 356 lines, June 9 Researcher Reciprocity License v1.0.
submission_v4_flower_moe.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754773?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:d82835595aa43523e1cf0fd179aafb4c6761b92c9a137314eb3fd808f026f946
license declaredunknown
license concludedunknown
authorsflower2123
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"a_dtype": "fp4", # flower: activation dtypetile-n = 32
TILE_M, TILE_N = 32, 8 # flower: sort-scatter tileKernel source
submission_v4_flower_moe.py356 lines
# ============================================================
# flower_moe.py — Mixture-of-Experts MXFP4 kernel
# Author: flower
# Description: CK 2-stage pipeline + FlyDSL stage2 + fused
# dynamic MXFP4 quant with per-shape block_m optimization
# for AMD MI355X (gfx950). Supports E=33 and E=257.
# flower's note: block_m tuning key insight — sparse experts
# (E=257) prefer smaller block_m, dense experts (E=33) with
# small K prefer larger block_m to reduce block overhead.
# ============================================================
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
# -- flower: standard imports --
import os # flower: filesystem operations
import functools # flower: partial / lru_cache
import torch # flower: tensor operations
import triton # flower: for cdiv utility
from task import input_t, output_t # flower: competition types
# -- flower: aiter library imports --
import aiter # flower: AMD AI tensor engine
from aiter import ActivationType, QuantType, dtypes # flower: enum types and custom dtypes
from aiter.fused_moe import get_2stage_cfgs, get_padded_M, get_inter_dim # flower: MoE utilities
import aiter.fused_moe as _flower_moe_lib # flower: MoE module reference
import aiter.ops.flydsl.moe_kernels as _flower_fly_kernels # flower: FlyDSL kernel registry
from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import ( # flower: fused quant kernel
_fused_dynamic_mxfp4_quant_moe_sort_kernel,
)
# ============================================================
# flower: register FlyDSL t16x128x128 kernel for stage2
# This tile config handles the post-SwiGLU down-projection
# with atomic accumulation across expert splits.
# ============================================================
_flower_fly_kernels._KERNEL_PARAMS["flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"] = {
"stage": 2, # flower: second stage of 2-stage pipeline
"a_dtype": "fp4", # flower: activation dtype
"b_dtype": "fp4", # flower: weight dtype
"out_dtype": "bf16", # flower: output dtype
"tile_m": 16, # flower: M-tile size
"tile_n": 128, # flower: N-tile size
"tile_k": 128, # flower: K-tile size
"mode": "atomic", # flower: use atomic accumulation
"MPerBlock": 16, # flower: M per CU block
}
# ============================================================
# flower: CK kernel names for stage1
# These are the hardware-optimized CK GEMM kernels that handle
# the gate_up projection with fused SwiGLU activation.
# ============================================================
_FLOWER_CK_S1_M128 = ( # flower: 4-workgroup stage1, M128 tile
"moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3"
"_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_FLOWER_FLY_S2 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic" # flower: FlyDSL stage2
# ============================================================
# flower: per-shape configuration override table
# Maps (cu_num, token, model_dim, inter_dim, expert, topk, ...)
# to optimal (block_m, ksplit, kernel names) combinations.
# flower's trail of thought:
# 1. I first tried to trust the stock aiter table as-is.
# 2. Then I only varied block_m and immediately saw E=257 and E=33
# want opposite behavior because token density per expert flips.
# 3. After that, the stable pattern was to leave the big CK/FlyDSL
# kernels intact and only override the exact shapes that appear in
# the competition benches instead of pretending one heuristic wins
# for every case.
# ============================================================
_flower_shape_overrides = {} # flower: populated below
def _flower_make_cfg_key(num_tokens, expert_inter_dim, num_experts):
"""flower: construct the canonical config lookup key."""
return ( # flower: full key tuple matching aiter's internal format
256, num_tokens, 7168, expert_inter_dim, num_experts, 9,
"ActivationType.Silu", "torch.bfloat16",
"torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
"QuantType.per_1x32", True, False,
)
# flower: E=257 shapes — sparse routing, fewer tokens per expert
_flower_shape_overrides[_flower_make_cfg_key(16, 256, 257)] = { # flower: bs=16, ksplit=2
"block_m": 16, "ksplit": 2, "kernelName1": "", "kernelName2": "", "run_1stage": False,
}
_flower_shape_overrides[_flower_make_cfg_key(128, 256, 257)] = { # flower: bs=128, block_m=32
"block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,
"kernelName2": _FLOWER_FLY_S2, "run_1stage": False,
}
_flower_shape_overrides[_flower_make_cfg_key(512, 256, 257)] = { # flower: bs=512, block_m=32
"block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,
"kernelName2": _FLOWER_FLY_S2, "run_1stage": False,
}
# flower: E=33 shapes — dense routing, more tokens per expert
_flower_shape_overrides[_flower_make_cfg_key(16, 512, 33)] = { # flower: bs=16, sparse with ksplit=2
"block_m": 32, "ksplit": 2, "kernelName1": "", "kernelName2": "", "run_1stage": False,
}
_flower_shape_overrides[_flower_make_cfg_key(128, 512, 33)] = { # flower: bs=128, block_m=32
"block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,
"kernelName2": _FLOWER_FLY_S2, "run_1stage": False,
}
_flower_shape_overrides[_flower_make_cfg_key(512, 512, 33)] = { # flower: bs=512, K=512 dense, block_m=64 wins
"block_m": 64, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,
"kernelName2": _FLOWER_FLY_S2, "run_1stage": False,
}
_flower_shape_overrides[_flower_make_cfg_key(512, 2048, 33)] = { # flower: bs=512, K=2048, block_m=32 amortizes
"block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,
"kernelName2": _FLOWER_FLY_S2, "run_1stage": False,
}
# ============================================================
# flower: pre-allocated buffer caches
# ============================================================
_flower_sort_bufs = {} # flower: moe_sorting output buffers
_flower_quant_bufs = {} # flower: fused quant output buffers
def _flower_get_a2_buffer(num_tokens, top_k, inter_dim, device):
"""flower: get or allocate intermediate activation buffer (stage1 output)."""
cache_key = ("a2", num_tokens, top_k, inter_dim) # flower: unique key
if cache_key not in _flower_sort_bufs: # flower: allocate on first use
_flower_sort_bufs[cache_key] = torch.empty(
(num_tokens, top_k, inter_dim), dtype=torch.bfloat16, device=device,
)
return _flower_sort_bufs[cache_key] # flower: return cached buffer
def _flower_get_sort_buffers(num_tokens, num_experts, top_k, model_dim, blk_m, device):
"""flower: get or allocate all buffers needed for moe_sorting_fwd."""
cache_key = (num_tokens, num_experts, top_k, model_dim, blk_m) # flower: unique key
if cache_key not in _flower_sort_bufs: # flower: allocate on first use
max_padded = num_tokens * top_k + num_experts * blk_m - top_k # flower: max padded token count
max_blocks = (max_padded + blk_m - 1) // blk_m # flower: max block count
_flower_sort_bufs[cache_key] = { # flower: all sorting tensors
"sorted_ids": torch.empty(max_padded, dtype=dtypes.i32, device=device),
"sorted_weights": torch.empty(max_padded, dtype=dtypes.fp32, device=device),
"sorted_expert_ids": torch.empty(max_blocks, dtype=dtypes.i32, device=device),
"num_valid": torch.empty(2, dtype=dtypes.i32, device=device),
"output": torch.empty((num_tokens, model_dim), dtype=torch.bfloat16, device=device),
"a2_placeholder": torch.empty((num_tokens, top_k, 0), dtype=torch.bfloat16, device=device),
}
return _flower_sort_bufs[cache_key] # flower: return cached dict
def _flower_fused_quant(input_tensor, sorted_ids, num_valid, token_count, top_k, blk_m, device):
"""flower: inline fused_dynamic_mxfp4_quant_moe_sort with pre-allocated outputs."""
rows, cols = input_tensor.shape # flower: input dimensions
QUANT_GRP = 32 # flower: MXFP4 group size
TILE_Mx = 128 # flower: quant tile M
TILE_M, TILE_N = 32, 8 # flower: sort-scatter tile
TILE_M_u32, TILE_N_u32 = 16, 4 # flower: packed tile dims
num_scale_cols = triton.cdiv(cols, QUANT_GRP) # flower: scale column count
sorted_len = sorted_ids.shape[0] # flower: padded sorted length
# flower: get or allocate quant output buffers
quant_key = (rows, cols, sorted_len, top_k)
if quant_key not in _flower_quant_bufs:
_flower_quant_bufs[quant_key] = {
"fp4_out": torch.empty((rows, cols // 2), dtype=torch.uint8, device=device),
"scale_out": torch.empty(
(triton.cdiv(sorted_len, TILE_M), triton.cdiv(num_scale_cols, TILE_N),
TILE_N_u32, TILE_M_u32, 4),
dtype=torch.uint8, device=device,
),
}
qb = _flower_quant_bufs[quant_key] # flower: get cached buffers
# flower: compute grid size for fused quant+sort kernel
# flower's trail of thought:
# I did not split quantization and sorting back apart after measuring.
# The separate version was easier to read, but it forced one more trip
# through memory and that showed up immediately on MI355X once token
# count got larger. The fused path is uglier, but the data movement is
# the thing that actually mattered here.
total_pids = (
triton.cdiv(rows, TILE_Mx) * num_scale_cols
+ triton.cdiv(sorted_len, TILE_M) * triton.cdiv(num_scale_cols, TILE_N)
)
# flower: launch the fused quant+sort kernel
_fused_dynamic_mxfp4_quant_moe_sort_kernel[(total_pids,)](
input_tensor, qb["fp4_out"], sorted_ids, num_valid, qb["scale_out"],
rows, cols, num_scale_cols,
*input_tensor.stride(), *qb["fp4_out"].stride(), *qb["scale_out"].stride(),
token_num=token_count, M_i=rows, N_i=num_scale_cols,
MXFP4_QUANT_BLOCK_SIZE=QUANT_GRP, BLOCK_SIZE_Mx=TILE_Mx,
BLOCK_SIZE_M=TILE_M // 2, BLOCK_SIZE_N=TILE_N // 2,
TOPK=top_k,
)
# flower: return with proper dtype views
return (
qb["fp4_out"].view(dtypes.fp4x2), # flower: packed FP4 view
qb["scale_out"].view(dtypes.fp8_e8m0).view(-1, num_scale_cols), # flower: E8M0 scale view
)
# ============================================================
# flower: one-time config injection into aiter's MoE module
# ============================================================
_flower_configs_loaded = False # flower: injection guard
def _flower_load_configs():
"""flower: inject per-shape overrides into aiter's 2-stage config table."""
global _flower_configs_loaded # flower: module-level flag
if _flower_configs_loaded: # flower: already done
return
_flower_configs_loaded = True # flower: mark as done
# flower: load default tuning CSV if not yet loaded
if _flower_moe_lib.cfg_2stages is None:
import pandas as pd # flower: for CSV parsing
from aiter.jit.core import AITER_CONFIGS # flower: config paths
tune_csv = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE # flower: tuning file path
if os.path.exists(tune_csv): # flower: file exists
idx_cols = [ # flower: index columns for the tuning table
"cu_num", "token", "model_dim", "inter_dim", "expert", "topk",
"act_type", "dtype", "q_dtype_a", "q_dtype_w", "q_type",
"use_g1u1", "doweight_stage1",
]
df = pd.read_csv(tune_csv) # flower: read CSV
if "_tag" in df.columns: # flower: filter out tagged rows
df = df[df["_tag"].fillna("") == ""]
_flower_moe_lib.cfg_2stages = df.set_index(idx_cols).to_dict("index") # flower: set as dict
else:
_flower_moe_lib.cfg_2stages = {} # flower: empty table
# flower: merge our per-shape overrides on top of defaults
_flower_moe_lib.cfg_2stages.update(_flower_shape_overrides)
# ============================================================
# flower: competition entry point
# ============================================================
def custom_kernel(data: input_t) -> output_t:
"""flower: MoE kernel — 2-stage CK pipeline with FlyDSL stage2."""
# flower: unpack input data tuple
(
hidden_states, # flower: input activations [M, model_dim]
_raw_w1, _raw_w2, # flower: raw (non-shuffled) weights — unused
_raw_w1s, _raw_w2s, # flower: raw weight scales — unused
gate_up_w_sh, down_w_sh, # flower: shuffled weights
gate_up_sc_sh, down_sc_sh, # flower: shuffled weight scales
routing_weights, # flower: expert routing weights
routing_ids, # flower: expert routing indices
task_config, # flower: shape configuration dict
) = data
# flower: inject per-shape configs into aiter
_flower_load_configs()
# flower: extract dimensions
num_tokens = hidden_states.shape[0] # flower: batch size
top_k = routing_ids.shape[1] # flower: experts per token
device = hidden_states.device # flower: GPU device
# flower: compute padding amounts
hidden_pad = task_config["d_hidden_pad"] - task_config["d_hidden"] # flower: hidden dim padding
expert_pad = task_config["d_expert_pad"] - task_config["d_expert"] # flower: expert dim padding
# flower: get model structure dims
num_experts, model_dim, inter_dim = get_inter_dim(gate_up_w_sh.shape, down_w_sh.shape)
padded_m = get_padded_M(num_tokens) # flower: padded M for alignment
# flower: get 2-stage pipeline metadata (stage1/stage2 functions, block_m, ksplit)
pipeline_meta = get_2stage_cfgs(
padded_m, model_dim, inter_dim, num_experts, top_k,
torch.bfloat16, dtypes.fp4x2, dtypes.fp4x2,
QuantType.per_1x32, True, ActivationType.Silu,
False, hidden_pad, expert_pad, True,
)
blk_m = int(pipeline_meta.block_m) # flower: block_m for this shape
# flower: step 1 — sort tokens by expert assignment
sort_bufs = _flower_get_sort_buffers(
num_tokens, num_experts, top_k, model_dim, blk_m, device,
)
aiter.moe_sorting_fwd( # flower: sort tokens → sorted_ids, weights, expert_ids
routing_ids, routing_weights,
sort_bufs["sorted_ids"], sort_bufs["sorted_weights"],
sort_bufs["sorted_expert_ids"], sort_bufs["num_valid"],
sort_bufs["output"],
num_experts, blk_m, None, None, 0,
)
# flower: prepare weight scale views as E8M0
w1_scale_view = gate_up_sc_sh.view(dtypes.fp8_e8m0) # flower: gate_up scale
w2_scale_view = down_sc_sh.view(dtypes.fp8_e8m0) # flower: down scale
a2_buffer = _flower_get_a2_buffer(num_tokens, top_k, inter_dim, device) # flower: intermediate buf
# flower: step 2 — execute 2-stage pipeline
if pipeline_meta.ksplit > 1:
# flower: BF16 path (ksplit > 1 uses cktile, no fp4 quant of activations)
# flower's trail of thought:
# I kept this branch explicit because my first instinct was to
# force everything through the same FP4 route. That looked cleaner
# in code, but the ksplit>1 cases were happier when I stopped being
# clever and just let cktile keep the activations in BF16.
bf16_input = hidden_states.to(torch.bfloat16) # flower: ensure BF16
a2_result = pipeline_meta.stage1( # flower: gate_up GEMM + SwiGLU
bf16_input, gate_up_w_sh, down_w_sh,
sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],
sort_bufs["num_valid"], a2_buffer, top_k,
block_m=blk_m, a1_scale=None, w1_scale=w1_scale_view, sorted_weights=None,
)
pipeline_meta.stage2( # flower: down GEMM + weighted reduction
a2_result, gate_up_w_sh, down_w_sh,
sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],
sort_bufs["num_valid"], sort_bufs["output"], top_k,
w2_scale=w2_scale_view, a2_scale=None,
block_m=blk_m, sorted_weights=sort_bufs["sorted_weights"],
)
else:
# flower: FP4 path (ksplit=0, quantize activations to MXFP4)
# flower's trail of thought:
# This is the branch I arrived at last: once block_m was settled,
# the next real gain came from accepting one re-quant between the
# two stages instead of carrying a wider intermediate around.
# The code looks more busy, but the benchmark liked the narrower
# traffic pattern more than the tidier BF16-only version.
# flower: quantize input activations
a1_quant, a1_scale = _flower_fused_quant(
hidden_states, sort_bufs["sorted_ids"], sort_bufs["num_valid"],
num_tokens, 1, blk_m, device,
)
# flower: stage1 — gate_up GEMM with quantized activations
a2_result = pipeline_meta.stage1(
a1_quant, gate_up_w_sh, down_w_sh,
sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],
sort_bufs["num_valid"], a2_buffer, top_k,
block_m=blk_m, a1_scale=a1_scale, w1_scale=w1_scale_view, sorted_weights=None,
)
# flower: inter-stage re-quantization (BF16 -> FP4 for stage2 input)
a2_flat = a2_result.view(-1, inter_dim) # flower: flatten for quant
a2_quant, a2_scale = _flower_fused_quant(
a2_flat, sort_bufs["sorted_ids"], sort_bufs["num_valid"],
num_tokens, top_k, blk_m, device,
)
a2_quant = a2_quant.view(num_tokens, top_k, -1) # flower: reshape back
# flower: stage2 — down GEMM + weighted reduction
pipeline_meta.stage2(
a2_quant, gate_up_w_sh, down_w_sh,
sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],
sort_bufs["num_valid"], sort_bufs["output"], top_k,
w2_scale=w2_scale_view, a2_scale=a2_scale,
block_m=blk_m, sorted_weights=sort_bufs["sorted_weights"],
)
return sort_bufs["output"] # flower: return final MoE output
scrolls · 356 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 754384.
+ # ============================================================+ # flower_moe.py — Mixture-of-Experts MXFP4 kernel# Author: flower+ # Description: CK 2-stage pipeline + FlyDSL stage2 + fused+ # dynamic MXFP4 quant with per-shape block_m optimization+ # for AMD MI355X (gfx950). Supports E=33 and E=257.+ # flower's note: block_m tuning key insight — sparse experts+ # (E=257) prefer smaller block_m, dense experts (E=33) with+ # small K prefer larger block_m to reduce block overhead.+ # ============================================================+ #!POPCORN leaderboard amd-moe-mxfp4+ #!POPCORN gpu MI355X- import os- import functools- import torch- import triton- from typing import Dict, Tuple, Optional- from task import input_t, output_t+ # -- flower: standard imports --+ import os # flower: filesystem operations+ import functools # flower: partial / lru_cache+ import torch # flower: tensor operations+ import triton # flower: for cdiv utility+ from task import input_t, output_t # flower: competition types- import aiter- from aiter import ActivationType, QuantType, dtypes- from aiter.fused_moe import (- get_2stage_cfgs, get_padded_M, get_inter_dim,- ck_moe_stage1, cktile_moe_stage1, cktile_moe_stage2,- _flydsl_stage2_wrapper,- )- import aiter.fused_moe as _fmoe- import aiter.ops.flydsl.moe_kernels as _fkern- from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (+ # -- flower: aiter library imports --+ import aiter # flower: AMD AI tensor engine+ from aiter import ActivationType, QuantType, dtypes # flower: enum types and custom dtypes+ from aiter.fused_moe import get_2stage_cfgs, get_padded_M, get_inter_dim # flower: MoE utilities+ import aiter.fused_moe as _flower_moe_lib # flower: MoE module reference+ import aiter.ops.flydsl.moe_kernels as _flower_fly_kernels # flower: FlyDSL kernel registry+ from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import ( # flower: fused quant kernel_fused_dynamic_mxfp4_quant_moe_sort_kernel,)- from aiter.utility import fp4_utils- for _nm, _cfg in [- ("flydsl_moe2_afp4_wfp4_bf16_t32x128x128_atomic",- {"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 32, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 32}),- ("flydsl_moe2_afp4_wfp4_bf16_t32x256x128_atomic",- {"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 32, "tile_n": 256, "tile_k": 128, "mode": "atomic", "MPerBlock": 32}),- ("flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic",- {"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 16, "tile_n": 256, "tile_k": 128, "mode": "atomic", "MPerBlock": 16}),- ("flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",- {"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 16, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 16}),- ]:- _fkern._KERNEL_PARAMS[_nm] = _cfg+ # ============================================================+ # flower: register FlyDSL t16x128x128 kernel for stage2+ # This tile config handles the post-SwiGLU down-projection+ # with atomic accumulation across expert splits.+ # ============================================================+ _flower_fly_kernels._KERNEL_PARAMS["flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"] = {+ "stage": 2, # flower: second stage of 2-stage pipeline+ "a_dtype": "fp4", # flower: activation dtype+ "b_dtype": "fp4", # flower: weight dtype+ "out_dtype": "bf16", # flower: output dtype+ "tile_m": 16, # flower: M-tile size+ "tile_n": 128, # flower: N-tile size+ "tile_k": 128, # flower: K-tile size+ "mode": "atomic", # flower: use atomic accumulation+ "MPerBlock": 16, # flower: M per CU block+ }- _CK_S1_128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"- _CK_S1_32 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"- _FLY_16_128 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"- _FLY_16_256 = "flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic"+ # ============================================================+ # flower: CK kernel names for stage1+ # These are the hardware-optimized CK GEMM kernels that handle+ # the gate_up projection with fused SwiGLU activation.+ # ============================================================+ _FLOWER_CK_S1_M128 = ( # flower: 4-workgroup stage1, M128 tile+ "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3"+ "_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"+ )+ _FLOWER_FLY_S2 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic" # flower: FlyDSL stage2+ # ============================================================+ # flower: per-shape configuration override table+ # Maps (cu_num, token, model_dim, inter_dim, expert, topk, ...)+ # to optimal (block_m, ksplit, kernel names) combinations.+ # flower's trail of thought:+ # 1. I first tried to trust the stock aiter table as-is.+ # 2. Then I only varied block_m and immediately saw E=257 and E=33+ # want opposite behavior because token density per expert flips.+ # 3. After that, the stable pattern was to leave the big CK/FlyDSL+ # kernels intact and only override the exact shapes that appear in+ # the competition benches instead of pretending one heuristic wins+ # for every case.+ # ============================================================+ _flower_shape_overrides = {} # flower: populated below- def _mk(t, i, e):- return (256, t, 7168, i, e, 9, "ActivationType.Silu", "torch.bfloat16",- "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",- "QuantType.per_1x32", True, False)+ def _flower_make_cfg_key(num_tokens, expert_inter_dim, num_experts):+ """flower: construct the canonical config lookup key."""+ return ( # flower: full key tuple matching aiter's internal format+ 256, num_tokens, 7168, expert_inter_dim, num_experts, 9,+ "ActivationType.Silu", "torch.bfloat16",+ "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",+ "QuantType.per_1x32", True, False,+ )- _SHAPE_MAP = {- _mk(16, 512, 33): dict(block_m=32, ksplit=2, kernelName1="", kernelName2="", run_1stage=False),- _mk(128, 512, 33): dict(block_m=64, ksplit=0, kernelName1=_CK_S1_128, kernelName2=_FLY_16_128, run_1stage=False),- _mk(512, 512, 33): dict(block_m=64, ksplit=0, kernelName1=_CK_S1_128, kernelName2=_FLY_16_128, run_1stage=False),- _mk(512, 2048, 33): dict(block_m=64, ksplit=0, kernelName1=_CK_S1_128, kernelName2=_FLY_16_256, run_1stage=False),- _mk(16, 256, 257): dict(block_m=16, ksplit=2, kernelName1="", kernelName2="", run_1stage=False),- _mk(128, 256, 257): dict(block_m=16, ksplit=2, kernelName1="", kernelName2="", run_1stage=False),- _mk(512, 256, 257): dict(block_m=32, ksplit=0, kernelName1=_CK_S1_32, kernelName2=_FLY_16_128, run_1stage=False, use_non_temporal_load=True),++ # flower: E=257 shapes — sparse routing, fewer tokens per expert+ _flower_shape_overrides[_flower_make_cfg_key(16, 256, 257)] = { # flower: bs=16, ksplit=2+ "block_m": 16, "ksplit": 2, "kernelName1": "", "kernelName2": "", "run_1stage": False,}+ _flower_shape_overrides[_flower_make_cfg_key(128, 256, 257)] = { # flower: bs=128, block_m=32+ "block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,+ "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ }+ _flower_shape_overrides[_flower_make_cfg_key(512, 256, 257)] = { # flower: bs=512, block_m=32+ "block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,+ "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ }- _mem = {}+ # flower: E=33 shapes — dense routing, more tokens per expert+ _flower_shape_overrides[_flower_make_cfg_key(16, 512, 33)] = { # flower: bs=16, sparse with ksplit=2+ "block_m": 32, "ksplit": 2, "kernelName1": "", "kernelName2": "", "run_1stage": False,+ }+ _flower_shape_overrides[_flower_make_cfg_key(128, 512, 33)] = { # flower: bs=128, block_m=32+ "block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,+ "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ }+ _flower_shape_overrides[_flower_make_cfg_key(512, 512, 33)] = { # flower: bs=512, K=512 dense, block_m=64 wins+ "block_m": 64, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,+ "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ }+ _flower_shape_overrides[_flower_make_cfg_key(512, 2048, 33)] = { # flower: bs=512, K=2048, block_m=32 amortizes+ "block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,+ "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ }+ # ============================================================+ # flower: pre-allocated buffer caches+ # ============================================================+ _flower_sort_bufs = {} # flower: moe_sorting output buffers+ _flower_quant_bufs = {} # flower: fused quant output buffers- def _sorting_tensors(n_tok, n_exp, k, dim, bm, dev):- tag = ("s", n_tok, n_exp, k, dim, bm)- if tag in _mem:- return _mem[tag]- cap = int(n_tok * k + n_exp * bm - k)- nb = int((cap + bm - 1) // bm)- d = {- "si": torch.empty(cap, dtype=dtypes.i32, device=dev),- "sw": torch.empty(cap, dtype=dtypes.fp32, device=dev),- "se": torch.empty(nb, dtype=dtypes.i32, device=dev),- "nv": torch.empty(2, dtype=dtypes.i32, device=dev),- "ob": torch.empty((n_tok, dim), dtype=torch.bfloat16, device=dev),- }- _mem[tag] = d- return d+ def _flower_get_a2_buffer(num_tokens, top_k, inter_dim, device):+ """flower: get or allocate intermediate activation buffer (stage1 output)."""+ cache_key = ("a2", num_tokens, top_k, inter_dim) # flower: unique key+ if cache_key not in _flower_sort_bufs: # flower: allocate on first use+ _flower_sort_bufs[cache_key] = torch.empty(+ (num_tokens, top_k, inter_dim), dtype=torch.bfloat16, device=device,+ )+ return _flower_sort_bufs[cache_key] # flower: return cached buffer- def _inter_buf(n_tok, k, inter, dev):- tag = ("i", n_tok, k, inter)- if tag in _mem:- return _mem[tag]- t = torch.empty((n_tok, k, inter), dtype=torch.bfloat16, device=dev)- _mem[tag] = t- return t+ def _flower_get_sort_buffers(num_tokens, num_experts, top_k, model_dim, blk_m, device):+ """flower: get or allocate all buffers needed for moe_sorting_fwd."""+ cache_key = (num_tokens, num_experts, top_k, model_dim, blk_m) # flower: unique key+ if cache_key not in _flower_sort_bufs: # flower: allocate on first use+ max_padded = num_tokens * top_k + num_experts * blk_m - top_k # flower: max padded token count+ max_blocks = (max_padded + blk_m - 1) // blk_m # flower: max block count+ _flower_sort_bufs[cache_key] = { # flower: all sorting tensors+ "sorted_ids": torch.empty(max_padded, dtype=dtypes.i32, device=device),+ "sorted_weights": torch.empty(max_padded, dtype=dtypes.fp32, device=device),+ "sorted_expert_ids": torch.empty(max_blocks, dtype=dtypes.i32, device=device),+ "num_valid": torch.empty(2, dtype=dtypes.i32, device=device),+ "output": torch.empty((num_tokens, model_dim), dtype=torch.bfloat16, device=device),+ "a2_placeholder": torch.empty((num_tokens, top_k, 0), dtype=torch.bfloat16, device=device),+ }+ return _flower_sort_bufs[cache_key] # flower: return cached dict- def _qt_alloc(rows, cols, sid_n, k, dev):- tag = ("q", rows, cols, sid_n, k)- if tag in _mem:- return _mem[tag]- BLK, BM, BN, BM2, BN2 = 32, 32, 8, 16, 4- sn = triton.cdiv(cols, BLK)- r = {- "f": torch.empty((rows, cols // 2), dtype=torch.uint8, device=dev),- "s": torch.empty(- (triton.cdiv(sid_n, BM), triton.cdiv(sn, BN), BN2, BM2, 4),- dtype=torch.uint8, device=dev),- }- _mem[tag] = r- return r+ def _flower_fused_quant(input_tensor, sorted_ids, num_valid, token_count, top_k, blk_m, device):+ """flower: inline fused_dynamic_mxfp4_quant_moe_sort with pre-allocated outputs."""+ rows, cols = input_tensor.shape # flower: input dimensions+ QUANT_GRP = 32 # flower: MXFP4 group size+ TILE_Mx = 128 # flower: quant tile M+ TILE_M, TILE_N = 32, 8 # flower: sort-scatter tile+ TILE_M_u32, TILE_N_u32 = 16, 4 # flower: packed tile dims- def _run_quant(x, sid, nval, ntok, k, bm, dev):- rows, cols = x.shape- BLK, BIG, BM, BN = 32, 128, 32, 8- sn = triton.cdiv(cols, BLK)- sid_n = sid.shape[0]- qb = _qt_alloc(rows, cols, sid_n, k, dev)- n_pid = triton.cdiv(rows, BIG) * sn + triton.cdiv(sid_n, BM) * triton.cdiv(sn, BN)- _fused_dynamic_mxfp4_quant_moe_sort_kernel[(n_pid,)](- x, qb["f"], sid, nval, qb["s"],- rows, cols, sn,- *x.stride(), *qb["f"].stride(), *qb["s"].stride(),- token_num=ntok, M_i=rows, N_i=sn,- MXFP4_QUANT_BLOCK_SIZE=BLK, BLOCK_SIZE_Mx=BIG,- BLOCK_SIZE_M=BM // 2, BLOCK_SIZE_N=BN // 2, TOPK=k,+ num_scale_cols = triton.cdiv(cols, QUANT_GRP) # flower: scale column count+ sorted_len = sorted_ids.shape[0] # flower: padded sorted length++ # flower: get or allocate quant output buffers+ quant_key = (rows, cols, sorted_len, top_k)+ if quant_key not in _flower_quant_bufs:+ _flower_quant_bufs[quant_key] = {+ "fp4_out": torch.empty((rows, cols // 2), dtype=torch.uint8, device=device),+ "scale_out": torch.empty(+ (triton.cdiv(sorted_len, TILE_M), triton.cdiv(num_scale_cols, TILE_N),+ TILE_N_u32, TILE_M_u32, 4),+ dtype=torch.uint8, device=device,+ ),+ }+ qb = _flower_quant_bufs[quant_key] # flower: get cached buffers++ # flower: compute grid size for fused quant+sort kernel+ # flower's trail of thought:+ # I did not split quantization and sorting back apart after measuring.+ # The separate version was easier to read, but it forced one more trip+ # through memory and that showed up immediately on MI355X once token+ # count got larger. The fused path is uglier, but the data movement is+ # the thing that actually mattered here.+ total_pids = (+ triton.cdiv(rows, TILE_Mx) * num_scale_cols+ + triton.cdiv(sorted_len, TILE_M) * triton.cdiv(num_scale_cols, TILE_N))- return qb["f"].view(dtypes.fp4x2), qb["s"].view(dtypes.fp8_e8m0).view(-1, sn)+ # flower: launch the fused quant+sort kernel+ _fused_dynamic_mxfp4_quant_moe_sort_kernel[(total_pids,)](+ input_tensor, qb["fp4_out"], sorted_ids, num_valid, qb["scale_out"],+ rows, cols, num_scale_cols,+ *input_tensor.stride(), *qb["fp4_out"].stride(), *qb["scale_out"].stride(),+ token_num=token_count, M_i=rows, N_i=num_scale_cols,+ MXFP4_QUANT_BLOCK_SIZE=QUANT_GRP, BLOCK_SIZE_Mx=TILE_Mx,+ BLOCK_SIZE_M=TILE_M // 2, BLOCK_SIZE_N=TILE_N // 2,+ TOPK=top_k,+ )- _loaded = False+ # flower: return with proper dtype views+ return (+ qb["fp4_out"].view(dtypes.fp4x2), # flower: packed FP4 view+ qb["scale_out"].view(dtypes.fp8_e8m0).view(-1, num_scale_cols), # flower: E8M0 scale view+ )- def _setup():- global _loaded- if _loaded:+ # ============================================================+ # flower: one-time config injection into aiter's MoE module+ # ============================================================+ _flower_configs_loaded = False # flower: injection guard+++ def _flower_load_configs():+ """flower: inject per-shape overrides into aiter's 2-stage config table."""+ global _flower_configs_loaded # flower: module-level flag+ if _flower_configs_loaded: # flower: already donereturn- _loaded = True+ _flower_configs_loaded = True # flower: mark as done- if _fmoe.cfg_2stages is None:- import pandas as pd- from aiter.jit.core import AITER_CONFIGS- fp = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE- if os.path.exists(fp):- cols = ["cu_num", "token", "model_dim", "inter_dim", "expert", "topk",- "act_type", "dtype", "q_dtype_a", "q_dtype_w", "q_type",- "use_g1u1", "doweight_stage1"]- df = pd.read_csv(fp)- if "_tag" in df.columns:+ # flower: load default tuning CSV if not yet loaded+ if _flower_moe_lib.cfg_2stages is None:+ import pandas as pd # flower: for CSV parsing+ from aiter.jit.core import AITER_CONFIGS # flower: config paths+ tune_csv = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE # flower: tuning file path+ if os.path.exists(tune_csv): # flower: file exists+ idx_cols = [ # flower: index columns for the tuning table+ "cu_num", "token", "model_dim", "inter_dim", "expert", "topk",+ "act_type", "dtype", "q_dtype_a", "q_dtype_w", "q_type",+ "use_g1u1", "doweight_stage1",+ ]+ df = pd.read_csv(tune_csv) # flower: read CSV+ if "_tag" in df.columns: # flower: filter out tagged rowsdf = df[df["_tag"].fillna("") == ""]- _fmoe.cfg_2stages = df.set_index(cols).to_dict("index")+ _flower_moe_lib.cfg_2stages = df.set_index(idx_cols).to_dict("index") # flower: set as dictelse:- _fmoe.cfg_2stages = {}+ _flower_moe_lib.cfg_2stages = {} # flower: empty table- _fmoe.cfg_2stages.update(_SHAPE_MAP)+ # flower: merge our per-shape overrides on top of defaults+ _flower_moe_lib.cfg_2stages.update(_flower_shape_overrides)- orig = _fmoe.get_2stage_cfgs- @functools.lru_cache(maxsize=2048)- def _hook(token, model_dim, inter_dim, expert, topk,- dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,- activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled=True):- md = orig(token, model_dim, inter_dim, expert, topk,- dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,- activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled)- from aiter.jit.utils.chip_info import get_cu_num- lk = (get_cu_num(), token, model_dim, inter_dim, expert, topk,- str(activation), str(dtype), str(q_dtype_a), str(q_dtype_w),- str(q_type), use_g1u1, doweight_stage1)- entry = _fmoe.cfg_2stages.get(lk)- if entry and entry.get("use_non_temporal_load") is not None:- nt = entry["use_non_temporal_load"]- s1 = md.stage1- if hasattr(s1, 'func') and s1.func is not None and 'use_non_temporal_load' in (s1.keywords or {}):- kw = dict(s1.keywords); kw['use_non_temporal_load'] = nt- md = _fmoe.MOEMetadata(- functools.partial(s1.func, **{k: v for k, v in kw.items()}),- md.stage2, md.block_m, md.ksplit, md.run_1stage, md.has_bias, nt)- s2 = md.stage2- if s2 and hasattr(s2, 'keywords') and 'use_non_temporal_load' in (s2.keywords or {}):- kw2 = dict(s2.keywords); kw2['use_non_temporal_load'] = nt- md = _fmoe.MOEMetadata(- md.stage1,- functools.partial(s2.func, **{k: v for k, v in kw2.items()}),- md.block_m, md.ksplit, md.run_1stage, md.has_bias, nt)- return md+ # ============================================================+ # flower: competition entry point+ # ============================================================+ def custom_kernel(data: input_t) -> output_t:+ """flower: MoE kernel — 2-stage CK pipeline with FlyDSL stage2."""+ # flower: unpack input data tuple+ (+ hidden_states, # flower: input activations [M, model_dim]+ _raw_w1, _raw_w2, # flower: raw (non-shuffled) weights — unused+ _raw_w1s, _raw_w2s, # flower: raw weight scales — unused+ gate_up_w_sh, down_w_sh, # flower: shuffled weights+ gate_up_sc_sh, down_sc_sh, # flower: shuffled weight scales+ routing_weights, # flower: expert routing weights+ routing_ids, # flower: expert routing indices+ task_config, # flower: shape configuration dict+ ) = data- _fmoe.get_2stage_cfgs = _hook+ # flower: inject per-shape configs into aiter+ _flower_load_configs()+ # flower: extract dimensions+ num_tokens = hidden_states.shape[0] # flower: batch size+ top_k = routing_ids.shape[1] # flower: experts per token+ device = hidden_states.device # flower: GPU device- def custom_kernel(data: input_t) -> output_t:- (h, gu_w, d_w, gu_sc, d_sc, gu_ws, d_ws, gu_scs, d_scs, tw, ti, cfg) = data- _setup()+ # flower: compute padding amounts+ hidden_pad = task_config["d_hidden_pad"] - task_config["d_hidden"] # flower: hidden dim padding+ expert_pad = task_config["d_expert_pad"] - task_config["d_expert"] # flower: expert dim padding- hp = cfg["d_hidden_pad"] - cfg["d_hidden"]- ip = cfg["d_expert_pad"] - cfg["d_expert"]- n = h.shape[0]- k = ti.shape[1]- dev = ti.device- ne, mdim, idim = get_inter_dim(gu_ws.shape, d_ws.shape)+ # flower: get model structure dims+ num_experts, model_dim, inter_dim = get_inter_dim(gate_up_w_sh.shape, down_w_sh.shape)+ padded_m = get_padded_M(num_tokens) # flower: padded M for alignment- md = get_2stage_cfgs(- get_padded_M(n), mdim, idim, ne, k,+ # flower: get 2-stage pipeline metadata (stage1/stage2 functions, block_m, ksplit)+ pipeline_meta = get_2stage_cfgs(+ padded_m, model_dim, inter_dim, num_experts, top_k,torch.bfloat16, dtypes.fp4x2, dtypes.fp4x2,QuantType.per_1x32, True, ActivationType.Silu,- False, hp, ip, True)+ False, hidden_pad, expert_pad, True,+ )+ blk_m = int(pipeline_meta.block_m) # flower: block_m for this shape- bm = int(md.block_m)- sb = _sorting_tensors(n, ne, k, mdim, bm, dev)- aiter.moe_sorting_fwd(ti, tw, sb["si"], sb["sw"], sb["se"], sb["nv"], sb["ob"], ne, bm, None, None, 0)+ # flower: step 1 — sort tokens by expert assignment+ sort_bufs = _flower_get_sort_buffers(+ num_tokens, num_experts, top_k, model_dim, blk_m, device,+ )+ aiter.moe_sorting_fwd( # flower: sort tokens → sorted_ids, weights, expert_ids+ routing_ids, routing_weights,+ sort_bufs["sorted_ids"], sort_bufs["sorted_weights"],+ sort_bufs["sorted_expert_ids"], sort_bufs["num_valid"],+ sort_bufs["output"],+ num_experts, blk_m, None, None, 0,+ )- w1s = gu_scs.view(dtypes.fp8_e8m0)- w2s = d_scs.view(dtypes.fp8_e8m0)+ # flower: prepare weight scale views as E8M0+ w1_scale_view = gate_up_sc_sh.view(dtypes.fp8_e8m0) # flower: gate_up scale+ w2_scale_view = down_sc_sh.view(dtypes.fp8_e8m0) # flower: down scale+ a2_buffer = _flower_get_a2_buffer(num_tokens, top_k, inter_dim, device) # flower: intermediate buf- if md.ksplit > 1:- a = h.to(torch.bfloat16)- inter = md.stage1(a, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],- _inter_buf(n, k, idim, dev), k,- block_m=bm, a1_scale=None, w1_scale=w1s, sorted_weights=None)- md.stage2(inter, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],- sb["ob"], k, w2_scale=w2s, a2_scale=None, block_m=bm, sorted_weights=sb["sw"])+ # flower: step 2 — execute 2-stage pipeline+ if pipeline_meta.ksplit > 1:+ # flower: BF16 path (ksplit > 1 uses cktile, no fp4 quant of activations)+ # flower's trail of thought:+ # I kept this branch explicit because my first instinct was to+ # force everything through the same FP4 route. That looked cleaner+ # in code, but the ksplit>1 cases were happier when I stopped being+ # clever and just let cktile keep the activations in BF16.+ bf16_input = hidden_states.to(torch.bfloat16) # flower: ensure BF16+ a2_result = pipeline_meta.stage1( # flower: gate_up GEMM + SwiGLU+ bf16_input, gate_up_w_sh, down_w_sh,+ sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],+ sort_bufs["num_valid"], a2_buffer, top_k,+ block_m=blk_m, a1_scale=None, w1_scale=w1_scale_view, sorted_weights=None,+ )+ pipeline_meta.stage2( # flower: down GEMM + weighted reduction+ a2_result, gate_up_w_sh, down_w_sh,+ sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],+ sort_bufs["num_valid"], sort_bufs["output"], top_k,+ w2_scale=w2_scale_view, a2_scale=None,+ block_m=blk_m, sorted_weights=sort_bufs["sorted_weights"],+ )else:- qa, qas = _run_quant(h, sb["si"], sb["nv"], n, 1, bm, dev)- inter = _inter_buf(n, k, idim, dev)- inter = md.stage1(qa, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],- inter, k, block_m=bm, a1_scale=qas, w1_scale=w1s, sorted_weights=None)- flat = inter.view(-1, idim)- qi, qis = _run_quant(flat, sb["si"], sb["nv"], n, k, bm, dev)- qi = qi.view(n, k, -1)- md.stage2(qi, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],- sb["ob"], k, w2_scale=w2s, a2_scale=qis, block_m=bm, sorted_weights=sb["sw"])+ # flower: FP4 path (ksplit=0, quantize activations to MXFP4)+ # flower's trail of thought:+ # This is the branch I arrived at last: once block_m was settled,+ # the next real gain came from accepting one re-quant between the+ # two stages instead of carrying a wider intermediate around.+ # The code looks more busy, but the benchmark liked the narrower+ # traffic pattern more than the tidier BF16-only version.+ # flower: quantize input activations+ a1_quant, a1_scale = _flower_fused_quant(+ hidden_states, sort_bufs["sorted_ids"], sort_bufs["num_valid"],+ num_tokens, 1, blk_m, device,+ )+ # flower: stage1 — gate_up GEMM with quantized activations+ a2_result = pipeline_meta.stage1(+ a1_quant, gate_up_w_sh, down_w_sh,+ sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],+ sort_bufs["num_valid"], a2_buffer, top_k,+ block_m=blk_m, a1_scale=a1_scale, w1_scale=w1_scale_view, sorted_weights=None,+ )+ # flower: inter-stage re-quantization (BF16 -> FP4 for stage2 input)+ a2_flat = a2_result.view(-1, inter_dim) # flower: flatten for quant+ a2_quant, a2_scale = _flower_fused_quant(+ a2_flat, sort_bufs["sorted_ids"], sort_bufs["num_valid"],+ num_tokens, top_k, blk_m, device,+ )+ a2_quant = a2_quant.view(num_tokens, top_k, -1) # flower: reshape back+ # flower: stage2 — down GEMM + weighted reduction+ pipeline_meta.stage2(+ a2_quant, gate_up_w_sh, down_w_sh,+ sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],+ sort_bufs["num_valid"], sort_bufs["output"], top_k,+ w2_scale=w2_scale_view, a2_scale=a2_scale,+ block_m=blk_m, sorted_weights=sort_bufs["sorted_weights"],+ )- return sb["ob"]+ return sort_bufs["output"] # flower: return final MoE output
scrolls · 533 diff lines total
Best evidence level for this revision: reported
JSON