Conversation
|
This pull request has merge conflicts that must be resolved before it can be |
|
This pull request has merge conflicts that must be resolved before it can be |
Resolve conflict in vllm/models/glm5next/nvidia/sparse_indexer.py (formerly model_executor/layers/sparse_attn_indexer_kpool.py, moved and AMD-split by vllm-project#55358): keep upstream's post-refactor NVIDIA code path and re-apply the TP row-shard slicing (row_start/row_end, sliced cu_seqlen_ks/ke, all_gatherv exchange) on top. Signed-off-by: Zheng Cai <8370601+zigzagcai@users.noreply.github.com>
|
Hi @zigzagcai, the pre-commit checks have failed. Please run: uv pip install pre-commit>=4.5.1
pre-commit install
pre-commit run --all-filesThen, commit the changes and push to your branch. For future commits, |
yewentao256
left a comment
There was a problem hiding this comment.
Thanks for the work!
Please do the following:
- shrink the diff, idealy < 500 LOC
- e2e accuracy using
lm_eval... - e2e benchmark using
vllm bench serve...
When doing e2e test, please attach with full output log in case some agents fake the data.
|
Hi @yewentao256 Thanks for the invaluable feedback!
The diff of this PR's functionality is actually <500 LOC, and the majority of diff lines is the added unit test file.
I have added the e2e accuracy/benchmarking details and logs, with the latest upstream baseline 3059155b47cee5269eaf446e24c9d0f8f908c7c9 and latest PR commit f0fc8269fbea535b9bae87eb8d021b37682e5cbb. E2E accuracy and benchmark (FlashInfer cubin enabled)I reran the Model Runner V2 + SpecDecode comparison after installing and enabling the FlashInfer cubin artifacts. Both revisions used the same GLM-5.3-Flash weights, 4x NVIDIA H200 GPUs, MTP with 5 speculative tokens,
Long-context accuracy (
|
| Context | Baseline exact match | Optimized exact match |
|---|---|---|
| 256K | 6/6 (1.0) | 6/6 (1.0) |
| 512K | 6/6 (1.0) | 6/6 (1.0) |
| 1M | 6/6 (1.0) | 6/6 (1.0) |
Full output logs and machine-readable results are available here: public gist. The gist contains full-output.log (complete server, benchmark, and lm_eval output), comparison.json, and this report. Download the complete archive, including per-request JSON files, benchmark server logs, task definitions, and lm_eval samples: pr54951-flashinfer-logs.tar.gz.
|
Thanks for this — it ports cleanly to ROCm, and it pays off there too, at a much lower threshold. ROCm port: jin-amd@cc0609a — one commit on top of this branch's head, +73/−26. On the threshold question: on MI325X at TP4 the replicated kpool indexer costs ~58 ms of a ~663 ms 16K-token prefill chunk ( Measured with
It also runs with MTP (k=5) on ROCm: long prompts shard and complete without errors. To keep this PR's diff down (per @yewentao256), I'll open it as a follow-up once this lands — unless you'd rather take it here. Two smaller notes:
|
|
@zigzagcai There are still a lot of code including unit tests could be simplified from my view. Also, as agents may generated fake results sometimes, please just copy paste full command line and the raw output here. |
Reduce redundant row-sharding tests and comments while retaining core correctness and gate coverage. Co-authored-by: OpenAI Codex <noreply@openai.com>
The metadata builder uses its class-level chunking implementation after the PCP/DCP refactor. Co-authored-by: OpenAI Codex <noreply@openai.com>
Hi @yewentao256 , thanks for the careful review! Firstly, this PR's optimziation was driven and discovered by my-developed end-to-end autonomous evolution system through a loop of profiling, experimentation, and verifier feedback. I understand your concern that agents may sometimes generate fake results. We've done a lot of work to prevent reward hacking from verification system(for example, adversarial agent to perform verification, since the implementation agent and adversarial agent usually have different training data distribution, and some human developed verification framework to improve benchmark numerical stability and detect code that might generate fake results, using the techniques like AST/roofline model/...). And we also integrate loopx framework to prevent goal drift in this kind of long-horizon tasks. Besides, I also manually reviewed the implementation logic behind the optimization to ensure that it wasn’t just AI-generated sloppy code, but a valuable optimization. Secondly, for the follow-up questions:
|
|
@yewentao256 , I also added table to compare computation and communication complexity, and flowchart to show the computation saving and communication overhead. Therefore the benefits of this PR can be better understood.
1. Before vs After (TP=4)2. Why it is faster — three key points
3. Prefill data flow (optimized path)Not an attention approximation. Every query row still attends to the full key history — only which rank computes which row changes. Decode kernels are untouched. 4. Why shard on query rows (not heads / KV)
The planner is cost-balanced, not a naive equal-row split: causal history lengths differ per row, so slices are chosen so each rank does roughly the same amount of scoring work (avoids a long tail before the In one sentence: replace |
|
I see the tests were done with H100, do these improvements also translate to Blackwell GPU's ? |
Hi @gaby ,the tests were done under H200. I don't have Blackwell GPU at hand, but I think these improvements are generic and should works well with Blackwell. |
yewentao256
left a comment
There was a problem hiding this comment.
Thanks! Please take a look at this AI generated comment:
[P2] Avoid row sharding when the indexer takes the short-prefill path
sparse_mla_force_mqa=True makes prefill_uses_mqa true even when max_prefill_seq_len <= index_topk, so this can create row_shard_sizes for a batch where the GLM indexer later takes short_prefill. In that path the scoring loop is skipped entirely and every rank just fills the causal indices locally, but shard_sizes != None still causes the 2051-wide all_gatherv.
For example, at TP4, 32 × 2048-token prefills reach the 65,536-row threshold and would exchange ~513 MiB/layer while saving zero indexer scoring work. sparse_mla_force_mqa is also a supported workaround on SM120, so this is reachable in practice. Could we suppress row_shard_sizes whenever prefill_max_seq_len <= index_topk, irrespective of whether MQA itself was forced?Short prefills fill all causal indices locally without MQA scoring. Ignore row-shard sizes on this path to avoid an unnecessary all-gather when shared metadata was built with forced MQA. Co-authored-by: Codex <noreply@openai.com> Signed-off-by: Codex <noreply@openai.com>
5bd642d to
eeb7409
Compare
Thanks for catching this! |
yewentao256
left a comment
There was a problem hiding this comment.
Thanks! Could you check if this PR would affect dsv4.1? eg. indexer_sparse_logits is True
Signed-off-by: Zheng Cai <8370601+zigzagcai@users.noreply.github.com>
1b73618 to
2e26eac
Compare
|
This pull request has merge conflicts that must be resolved before it can be |
TL;DR
How it works / why it is faster: Instead of repeating sparse-indexer prefill MQA scoring and top-k on every TP rank, partition independent query rows into cost-balanced contiguous slices. Each rank scores its rows against the full available key history;
all_gathervreassembles the final int32 indices. This removes redundant computation in exchange for fixed-width index communication, so savings grow with context length. There is no attention approximation. Decode kernels are unchanged; smaller batches retain the existing path (activation requires at least16,384 × TPscheduled prefill query rows).Performance, baseline → optimized: main parent
5893426b88f7→ PR head16934c6e76e5; GLM-5.3-Flash on 4× H200, TP4, C1, 256 output tokens, 65,536-token batch budget, legacy runner + explicit FA3 backend. Four fresh-process A/B pairs in ABBA/BAAB order; warmups excluded. Throughput below counts input + output tokens at C1, not saturated serving throughput.Server-side TPOT at 1M: 6.631 → 6.633 ms/token. Across the full 1K–1M scan, absolute TPOT percentage-change point estimates are <0.04%; every paired 95% interval includes zero. This does not prove equivalence. The 2K TTFT point estimate is 2.84% slower (interval crosses zero); we do not claim blanket short-input non-regression.
Correctness / targeted accuracy, baseline → optimized: Across 33 distinct retrieval probes spanning 1K–1M, repeated over four processes per arm, expected-answer presence is 132/132 → 132/132; strict answer-only accuracy is 127/132 → 129/132. Exact generated-token agreement is 130/132 A/B (baseline A/A: 98/99). Both observed A/B differences contain the same correct code, with an additional explanatory sentence in baseline. These are targeted retrieval checks, not a broad quality benchmark or a claim of bitwise equivalence. The 54 focused PR tests also pass.
Full eleven-length results, paired confidence intervals, implementation scope, and reproducibility details follow.
Summary
Shard replicated sparse-indexer prefill query rows across tensor-parallel ranks. Each rank computes MQA logits and top-k for a cost-balanced contiguous slice, then one variable-size
all_gathervcall per indexer layer restores the original row layout. This removes redundant scoring; it does not approximate attention or merge independent top-k candidates. The API call is not a claim of one underlying physical NCCL operation.For GLM-5.3-Flash, the exchange includes the incomplete k-pool tail:
index_topk + index_kpool - 1 = 2048 + 4 - 1 = 2,051int32 columns. KV gathering remains unconditional because subsequent query chunks can reuse that workspace.The avoided MQA scoring work grows with both query rows and the available pooled-key history, whereas the exchanged index width per query row stays fixed. Longer histories therefore offer more computation to save per byte exchanged. The cost-balanced partition also accounts for different causal-history lengths across rows. This targets prefill TTFT and prefill-dominated total-token throughput; the row partition and exchange are outside the decode branch, so a direct decode-kernel speedup is not expected. Indirect TPOT changes are measured and reported below.
Scope and activation
The diff against main is four files: the shared DSA indexer metadata/planner, the generic sparse indexer, the GLM k-pool indexer, and their focused tests. #53906 is already in the baseline; this PR does not reintroduce its model support or kernel changes.
16,384 × TPscheduled prefill query rows in the current batch.Important configuration constraint: at TP4,
max_num_batched_tokens=8192cannot activate this implementation, even with a 1M-token accumulated context. Both arms below use 65,536, so C1 inputs of 1K–32K exercise the fallback, and full 64K prefill chunks can exercise sharding. The gate is a batch-row count, not the request's context length; several shorter concurrent requests can also reach it.Exact-head A/B rescan
These measurements replace the previous historical tables. Those tables used older PR/runtime overlays and an 8,192-token batch budget; their gains cannot be attributed to the current head and gate.
5893426b88f7b3cd21101d194eb1c6f0a6f0e27b, the optimized commit's parent16934c6e76e508b22037d5d4040488d490090044Both arms load their respective clean source trees, the same model snapshot (
3f1971b7b5f7a528c9c4ef6212c8785298a8c24a), and identical compiled vLLM extensions from the official wheel for the baseline commit. Worker-level provenance checks verify the loaded vLLM paths for all four TP ranks. The historical vendor-vLLM Python overlay is not used. The PR range contains one commit.Configuration and stability protocol
VLLM_USE_V2_MODEL_RUNNER=0) and explicitFLASH_ATTN_MLA_SPARSEattention backend, chunked prefill,max_model_len=1,048,576,max_num_batched_tokens=65,536,max_num_seqs=8, GPU-memory utilization 0.90.FULL_AND_PIECEWISECUDA-graph configuration with main's default breakable-graph behavior. The effective engine configuration reportsCompilationMode.NONEand CUDA-graph capture sizes[1, 2, 4, 8, 16]; this is not a torch.compile-enabled prefill benchmark. Torch 2.13.0+cu130, Triton 3.7.1, FlashInfer 0.6.18, Transformers 5.12.1. Same process environment/cache policy in both arms; OMP/MKL threads fixed at 1. Synchronous CUDA debugging is off for timed runs. These explicit runner/backend settings avoid failures observed in the unmodified baseline; see limitations below.ignore_eos=true. 1M means 1,000,000 input tokens, leaving room for generation within the server limit. Other K labels are powers of 1,024.1 − optimized/baseline) or higher throughput (optimized/baseline − 1). Confidence intervals resample the four restart pairs, not individual requests (20,000 bootstrap draws). Four pairs give limited statistical resolution; an interval crossing zero is not proof of equivalence.Performance
At 1M, measured TTFT improves by +25.42% and C1 total-token throughput by +32.59%. Across all lengths, TPOT improvement point estimates range from -0.03% to +0.01%. The fallback 1K–32K TTFT changes range from -2.84% to +0.43%. The intervals below and individual restart results, not just the sign of a small point estimate, determine how confidently an improvement or regression can be distinguished from run-to-run variation. These results do not prove zero overhead outside the tested configuration. Every TPOT interval includes zero: this scan does not establish a systematic TPOT change, nor does it prove exact equivalence.
Short-input caveat: the most negative fallback TTFT point is 2K: 151.39 → 155.69 ms (-2.84% improvement; paired 95% interval [-8.18, +2.24]%). This observation is retained, not discarded or declared harmless solely because the new collective is gated off. The scan does not establish blanket short-input non-regression or a causal explanation for this timing difference.
95% paired-restart bootstrap intervals for percentage improvements
The four individual adjacent-restart improvements are shown in acquisition order (AB, BA, BA, AB):
Client-side latency cross-check (same samples)
Client TTFT includes request upload/processing and transport. Server-side TPOT is the primary decode measure because clients can receive multiple generated tokens in a single SSE chunk, especially after very long prefills.
Correctness / targeted accuracy A/B
This is a targeted long-context retrieval and numerical-equivalence evaluation, not a general model-quality benchmark. Three distinct generated record probes per length place the answer at 10%, 50%, and 90% of the prompt. Each is repeated in the four independent processes per arm: 12 observations per arm per length, but only three distinct tasks. The prompt has a non-thinking assistant boundary; greedy decoding stops normally with at most 64 output tokens.
The table reports strict expected-answer accuracy, whether the expected answer is present, and exact generated-token equality for matched prompts. Answer containment alone is a weaker check: a response can contain the correct code yet fail strict answer-only formatting by adding an explanatory sentence. Baseline A/A compares adjacent independent baseline restarts, providing a control for process-to-process repeatability. HTTP success and valid token counts alone are not treated as accuracy.
Matched A/B output-token sequences agree in 130/132 comparisons; baseline A/A agrees in 98/99. The expected answer is present in 132/132 baseline and 132/132 optimized responses. Strict answer-only format passes 127/132 baseline and 129/132 optimized responses. Answer containment does not excuse extra or conflicting text; strict answer-only formatting and full token equality are reported separately. Repeated observations of the same three tasks are not independent evidence of broad model accuracy. Multimodal inputs, MTP, other DSA models, and general reasoning quality have not been evaluated by this scan.
Inspection of both A/B mismatches (the 90%-depth task at 512K and 1M) shows the same correct access code in both responses: baseline adds an explanatory sentence, while optimized returns only the code. The baseline A/A mismatch at 1M has the same formatting distinction; optimized's own restarts also differ in formatting at 512K. No different retrieved code was observed in these probes. This is an inspection of the observed mismatches, not a claim that all possible numerical differences are formatting-only.
Additional checks:
[65,536, 2,051]gathered result. Its timings are excluded from the performance table.Runtime qualification / limitations: all 704 measured performance requests complete with exactly 256 output tokens; all 264 separate retrieval requests complete. This is not an error-free-runtime claim: the common FA3 runtime emits TMA-descriptor diagnostics, which are preserved in the server logs. Two separate baseline failures were investigated before fixing the common benchmark configuration: (1) the automatically selected FlashInfer SM90 sparse MLA backend fails during 65,536-row startup/autotune; a standalone NoPE MLA probe also fails at 65,536 rows while passing 64 rows; (2) with FA3, the V2 runner fails at 256K in its slot-mapping kernel, reproduced with eager execution and CUDA_LAUNCH_BLOCKING=1. The legacy-runner/FA3 diagnostic then passes 256K, 512K, and 1M. Failed and synchronous diagnostic admissions are excluded from the performance table. No unrelated baseline source fix is overlaid. The common user environment also registers an unrelated SGLang operator at Python startup; neither pinned vLLM source tree has a callsite for that operator. These results qualify only the explicit legacy-runner/FA3 configuration, not default V2 or auto-selected FlashInfer SM90 MLA.
Reproduction
Use clean checkouts of the two pinned commits and one matched environment. Pin the compiled artifacts to the parent-main wheel (rather than allowing each editable installation to choose a different nightly):
Wheel SHA256:
93097b58f56b838b3d18928753e00b5e18624fb9084d14928e748e5fdc176dc3. The measurements used the same extracted wheel extensions in both source worktrees and checked the four workers' module paths and extension hashes.--no-depsabove assumes the explicitly pinned common runtime is already installed; it is not a complete environment bootstrap.Set
MODEL_DIRto the local model snapshot, andPYTHONto the common environment's Python. Start only one arm at a time; wait for health before running requests:Use fresh process tags
00-baseline,01-optimized,02-optimized,03-baseline,04-optimized,05-baseline,06-baseline,07-optimized. For each, run performance and then correctness with the client below; stop that owned server before starting the next. Keep the same shared compilation caches and do not include server startup in timing.Pair adjacent fresh processes and aggregate as specified above. Do not merge any diagnostic pilot or failed baseline admission into these samples. The source trees are unchanged by the benchmark; benchmark helpers are not additional PR commits.
Exact-token client and synthetic retrieval prompt generator (save as pr54951_measure.py)
Only the local model-location setting is made portable here. The tokenizer, prompt construction, seeds, token accounting, and measurement logic are the same as the measured campaign.
AI assistance and ownership
AI assistance was used for implementation and benchmark analysis. The author owns validation, review follow-up, and maintenance. This PR optimizes the already-merged GLM/DSA prefill path; it is not a duplicate model-support submission.