[Feature] Complete Kimi SM120 FlashInfer MXFP4 integration

This commit is contained in:
Zhiyi Hong 2026-08-17 14:58:52 +08:00
parent a1c18d736b
commit 6493798ad5
3 changed files with 418 additions and 44 deletions

View File

@ -304,20 +304,98 @@ FLASHINFER_DISABLE_VERSION_CHECK=1 \
### NaN/Inf 决策记录 ### NaN/Inf 决策记录
MXFP8 quantize/dequantize 后,注入的 NaN 被量化为有限值,而 Inf 保留。相同 Inf 行在 PyTorch GEMM 与 CUTLASS reduction 中产生不同的逐元素 NaN/Inf 混合,但两者影响同一 token row未受影响行继续满足普通数值阈值。IEEE reduction 顺序不保证逐元素 NaN/Inf 分类一致,因此 harness 要求: 2026-08-17 的统一 gate 发现旧 oracle 不成立:原始 BF16 NaN/Inf 经
`mxfp8_quantize`E4M3 payload 全部有限scale 最大值为 `0xFE`
NVIDIA E8M0 定义中 `0x00..0xFE` 表示有限的二次幂,只有 `0xFF` 表示
NaN转换采用 saturating finite 行为。旧 reference 把 `0xFE` scale
重新展开到 BF16 时发生溢出,再要求 Tensor Core candidate 传播相同
non-finite row不是有效的 MXFP8 oracle。
1. 相同 token rows 被标记为 non-finite 修正后的 nonfinite case 同时要求:
2. 未受影响 rows 继续满足 `rtol=0.2, atol=0.2`
3. candidate 不崩溃,且三次结果确定。
这不是放宽有限值正确性,有限 case 仍采用与 FlashInfer 官方 MXFP8×MXFP4 测试一致的阈值。 1. 量化后的 E4M3 payload 全部有限;
2. E8M0 scale 不得出现 `0xFF`
3. candidate 输出全部有限且三次 bitwise deterministic
4. BF16 dequant reference 的 overflow rows 继续记录为 diagnostic不参与数值判定。
## Phase 3 硬门槛 这项修正没有改变任何 finite case 的 reference、容差或性能计时。有限值仍
使用同一量化输入、权重、scale 与 routing并满足 `rtol=0.2, atol=0.2`
下一阶段只做最小实现: 官方格式依据:
`https://docs.nvidia.com/cuda/cuda-math-api/cuda_math_api/struct____nv__fp8__e8m0.html`
1. FlashInfer C++ 增加与 Python 对齐的 `Situ` enum、gated classification 和 SiTU adaptor复用已有 per-expert `alpha/beta` 参数与 grouped GEMM。 ## Phase 3 最小实现
2. SGLang SM120 processor 按 `gate_up_interleaved=False` 处理 Kimi 连续半区,生成 CUTLASS `[up; gate]`,不得复用 GPT-OSS 奇偶行拆分。
3. SGLang runner 从 `runner_config` 透传 `ActivationType.Situ` 和 Kimi `4/25` 参数。 状态:**完成;有限值正确性与定义明确的 MXFP8 nonfinite 边界通过,可以进入 Phase 4 profiling。尚未形成整模型性能结论。**
4. 未满足支持条件时在服务初始化阶段明确选择 Marlin不得在 kernel 失败后热切换,不修改 checkpoint不使用 runtime monkeypatch。
5. 完成后必须以 `--require-situ --suite all` 通过 Phase 2 全矩阵,才能进入任何吞吐比较。 ### 源码提交
| 仓库 | 分支 | 提交 |
|---|---|---|
| FlashInfer | `hzy/kimi-k3-sm120-situ-mxfp4` | `7d1a4ee` `[Feature] Add CUTLASS SiTU fused MoE support on SM120` |
| SGLang | `hzy/kimi-k3-sm120-flashinfer-mxfp4` | `fb929bb` `[Feature] Support Kimi-K3 FlashInfer MXFP4 MoE on SM120` |
FlashInfer 的最小增量:
1. C++ `ActivationType` 增加与 Python ABI 对齐的 `Situ`
2. 把 SiTU 标为 gated activation
3. 增加 `4*tanh(gate/4)*sigmoid(gate) * 25*tanh(up/25)` adaptor复用已有 grouped GEMM 和 per-expert 参数;
4. Python 暴露显式 capability marker避免仅凭早已存在的 enum 误判 kernel 可用。
SGLang 的最小增量:
1. 同时支持 GPT-OSS 交错 `[gate_i, up_i]` 与 Kimi 连续 `[gate; up]` checkpoint 布局,统一生成 FlashInfer `[up; gate]`
2. 从 `MoeRunnerConfig` 透传 `Situ`、alpha `4`、beta `25`,不再硬编码 SwiGLU
3. bias-free Kimi 路径不向 kernel 传零占位 bias
4. FlashInfer 缺少 SiTU capability marker 时在 runner 初始化阶段明确报错,并提示以 `--moe-runner-backend marlin` 重启;不做 kernel 失败后的热切换。
### 构建产物
```text
/data/hzy/artifacts/flashinfer-kimi-sm120-situ-phase3-v3/flashinfer_python-0.6.18-py3-none-any.whl
SHA256 9298d34b6f9cec9b8bf535d1688fc6e8c73f7db65f64e0c496580e881c3d997b
```
### 正确性与集成结果
统一 Phase 2/3 gate 命令:
```bash
cd /data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4
CUDA_VISIBLE_DEVICES=6 \
PYTHONPATH=/tmp/hzy-fi-situ-phase3-v3 \
FLASHINFER_DISABLE_VERSION_CHECK=1 \
/data/hzy/envs/sparse-opd-sm120/bin/python validate_kimi_sm120_mxfp4.py \
--suite all --activation situ --require-situ \
--warmup 2 --benchmark-iterations 3 \
--output results/phase3_target_situ_all_gpu6_v3.json
```
结果:`8/8 PASS`
| 覆盖 | 结果 |
|---|---|
| M=1/2/8/32uniform/skewed/empty/masked | PASS |
| EP2 global/local IDs | PASS |
| MXFP8 NaN/Inf saturation boundary | PASScandidate finite、3-run bitwise |
| CUDA Graph capture/replay | PASSreplay 与 eager `max_abs=0` |
| Kimi decode `M=1, H=3584, N=3072, E=28, topK=16` | cosine `0.9996223`NRMSE `2.748%` |
| Kimi prefill `M=128` 同真尺寸 | cosine `0.9996152`NRMSE `2.774%` |
SGLang 集成测试:`6 passed, 18 warnings in 563.41s`。warnings 为 asyncio
配置、PyTorch deprecation 和只读 pytest cache不含断言失败。
原始证据:
```text
/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/results/phase3_target_situ_all_gpu6_v3.log
/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/results/phase3_target_situ_all_gpu6_v3.json
/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_sglang_sm120_flashinfer_mxfp4/results/phase3_sglang_integration_gpu6_final.log
```
### Phase 3 决策
1. 当前缺口同时位于 FlashInfer SiTU activation 与 SGLang Kimi layout/参数集成;两个补丁都必要。
2. 没有修改 checkpoint没有 runtime monkeypatch没有通过放宽 finite tolerance 获得通过。
3. microbenchmark latency 仅用于检查 candidate 可执行,不作为生产性能结论。
4. 下一阶段必须用 Nsight Systems/Compute 和三轮交错 A/B 区分 wrapper、同步、权重搬运、dispatch 与 kernel 本体;在此之前不宣称优于 Marlin。

View File

@ -0,0 +1,296 @@
{
"arguments": {
"activation": "situ",
"benchmark_iterations": 3,
"output": "results/phase3_target_situ_all_gpu6_v3.json",
"require_situ": true,
"seed": 20260817,
"suite": "all",
"warmup": 2
},
"compute_capability": [
12,
0
],
"created_at": "2026-08-17T14:54:33.526558+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
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 1.2600979804992676,
"latency_ms": {
"candidate": 0.24848000208536783,
"reference": 0.7109653155008951
},
"metrics": {
"cosine": 0.9996331930160522,
"finite_elements": 128,
"max_abs": 0.046875,
"nrmse": 0.038191329687833786,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"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
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 0.041979312896728516,
"latency_ms": {
"candidate": 0.24086399873097739,
"reference": 0.6774506568908691
},
"metrics": {
"cosine": 0.9996317625045776,
"finite_elements": 256,
"max_abs": 0.046875,
"nrmse": 0.03233937919139862,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"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
},
"cuda_graph": {
"replay_matches_eager": {
"cosine": 0.9999998807907104,
"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.0908958911895752,
"latency_ms": {
"candidate": 0.23886932929356894,
"reference": 0.7901653448740641
},
"metrics": {
"cosine": 0.9996166229248047,
"finite_elements": 1024,
"max_abs": 0.046875,
"nrmse": 0.02830333448946476,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"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
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 0.02089524269104004,
"latency_ms": {
"candidate": 0.22908800840377808,
"reference": 1.0701546669006348
},
"metrics": {
"cosine": 0.9996225237846375,
"finite_elements": 4096,
"max_abs": 0.03125,
"nrmse": 0.027592986822128296,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"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
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 0.07790017127990723,
"latency_ms": {
"candidate": 0.23056000471115112,
"reference": 1.2077759901682537
},
"metrics": {
"cosine": 0.9995924830436707,
"finite_elements": 1024,
"max_abs": 0.0166015625,
"nrmse": 0.02855238877236843,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"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
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 0.008931636810302734,
"metrics": {
"candidate_inf_elements": 0,
"candidate_nan_elements": 0,
"candidate_nonfinite_rows": [
false,
false
],
"contract": "mxfp8_satfinite_input_and_finite_candidate",
"input_payload_finite": true,
"max_scale_code": 254,
"reference_nonfinite_rows_diagnostic_only": [
false,
true
],
"scale_has_nan_code": false
},
"status": "PASS"
},
{
"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
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 12.66703748703003,
"latency_ms": {
"candidate": 0.295904000600179,
"reference": 2.942485491434733
},
"metrics": {
"cosine": 0.9996223449707031,
"finite_elements": 3584,
"max_abs": 0.009765625,
"nrmse": 0.027483593672513962,
"same_nonfinite_mask": true
},
"status": "PASS"
},
{
"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
},
"determinism": {
"bitwise_equal": true,
"runs": 3
},
"elapsed_s": 12.947205066680908,
"latency_ms": {
"candidate": 0.6200213432312012,
"reference": 8.530954360961914
},
"metrics": {
"cosine": 0.9996151924133301,
"finite_elements": 458752,
"max_abs": 0.013671875,
"nrmse": 0.02774309180676937,
"same_nonfinite_mask": true
},
"status": "PASS"
}
],
"schema_version": 1,
"torch_version": "2.11.0+cu130"
}

View File

@ -378,42 +378,42 @@ def assert_matches(reference: torch.Tensor, candidate: torch.Tensor) -> dict:
def assert_nonfinite_contract( def assert_nonfinite_contract(
reference: torch.Tensor, candidate: torch.Tensor data: QuantizedCase,
reference: torch.Tensor,
candidate: torch.Tensor,
) -> dict: ) -> dict:
"""Check stable row-level propagation without over-specifying reduction order. """Validate the defined MXFP8 boundary for non-finite BF16 inputs.
An Inf entering a GEMM may produce a different element-wise mix of Inf and ``mxfp8_quantize`` uses saturating conversion: the generated E4M3 payload
NaN when the reduction tree changes. Requiring an identical element mask is finite and E8M0 codes 0x00..0xFE denote finite powers of two; only 0xFF
would therefore be stricter than IEEE arithmetic permits. We require the denotes NaN. A BF16 dequantized reference can overflow for the 0xFE scale,
same affected rows, finite unaffected rows, and ordinary numeric closeness so its NaN/Inf rows are diagnostic only and are not a valid Tensor Core
on those unaffected rows. oracle. The kernel contract here is finite, deterministic output from the
finite quantized representation.
""" """
input_payload_finite = bool(torch.isfinite(data.input_q.float()).all())
scale_has_nan_code = bool((data.input_sf == 0xFF).any())
candidate_finite = bool(torch.isfinite(candidate).all())
if not input_payload_finite or scale_has_nan_code or not candidate_finite:
raise AssertionError(
"MXFP8 saturation contract failed: "
f"input_payload_finite={input_payload_finite} "
f"scale_has_nan_code={scale_has_nan_code} "
f"candidate_finite={candidate_finite}"
)
reference_rows = (~torch.isfinite(reference)).any(dim=1) reference_rows = (~torch.isfinite(reference)).any(dim=1)
candidate_rows = (~torch.isfinite(candidate)).any(dim=1) candidate_rows = (~torch.isfinite(candidate)).any(dim=1)
if not torch.equal(reference_rows, candidate_rows): return {
raise AssertionError( "contract": "mxfp8_satfinite_input_and_finite_candidate",
"nonfinite propagation reached different token rows: " "input_payload_finite": input_payload_finite,
f"reference={reference_rows.tolist()} candidate={candidate_rows.tolist()}" "scale_has_nan_code": scale_has_nan_code,
) "max_scale_code": int(data.input_sf.max().item()),
unaffected = ~reference_rows "reference_nonfinite_rows_diagnostic_only": reference_rows.tolist(),
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_nonfinite_rows": candidate_rows.tolist(),
"candidate_nan_elements": int(torch.isnan(candidate).sum().item()), "candidate_nan_elements": int(torch.isnan(candidate).sum().item()),
"candidate_inf_elements": int(torch.isinf(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: def timed_ms(fn: Callable[[], torch.Tensor], warmup: int, iterations: int) -> float:
@ -549,7 +549,7 @@ def run_one(
use_fused_finalize=False, use_fused_finalize=False,
) )
result["metrics"] = ( result["metrics"] = (
assert_nonfinite_contract(reference, candidate) assert_nonfinite_contract(data, reference, candidate)
if case.inject_nonfinite if case.inject_nonfinite
else assert_matches(reference, candidate) else assert_matches(reference, candidate)
) )