[Test] Add Kimi SM120 MXFP4 correctness matrix

This commit is contained in:
Zhiyi Hong 2026-08-17 12:08:06 +08:00
parent dac1bb652d
commit a1c18d736b
7 changed files with 1259 additions and 10 deletions

View File

@ -1,6 +1,6 @@
# sskj — 多平台大模型推理性能基准测试项目 # sskj — 多平台大模型推理性能基准测试项目
**更新2026-08-14 17:08:00 CST**:完成 Kimi-K3 / RTX PRO 6000D / SGLang FlashInfer MXFP4 MoE Phase 1 数据契约审计;已从真实 checkpoint 追到 SM120 CUTLASS 调用,确认 SiTU 激活、gate/up 权重变换和参数透传三处独立缺口,并在 6000D 上复现 C++ binding 的门控 shape 首错,详见 `experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/README.md` **更新2026-08-17 11:56:47 CST**:完成 Kimi-K3 / RTX PRO 6000D / SGLang FlashInfer MXFP4 MoE Phase 2 correctness harnessSwiGLU 控制矩阵 8/8 通过,覆盖 edge routing、EP global ID、NaN/Inf、确定性、CUDA Graph 和 Kimi 真尺寸SiTU 目标矩阵 8/8 稳定复现同一 C++ 缺口,详见 `experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/README.md`
> 历史更新见 `git log`。项目目的与工作流见下方。 > 历史更新见 `git log`。项目目的与工作流见下方。
> **项目目的**当新显卡GPU/NPU到货时用最短时间完成大模型在该平台上的推理性能评估与部署配置选型。 > **项目目的**当新显卡GPU/NPU到货时用最短时间完成大模型在该平台上的推理性能评估与部署配置选型。

View File

@ -231,15 +231,93 @@ FLASHINFER_DISABLE_VERSION_CHECK=1 \
4. 不把旧 `flashinfer_cutedsl` 实验当作本路径证据。 4. 不把旧 `flashinfer_cutedsl` 实验当作本路径证据。
5. 不进入整模型性能测试,直到 Phase 2 的 BF16/PyTorch 或 Marlin 参考正确性通过。 5. 不进入整模型性能测试,直到 Phase 2 的 BF16/PyTorch 或 Marlin 参考正确性通过。
## 下一阶段硬门槛 ## Phase 2 correctness harness
Phase 2 建立独立、对称的 correctness harness至少覆盖 状态:**完成**。本阶段建立验证框架并校准 reference但 SiTU candidate 按预期尚未通过;未产生性能晋级结论。
- token 数 1/2/8/32 以及 Kimi 真实 prefill/decode 代表 shape 唯一入口:
- 均匀、倾斜、空 expert 和 masked/global expert id
- SiTU `4/25`、MXFP8 activation、MXFP4 packed weight 与 E8M0 scale
- NaN/Inf、确定性、CUDA Graph capture/replay
- TP1/EP1 micro case 与 TP32/EP32 的 local-expert 映射;
- candidate 与 reference 使用同一 routing、同一量化输入和同一 finalize 口径。
正确性未通过前,不做吞吐结论,不修改 checkpoint不使用运行时 monkeypatch。 ```text
validate_kimi_sm120_mxfp4.py
```
### 对称性约束
- reference 与 candidate 共用同一 `Case`、随机种子、routing IDs/weights、量化后的 MXFP8 input、packed MXFP4 weights 和 E8M0 scalesreference 不读取量化前的 BF16 权重。
- FC1 在生成时显式采用 CUTLASS oracle 布局 `[up; gate]`。这使 harness 同时约束 SGLang 后续必须把 Kimi checkpoint 的 `[gate; up]` 正确交换。
- SwiGLU 作为当前 kernel 已支持的控制组SiTU `(4, 25)` 使用完全相同的数据与调用,只改变 activation contract。
- correctness 固定 `use_fused_finalize=False`,避免非结合 atomic reduction 干扰确定性benchmark 也由同一个 CUDA Event 计时函数驱动,但本阶段设置 iterations=0不在共享 GPU 负载下生成性能数字。
覆盖矩阵:
| 类别 | 覆盖 |
|---|---|
| token shape | `M=1/2/8/32` |
| routing | uniform、skewed、空 expert、`-1` masked slot |
| EP | `EP2` global/local ID 过滤Kimi `EP32/rank7/E_local=28` |
| 数值 | finite、NaN/Inf row-level propagation、三次 bitwise determinism |
| execution | eager、CUDA Graph capture、两次 replay |
| Kimi 真尺寸 | `H=3584, N=3072, E_local=28, topK=16`decode `M=1`、prefill `M=128` |
### 执行命令
```bash
cd /data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4
# 已支持的 SwiGLU 控制组;本阶段最终证据统一使用 GPU 6
CUDA_VISIBLE_DEVICES=6 PYTHONPATH=/tmp/hzy-fi-main-phase1 \
FLASHINFER_DISABLE_VERSION_CHECK=1 \
/data/hzy/envs/sparse-opd-sm120/bin/python validate_kimi_sm120_mxfp4.py \
--suite smoke --activation swiglu \
--output results/phase2_control_swiglu_smoke.json
# Kimi SiTU 目标真尺寸
CUDA_VISIBLE_DEVICES=6 PYTHONPATH=/tmp/hzy-fi-main-phase1 \
FLASHINFER_DISABLE_VERSION_CHECK=1 \
/data/hzy/envs/sparse-opd-sm120/bin/python validate_kimi_sm120_mxfp4.py \
--suite real --activation situ \
--output results/phase2_target_situ_real_gpu6.json
```
实现完成后的 CI-style gate 必须增加 `--require-situ`;若仍命中已知缺口,入口返回非零。
### 结果与证据
| 组别 | 结果 | 关键证据 |
|---|---|---|
| SwiGLU smoke control | 6/6 PASS | cosine `0.9996250.999812`NRMSE `1.50%2.75%` |
| SwiGLU Kimi real shape | 2/2 PASS | decode/prefill cosine `0.999597/0.999621`NRMSE `2.85%/2.75%` |
| SiTU smoke target | 6/6 EXPECTED_MISSING_SITU | 全部命中 binding `256 vs. 128` 首错 |
| SiTU Kimi real shape | 2/2 EXPECTED_MISSING_SITU | GPU6 上 decode/prefill 均命中同一首错 |
| determinism | PASS | non-fused finalize 三次 bitwise equal |
| CUDA Graph | PASS | control replay 与 eager `max_abs=0`,两次 replay bitwise equal |
原始证据:
```text
/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/phase2_control_swiglu_smoke_gpu6.log
/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/phase2_target_situ_smoke_gpu6.log
/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/phase2_control_swiglu_real_gpu6.log
/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/phase2_target_situ_real_gpu6.log
/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/results/phase2_*.json
```
### NaN/Inf 决策记录
MXFP8 quantize/dequantize 后,注入的 NaN 被量化为有限值,而 Inf 保留。相同 Inf 行在 PyTorch GEMM 与 CUTLASS reduction 中产生不同的逐元素 NaN/Inf 混合,但两者影响同一 token row未受影响行继续满足普通数值阈值。IEEE reduction 顺序不保证逐元素 NaN/Inf 分类一致,因此 harness 要求:
1. 相同 token rows 被标记为 non-finite
2. 未受影响 rows 继续满足 `rtol=0.2, atol=0.2`
3. candidate 不崩溃,且三次结果确定。
这不是放宽有限值正确性,有限 case 仍采用与 FlashInfer 官方 MXFP8×MXFP4 测试一致的阈值。
## Phase 3 硬门槛
下一阶段只做最小实现:
1. FlashInfer C++ 增加与 Python 对齐的 `Situ` enum、gated classification 和 SiTU adaptor复用已有 per-expert `alpha/beta` 参数与 grouped GEMM。
2. SGLang SM120 processor 按 `gate_up_interleaved=False` 处理 Kimi 连续半区,生成 CUTLASS `[up; gate]`,不得复用 GPT-OSS 奇偶行拆分。
3. SGLang runner 从 `runner_config` 透传 `ActivationType.Situ` 和 Kimi `4/25` 参数。
4. 未满足支持条件时在服务初始化阶段明确选择 Marlin不得在 kernel 失败后热切换,不修改 checkpoint不使用 runtime monkeypatch。
5. 完成后必须以 `--require-situ --suite all` 通过 Phase 2 全矩阵,才能进入任何吞吐比较。

View File

@ -0,0 +1,81 @@
{
"arguments": {
"activation": "swiglu",
"benchmark_iterations": 0,
"output": "results/phase2_control_swiglu_real_gpu6.json",
"require_situ": false,
"seed": 20260817,
"suite": "real",
"warmup": 3
},
"compute_capability": [
12,
0
],
"created_at": "2026-08-17T12:02:07.099079+08:00",
"cuda_version": "13.0",
"cuda_visible_devices": "6",
"device": "NVIDIA RTX 6000D",
"flashinfer_version": "0.6.18",
"logical_cuda_device": 0,
"results": [
{
"activation": "swiglu",
"case": {
"ep_rank": 7,
"ep_size": 32,
"hidden": 3584,
"inject_nonfinite": false,
"intermediate": 3072,
"local_experts": 28,
"name": "kimi_decode_m1_ep32",
"routing": "ep_mixed",
"tokens": 1,
"top_k": 16
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 13.75742506980896,
"metrics": {
"cosine": 0.9995965957641602,
"finite_elements": 3584,
"max_abs": 0.01025390625,
"nrmse": 0.028520602732896805,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"activation": "swiglu",
"case": {
"ep_rank": 7,
"ep_size": 32,
"hidden": 3584,
"inject_nonfinite": false,
"intermediate": 3072,
"local_experts": 28,
"name": "kimi_prefill_m128_ep32",
"routing": "ep_mixed",
"tokens": 128,
"top_k": 16
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 12.48392629623413,
"metrics": {
"cosine": 0.9996211528778076,
"finite_elements": 458752,
"max_abs": 0.013671875,
"nrmse": 0.027537915855646133,
"same_nonfinite_mask": true
},
"status": "PASS"
}
],
"schema_version": 1,
"torch_version": "2.11.0+cu130"
}

View File

@ -0,0 +1,214 @@
{
"arguments": {
"activation": "swiglu",
"benchmark_iterations": 0,
"output": "results/phase2_control_swiglu_smoke_gpu6.json",
"require_situ": false,
"seed": 20260817,
"suite": "smoke",
"warmup": 3
},
"compute_capability": [
12,
0
],
"created_at": "2026-08-17T12:01:32.407549+08:00",
"cuda_version": "13.0",
"cuda_visible_devices": "6",
"device": "NVIDIA RTX 6000D",
"flashinfer_version": "0.6.18",
"logical_cuda_device": 0,
"results": [
{
"activation": "swiglu",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 4,
"name": "decode_m1_uniform",
"routing": "uniform",
"tokens": 1,
"top_k": 2
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 1.2227470874786377,
"metrics": {
"cosine": 0.9998119473457336,
"finite_elements": 128,
"max_abs": 0.03125,
"nrmse": 0.019757770001888275,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"activation": "swiglu",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 4,
"name": "decode_m2_skewed",
"routing": "skewed",
"tokens": 2,
"top_k": 2
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 0.03631448745727539,
"metrics": {
"cosine": 0.9997045397758484,
"finite_elements": 256,
"max_abs": 0.03125,
"nrmse": 0.024355093017220497,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"activation": "swiglu",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 8,
"name": "batch_m8_empty_experts",
"routing": "skewed",
"tokens": 8,
"top_k": 2
},
"cuda_graph": {
"replay_matches_eager": {
"cosine": 0.9999999403953552,
"finite_elements": 1024,
"max_abs": 0.0,
"nrmse": 0.0,
"same_nonfinite_mask": true
},
"replays_bitwise_equal": true
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 0.08502435684204102,
"metrics": {
"cosine": 0.9996250867843628,
"finite_elements": 1024,
"max_abs": 0.0546875,
"nrmse": 0.02746359072625637,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"activation": "swiglu",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 8,
"name": "batch_m32_masked_slot",
"routing": "masked",
"tokens": 32,
"top_k": 4
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 0.013071775436401367,
"metrics": {
"cosine": 0.9996730089187622,
"finite_elements": 4096,
"max_abs": 0.0390625,
"nrmse": 0.02586047165095806,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"activation": "swiglu",
"case": {
"ep_rank": 0,
"ep_size": 2,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 4,
"name": "ep2_global_ids",
"routing": "ep_mixed",
"tokens": 8,
"top_k": 4
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 0.0670325756072998,
"metrics": {
"cosine": 0.9996550679206848,
"finite_elements": 1024,
"max_abs": 0.01953125,
"nrmse": 0.02657792530953884,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"activation": "swiglu",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": true,
"intermediate": 128,
"local_experts": 4,
"name": "nonfinite_activation",
"routing": "uniform",
"tokens": 2,
"top_k": 2
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 0.009414196014404297,
"metrics": {
"candidate_inf_elements": 68,
"candidate_nan_elements": 60,
"candidate_nonfinite_rows": [
false,
true
],
"contract": "same_affected_rows_and_close_unaffected_rows",
"cosine": 0.9998903870582581,
"finite_elements": 128,
"max_abs": 0.025146484375,
"nrmse": 0.015044113621115685,
"reference_nonfinite_rows": [
false,
true
],
"same_nonfinite_mask": false
},
"status": "PASS"
}
],
"schema_version": 1,
"torch_version": "2.11.0+cu130"
}

View File

@ -0,0 +1,61 @@
{
"arguments": {
"activation": "situ",
"benchmark_iterations": 0,
"output": "results/phase2_target_situ_real_gpu6.json",
"require_situ": false,
"seed": 20260817,
"suite": "real",
"warmup": 3
},
"compute_capability": [
12,
0
],
"created_at": "2026-08-17T12:04:05.231369+08:00",
"cuda_version": "13.0",
"cuda_visible_devices": "6",
"device": "NVIDIA RTX 6000D",
"flashinfer_version": "0.6.18",
"logical_cuda_device": 0,
"results": [
{
"activation": "situ",
"case": {
"ep_rank": 7,
"ep_size": 32,
"hidden": 3584,
"inject_nonfinite": false,
"intermediate": 3072,
"local_experts": 28,
"name": "kimi_decode_m1_ep32",
"routing": "ep_mixed",
"tokens": 1,
"top_k": 16
},
"elapsed_s": 15.5287344455719,
"error": "Check failed: fc1_expert_weights.size(1) == fc2_expert_weights.size(2) * mInnerDimMultiplier (6144 vs. 3072) : fc1_expert_weights inter size must be equal to fc2_expert_weights inter size.",
"status": "EXPECTED_MISSING_SITU"
},
{
"activation": "situ",
"case": {
"ep_rank": 7,
"ep_size": 32,
"hidden": 3584,
"inject_nonfinite": false,
"intermediate": 3072,
"local_experts": 28,
"name": "kimi_prefill_m128_ep32",
"routing": "ep_mixed",
"tokens": 128,
"top_k": 16
},
"elapsed_s": 11.09748888015747,
"error": "Check failed: fc1_expert_weights.size(1) == fc2_expert_weights.size(2) * mInnerDimMultiplier (6144 vs. 3072) : fc1_expert_weights inter size must be equal to fc2_expert_weights inter size.",
"status": "EXPECTED_MISSING_SITU"
}
],
"schema_version": 1,
"torch_version": "2.11.0+cu130"
}

View File

@ -0,0 +1,133 @@
{
"arguments": {
"activation": "situ",
"benchmark_iterations": 0,
"output": "results/phase2_target_situ_smoke_gpu6.json",
"require_situ": false,
"seed": 20260817,
"suite": "smoke",
"warmup": 3
},
"compute_capability": [
12,
0
],
"created_at": "2026-08-17T12:00:47.947426+08:00",
"cuda_version": "13.0",
"cuda_visible_devices": "6",
"device": "NVIDIA RTX 6000D",
"flashinfer_version": "0.6.18",
"logical_cuda_device": 0,
"results": [
{
"activation": "situ",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 4,
"name": "decode_m1_uniform",
"routing": "uniform",
"tokens": 1,
"top_k": 2
},
"elapsed_s": 2.840806484222412,
"error": "Check failed: fc1_expert_weights.size(1) == fc2_expert_weights.size(2) * mInnerDimMultiplier (256 vs. 128) : fc1_expert_weights inter size must be equal to fc2_expert_weights inter size.",
"status": "EXPECTED_MISSING_SITU"
},
{
"activation": "situ",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 4,
"name": "decode_m2_skewed",
"routing": "skewed",
"tokens": 2,
"top_k": 2
},
"elapsed_s": 0.03530144691467285,
"error": "Check failed: fc1_expert_weights.size(1) == fc2_expert_weights.size(2) * mInnerDimMultiplier (256 vs. 128) : fc1_expert_weights inter size must be equal to fc2_expert_weights inter size.",
"status": "EXPECTED_MISSING_SITU"
},
{
"activation": "situ",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 8,
"name": "batch_m8_empty_experts",
"routing": "skewed",
"tokens": 8,
"top_k": 2
},
"elapsed_s": 0.016995668411254883,
"error": "Check failed: fc1_expert_weights.size(1) == fc2_expert_weights.size(2) * mInnerDimMultiplier (256 vs. 128) : fc1_expert_weights inter size must be equal to fc2_expert_weights inter size.",
"status": "EXPECTED_MISSING_SITU"
},
{
"activation": "situ",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 8,
"name": "batch_m32_masked_slot",
"routing": "masked",
"tokens": 32,
"top_k": 4
},
"elapsed_s": 0.01120138168334961,
"error": "Check failed: fc1_expert_weights.size(1) == fc2_expert_weights.size(2) * mInnerDimMultiplier (256 vs. 128) : fc1_expert_weights inter size must be equal to fc2_expert_weights inter size.",
"status": "EXPECTED_MISSING_SITU"
},
{
"activation": "situ",
"case": {
"ep_rank": 0,
"ep_size": 2,
"hidden": 128,
"inject_nonfinite": false,
"intermediate": 128,
"local_experts": 4,
"name": "ep2_global_ids",
"routing": "ep_mixed",
"tokens": 8,
"top_k": 4
},
"elapsed_s": 0.007741689682006836,
"error": "Check failed: fc1_expert_weights.size(1) == fc2_expert_weights.size(2) * mInnerDimMultiplier (256 vs. 128) : fc1_expert_weights inter size must be equal to fc2_expert_weights inter size.",
"status": "EXPECTED_MISSING_SITU"
},
{
"activation": "situ",
"case": {
"ep_rank": 0,
"ep_size": 1,
"hidden": 128,
"inject_nonfinite": true,
"intermediate": 128,
"local_experts": 4,
"name": "nonfinite_activation",
"routing": "uniform",
"tokens": 2,
"top_k": 2
},
"elapsed_s": 0.007404804229736328,
"error": "Check failed: fc1_expert_weights.size(1) == fc2_expert_weights.size(2) * mInnerDimMultiplier (256 vs. 128) : fc1_expert_weights inter size must be equal to fc2_expert_weights inter size.",
"status": "EXPECTED_MISSING_SITU"
}
],
"schema_version": 1,
"torch_version": "2.11.0+cu130"
}

View File

@ -0,0 +1,682 @@
#!/usr/bin/env python3
"""Validate Kimi-K3 SiTU semantics for FlashInfer SM120 MXFP8 x MXFP4 MoE.
The control and target paths share the same generated case, quantized tensors,
routing metadata, CUTLASS ABI wrapper, synchronization, and timing method. The
only intentional difference is the activation contract: supported SwiGLU is a
harness control, while SiTU is the Kimi-K3 target.
"""
from __future__ import annotations
import argparse
from dataclasses import asdict, dataclass
from datetime import datetime
import importlib.metadata
import json
import math
import os
from pathlib import Path
import sys
import time
from typing import Callable, Literal
import torch
import torch.nn.functional as F
from flashinfer import (
mxfp4_dequantize,
mxfp4_quantize,
mxfp8_dequantize_host,
mxfp8_quantize,
)
from flashinfer.fused_moe import cutlass_fused_moe
from flashinfer.tllm_enums import ActivationType
ActivationName = Literal["swiglu", "situ"]
RoutingName = Literal["uniform", "skewed", "masked", "ep_mixed"]
KNOWN_MISSING_SITU_ERROR = (
"fc1_expert_weights inter size must be equal to "
"fc2_expert_weights inter size"
)
@dataclass(frozen=True)
class Case:
name: str
tokens: int
hidden: int
intermediate: int
local_experts: int
top_k: int
routing: RoutingName
ep_size: int = 1
ep_rank: int = 0
inject_nonfinite: bool = False
@property
def total_experts(self) -> int:
return self.local_experts * self.ep_size
@property
def local_expert_start(self) -> int:
return self.local_experts * self.ep_rank
@dataclass
class QuantizedCase:
case: Case
input_q: torch.Tensor
input_sf: torch.Tensor
input_dq: torch.Tensor
fc1_q: torch.Tensor
fc1_sf: torch.Tensor
fc1_dq: torch.Tensor
fc2_q: torch.Tensor
fc2_sf: torch.Tensor
fc2_dq: torch.Tensor
topk_ids: torch.Tensor
topk_weights: torch.Tensor
expert_scale: torch.Tensor
situ_gate_beta: torch.Tensor
situ_linear_beta: torch.Tensor
def quantize_mxfp4_experts(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
packed = []
scales = []
for expert_weight in weight:
expert_packed, expert_scale = mxfp4_quantize(expert_weight)
packed.append(expert_packed)
scales.append(expert_scale)
return torch.stack(packed), torch.stack(scales)
def dequantize_mxfp4_experts(
packed: torch.Tensor, scales: torch.Tensor, dtype: torch.dtype
) -> torch.Tensor:
outputs = []
flat_scales = scales.reshape(scales.shape[0], -1)
for expert in range(packed.shape[0]):
outputs.append(
mxfp4_dequantize(
packed[expert].cpu().view(torch.uint8),
flat_scales[expert].cpu().view(torch.uint8),
)
)
return torch.stack(outputs).to(device=packed.device, dtype=dtype)
def dequantize_mxfp8_input(
packed: torch.Tensor, scales: torch.Tensor, dtype: torch.dtype
) -> torch.Tensor:
return (
mxfp8_dequantize_host(
packed.cpu().view(torch.uint8),
scales.cpu().view(torch.uint8).reshape(-1),
True,
)
.to(device=packed.device, dtype=dtype)
.reshape(packed.shape)
)
def make_routing(case: Case, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
ids = torch.empty(case.tokens, case.top_k, dtype=torch.int32, device=device)
local_start = case.local_expert_start
local_end = local_start + case.local_experts
if case.routing == "uniform":
for token in range(case.tokens):
for slot in range(case.top_k):
ids[token, slot] = (token * case.top_k + slot) % case.total_experts
elif case.routing == "skewed":
chosen = torch.arange(case.top_k, device=device) % case.total_experts
ids.copy_(chosen.expand(case.tokens, -1))
elif case.routing == "masked":
chosen = torch.arange(case.top_k, device=device) % case.total_experts
ids.copy_(chosen.expand(case.tokens, -1))
ids[:, -1] = -1
elif case.routing == "ep_mixed":
remote_start = local_end % case.total_experts
for token in range(case.tokens):
for slot in range(case.top_k):
if slot % 2 == 0:
ids[token, slot] = local_start + (token + slot) % case.local_experts
else:
ids[token, slot] = remote_start + (token + slot) % case.local_experts
else:
raise ValueError(f"unsupported routing mode: {case.routing}")
logits = torch.linspace(
0.25,
1.25,
case.top_k,
dtype=torch.float32,
device=device,
).expand(case.tokens, -1)
weights = torch.softmax(logits, dim=-1)
weights = torch.where(ids >= 0, weights, torch.zeros_like(weights))
weights /= weights.sum(dim=-1, keepdim=True)
return ids, weights.contiguous()
def prepare_case(case: Case, seed: int, device: torch.device) -> QuantizedCase:
generator = torch.Generator(device=device).manual_seed(seed)
dtype = torch.bfloat16
x = torch.randn(
case.tokens,
case.hidden,
generator=generator,
dtype=dtype,
device=device,
)
if case.inject_nonfinite:
x[0, 0] = float("nan")
if case.hidden > 1:
x[-1, 1] = float("inf")
std_fc1 = 1.0 / math.sqrt(case.hidden)
std_fc2 = 1.0 / math.sqrt(case.intermediate)
gate = torch.randn(
case.local_experts,
case.intermediate,
case.hidden,
generator=generator,
dtype=dtype,
device=device,
) * std_fc1
up = torch.randn(
case.local_experts,
case.intermediate,
case.hidden,
generator=generator,
dtype=dtype,
device=device,
) * std_fc1
down = torch.randn(
case.local_experts,
case.hidden,
case.intermediate,
generator=generator,
dtype=dtype,
device=device,
) * std_fc2
# CUTLASS consumes [up; gate]. Kimi checkpoints load as [gate; up], so
# this is also the oracle layout required from the SGLang weight processor.
fc1 = torch.cat([up, gate], dim=1).contiguous()
input_q, input_sf = mxfp8_quantize(x, True, 32)
fc1_q, fc1_sf = quantize_mxfp4_experts(fc1)
fc2_q, fc2_sf = quantize_mxfp4_experts(down)
input_dq = dequantize_mxfp8_input(input_q, input_sf, dtype)
fc1_dq = dequantize_mxfp4_experts(fc1_q, fc1_sf, dtype)
fc2_dq = dequantize_mxfp4_experts(fc2_q, fc2_sf, dtype)
topk_ids, topk_weights = make_routing(case, device)
return QuantizedCase(
case=case,
input_q=input_q,
input_sf=input_sf,
input_dq=input_dq,
fc1_q=fc1_q,
fc1_sf=fc1_sf,
fc1_dq=fc1_dq,
fc2_q=fc2_q,
fc2_sf=fc2_sf,
fc2_dq=fc2_dq,
topk_ids=topk_ids,
topk_weights=topk_weights,
expert_scale=torch.ones(
case.local_experts, dtype=torch.float32, device=device
),
situ_gate_beta=torch.full(
(case.local_experts,), 4.0, dtype=torch.float32, device=device
),
situ_linear_beta=torch.full(
(case.local_experts,), 25.0, dtype=torch.float32, device=device
),
)
def apply_activation(
gate: torch.Tensor, up: torch.Tensor, activation: ActivationName
) -> torch.Tensor:
if activation == "swiglu":
return F.silu(gate) * up
if activation == "situ":
gate = 4.0 * torch.tanh(gate / 4.0) * torch.sigmoid(gate)
up = 25.0 * torch.tanh(up / 25.0)
return gate * up
raise ValueError(f"unsupported activation: {activation}")
def run_reference(data: QuantizedCase, activation: ActivationName) -> torch.Tensor:
case = data.case
output = torch.zeros(
case.tokens,
case.hidden,
dtype=torch.bfloat16,
device=data.input_q.device,
)
local_start = case.local_expert_start
for local_expert in range(case.local_experts):
global_expert = local_start + local_expert
selected = data.topk_ids == global_expert
if not bool(selected.any()):
continue
token_ids, slots = torch.where(selected)
up_weight, gate_weight = torch.chunk(data.fc1_dq[local_expert], 2, dim=0)
expert_input = data.input_dq[token_ids]
gate = expert_input @ gate_weight.transpose(0, 1)
up = expert_input @ up_weight.transpose(0, 1)
hidden = apply_activation(gate, up, activation)
expert_output = hidden @ data.fc2_dq[local_expert].transpose(0, 1)
output[token_ids] += (
data.topk_weights[token_ids, slots, None].to(expert_output.dtype)
* expert_output
)
return output
def run_candidate(
data: QuantizedCase,
activation: ActivationName,
*,
use_fused_finalize: bool,
output: torch.Tensor | None = None,
) -> torch.Tensor:
case = data.case
if output is None:
output = torch.empty(
case.tokens,
case.hidden,
dtype=torch.bfloat16,
device=data.input_q.device,
)
activation_type = (
ActivationType.Swiglu if activation == "swiglu" else ActivationType.Situ
)
alpha = data.situ_gate_beta if activation == "situ" else None
beta = data.situ_linear_beta if activation == "situ" else None
cutlass_fused_moe(
input=data.input_q,
token_selected_experts=data.topk_ids,
token_final_scales=data.topk_weights,
fc1_expert_weights=data.fc1_q.contiguous().view(torch.int64),
fc2_expert_weights=data.fc2_q.contiguous().view(torch.int64),
output_dtype=torch.bfloat16,
quant_scales=[
data.fc1_sf.view(torch.int32),
data.expert_scale,
data.fc2_sf.view(torch.int32),
data.expert_scale,
],
input_sf=data.input_sf,
swiglu_alpha=alpha,
swiglu_beta=beta,
swiglu_limit=None,
tp_size=1,
tp_rank=0,
ep_size=case.ep_size,
ep_rank=case.ep_rank,
use_mxfp8_act_scaling=True,
activation_type=activation_type,
tune_max_num_tokens=1 << (case.tokens - 1).bit_length(),
output=output,
use_fused_finalize=use_fused_finalize,
)
return output
def tensor_metrics(reference: torch.Tensor, candidate: torch.Tensor) -> dict:
ref_nan = torch.isnan(reference)
got_nan = torch.isnan(candidate)
ref_inf = torch.isinf(reference)
got_inf = torch.isinf(candidate)
same_nonfinite = bool(torch.equal(ref_nan, got_nan) and torch.equal(ref_inf, got_inf))
finite = torch.isfinite(reference) & torch.isfinite(candidate)
if not bool(finite.any()):
return {
"same_nonfinite_mask": same_nonfinite,
"finite_elements": 0,
"max_abs": None,
"nrmse": None,
"cosine": None,
}
ref = reference[finite].float()
got = candidate[finite].float()
diff = got - ref
denominator = torch.sqrt(torch.mean(ref.square())).clamp_min(1e-12)
if bool(torch.count_nonzero(ref)) and bool(torch.count_nonzero(got)):
cosine = float(F.cosine_similarity(ref, got, dim=0).item())
else:
cosine = 1.0 if torch.equal(ref, got) else 0.0
return {
"same_nonfinite_mask": same_nonfinite,
"finite_elements": int(finite.sum().item()),
"max_abs": float(diff.abs().max().item()),
"nrmse": float((torch.sqrt(torch.mean(diff.square())) / denominator).item()),
"cosine": cosine,
}
def assert_matches(reference: torch.Tensor, candidate: torch.Tensor) -> dict:
metrics = tensor_metrics(reference, candidate)
if not metrics["same_nonfinite_mask"]:
raise AssertionError(f"NaN/Inf masks differ: {metrics}")
torch.testing.assert_close(
candidate,
reference,
rtol=0.2,
atol=0.2,
equal_nan=True,
)
return metrics
def assert_nonfinite_contract(
reference: torch.Tensor, candidate: torch.Tensor
) -> dict:
"""Check stable row-level propagation without over-specifying reduction order.
An Inf entering a GEMM may produce a different element-wise mix of Inf and
NaN when the reduction tree changes. Requiring an identical element mask
would therefore be stricter than IEEE arithmetic permits. We require the
same affected rows, finite unaffected rows, and ordinary numeric closeness
on those unaffected rows.
"""
reference_rows = (~torch.isfinite(reference)).any(dim=1)
candidate_rows = (~torch.isfinite(candidate)).any(dim=1)
if not torch.equal(reference_rows, candidate_rows):
raise AssertionError(
"nonfinite propagation reached different token rows: "
f"reference={reference_rows.tolist()} candidate={candidate_rows.tolist()}"
)
unaffected = ~reference_rows
if bool(unaffected.any()):
torch.testing.assert_close(
candidate[unaffected],
reference[unaffected],
rtol=0.2,
atol=0.2,
)
metrics = tensor_metrics(reference, candidate)
metrics.update(
{
"contract": "same_affected_rows_and_close_unaffected_rows",
"reference_nonfinite_rows": reference_rows.tolist(),
"candidate_nonfinite_rows": candidate_rows.tolist(),
"candidate_nan_elements": int(torch.isnan(candidate).sum().item()),
"candidate_inf_elements": int(torch.isinf(candidate).sum().item()),
}
)
return metrics
def timed_ms(fn: Callable[[], torch.Tensor], warmup: int, iterations: int) -> float:
for _ in range(warmup):
fn()
torch.cuda.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(iterations):
fn()
end.record()
end.synchronize()
return float(start.elapsed_time(end) / iterations)
def validate_determinism(data: QuantizedCase, activation: ActivationName) -> dict:
outputs = [
run_candidate(data, activation, use_fused_finalize=False).clone()
for _ in range(3)
]
exact = True
for item in outputs[1:]:
try:
torch.testing.assert_close(
outputs[0], item, rtol=0.0, atol=0.0, equal_nan=True
)
except AssertionError:
exact = False
break
if not exact:
raise AssertionError("non-fused finalize is not bitwise deterministic")
return {"runs": len(outputs), "bitwise_equal": exact}
def validate_cuda_graph(data: QuantizedCase, activation: ActivationName) -> dict:
eager = run_candidate(data, activation, use_fused_finalize=False).clone()
static_output = torch.empty_like(eager)
# Warm module loading and tactic selection before capture.
run_candidate(
data,
activation,
use_fused_finalize=False,
output=static_output,
)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
run_candidate(
data,
activation,
use_fused_finalize=False,
output=static_output,
)
graph.replay()
torch.cuda.synchronize()
first = static_output.clone()
graph.replay()
torch.cuda.synchronize()
second = static_output.clone()
metrics = assert_matches(eager, first)
if not torch.equal(first, second):
raise AssertionError("CUDA Graph replay is not bitwise deterministic")
return {"replay_matches_eager": metrics, "replays_bitwise_equal": True}
def smoke_cases() -> list[Case]:
return [
Case("decode_m1_uniform", 1, 128, 128, 4, 2, "uniform"),
Case("decode_m2_skewed", 2, 128, 128, 4, 2, "skewed"),
Case("batch_m8_empty_experts", 8, 128, 128, 8, 2, "skewed"),
Case("batch_m32_masked_slot", 32, 128, 128, 8, 4, "masked"),
Case("ep2_global_ids", 8, 128, 128, 4, 4, "ep_mixed", ep_size=2),
Case(
"nonfinite_activation",
2,
128,
128,
4,
2,
"uniform",
inject_nonfinite=True,
),
]
def real_cases() -> list[Case]:
# One TP32/EP32 rank owns 28 of Kimi-K3's 896 experts. Global expert IDs
# intentionally mix local and remote routes, matching the per-rank kernel ABI.
return [
Case(
"kimi_decode_m1_ep32",
1,
3584,
3072,
28,
16,
"ep_mixed",
ep_size=32,
ep_rank=7,
),
Case(
"kimi_prefill_m128_ep32",
128,
3584,
3072,
28,
16,
"ep_mixed",
ep_size=32,
ep_rank=7,
),
]
def run_one(
case: Case,
activation: ActivationName,
args: argparse.Namespace,
) -> dict:
started = time.time()
data = prepare_case(case, args.seed, torch.device("cuda"))
reference = run_reference(data, activation)
result = {
"case": asdict(case),
"activation": activation,
"status": "UNKNOWN",
}
try:
candidate = run_candidate(
data,
activation,
use_fused_finalize=False,
)
result["metrics"] = (
assert_nonfinite_contract(reference, candidate)
if case.inject_nonfinite
else assert_matches(reference, candidate)
)
result["determinism"] = validate_determinism(data, activation)
if case.name == "batch_m8_empty_experts":
result["cuda_graph"] = validate_cuda_graph(data, activation)
if args.benchmark_iterations > 0 and not case.inject_nonfinite:
result["latency_ms"] = {
"reference": timed_ms(
lambda: run_reference(data, activation),
args.warmup,
args.benchmark_iterations,
),
"candidate": timed_ms(
lambda: run_candidate(
data,
activation,
use_fused_finalize=False,
),
args.warmup,
args.benchmark_iterations,
),
}
result["status"] = "PASS"
except RuntimeError as error:
if activation == "situ" and KNOWN_MISSING_SITU_ERROR in str(error):
result["status"] = "EXPECTED_MISSING_SITU"
result["error"] = str(error)
if args.require_situ:
raise
else:
result["status"] = "ERROR"
result["error"] = repr(error)
raise
finally:
result["elapsed_s"] = time.time() - started
return result
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument(
"--suite",
choices=("smoke", "real", "all"),
default="smoke",
)
parser.add_argument(
"--activation",
choices=("swiglu", "situ", "both"),
default="both",
)
parser.add_argument(
"--require-situ",
action="store_true",
help="Fail instead of recording the known missing-SiTU boundary.",
)
parser.add_argument("--seed", type=int, default=20260817)
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--benchmark-iterations", type=int, default=0)
parser.add_argument("--output", type=Path)
return parser.parse_args()
def main() -> None:
args = parse_args()
if not torch.cuda.is_available():
raise SystemExit("CUDA is required")
major, minor = torch.cuda.get_device_capability()
if major != 12:
raise SystemExit(f"SM120/SM121 is required, got SM{major}{minor}")
cases = []
if args.suite in ("smoke", "all"):
cases.extend(smoke_cases())
if args.suite in ("real", "all"):
cases.extend(real_cases())
activations: list[ActivationName] = (
["swiglu", "situ"] if args.activation == "both" else [args.activation]
)
report = {
"schema_version": 1,
"created_at": datetime.now().astimezone().isoformat(),
"torch_version": torch.__version__,
"cuda_version": torch.version.cuda,
"flashinfer_version": importlib.metadata.version("flashinfer-python"),
"device": torch.cuda.get_device_name(),
"cuda_visible_devices": os.environ.get("CUDA_VISIBLE_DEVICES"),
"logical_cuda_device": torch.cuda.current_device(),
"compute_capability": [major, minor],
"arguments": vars(args) | {"output": str(args.output) if args.output else None},
"results": [],
}
failed = False
for activation in activations:
for case in cases:
print(f"[RUN] activation={activation} case={case.name}", flush=True)
try:
result = run_one(case, activation, args)
except Exception as error:
failed = True
result = {
"case": asdict(case),
"activation": activation,
"status": "ERROR",
"error": repr(error),
}
report["results"].append(result)
print(
f"[{result['status']}] activation={activation} case={case.name}",
flush=True,
)
output = args.output
if output is None:
stamp = datetime.now().strftime("%Y%m%d-%H%M%S")
output = Path("results") / f"correctness_{args.suite}_{stamp}.json"
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(json.dumps(report, indent=2, sort_keys=True), encoding="utf-8")
print(f"report={output.resolve()}")
statuses = [item["status"] for item in report["results"]]
counts = {status: statuses.count(status) for status in sorted(set(statuses))}
print("summary=" + json.dumps(counts))
if failed:
raise SystemExit(1)
if __name__ == "__main__":
main()