Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
108.7µs
#16 of 782
2026-04-07

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 dtype
tile-n = 32TILE_M, TILE_N = 32, 8 # flower: sort-scatter tile

Kernel 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 done
return
- _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 rows
df = 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 dict
else:
- _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