Skip to content

[FlyDSL][DSv4] Prototype FP4 MQA streaming TopK - #5282

Draft
AMD-yanfeiwang wants to merge 4 commits into
ROCm:mainfrom
AMD-yanfeiwang:perf/fp4-mqa-streaming-topk
Draft

AMD-yanfeiwang wants to merge 4 commits into
ROCm:mainfrom
AMD-yanfeiwang:perf/fp4-mqa-streaming-topk

Conversation

@AMD-yanfeiwang

Copy link
Copy Markdown
Contributor

Motivation

The DeepSeek-V4 FP4 indexer currently materializes an FP32 [query_rows, c4_context] score matrix before TopK. At 768K full-token context this is about 0.75 MiB per query row; a 1024-row local prefill is about 768 MiB per C4 layer invocation. SGLang currently avoids allocator fragmentation with a persistent 2 GiB slab (sgl-project/sglang#37660), but that reserves context-sized storage instead of removing the intermediate.

This draft adds an AITER operator that computes exact TopK without materializing the full logits matrix.

Changes

  • Add flydsl_pa_mqa_topk_fp4_prefill for gfx950 FP4 paged MQA.
  • Fuse score production with CTA-local exact TopK:
    • scores stream through a bounded LDS pool;
    • four-byte score radix plus four-byte logical-index radix implements total order (score descending, logical index ascending);
    • NaNs sort below numeric values and +0 > -0 is preserved;
    • each CTA emits only (candidate_value, raw_index, valid_count).
  • Merge split candidates exactly, then map raw C4 positions through the page table.
  • Preserve the existing short-row contract: when length <= topk, return every logical position sequentially and pad with -1.
  • Add a bounded context-tiled implementation as a correctness/reference fallback.
  • Keep the existing full-logits API source-compatible.

For the common SGLang prefill shape rows=1024, parallel_unit_num=1024, topk=1024, internal candidate + merge scratch is about 20 MiB, independent of context length, versus roughly 768 MiB of logits at 768K context. No device-global context-sized allocation is introduced.

Scope

Initial production specialization:

  • gfx950 / wave64
  • H=64, D=128
  • FP4 E2M1 payload with shuffled UE8M0 scales
  • KV page size 64, score block 256
  • TopK 512 or 1024
  • ragged prefill windows

Graph capture is rejected explicitly in this first version; SGLang will retain its existing graph path as fallback. A dependent SGLang draft will capability-probe this API and keep old AITER compatibility. sgl-project/sglang#38086 remains an independent workspace-based alternative.

Validation

Completed:

  • Black / Ruff / Python syntax checks
  • CPU reference/property tests for exact candidate union and special-value ordering
  • offline gfx950 FlyDSL compilation for TopK 512/1024 and the legacy scorer
  • safety review for empty windows, nonzero starts, speculative page lookahead, stream affinity, and launcher lifetime

Queued on MI355X:

  • runtime fused vs full-logits correctness for TopK 512/1024
  • shuffled page tables, empty/short/ragged windows, and non-current streams
  • fused/tiled/full-logits latency and peak-memory benchmarks
  • dependent SGLang adapter tests

This stays draft until those GPU and end-to-end long-context results are attached.

@github-actions

github-actions Bot commented Sep 5, 2026

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every PR:

  • ✅ Pre-checks (submodule verification, code formatting)
  • ✅ Aiter op tests (gfx942 + gfx950)
  • ✅ Triton tests on MI35X (only when aiter/ops/triton/** or related paths are changed)

Extended tests (opt-in via labels):

Label Tests
ci:gfx1250-ffm-triton Run the five-shard gfx1250 FFM Triton test suite
ci:triton-300x Run an additional Triton test job on MI300X in PRs; main branch always runs both MI35X and MI300X
multigpu Aiter multi-GPU tests on the 8-GPU runner
ci:sglang SGLang integration tests: DeepSeek-R1-MXFP4 accuracy, Qwen 3.5 accuracy
ci:atom ATOM benchmark: DeepSeek-R1-0528, GPT-OSS-120B
ci:atom_full ATOM accuracy suite for PR and main models from ATOM models_accuracy.json
ci:vllm vLLM benchmark: GPT-OSS-120B, DeepSeek-R1-0528, Kimi-K2.5
ci:all All standard extended tests (excludes ci:atom_full)

Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 5282 --add-label <label>

PR title tags & labels:
Component tags ([Triton/Gluon], [HIP], [CK], [ASM], ...) are added to the PR title and as PR labels automatically from the changed files and re-synced on every push — change-type tags like [fix]/[Perf], op tags like [MLA], and human labels (ci:*) are left untouched. Add the no-auto-title label to opt this PR out.

@AMD-yanfeiwang

Copy link
Copy Markdown
Contributor Author

MI355X runtime update (gfx950, ROCm 7.2):

  • focused AITER suite: 18 passed (K=512/1024, exact candidate merge, empty/ragged windows, streams)
  • dependent SGLang adapter suite: 5 passed
  • bounded and true-fused outputs match the full-logits selected set for K=512/1024

Production-grid microbenchmark (rows=1024, parallel_unit_num=1024, score_batch_chunks=16):

  • C4 width 196,608 (~768K full context), K=512: legacy 1.908 ms / 868 MiB path storage; fused 4.522 ms / 16 MiB
  • C4 width 196,608, K=1024: legacy 1.935 ms / 872 MiB; fused 4.810 ms / 32 MiB
  • C4 width 262,144 (~1M full context), K=512: legacy 2.505 ms / 1,156 MiB; fused 6.059 ms / 16 MiB
  • C4 width 262,144, K=1024: legacy 2.537 ms / 1,160 MiB; fused 6.438 ms / 32 MiB

So the prototype removes 96–99% of score-path storage, but is currently about 2.4–2.5x slower at the core operator. That does not meet the default-path performance gate yet; this PR remains Draft while the 8-GPU end-to-end run finishes. The result supports keeping sgl-project/sglang#38086 as the near-term alternative unless further local-selection optimization closes the gap.

@AMD-yanfeiwang AMD-yanfeiwang changed the title [FlyDSL][DSv4] Fuse FP4 MQA scoring with streaming TopK [FlyDSL][DSv4] Prototype FP4 MQA streaming TopK Sep 5, 2026
@AMD-yanfeiwang
AMD-yanfeiwang force-pushed the perf/fp4-mqa-streaming-topk branch from e369bec to b400855 Compare September 6, 2026 13:24
zufayu added a commit that referenced this pull request Sep 7, 2026
…y looked right

Pointing the FlyDSL collectors at $WORK/head was correct and insufficient: that
tree was populated only inside `if grep -q "invariant-removed\|api-signature"`.
On a PR deriving neither family it is an empty directory, and a collector that
reads files finds nothing to read.

#5207 has six buffer calls in its diff -- five bounded, one a bare
`make_buffer_tensor(W_scale, max_size=True)` -- and flydslbounds returned zero.
#5301, which the previous commit was verified against, deletes a guard and so
derives invariant-removed. That is the only reason the fix appeared to work. A
fixture cannot see this: every test here builds its own root and hands it over,
so the tree is always populated.

Materialising head is now unconditional; only the `evidence` call stays behind
the family test. The cost is a few `git show` calls against objects already
fetched at line 125.

On #5207 B8 now reports exactly one candidate and it is worth asking about:
`W_scale` is bound with `num_records_bytes=n_heads * self.n_per_head * 4` at
line 184 of the same file, and with `max_size=True` at line 92, inside a
BlockScale whose constructor already receives head and n_scale_cols. Same
tensor, two treatments, one of them unbounded.

Also verified silent for the right reason on #5146 (its one added buffer call
is buffer_ops.py forwarding variables, not a kernel binding a tensor) and on
#5282 (no buffer construction in the diff at all).

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

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant