Skip to content

[NPU] Add opt-in MLA prefix FIA and Gemma norm fallback - #40131

Open
McZyWu wants to merge 10 commits into
sgl-project:mainfrom
McZyWu:fix/ascend-mla-prefix-fia-opt-in
Open

McZyWu wants to merge 10 commits into
sgl-project:mainfrom
McZyWu:fix/ascend-mla-prefix-fia-opt-in

Conversation

@McZyWu

@McZyWu McZyWu commented Sep 18, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

MLA prefix-cache prefill currently requires ATB npu_ring_mla, even when ASCEND_USE_FIA=1. This prevents the prefix path from running in environments where RingMLA is unavailable.

Allow the existing ASCEND_USE_FIA flag to select CANN FIA V2 for this path. With the flag unset or disabled, retain the original ATB calls and RingMLA mask.

Honor SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM in Gemma3RMSNorm on NPU.

Modifications

  • Apply the existing native Gemma RMSNorm flag to Gemma3RMSNorm.forward_npu. The flag remains disabled by default; Gemma4RMSNorm already falls back to native on NPU.
  • Select FIA or ATB once for MLA prefix prefill. The complete FIA branch performs current-token attention, NZ-aware prefix gather/projection, prefix attention and FP32 output merge. The complete ATB branch performs both RingMLA calls with prefix gather/projection between them. Input preparation and output padding remain common.
  • Gather cached latent and RoPE pages through gather_mla_cache_pages in the two-FIA branch, restoring logical token order for PA-NZ storage before projection and attention. Ordinary ND cache behavior is unchanged.
  • The FIA path runs causal attention on current tokens and unmasked attention on the prefix, then merges their outputs with npu_attention_update in FP32 before restoring the query dtype.
  • Use cumulative TND sequence lengths, contiguous inputs and consistent token/head ordering for LSE. Set empty-prefix output/LSE to zero/negative infinity before merging.
  • Add numerical tests against full prefix-first dense SDPA and verify that FIA enabled/disabled selects only the intended operators.

Enable the new path with:

export ASCEND_USE_FIA=1

This reuses the existing backend-wide flag; its default remains disabled. FIA V2 is used because the TorchNPU 26.1 implementation dispatches to CANN FIA V4 on A2/A3 and V5 on Ascend 950. The approach follows the FIA + AttentionUpdate composition used by vLLM Ascend.

To select the native Gemma normalization implementation on NPU:

export SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM=1

Accuracy Tests

Prefix dispatch refactor validation (2026-09-28)

Commit ee3dfe96a1bcc93bcdf3d6f4c04eb489b659de28 reorganizes the MLA prefix path into one complete if self.use_fia / else block. For each fixed value of self.use_fia, the old and new Python ASTs contain the same executable statements in the same order, including operator arguments, NZ cache handling, empty-prefix handling and output padding. Ruff 0.15.1 formatting and lint (F401,F821,UP037), Python syntax parsing and git diff --check passed. No new unit tests or NPU model evaluations were run for this structural change; the numerical and performance results below retain their recorded commit attribution.

NZ cache validation (2026-09-24)

Commit 0c8ebbce034c11b5e0b9c2fec3193def9de2aa74 passed the original full HiCache MLA GSM8K case on A3/151 with FIA enabled in both runs: NZ disabled 473/1319 (35.860500%), NZ enabled 465/1319 (35.253980%). Both exceed the 34% case threshold. The 18 operator-path cases and 6 independent reruns also passed, with fixed NZ outputs bitwise equal to the ND baseline. See the NZ validation section under Speed Tests and Profiling for the complete environment, comparison, numerical checks, performance table and limitations. Python syntax parsing and git diff --check passed; no new unit tests were added.

Latest validation (2026-09-24)

At commit ef8a80be18d4d1e6efa84252cca07dd31284850a, Python syntax parsing and git diff --check passed for the three changed files. The quantization implementation and its test file match the PR base; the remaining PR changes concern MLA prefix attention and Gemma normalization. No additional unit tests or model evaluations were run for this update. The accuracy and performance measurements below remain attributed to their explicitly recorded earlier commits.

Gemma native fallback validation (2026-09-22)

Commit ec72a1dd7c4292477bdd8e317a55f87c44f04357 changes only layernorm.py. Python syntax compilation and git diff --check passed. No additional unit tests or NPU model evaluations were run for this change. The model accuracy measurements below belong to their explicitly recorded earlier commits.

Validation of updated PR head 3ba07dd on machine 216 (2026-09-21)

After the author merged community main, tested the actual PR head 3ba07ddc43d55782e2195f5cd2457297b1eef9fd, which includes main 62ba9648482e1b0a253187c8dc6ac5d23bbbd403. Main is already an ancestor of the PR; git merge-tree succeeds and GitHub reports MERGEABLE. This rerun uses the published PR commit directly, rather than the earlier temporary integration snapshot.

Ran its original test/registered/npu/basic_function/HiCache/test_npu_hicache_mla.py on machine 216, physical A3 devices 0–3, with CANN 9.1.0 / torch_npu 2.10.0.post4 / PyTorch 2.10.0 / Transformers 5.12.1. Reused the same full DeepSeek-V2-Lite-W8A8 checkpoint and GSM8K dataset. Both modes used TP=4, mem_fraction_static=0.8, HiCache enabled (current defaults: ratio=2.0, host memory fraction=0.8), all 1319 questions, 5-shot, temperature=0, max_new_tokens=512, concurrency=128, and seed=42. The PR smoke shortcut was disabled. No additional unit tests were added or run.

ASCEND_USE_FIA Correct / total Accuracy Invalid answers Case result
unset 472 / 1319 35.784685% 5 passed
1 469 / 1319 35.557240% 5 passed

FIA minus baseline: -3 correct answers / -0.227445 percentage points. 681 final numerical predictions differed: 111 changed from correct to incorrect, and 108 from incorrect to correct. This single pair does not establish lossless replacement. The existing backend-wide ASCEND_USE_FIA flag also affects other attention paths, so this does not isolate the new prefix branch alone.

Source hashes (including files updated by the merged main), dataset, devices, model path, evaluation settings, and actual server argument dictionaries were checked for consistency. Device-prefix cache hits: 1319/1319 unset and 1319/1319 enabled, 1012992 and 1012992 cached tokens respectively. Host-cache reload hits: 0 and 0; this case does not force host-cache eviction/reload. Full logs, metrics and per-question answers were retained.

Observed accuracy-evaluation runtime / output throughput: unset 198.626 s / 794.967 token/s, enabled 207.717 s / 755.773 token/s. These are not controlled performance measurements. A2/A5 and CANN 9.2 remain untested. Earlier validation records are preserved below.

Latest-main integration rerun on machine 216 (2026-09-21)

Community main 176dbcb85d3b7737564e4911947035cc0af65f0a and PR head 1344e1e2a5ca3e07bf6140b8c606e38465fbc55c merged without conflicts in both git merge-tree and an actual isolated merge. GitHub reports MERGEABLE; the CI gate is blocked by the missing run-ci label, independently of merge conflicts. Tested the resulting temporary merge 2b388ce85710f74030b266f8446a01b91c556d48 without pushing it to the PR branch.

During the run, main advanced to ab03a8e7eb82d35907bfbeb645bf0af8b4ce290e (#40201, import/startup changes). A final git merge-tree check against that main also succeeded and GitHub still reports MERGEABLE. It does not modify this PR's NPU attention/quantization files or the HiCache MLA case. The accuracy results below belong to the fixed merge snapshot above; the full model test was not repeated on this newer main commit.

Ran the current merged-source test_npu_hicache_mla.py on machine 216, four Ascend A3 devices (physical 0–3), CANN 9.1.0 / torch_npu 2.10.0.post4 / PyTorch 2.10.0 / Transformers 5.12.1. Used the same full DeepSeek-V2-Lite-W8A8 checkpoint as the previous runs, copied from machine 209, with TP=4, mem_fraction_static=0.8, HiCache enabled, all 1319 GSM8K questions, 5-shot, temperature 0, max 512 new tokens, concurrency 128, and seed 42 for both modes. The updated main test no longer explicitly sets hicache_ratio=1.2; this rerun preserved its defaults, resolving to ratio 2.0 / host_memory_fraction 0.8. No additional unit tests were added or run for this rerun.

ASCEND_USE_FIA Correct / total Accuracy Invalid answers Case result
unset 473 / 1319 35.860500% 3 passed
1 484 / 1319 36.694466% 4 passed

FIA minus baseline: +11 correct answers / +0.833965 percentage points. This single paired run does not establish lossless replacement or isolate the new prefix branch: ASCEND_USE_FIA affects other existing attention paths too. 661 final numerical predictions differed; 102 changed from correct to incorrect and 113 from incorrect to correct.

Both runs used matching source, dataset, devices, model path, evaluation parameters and actual server argument dictionaries. All 1319 requests in each run reused 768 device-cached prefix tokens; neither run forced host-cache reload. Full logs, exact metrics, source hashes and per-question answers were retained. Observed evaluation runtime / output throughput: unset 227.836 s / 691.215 token/s, enabled 205.004 s / 768.774 token/s. Other devices on the host had existing workloads; these observations are not a controlled performance benchmark. A2, A5 and CANN 9.2 remain untested.

Historical validation on machine 209 (2026-09-18)

Ran the existing test/registered/npu/basic_function/HiCache/test_npu_hicache_mla.py at commit 1344e1e2a5ca3e07bf6140b8c606e38465fbc55c on four Ascend910_9362 (A3) devices with CANN 9.1.0, PyTorch 2.10.0, torch_npu 2.10.0.post4 and Transformers 5.12.1.

Used the full vllm-ascend/DeepSeek-V2-Lite-W8A8 checkpoint, adapting its local filesystem path. Each run used TP=4, HiCache enabled, hicache_ratio=1.2, mem_fraction_static=0.8, all 1319 GSM8K questions, 5-shot prompts, temperature 0, max_new_tokens=512 and concurrency 128. The PR-pipeline single-request smoke shortcut was disabled.

Run pair ASCEND_USE_FIA Correct / total Accuracy Invalid answers
Original case defaults unset 477 / 1319 36.163760% 5
Original case defaults 1 469 / 1319 35.557240% 7
Both servers use --random-seed 42 unset 481 / 1319 36.467020% 6
Both servers use --random-seed 42 1 470 / 1319 35.633055% 7

All four runs passed the case's 34% accuracy threshold. FIA scored 0.606520 percentage points lower in the first pair and 0.833965 percentage points lower in the fixed-seed pair. These measurements do not establish accuracy parity or lossless replacement; the ATB default and opt-in gate are retained.

In the fixed-seed pair, 686 final numerical predictions differed: 123 questions changed from correct to incorrect and 112 from incorrect to correct. Both actual server argument dictionaries matched, including seed 42. The first pair used automatically generated server seeds; both pairs are reported. The ATB runs themselves differed on 625 predictions between these two runs, so per-question changes cannot all be attributed to FIA from this experiment alone.

All 1319 requests in each run reused 768 device-cached prefix tokens, exercising prefix prefill with HiCache attached. This workload did not force host-cache eviction/reload. Source hashes, full server logs, exact metrics and per-question answers were retained. ASCEND_USE_FIA is an existing backend-wide flag, so this model comparison includes every path affected by the flag rather than isolating only the new prefix branch.

Additional validation performed before the model runs:

  • CPU regression: 3 test methods / 11 parameter cases passed, covering FP16/BF16, variable lengths, mixed empty prefixes, noncontiguous inputs, padding, large logits and both operator routes. These tests use numerical substitutes for NPU kernels and compare with full dense SDPA.
  • NPU quantization unit suite: all 25 tests passed in the torch_npu 2.10.0.post4 container. These are CPU/mocked tests, including shape, bias and dtype behavior for MXFP8 linear operations.
  • Re-ran the gated production forward_extend prefix branch on Ascend910_9362 (A3), CANN 9.1.0, PyTorch 2.10.0 and torch_npu 2.10.0.post4: 18/18 cases passed. Cache gather, KV projection, FIA, merge and padding execute on NPU; the reference attention executes on CPU. The branch is loaded unchanged from source with unrelated imports isolated; this is not a complete serving/model test.
  • NPU cases cover page sizes 1/128, heads 2/4/8/16/128, query lengths up to 257, prefixes up to 2048, mixed empty prefixes and zero/nonzero padding. NoPE/RoPE/V dimensions are 128/64/128.
dtype Maximum absolute error Maximum per-case RMSE atol / rtol
FP16 0.001953125 0.000209168 0.003 / 0.003
BF16 0.015625 0.001710953 0.02 / 0.02

All NPU branch cases also check output device, shape, dtype, finite values and zero padding. A2/A5 and CANN 9.2 have not been hardware tested.

Speed Tests and Profiling

NZ prefix-cache validation on A3 (2026-09-24)

Commit 0c8ebbce034c11b5e0b9c2fec3193def9de2aa74 uses the existing gather_mla_cache_pages helper from #39589 for both latent and RoPE prefix reads in the two-FIA branch. It restores logical token order before projection and TND attention when the KV cache uses PA-NZ storage. The helper is byte-identical to #39589 after newline normalization. ND behavior is preserved.

Environment: machine 151, Ascend910_9362 (A3), CANN 9.1.0, PyTorch 2.10.0, torch_npu 2.10.0.post6. The patch changes only the two cache-gather calls; it does not reintroduce the reverted quantization API changes.

Operator-path validation and performance: 18 cases plus 6 independent reruns passed. The fixture uses the actual NPUMLATokenToKVPool._set_fia_nz_kv_buffer and its index helper, executing real NPU scatter with randomized token write order. Both latent and RoPE readback are bitwise equal to the logical ND cache. Across all cases, the fixed NZ branch and fixed ND branch outputs are bitwise equal to the original ND branch output. Sampled CPU FP32 reference checks also pass (maximum RMS relative error 0.2995%). The unfixed NZ branch is retained as a negative control and shows RMS relative errors of 6.55%-119.16% on these synthetic cases.

The independent rerun below measures median amortized wall time for cache gather, BF16 projection, attention, copies and merge. QK=192, V=128, latent=512, page=128, H=4; BF16/eager. It uses 20 warmups and 9 rounds of 30 calls per method with randomized method order. NZ BSND explicitly selects the existing per-request branch for the same inputs, without changing production dispatch. NZ here describes cache storage; the two attention calls still use TND inputs after gathering.

Case ND before (ms) ND fixed (ms) NZ fixed (ms) NZ overhead vs ND NZ BSND (ms)
b1_q100_p768_h4 0.507 0.509 0.567 +11.39% 0.440
b8_q100_p768_h4 0.563 0.565 0.627 +11.04% 2.061
b32_q100_p768_h4 0.630 0.629 0.712 +13.16% 6.819
b1_q1024_p8192_h4 0.550 0.552 0.611 +10.60% 0.494
b4_q1024_p8192_h4 1.746 1.745 1.833 +5.03% 1.718
b8_q1024_p8192_h4 3.365 3.364 3.502 +4.09% 3.518

NZ layout restoration adds about 4%-13% to the two-FIA prefix branch in this rerun. With NZ disabled, the helper adds no layout conversion; the first sweep measured ND before/after differences within approximately -0.34% to +0.65%. Short-query batched two-FIA remains faster than per-request BSND. These are branch measurements and do not establish whole-model speedup.

Full model accuracy: ran the original test_npu_hicache_mla.py on the same physical devices 0-3 with DeepSeek-V2-Lite-W8A8, TP=4, HiCache enabled, radix enabled, seed=42, all 1319 GSM8K questions, 5-shot, temperature=0, max_new_tokens=512 and concurrency=128. Both runs set ASCEND_USE_FIA=1; only SGLANG_USE_FIA_NZ changes. The source, model path, dataset, devices, launch arguments and evaluation settings match.

SGLANG_USE_FIA_NZ Correct / total Accuracy Invalid Result Eval runtime (s) Output token/s
0 473 / 1319 35.860500% 6 passed 229.363 701.845
1 465 / 1319 35.253980% 5 passed 246.641 652.957

NZ minus ND: -0.606520 percentage points. Final numerical predictions differ on 581 questions (96 correct-to-incorrect, 88 incorrect-to-correct). These full-model runs exercise other NZ-dependent cache/decode paths as well, so model-level differences cannot be attributed to the prefix gather alone. Evaluation runtime and token/s are observations from accuracy runs with variable generation lengths, not a controlled serving performance benchmark. The host is shared; another process was observed on device 3 at the end of the NZ run, so model-run timing is not a device-exclusive measurement.

Device-prefix cache hits: ND 1319/1319, NZ 1319/1319; cached tokens: ND 1012992, NZ 1012992. Host-cache reload hits: ND 0, NZ 0. This does not establish host-cache eviction/reload correctness when those counts are zero. A2/A5 and CANN 9.2 are not hardware-validated by this run. Full logs, inputs, source hashes, timing samples and per-question predictions were retained.

Controlled prefix-branch comparison on A3 (2026-09-22)

Measured commit ec72a1dd7c4292477bdd8e317a55f87c44f04357 on 2026-09-22, machine 151: one Ascend910_9362 (A3), CANN 9.1.0 (B070 container), PyTorch 2.10.0, torch_npu 2.10.0.post6, BF16, NZ disabled, eager execution. The main shape uses QK=192 (NoPE=128 + RoPE=64), V=128, latent=512, page size=128 and H=4 per rank (the DeepSeek-V2-Lite TP=4 head shape); no TP communication is executed.

The comparison is between two FIA implementations of MLA prefix prefill:

  • Two FIA V2 + merge: this PR's batched TND current-token attention and prefix attention, followed by npu_attention_update.
  • Concat + BSND FIA: the existing if layer.qk_head_dim == layer.v_head_dim branch, concatenating prefix/current K/V and issuing one BSND FIA call per request. B requests therefore issue B FIA calls.

The benchmark extracts both branch bodies directly from the recorded source and uses identical Q/K/V, cache, projection weights and sequence metadata. To compare the same DeepSeek input, it explicitly selects the BSND body for QK=192/V=128 despite the production dispatch condition; this CANN build accepted that shape and passed the numerical checks. The production condition and kernel arguments were not changed. This is not a comparison with ATB or an end-to-end comparison of ASCEND_USE_FIA enabled/disabled.

Times include cache gather, BF16 kv_b_proj, FIA, concatenation/contiguous copies, dtype conversions and merge. They exclude the rest of the model, QKV preparation, scheduling, decode and TP communication, and do not reproduce W8A8 projection. Each value is the median amortized wall time per complete prefix branch, with synchronization at each measurement block. Method order is randomized with a fixed seed. The first sweep uses 12 warmups and 7 rounds of 20 calls per method; an independent rerun uses 20 warmups and 9 rounds of 30 calls. Individually synchronized calls and NPU event spans were also recorded; event spans can include host submission gaps.

Independent rerun (8 cases). B = requests, Q = new query tokens per request, P = cached prefix tokens per request, H = heads per rank. A BSND / two-FIA time ratio above 1 favors the PR's two-FIA path.

B Q per request Prefix per request H Two FIA V2 + merge (ms) Concat + BSND FIA (ms) BSND / two-FIA time
1 100 768 4 0.513 0.395 0.77
2 100 768 4 0.520 0.629 1.21
8 100 768 4 0.524 1.862 3.55
32 100 768 4 0.619 6.677 10.80
128 100 768 4 2.434 26.031 10.70
1 1024 8192 4 0.542 0.457 0.84
4 1024 8192 4 1.715 1.645 0.96
8 1024 8192 4 3.306 3.374 1.02

Result: for Q=100/P=768, the two-FIA branch wins from B=2 in this sweep, reaching 3.55x at B=8 and 10.80x at B=32 in the rerun. At B=1, concatenation is faster: about 23% lower latency for Q=100/P=768 and 16% lower for Q=1024/P=8192. With Q=1024/P=8192, B=4 favors BSND by about 4%, while B=8 is near parity; a 1-2% gap should not be treated as a significant stable win. These are prefix-branch ratios, not whole-model speedups.

The source structure is consistent with the trend: per-request concatenation and B FIA submissions become expensive for short-query batches, while the TND path batches requests into two FIA calls plus one merge. B=1 avoids the second call and merge with BSND. This interpretation has not been confirmed by kernel-level profiling. The results support retaining the batched two-FIA path; a B=1 specialization is a possible follow-up, subject to device/version and full-model validation.

Complete first sweep: 18 cases
B Q per request Prefix per request H Two FIA V2 + merge (ms) Concat + BSND FIA (ms) BSND / two-FIA time
1 100 768 4 0.526 0.386 0.73
2 100 768 4 0.560 0.661 1.18
4 100 768 4 0.611 1.204 1.97
8 100 768 4 0.595 2.044 3.44
16 100 768 4 0.595 3.745 6.30
32 100 768 4 0.618 6.667 10.79
64 100 768 4 1.277 12.902 10.11
128 100 768 4 2.443 26.060 10.67
1 256 2048 4 0.500 0.381 0.76
4 256 2048 4 0.500 0.986 1.97
16 256 2048 4 0.939 3.446 3.67
32 256 2048 4 1.828 6.595 3.61
1 1024 8192 4 0.549 0.457 0.83
4 1024 8192 4 1.695 1.635 0.96
8 1024 8192 4 3.427 3.459 1.01
16 100 768 16 1.181 3.011 2.55
32 100 768 / 0 alternating 4 0.844 6.740 7.98
10 mixed (see below) 768, then nine zeros 4 0.671 1.906 2.84

The mixed B=10 case uses query lengths [101,870,876,874,869,864,904,863,837,768] and prefix lengths [768,0,0,0,0,0,0,0,0,0], matching an earlier observed batch shape with synthetic tensor values. It favors two FIA by 2.84x. B=128/Q=100 contains 12,800 query tokens and is an extended sweep point, not a claim that default chunk limits produce that batch.

Numerical checks: all 18 sweep cases and all 8 rerun cases passed full-output pairwise comparison (atol=0.02, rtol=0.02) and sampled CPU FP32 full-attention reference checks. Maximum full-output absolute difference between methods was 0.001953125; maximum per-case pairwise RMS relative difference was 0.3022%. Maximum sampled RMS relative error against CPU FP32 was 0.2995% for two FIA and 0.2296% for BSND. The CPU reference uses the same projected K/V. This checks the sampled operator paths; it is not a new GSM8K accuracy run. Hardware conclusions here are limited to this A3/CANN/torch_npu configuration.

Historical accuracy-run timings

The fixed-seed full GSM8K runs measured 193.304 s / 814.391 output tokens per second with FIA unset, and 196.474 s / 801.729 output tokens per second with FIA enabled. The first pair measured 225.281 s / 697.903 and 195.270 s / 802.929 respectively. These are accuracy-run measurements with differing generation lengths and dynamic batches, not a controlled performance benchmark; no speedup is claimed. The FIA path adds a separate merge and FP32 intermediate outputs and remains opt-in.

Checklist

  • Format code and run the applicable checks (isort, Ruff formatting/lint, codespell, AST, debug statements, merge markers, whitespace and EOF).
  • Add unit tests for numerical correctness and both operator routes.
  • Document the existing flag's new behavior and validation scope in this PR.
  • Run complete model-level accuracy comparisons, including a fixed-seed pair.
  • Run a dedicated prefix-branch performance benchmark on A3 (full-model performance remains unmeasured).
  • Follow the existing Ascend backend style.

CI States

Latest PR Test (Base): ✅ Run #36403090220
Latest PR Test (Extra): ❌ Run #36403089905
Latest PR Test (AMD ROCm 10): ❌ Run #36403090380

@McZyWu McZyWu changed the title [NPU] Add opt-in FIA for MLA prefix prefill [NPU] Add opt-in FIA for MLA prefix prefill and use torch_npu quant APIs Sep 18, 2026
@github-actions github-actions Bot added the quant LLM Quantization label Sep 18, 2026
Comment thread python/sglang/srt/hardware_backend/npu/moe/quant.py
@McZyWu McZyWu changed the title [NPU] Add opt-in FIA for MLA prefix prefill and use torch_npu quant APIs [NPU] Add opt-in MLA prefix FIA, torch_npu quant APIs, and Gemma norm fallback Sep 22, 2026
@McZyWu

McZyWu commented Sep 22, 2026

Copy link
Copy Markdown
Contributor Author

A3 performance results: two FIA + merge vs. per-request concat + BSND FIA

Measured commit ec72a1dd7c4292477bdd8e317a55f87c44f04357 on 2026-09-22, machine 151: one Ascend910_9362 (A3), CANN 9.1.0 (B070 container), PyTorch 2.10.0, torch_npu 2.10.0.post6, BF16, NZ disabled, eager execution. The main shape uses QK=192 (NoPE=128 + RoPE=64), V=128, latent=512, page size=128 and H=4 per rank (the DeepSeek-V2-Lite TP=4 head shape); no TP communication is executed.

The comparison is between two FIA implementations of MLA prefix prefill:

  • Two FIA V2 + merge: this PR's batched TND current-token attention and prefix attention, followed by npu_attention_update.
  • Concat + BSND FIA: the existing if layer.qk_head_dim == layer.v_head_dim branch, concatenating prefix/current K/V and issuing one BSND FIA call per request. B requests therefore issue B FIA calls.

The benchmark extracts both branch bodies directly from the recorded source and uses identical Q/K/V, cache, projection weights and sequence metadata. To compare the same DeepSeek input, it explicitly selects the BSND body for QK=192/V=128 despite the production dispatch condition; this CANN build accepted that shape and passed the numerical checks. The production condition and kernel arguments were not changed. This is not a comparison with ATB or an end-to-end comparison of ASCEND_USE_FIA enabled/disabled.

Completed 18 shape cases and an independent 8-case rerun. The table below shows the rerun; the PR description now includes the complete 18-case sweep as well. B = request count, Q/P = query/prefix tokens per request, H = heads per rank; a ratio above 1 favors two FIA.

B Q per request Prefix per request H Two FIA V2 + merge (ms) Concat + BSND FIA (ms) BSND / two-FIA time
1 100 768 4 0.513 0.395 0.77
2 100 768 4 0.520 0.629 1.21
8 100 768 4 0.524 1.862 3.55
32 100 768 4 0.619 6.677 10.80
128 100 768 4 2.434 26.031 10.70
1 1024 8192 4 0.542 0.457 0.84
4 1024 8192 4 1.715 1.645 0.96
8 1024 8192 4 3.306 3.374 1.02

Result: for Q=100/P=768, the two-FIA branch wins from B=2 in this sweep, reaching 3.55x at B=8 and 10.80x at B=32 in the rerun. At B=1, concatenation is faster: about 23% lower latency for Q=100/P=768 and 16% lower for Q=1024/P=8192. With Q=1024/P=8192, B=4 favors BSND by about 4%, while B=8 is near parity; a 1-2% gap should not be treated as a significant stable win. These are prefix-branch ratios, not whole-model speedups.

The source structure is consistent with the trend: per-request concatenation and B FIA submissions become expensive for short-query batches, while the TND path batches requests into two FIA calls plus one merge. B=1 avoids the second call and merge with BSND. This interpretation has not been confirmed by kernel-level profiling. The results support retaining the batched two-FIA path; a B=1 specialization is a possible follow-up, subject to device/version and full-model validation.

Times include cache gather, BF16 kv_b_proj, FIA, concatenation/contiguous copies, dtype conversions and merge. They exclude the rest of the model, QKV preparation, scheduling, decode and TP communication, and do not reproduce W8A8 projection. Each value is the median amortized wall time per complete prefix branch, with synchronization at each measurement block. Method order is randomized with a fixed seed. The first sweep uses 12 warmups and 7 rounds of 20 calls per method; an independent rerun uses 20 warmups and 9 rounds of 30 calls. Individually synchronized calls and NPU event spans were also recorded; event spans can include host submission gaps.

Numerical checks: all 18 sweep cases and all 8 rerun cases passed full-output pairwise comparison (atol=0.02, rtol=0.02) and sampled CPU FP32 full-attention reference checks. Maximum full-output absolute difference between methods was 0.001953125; maximum per-case pairwise RMS relative difference was 0.3022%. Maximum sampled RMS relative error against CPU FP32 was 0.2995% for two FIA and 0.2296% for BSND. The CPU reference uses the same projected K/V. This checks the sampled operator paths; it is not a new GSM8K accuracy run. Hardware conclusions here are limited to this A3/CANN/torch_npu configuration.

@McZyWu McZyWu changed the title [NPU] Add opt-in MLA prefix FIA, torch_npu quant APIs, and Gemma norm fallback [NPU] Add opt-in MLA prefix FIA and Gemma norm fallback Sep 24, 2026
@McZyWu

McZyWu commented Sep 24, 2026

Copy link
Copy Markdown
Contributor Author

NZ prefix-cache validation on A3 (2026-09-24)

Commit 0c8ebbce034c11b5e0b9c2fec3193def9de2aa74 uses the existing gather_mla_cache_pages helper from #39589 for both latent and RoPE prefix reads in the two-FIA branch. It restores logical token order before projection and TND attention when the KV cache uses PA-NZ storage. The helper is byte-identical to #39589 after newline normalization. ND behavior is preserved.

Environment: machine 151, Ascend910_9362 (A3), CANN 9.1.0, PyTorch 2.10.0, torch_npu 2.10.0.post6. The patch changes only the two cache-gather calls; it does not reintroduce the reverted quantization API changes.

Operator-path validation and performance: 18 cases plus 6 independent reruns passed. The fixture uses the actual NPUMLATokenToKVPool._set_fia_nz_kv_buffer and its index helper, executing real NPU scatter with randomized token write order. Both latent and RoPE readback are bitwise equal to the logical ND cache. Across all cases, the fixed NZ branch and fixed ND branch outputs are bitwise equal to the original ND branch output. Sampled CPU FP32 reference checks also pass (maximum RMS relative error 0.2995%). The unfixed NZ branch is retained as a negative control and shows RMS relative errors of 6.55%-119.16% on these synthetic cases.

The independent rerun below measures median amortized wall time for cache gather, BF16 projection, attention, copies and merge. QK=192, V=128, latent=512, page=128, H=4; BF16/eager. It uses 20 warmups and 9 rounds of 30 calls per method with randomized method order. NZ BSND explicitly selects the existing per-request branch for the same inputs, without changing production dispatch. NZ here describes cache storage; the two attention calls still use TND inputs after gathering.

Case ND before (ms) ND fixed (ms) NZ fixed (ms) NZ overhead vs ND NZ BSND (ms)
b1_q100_p768_h4 0.507 0.509 0.567 +11.39% 0.440
b8_q100_p768_h4 0.563 0.565 0.627 +11.04% 2.061
b32_q100_p768_h4 0.630 0.629 0.712 +13.16% 6.819
b1_q1024_p8192_h4 0.550 0.552 0.611 +10.60% 0.494
b4_q1024_p8192_h4 1.746 1.745 1.833 +5.03% 1.718
b8_q1024_p8192_h4 3.365 3.364 3.502 +4.09% 3.518

NZ layout restoration adds about 4%-13% to the two-FIA prefix branch in this rerun. With NZ disabled, the helper adds no layout conversion; the first sweep measured ND before/after differences within approximately -0.34% to +0.65%. Short-query batched two-FIA remains faster than per-request BSND. These are branch measurements and do not establish whole-model speedup.

Full model accuracy: ran the original test_npu_hicache_mla.py on the same physical devices 0-3 with DeepSeek-V2-Lite-W8A8, TP=4, HiCache enabled, radix enabled, seed=42, all 1319 GSM8K questions, 5-shot, temperature=0, max_new_tokens=512 and concurrency=128. Both runs set ASCEND_USE_FIA=1; only SGLANG_USE_FIA_NZ changes. The source, model path, dataset, devices, launch arguments and evaluation settings match.

SGLANG_USE_FIA_NZ Correct / total Accuracy Invalid Result Eval runtime (s) Output token/s
0 473 / 1319 35.860500% 6 passed 229.363 701.845
1 465 / 1319 35.253980% 5 passed 246.641 652.957

NZ minus ND: -0.606520 percentage points. Final numerical predictions differ on 581 questions (96 correct-to-incorrect, 88 incorrect-to-correct). These full-model runs exercise other NZ-dependent cache/decode paths as well, so model-level differences cannot be attributed to the prefix gather alone. Evaluation runtime and token/s are observations from accuracy runs with variable generation lengths, not a controlled serving performance benchmark. The host is shared; another process was observed on device 3 at the end of the NZ run, so model-run timing is not a device-exclusive measurement.

Device-prefix cache hits: ND 1319/1319, NZ 1319/1319; cached tokens: ND 1012992, NZ 1012992. Host-cache reload hits: ND 0, NZ 0. This does not establish host-cache eviction/reload correctness when those counts are zero. A2/A5 and CANN 9.2 are not hardware-validated by this run. Full logs, inputs, source hashes, timing samples and per-question predictions were retained.

@McZyWu

McZyWu commented Sep 24, 2026

Copy link
Copy Markdown
Contributor Author

Kimi-K3 accuracy validation with FIA + NZ — 2026-09-24

Completed one full accuracy run on the latest-main integration snapshot recorded below. GSM8K met the requested target; GPQA Diamond finished at 184/198 (92.9293%), one correct answer short of the requested 93.4% target.

Evaluation Correct / total Accuracy Requested target Result
GSM8K, 50 questions 50 / 50 100% 98–100% Met
GPQA Diamond, full dataset 184 / 198 92.9293% ≥93.4% Below target; 185 / 198 would be 93.4343%

Exact source snapshot

  • Main: 261cb826e6c63b79ae984af6a8d40c10bcfd8571
  • PR40131 head: 0c8ebbce034c11b5e0b9c2fec3193def9de2aa74 (includes the NZ logical-prefix-page gather fix)
  • Local merged validation commit: 7e7f973013841a08283197737e418194335f7c89

Environment and serving configuration

Four Ascend nodes, TP=64 / DP=4, container hanwlax-cann910-0922, CANN 9.1.0, PyTorch 2.10.0, torch_npu 2.10.0.post6, Transformers 5.12.1, and the existing sgl-kernel-npu 2026.9.0 package. Model: Kimi-K3-w4a8-int-moe (modelslim, BF16); DSPARK draft: Kimi-K3-DSpark, block size 7.

Enabled ASCEND_USE_FIA=1 and SGLANG_USE_FIA_NZ=1. The launcher retained the active parameters of the container's original run.sh, including NPU graph and DSPARK. An isolated host checkout was selected through PYTHONPATH. Container source, the original run.sh, and the installed operator package were not modified. All four source checkouts remained clean, and the installed operator library hashes matched across nodes.

Evaluation settings

  • GSM8K: original sglang.test.few_shot_gsm8k, 50 questions, 5-shot, temperature 0, max 512 new tokens, client concurrency 128.
  • GPQA Diamond: EvalScope 1.12.0, all 198 questions, batch size 32, temperature 1.0, top_p 0.95, max_tokens 131072, reasoning_effort=max, evaluation seed 42; one complete run.

Validation checks

Both evaluations exited with code 0. GSM8K had no invalid answers; independently checking the saved per-question outputs reproduced 50/50. GPQA completed all 198 requests with zero API or review errors. All 198 responses had stop_reason=stop and a valid extracted A/B/C/D option. The 14 incorrect answers were not request failures, output-length truncations, or answer-extraction failures. The longest response contained 83,668 output tokens, below the configured 131,072 limit. Per-question scores agree with the EvalScope aggregate report. All four serving nodes remained healthy, with no fatal server errors recorded.

No baseline was run, and the evaluation was not repeated with adjusted settings to reach the target. This run shows GPQA slightly below the requested threshold; it does not establish an accuracy delta versus main or attribute that difference to this PR. Full logs, dataset hashes, actual server arguments, predictions, reviews, and the 14 incorrect cases are retained with the validation artifacts.

@McZyWu

McZyWu commented Sep 24, 2026

Copy link
Copy Markdown
Contributor Author

Kimi-K3 128K/1K cached-prefix TTFT — main vs this PR

Four-node end-to-end serving measurement: mean TTFT changed from 8.629 s to 7.401 s (14.22% lower; main/PR TTFT ratio 1.166×) across three runs of 32 requests per variant. The observed average improved, but the small sample and visible run-to-run variation should be considered when interpreting the size of the gain.

Run Main mean TTFT (s) Main + PR40131 mean TTFT (s)
1 10.267 10.418
2 5.236 5.918
3 10.383 5.867
All 96 requests per variant 8.629 7.401

The observed mean-TTFT difference is preliminary and sensitive to run-to-run variation. P99 TTFT did not improve. Additional measured metrics are included for completeness:

Metric Main Main + PR40131
Pooled P99 TTFT, 96 requests (s) 10.578 10.725
Mean TPOT, three-run average (ms) 24.027 28.931
Output throughput, mean of three runs (tokens/s) 839.42 813.25

All six measured runs completed 32/32 requests, with exactly 128000 input tokens and 1000 output tokens per request, and no request errors. Each measured request reported 127872 cached tokens / 128000 = 99.9% cache hit. The dataset is named 100cache and its prompts share all 128000 tokens; the serving path recomputes the final 128-token page, so the observed hit rate is reported as 99.9%, not rounded up to 100%.

Snapshots and configuration

  • Main: 28ec6700dade98786e948603fa80b472852a4b74.
  • PR40131 head: 0c8ebbce034c11b5e0b9c2fec3193def9de2aa74.
  • Merged validation commit: afbcf107613460050fbbe9ae55608e20c1045602.
  • Four Ascend nodes, TP=64 / DP=4; container hanwlax-cann910-0922; existing operator package unchanged.
  • Kimi-K3-w4a8-int-moe, ModelSlim/BF16, DSPARK block size 7; original active run.sh serving parameters retained.
  • ASCEND_USE_FIA=1 and SGLANG_USE_FIA_NZ=1 in both variants. Main uses the existing ring-MLA prefix implementation; the PR selects its FIA prefix implementation and NZ logical-page gather fix. Thus this is a whole-PR comparison against main, not a FIA-toggle-only ablation.
  • Effective server configuration matched, including each node's server random seed. All eight source checkouts were clean and all operator-library hashes matched the original installed library.

Workload and timing method

Used the existing run_exact_benchmark.sh 128K/1K 100cache workload: sglang.bench_serving, random dataset, seed 42, range ratio 1, 32 requests, max concurrency 40 (at most 32 requests in this workload), no implicit benchmark warmup. The exact dataset SHA256 is e234a39cffcacbd5b90296e62b046fc834ede3a2c48712bdaf115ff42e91df81.

Each measured round starts with cache flush, four 128000/1 priming requests, then four unmeasured 128000/1000 stabilization requests to populate the reusable cache state on all four DP ranks. The initial main-side check with only the original 1-token priming reached 98.92% overall cache hit; it was retained as a setup diagnostic and excluded before the matched hot-cache protocol was run on either variant. All six included rounds have identical per-request cache counts and output lengths.

Reported TTFT comes from the streaming benchmark client and includes serving/queueing overhead. All three rounds are retained; the main per-run mean ranges from 5.236–10.383 s, and PR from 5.867–10.418 s. This is a small end-to-end comparison with visible run-to-run variation, not an isolated kernel timing. Full per-request timings, cache reports, actual server arguments, logs, and environment attestations are retained.

This performance comparison uses the newer main snapshot above. The previous accuracy results used main 261cb826e6c63b79ae984af6a8d40c10bcfd8571: GSM8K 50/50 (100%) and GPQA Diamond 184/198 (92.9293%). Accuracy was not rerun as part of this performance comparison.

@sglang-npu-bot

Copy link
Copy Markdown
Collaborator

/tag-and-rerun-ci

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

npu quant LLM Quantization run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants