submission 576432
Koren · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 687 lines, June 9 Researcher Reciprocity License v1.0.
submission_2_aiter.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-576432?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:3abdb224bc85bcade4ce7e7d58e12e4b5a98f719a3d6ca023f33d8b3bc112ba0
license declaredunknown
license concludedunknown
authorsKoren
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
_LOG_PREFIX = "[mxfp4-moe-prof]"Kernel source
submission_2_aiter.py687 lines
import atexit
from collections import deque
import math
import sys
import time
from typing import Dict, Optional
import torch
from task import input_t, output_t
import aiter
from aiter import ActivationType, QuantType, dtypes, fp4_utils, get_hip_quant
from aiter import get_hip_quant as get_quant
from aiter.fused_moe import get_2stage_cfgs, get_inter_dim, get_padded_M, moe_sorting
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
from aiter.jit.utils.torch_guard import torch_compile_guard
# Profiler controls (edit these globals directly in this file).
_SUBMISSION_DEBUG = False
_SUBMISSION_SAMPLE_EVERY = 1
_SUBMISSION_EMIT_EVERY = 99
_SUBMISSION_TAIL_SIZE = 64
_LOG_PREFIX = "[mxfp4-moe-prof]"
_STAGE_NAMES = ("sort", "quant1", "stage1", "quant2", "stage2", "two_stage", "e2e")
_CASE_FIELDS = (
"bs",
"d_hidden",
"d_expert",
"n_routed_experts",
"n_shared_experts",
"n_experts_per_token",
"total_top_k",
"M",
"E_runtime",
"topk_runtime",
"dtype_hidden",
"dtype_w1",
)
_DEBUG = bool(_SUBMISSION_DEBUG)
_SAMPLE_EVERY = max(1, int(_SUBMISSION_SAMPLE_EVERY))
_EMIT_EVERY = max(1, int(_SUBMISSION_EMIT_EVERY))
_TAIL_SIZE = max(1, int(_SUBMISSION_TAIL_SIZE))
_stats: Dict[str, object] = {
"calls_total": 0,
"calls_timed": 0,
"by_case": {},
}
_active_profile: Optional[Dict[str, object]] = None
def _log(msg: str) -> None:
if _DEBUG:
print(f"{_LOG_PREFIX} {msg}", file=sys.stderr, flush=True)
def _sample_std(total: float, sumsq: float, n: int) -> float:
if n <= 1:
return 0.0
variance = (sumsq - (total * total) / float(n)) / float(n - 1)
return math.sqrt(max(variance, 0.0))
def _stage_bucket() -> Dict[str, object]:
return {
"total": 0.0,
"sum_sq": 0.0,
"min": float("inf"),
"max": 0.0,
"tail": deque(maxlen=_TAIL_SIZE),
}
def _new_case_stats() -> Dict[str, object]:
return {
"calls_total": 0,
"calls_timed": 0,
"gpu_us": {stage: _stage_bucket() for stage in _STAGE_NAMES},
"cpu_us": {stage: _stage_bucket() for stage in _STAGE_NAMES},
}
def _cfg_int(config: dict, primary: str, alternate: str, default: int = -1) -> int:
if primary in config:
return int(config[primary])
if alternate in config:
return int(config[alternate])
return default
def _make_case_key(config: dict, hidden_states: torch.Tensor, w1: torch.Tensor, topk_ids: torch.Tensor) -> tuple:
return (
_cfg_int(config, "bs", "bs"),
_cfg_int(config, "d_hidden", "dhidden"),
_cfg_int(config, "d_expert", "dexpert"),
_cfg_int(config, "n_routed_experts", "nroutedexperts"),
_cfg_int(config, "n_shared_experts", "nsharedexperts"),
_cfg_int(config, "n_experts_per_token", "nexpertspertoken"),
_cfg_int(config, "total_top_k", "totaltopk", int(topk_ids.shape[1])),
int(hidden_states.shape[0]),
int(w1.shape[0]),
int(topk_ids.shape[1]),
str(hidden_states.dtype),
str(w1.dtype),
)
def _format_case_key(case_key: tuple) -> str:
parts = []
for idx, field in enumerate(_CASE_FIELDS):
parts.append(f"{field}={case_key[idx]}")
return " ".join(parts)
def _ensure_case_stats(case_key: tuple) -> Dict[str, object]:
by_case = _stats["by_case"]
if case_key not in by_case:
by_case[case_key] = _new_case_stats()
return by_case[case_key]
def _record_case_call(case_key: tuple) -> None:
_stats["calls_total"] = int(_stats["calls_total"]) + 1
case_stats = _ensure_case_stats(case_key)
case_stats["calls_total"] = int(case_stats["calls_total"]) + 1
def _update_stage_bucket(case_stats: Dict[str, object], mode: str, stage: str, value_us: float) -> None:
bucket = case_stats[mode][stage]
bucket["total"] += value_us
bucket["sum_sq"] += value_us * value_us
bucket["min"] = min(bucket["min"], value_us)
bucket["max"] = max(bucket["max"], value_us)
bucket["tail"].append(value_us)
def _begin_profile(case_key: tuple) -> None:
global _active_profile
_active_profile = {
"case_key": case_key,
"cpu_us": {stage: 0.0 for stage in _STAGE_NAMES},
"events": [],
}
def _profile_active() -> bool:
return _DEBUG and _active_profile is not None
def _time_stage(stage: str, fn):
if not _profile_active():
return fn()
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
cpu_t0 = time.perf_counter_ns()
start_event.record()
out = fn()
end_event.record()
cpu_t1 = time.perf_counter_ns()
_active_profile["cpu_us"][stage] += (cpu_t1 - cpu_t0) / 1000.0
_active_profile["events"].append((stage, start_event, end_event))
return out
def _tail_mean_and_std(bucket: Dict[str, object]) -> tuple[float, float]:
tail = bucket["tail"]
tail_n = len(tail)
if tail_n == 0:
return 0.0, 0.0
total = float(sum(tail))
sumsq = float(sum(x * x for x in tail))
return total / float(tail_n), _sample_std(total, sumsq, tail_n)
def _tail_stats(bucket: Dict[str, object]) -> Dict[str, float]:
tail = bucket["tail"]
tail_n = len(tail)
if tail_n == 0:
return {"n": 0.0, "avg": 0.0, "std": 0.0}
avg, std = _tail_mean_and_std(bucket)
return {
"n": float(tail_n),
"avg": avg,
"std": std,
}
def _emit_case_line(case_key: tuple, case_stats: Dict[str, object]) -> None:
if int(case_stats["calls_timed"]) <= 0:
return
gpu = case_stats["gpu_us"]
cpu = case_stats["cpu_us"]
gpu_e2e = _tail_stats(gpu["e2e"])
gpu_sort = _tail_stats(gpu["sort"])
gpu_quant1 = _tail_stats(gpu["quant1"])
gpu_quant2 = _tail_stats(gpu["quant2"])
gpu_stage1 = _tail_stats(gpu["stage1"])
gpu_stage2 = _tail_stats(gpu["stage2"])
gpu_two_stage = _tail_stats(gpu["two_stage"])
cpu_e2e = _tail_stats(cpu["e2e"])
cpu_sort = _tail_stats(cpu["sort"])
cpu_quant1 = _tail_stats(cpu["quant1"])
cpu_quant2 = _tail_stats(cpu["quant2"])
cpu_stage1 = _tail_stats(cpu["stage1"])
cpu_stage2 = _tail_stats(cpu["stage2"])
cpu_two_stage = _tail_stats(cpu["two_stage"])
gpu_quant_avg = gpu_quant1["avg"] + gpu_quant2["avg"]
cpu_quant_avg = cpu_quant1["avg"] + cpu_quant2["avg"]
stage_share = {"sort": 0.0, "quant_total": 0.0, "stage1": 0.0, "stage2": 0.0}
if gpu_e2e["avg"] > 0.0:
stage_share["sort"] = 100.0 * gpu_sort["avg"] / gpu_e2e["avg"]
stage_share["quant_total"] = 100.0 * gpu_quant_avg / gpu_e2e["avg"]
stage_share["stage1"] = 100.0 * gpu_stage1["avg"] / gpu_e2e["avg"]
stage_share["stage2"] = 100.0 * gpu_stage2["avg"] / gpu_e2e["avg"]
ranked = sorted(stage_share.items(), key=lambda item: item[1], reverse=True)
bottleneck, bottleneck_share = ranked[0]
runner_up, runner_up_share = ranked[1]
_log(
f"case[{_format_case_key(case_key)}]\n"
f" calls_total={case_stats['calls_total']} calls_timed={case_stats['calls_timed']} "
f"tail_window_n={int(gpu_e2e['n'])} (all values are tail-window aggregate avg/std in us)\n"
f" GPU avg: e2e={gpu_e2e['avg']:.2f}us\n\ttwo_stage={gpu_two_stage['avg']:.2f}us"
f"\n\tsort={gpu_sort['avg']:.2f}us\n\tquant1={gpu_quant1['avg']:.2f}us quant2={gpu_quant2['avg']:.2f}us "
f"quant_total={gpu_quant_avg:.2f}us\n\tstage1={gpu_stage1['avg']:.2f}us stage2={gpu_stage2['avg']:.2f}us\n"
f"\n GPU std: e2e={gpu_e2e['std']:.2f}us\n\tsort={gpu_sort['std']:.2f}us "
f"\n\tquant_total={math.sqrt(max(gpu_quant1['std'] * gpu_quant1['std'] + gpu_quant2['std'] * gpu_quant2['std'], 0.0)):.2f}us "
f"\n\tstage1={gpu_stage1['std']:.2f}us stage2={gpu_stage2['std']:.2f}us\n"
f"\n\n CPU avg: e2e={cpu_e2e['avg']:.2f}us\n\ttwo_stage={cpu_two_stage['avg']:.2f}us "
f"\n\tsort={cpu_sort['avg']:.2f}us\n\tquant1={cpu_quant1['avg']:.2f}us quant2={cpu_quant2['avg']:.2f}us "
f"quant_total={cpu_quant_avg:.2f}us\n\tstage1={cpu_stage1['avg']:.2f}us stage2={cpu_stage2['avg']:.2f}us\n"
f"\n CPU std: e2e={cpu_e2e['std']:.2f}us\n\tsort={cpu_sort['std']:.2f}us "
f"\n\tquant_total={math.sqrt(max(cpu_quant1['std'] * cpu_quant1['std'] + cpu_quant2['std'] * cpu_quant2['std'], 0.0)):.2f}us "
f"\n\tstage1={cpu_stage1['std']:.2f}us stage2={cpu_stage2['std']:.2f}us\n"
f"\n\n bottleneck_gpu={bottleneck}({bottleneck_share:.1f}% of e2e) "
f"second={runner_up}({runner_up_share:.1f}% of e2e)"
)
def _finalize_profile() -> None:
global _active_profile
ctx = _active_profile
if ctx is None:
return
gpu_us_by_stage = {stage: 0.0 for stage in _STAGE_NAMES}
try:
torch.cuda.synchronize()
for stage, start_event, end_event in ctx["events"]:
gpu_us_by_stage[stage] += start_event.elapsed_time(end_event) * 1000.0
except Exception as err: # keep profiling best-effort only
_log(f"profiling synchronize/error: {err}")
case_key = ctx["case_key"]
case_stats = _ensure_case_stats(case_key)
case_stats["calls_timed"] = int(case_stats["calls_timed"]) + 1
_stats["calls_timed"] = int(_stats["calls_timed"]) + 1
for stage in _STAGE_NAMES:
cpu_value = float(ctx["cpu_us"].get(stage, 0.0))
if cpu_value > 0.0:
_update_stage_bucket(case_stats, "cpu_us", stage, cpu_value)
gpu_value = float(gpu_us_by_stage.get(stage, 0.0))
if gpu_value > 0.0:
_update_stage_bucket(case_stats, "gpu_us", stage, gpu_value)
calls_timed = int(case_stats["calls_timed"])
if calls_timed % _EMIT_EVERY == 0:
_emit_case_line(case_key, case_stats)
_active_profile = None
def _final_report() -> None:
if not _DEBUG:
return
if int(_stats["calls_timed"]) == 0:
_log("no timed calls recorded")
return
by_case = _stats["by_case"]
for case_key, case_stats in sorted(
by_case.items(),
key=lambda entry: int(entry[1]["calls_timed"]),
reverse=True,
):
if int(case_stats["calls_timed"]) > 0:
_emit_case_line(case_key, case_stats)
atexit.register(_final_report)
def fused_moe_2stages(
hidden_states,
w1, # [expert(local_expert:EP), inter_dim*2, dim] N,K
w2, # [expert(local_expert:EP), dim, inter_dim]
topk,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_out,
isG1U1,
block_size_M,
activation=ActivationType.Silu,
quant_type=QuantType.No,
doweight_stage1=False,
# following for quant
q_dtype_a=None,
q_dtype_w=None,
w1_scale=None, # [expert(local_expert:EP), inter_dim, 1]
w2_scale=None, # [expert(local_expert:EP), model_dim, 1]
a1_scale=None, # [expert(local_expert:EP), 1, model_dim]
a2_scale=None, # [expert(local_expert:EP), 1, inter_dim]
num_local_tokens: Optional[torch.tensor] = None,
# following for cktile support
hidden_pad=0,
intermediate_pad=0,
bias1=None,
bias2=None,
metadata=None,
):
token_num, _ = hidden_states.shape
E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
dtype = moe_out.dtype
device = hidden_states.device
is_shuffled = getattr(w1, "is_shuffled", False)
if metadata is None:
metadata = get_2stage_cfgs(
get_padded_M(token_num), # consider token_num > 1024 as prefill
model_dim,
inter_dim,
E,
topk,
dtype,
q_dtype_a,
q_dtype_w,
quant_type,
isG1U1,
activation,
doweight_stage1,
hidden_pad,
intermediate_pad,
is_shuffled,
)
assert quant_type == QuantType.per_1x32, "only per_1x32 quant type is supported"
a1, a1_scale = _time_stage(
"quant1",
lambda: fused_dynamic_mxfp4_quant_moe_sort(
hidden_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=1,
block_size=block_size_M,
),
)
a2 = torch.empty(
(token_num, topk, inter_dim),
dtype=dtype,
device=device,
)
extra_stage1_args = {}
extra_stage2_args = {}
# print(f"\nRunning stage1 {metadata.stage1.func.__name__}\n", file=sys.stderr, flush=True)
a2 = _time_stage(
"stage1",
lambda: metadata.stage1(
a1,
w1,
w2,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
a2,
topk,
block_m=block_size_M,
a1_scale=a1_scale,
w1_scale=(
w1_scale.view(dtypes.fp8_e8m0) if w1.dtype == dtypes.fp4x2 else w1_scale
),
sorted_weights=sorted_weights if doweight_stage1 else None,
**extra_stage1_args,
),
)
a2 = a2.view(-1, inter_dim)
a2, a2_scale = _time_stage(
"quant2",
lambda: fused_dynamic_mxfp4_quant_moe_sort(
a2,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=topk,
block_size=block_size_M,
),
)
a2 = a2.view(token_num, topk, -1)
# print(f"\nRunning stage2 {metadata.stage2.func.__name__}\n", file=sys.stderr, flush=True)
_time_stage(
"stage2",
lambda: metadata.stage2(
a2,
w1,
w2,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
moe_out,
topk,
w2_scale=(
w2_scale.view(dtypes.fp8_e8m0) if w2.dtype == dtypes.fp4x2 else w2_scale
),
a2_scale=a2_scale,
block_m=block_size_M,
sorted_weights=sorted_weights if not doweight_stage1 else None,
**extra_stage2_args,
),
)
return moe_out
def fused_moe_fake_submission(
hidden_states: torch.Tensor,
w1: torch.Tensor, # [expert(local_expert:EP), inter_dim*2, dim] N,K
w2: torch.Tensor, # [expert(local_expert:EP), dim, inter_dim]
topk_weight: torch.Tensor,
topk_ids: torch.Tensor,
expert_mask: Optional[torch.Tensor] = None, # EP
activation: int = ActivationType.Silu.value,
quant_type: int = QuantType.No.value,
doweight_stage1: bool = False,
# following for quant
w1_scale: Optional[torch.Tensor] = None, # [expert(local_expert:EP), inter_dim, 1]
w2_scale: Optional[torch.Tensor] = None, # [expert(local_expert:EP), model_dim, 1]
a1_scale: Optional[torch.Tensor] = None, # [expert(local_expert:EP), 1, model_dim]
a2_scale: Optional[torch.Tensor] = None, # [expert(local_expert:EP), 1, inter_dim]
# following for tuning
block_size_M: int = -1,
num_local_tokens: Optional[torch.Tensor] = None,
moe_sorting_dispatch_policy: bool = 0,
dtype: Optional[torch.dtype] = None,
hidden_pad: int = 0,
intermediate_pad: int = 0,
bias1: Optional[torch.Tensor] = None,
bias2: Optional[torch.Tensor] = None,
) -> torch.Tensor:
device = topk_ids.device
M, topk = topk_ids.shape
dtype = hidden_states.dtype if dtype is None else dtype
model_dim = w2.shape[1]
moe_buf = torch.empty((M, model_dim), dtype=dtype, device=device)
return moe_buf
quant_remap = {QuantType.per_128x128: QuantType.per_1x128}
@torch_compile_guard(gen_fake=fused_moe_fake_submission)
def fused_moe_submission(
hidden_states: torch.Tensor,
w1: torch.Tensor, # [expert(local_expert:EP), inter_dim*2, dim] N,K
w2: torch.Tensor, # [expert(local_expert:EP), dim, inter_dim]
topk_weight: torch.Tensor,
topk_ids: torch.Tensor,
expert_mask: Optional[torch.Tensor] = None, # EP
activation: int = ActivationType.Silu.value,
quant_type: int = QuantType.No.value,
doweight_stage1: bool = False,
# following for quant
w1_scale: Optional[torch.Tensor] = None, # [expert(local_expert:EP), inter_dim, 1]
w2_scale: Optional[torch.Tensor] = None, # [expert(local_expert:EP), model_dim, 1]
a1_scale: Optional[torch.Tensor] = None, # [expert(local_expert:EP), 1, model_dim]
a2_scale: Optional[torch.Tensor] = None, # [expert(local_expert:EP), 1, inter_dim]
# following for tuning
block_size_M: int = -1,
num_local_tokens: Optional[torch.Tensor] = None,
moe_sorting_dispatch_policy: bool = 0,
dtype: Optional[torch.dtype] = None,
hidden_pad: int = 0,
intermediate_pad: int = 0,
bias1: Optional[torch.Tensor] = None,
bias2: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# We do such convert since custom_op schema restriction on block_size_M, and Enum type
activation = ActivationType(activation)
quant_type = QuantType(quant_type)
if block_size_M == -1:
block_size_M = None
"""user API"""
M, topk = topk_ids.shape
E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
assert w1.shape[1] in [
inter_dim,
inter_dim * 2,
], f"Invalid MoE weight: {w1.shape=} {w2.shape=}"
isG1U1 = inter_dim != w1.shape[1]
isShuffled = getattr(w1, "is_shuffled", False)
global_E = E
if expert_mask is not None:
global_E = expert_mask.numel()
dtype = hidden_states.dtype if dtype is None else dtype
assert dtype in [
dtypes.fp16,
dtypes.bf16,
], f"Fused_moe unsupported out dtype: {dtype}"
quant_type = quant_remap.get(quant_type, quant_type)
q_dtype_w = w1.dtype
q_dtype_a = w1.dtype if w1.dtype != torch.uint32 else dtypes.fp8
bf16_fp8_bound = 512
if quant_type == QuantType.per_1x32:
if activation == ActivationType.Swiglu:
if get_gfx() != "gfx950" or M < bf16_fp8_bound:
q_dtype_a = dtypes.bf16
elif M >= bf16_fp8_bound:
q_dtype_a = dtypes.fp8
else:
q_dtype_a = dtypes.fp4x2
metadata = get_2stage_cfgs(
get_padded_M(M), # consider token_num > 1024 as prefill
model_dim,
inter_dim,
E,
topk,
dtype,
q_dtype_a,
q_dtype_w,
quant_type,
isG1U1,
activation,
doweight_stage1,
hidden_pad,
intermediate_pad,
isShuffled,
)
block_size_M = metadata.block_m if block_size_M is None else block_size_M
# Ensure block_size_M is int (metadata.block_m from CSV may be float)
if block_size_M is not None:
block_size_M = int(block_size_M)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = _time_stage(
"sort",
lambda: moe_sorting(
topk_ids,
topk_weight,
global_E,
model_dim,
dtype,
block_size_M,
expert_mask,
num_local_tokens,
moe_sorting_dispatch_policy,
),
)
assert not metadata.run_1stage, "Expected 2 stage kernel"
return _time_stage(
"two_stage",
lambda: fused_moe_2stages(
hidden_states,
w1,
w2,
topk,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_buf,
isG1U1,
block_size_M,
activation=activation,
quant_type=quant_type,
doweight_stage1=doweight_stage1,
q_dtype_a=q_dtype_a,
q_dtype_w=q_dtype_w,
w1_scale=w1_scale,
w2_scale=w2_scale,
a1_scale=a1_scale,
a2_scale=a2_scale,
num_local_tokens=num_local_tokens,
# following for cktile support
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
bias1=bias1,
bias2=bias2,
metadata=metadata,
)
)
def custom_kernel(data: input_t) -> output_t:
"""
Submission template for DeepSeek-R1 MXFP4 MoE kernel.
Input data tuple:
hidden_states: [M, d_hidden] bf16
gate_up_weight: [E, 2*d_expert_pad, d_hidden_pad//2] fp4x2 (raw)
down_weight: [E, d_hidden_pad, d_expert_pad//2] fp4x2 (raw)
gate_up_weight_scale: [E, 2*d_expert_pad, scale_K] e8m0 (raw)
down_weight_scale: [E, d_hidden_pad, scale_K] e8m0 (raw)
gate_up_weight_shuffled: [E, 2*d_expert_pad, d_hidden_pad//2] fp4x2 (shuffled)
down_weight_shuffled: [E, d_hidden_pad, d_expert_pad//2] fp4x2 (shuffled)
gate_up_weight_scale_shuffled:[padded, flat] e8m0 (shuffled)
down_weight_scale_shuffled: [padded, flat] e8m0 (shuffled)
topk_weights: [M, total_top_k] float32
topk_ids: [M, total_top_k] int32
config: dict
Returns:
output: [M, d_hidden] bf16
"""
(
hidden_states,
gate_up_weight,
down_weight,
gate_up_weight_scale,
down_weight_scale,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
) = data
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
case_key = _make_case_key(config, hidden_states, gate_up_weight_shuffled, topk_ids)
_record_case_call(case_key)
should_time = _DEBUG and (int(_stats["calls_total"]) % _SAMPLE_EVERY == 0)
if should_time:
_begin_profile(case_key)
try:
output = _time_stage(
"e2e",
lambda: fused_moe_submission(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
expert_mask=None,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
doweight_stage1=False,
w1_scale=gate_up_weight_scale_shuffled,
w2_scale=down_weight_scale_shuffled,
a1_scale=None,
a2_scale=None,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
block_size_M=-1,
),
)
finally:
if should_time:
_finalize_profile()
return output
scrolls · 687 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON