Skip to content

[FlyDSL] [gfx950] Add explicit page strides and 64-bit rebase for FP4 MQA logits - #5518

Open
LiuYinfeng01 wants to merge 5 commits into
ROCm:mainfrom
LiuYinfeng01:dsv4-fp4-strided-pages
Open

LiuYinfeng01 wants to merge 5 commits into
ROCm:mainfrom
LiuYinfeng01:dsv4-fp4-strided-pages

Conversation

@LiuYinfeng01

@LiuYinfeng01 LiuYinfeng01 commented Sep 15, 2026 •

Copy link
Copy Markdown
Contributor

Summary

Extend the existing gfx950 FlyDSL paged MXFP4 MQA scorers with shared-pool page strides and 64-bit physical-page rebasing. This changes addressing, not FP4 arithmetic or the default dense layout. Consumer: vllm-project/vllm#57517.

The public runtime wrappers read kv_cache.stride(0), kv_scale.stride(0), and block_tables.stride(0); optional stride arguments belong to the build/compile APIs. Cache strides are bytes for uint8 tensors; table stride is int32 elements. _i32_buffer(..., byte_offset=...) rebases the descriptor before page-local loads.

The scorers now also accept 128-row pages while retaining the 256-token compute chunk. Scale writers/readers use four equal token groups instead of a hard-coded 16-row group, preserving the existing 64-row ABI and correctly packing the two 64-token warp regions of a 128-row page. Both 64- and 128-row decode and prefill cases match the exact FP4-dequant reference at cosine 1.000000. The default decode and prefill correctness sweeps now include 128-row page cases.

The original offset is 4,608,000,000 bytes = 4.608 GB = 4.292 GiB, not 4.6 GiB. #4609 migrated existing scorers to the fx API; it did not first introduce FP4 scoring.

FP8 -> FP4 indexer micro-benchmark

This is the primary performance comparison for the scorer: the previous paged
FP8 indexer pipeline gathers paged K/scales into contiguous buffers and then
runs FP8 MQA logits; the FP4 path reads the paged cache directly in one kernel.
Both operate on the same synthetic indexer shape and exclude model GEMMs,
attention, scheduler and CPU/server overhead.

MI355X/gfx950, H=64, D=128, page=64. These results were measured inside the
Sep-10 ROCm 10 nightly (sha256:72e90adf360c...), with Torch
2.12.0+rocm10.0.0 and Triton
3.8.0+git4cff872c.rocm10.0.0. The FP8 arm uses the exact merged #5603 source
(a84bd368); the FP4 arm adds only this PR's three runtime scorer files. Each
shape uses five independent processes, with 10 warmups and 50 timed iterations
per process; medians are shown. FP8 total includes paged gather and FP8 logits;
FP4 is the single direct-paged scorer.

Shape B Query rows / request Context FP4 us FP8 gather us FP8 total us FP8 / FP4 FP4 vs BF16 cos FP8 vs BF16 cos FP4 vs FP8 cos
1x512x64K 1 512 65,536 259.82 8.41 315.00 1.21x 0.987314 0.999301 0.985827
1x1Kx64K 1 1,024 65,536 450.64 8.48 591.55 1.31x 0.987140 0.999295 0.985643
1x1Kx128K 1 1,024 131,072 864.12 13.90 1166.32 1.35x 0.987143 0.999293 0.985628
4x256x16K 4 256 16,384 126.85 8.35 184.82 1.47x 0.987142 0.999295 0.985647
8x128x8K 8 128 8,192 63.57 8.99 111.17 1.75x 0.987139 0.999295 0.985643
1x8Kx8K 1 8,192 8,192 414.53 4.31 536.86 1.30x 0.987229 0.999298 0.985725

Across this exact-stack sweep, FP4 is 1.21x–1.75x faster than FP8
gather+score. FP4-vs-FP8 is computed directly from both output vectors over the
same valid windows. Every shape also reports cos=1.000000 against the exact
FP4-dequant reference, separating kernel correctness from the expected
quantization difference.

for spec in \
  "1 512 65536" "1 1024 65536" "1 1024 131072" \
  "4 256 16384" "8 128 8192" "1 8192 8192"; do
  read -r bs nq ctx <<< "$spec"
  for repeat in {0..4}; do
    python - "$bs" "$nq" "$ctx" "$repeat" <<'PY'
import sys
from op_tests.test_flydsl_pa_mqa_logits_fp4_prefill import run_case
bs, nq, ctx, repeat = map(int, sys.argv[1:])
run_case(
    bs,
    [[ctx] * nq for _ in range(bs)],
    seed=42 + repeat,
    bench=True,
    iters=50,
    warmup=10,
    parallel_unit_num=max(512, bs * nq),
)
PY
  done
done

The existing op-test calls the FP8 reference pipeline ATOM cp_gather+fp8_logits; that test name and implementation remain unchanged.
Decode uses H=64, D=128, page=64 and the new graph-stable auto-CTA schedule.
The search covered 16–4096 CTAs at 8K/32K/64K/128K; each final row is the
median of five independent processes with 20 warmups and 100 timed iterations.
FP4 and FP8 are measured in the same process on identical inputs.

Context Decode batch Auto CTA FP4 us (range) FP8 us FP8 / FP4
8K 1 64 4.21 (3.62–4.27) 4.17 0.99x
8K 2 96 4.34 (3.75–4.43) 4.31 0.99x
8K 4 256 4.48 (3.95–4.58) 4.35 0.97x
8K 8 256 4.54 (3.94–4.58) 4.50 0.99x
8K 16 512 4.84 (4.80–4.91) 4.82 1.00x
8K 64 512 8.36 (8.09–8.62) 11.89 1.42x
128K 1 256 4.04 (3.99–4.59) 4.14 1.03x
128K 2 512 5.45 (4.86–5.50) 5.50 1.01x
128K 4 1280 7.37 (7.26–8.05) 12.36 1.68x
128K 8 1280 11.87 (11.07–12.04) 20.79 1.75x
128K 16 2048 17.60 (17.37–17.99) 33.27 1.89x
128K 64 4096 56.53 (55.46–57.02) 165.55 2.93x

At 8K, the tuned grid removes the fixed-512 low-batch penalty: B1–B16 are
within 3% of FP8 and B64 is 1.42x faster. At 128K, B1/B2 are at parity and
B4–B64 are 1.68x–2.93x faster; B64 drops from 70.44 us with 512 CTAs to
56.53 us with the tuned grid. The selection uses only host-known row and
maximum-context buckets, so it adds no device-to-host synchronization and is
stable under graph replay. Explicit parallel_unit_num remains an override.

Correctness / regression

Local MI355X, gfx950, GPU 7. Four-commit branch: functional kernel commit d6807379, consolidated test/benchmark commit 0030d058, CI coverage fix 063f3eb5, and auto-CTA tuning commit 6ab04df3; all are authored by LiuYinfeng01. Torch 2.12.0+git6bbd260, HIP 7.2.53211; original validation used image Triton 3.7.1; the complete correctness suite and standard shape sweep were rerun with the upstream CI Triton 3.8.0+amd.rocm7.2.0.git111ff227 pin. The strided-page test now enables graph replay and page-map/scale sensitivity checks by default, while non-gfx950 runners print [skip] and exit 0. No NVIDIA validation is implied.

python op_tests/test_flydsl_pa_mqa_logits_fp4.py
python op_tests/test_flydsl_pa_mqa_logits_fp4_prefill.py
python op_tests/test_flydsl_pa_mqa_logits_fp4_strided_pages.py \
  --json-output address-results.json
Check Result
Existing decode sweep PASS; cosine 1.000000, existing error/mask checks
Existing prefill/variable-qlen sweep PASS; cosine against FP4 reference 1.000000, BF16 comparison and -inf masks retained
Random dense vs shared-pool PASS; identical payloads, separately contiguous dense K/scales, bit-identical dense/strided outputs
Independent FP4-dequant oracle PASS; rtol=2e-4, atol=2e-3; nonuniform signed Q/K and per-group scales
Small, 2 GiB and 4 GiB boundary, original 4.608 GB offset PASS; sequentially bounded allocations, nonzero layer/Q/scale offsets
Multiple requests, shuffled pages, padded table rows PASS
Nonzero/empty windows, tail contexts, output guards PASS
Prefill and decode graph capture/replay PASS; logical outputs poisoned between replays
Deliberately wrong page map / scale permutation Detected; restored inputs pass
Total recorded layout/path/execution comparisons 40 PASS; includes original constant 8192.0 regression
Non-gfx950 default invocation SKIP; prints [skip] and exits 0
CTA tuning and paired final A/B 1,792 correctness-passing records; 5-process final medians

Negative control: running the new small strided case against upstream scorer files from 246efe9105ba produced a GPU memory-access fault (exit 134). The patched source passes the same test. Do not run this negative control in a shared serving process.

Appendix: stride/rebase performance non-regression

Following the same-change isolation used in #5285: only the three scorer files differ between baseline 246efe9105ba and patched source; same dependencies, GPU and test. Three alternating process-level repetitions, 20 warmups, 100 timed iterations each; table reports the median of the three event-time averages. Dense K/scales and tables are genuinely contiguous and valid on both versions.

Scope: public scorer API including internal schedule and output reset, CUDA graph replay; not kernel-only and not vLLM. Rows=32, requests=3, context upper bound=8192 with ragged/empty windows, H=64, D=128, page=64.

Path Before us After us Latency change
Prefill 36.564 36.442 -0.33%
Decode 37.704 37.779 +0.20%

Observed ranges overlap (prefill before 36.277–36.635, after 36.070–36.681 us; decode before 37.628–37.764, after 37.197–37.856 us). This auxiliary check only shows that shared-pool addressing did not materially regress the already-existing FP4 scorer; it is not the FP8 -> FP4 speedup table above.

# Run against each source revision, interleaving baseline/patched runs.
python op_tests/test_flydsl_pa_mqa_logits_fp4_strided_pages.py \
  --case small --layout dense --rows 32 --context 8192 \
  --graph --benchmark --iters 100 --warmup 20 --json-output dense-results.json

Artifacts retained locally under results/local-kernel-validation/: aiter-address-full.json, aiter-before-after.json, aiter-main-repeat*.json, aiter-patched-repeat*.json. Black and Ruff pass on changed files.

AI assistance

AI assistance was used at the author's request. Human review remains required before marking ready.

Signed-off-by: LiuYinfeng01 LiuYinfeng01@users.noreply.github.com

@github-actions

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 5518 --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.

@LiuYinfeng01
LiuYinfeng01 force-pushed the dsv4-fp4-strided-pages branch 4 times, most recently from 95fe436 to 563488a Compare September 15, 2026 11:45
@LiuYinfeng01
LiuYinfeng01 marked this pull request as ready for review September 15, 2026 11:53
@LiuYinfeng01
LiuYinfeng01 requested a review from a team September 15, 2026 11:53
…MQA logits

Teach pa_mqa_logits_fp4 decode/prefill to honor caller-provided K/scale/block-table
page strides and rebase each physical cache page in 64-bit address space before
building buffer descriptors. Fixes wrong logits when shared KV pools place pages
beyond the 4 GiB range of buffer-instruction byte offsets (reproduced at 4.6 GiB
in DeepSeek V4 BLHNC layouts).

Assisted-by: OpenAI API coding assistant

Signed-off-by: LiuYinfeng01 <LiuYinfeng01@users.noreply.github.com>
Signed-off-by: LiuYinfeng01 <LiuYinfeng01@users.noreply.github.com>
@github-actions github-actions Bot changed the title [FlyDSL][gfx950] Add explicit page strides and 64-bit rebase for FP4 MQA logits [FlyDSL] [gfx950] Add explicit page strides and 64-bit rebase for FP4 MQA logits Sep 15, 2026
@zufayu
zufayu requested a review from coderfeli September 16, 2026 01:52
@zufayu

zufayu commented Sep 17, 2026

Copy link
Copy Markdown
Contributor

Advisory review (static + hand-run; not a merge gate). Validation/Perf ran no GPU stage — reasons are on their lines below. Findings tagged [verified] are traced to code/repro.

Extends the gfx950 FlyDSL paged MXFP4 MQA scorers (decode + ragged prefill) to take the KV-cache, scale and block-table page strides from the real tensors and to rebase each physical page in 64-bit before the within-page buffer loads, so shared BLHNC pools whose pages sit beyond the 4 GiB range of a buffer instruction's byte offset address correctly.

Review (advisory): ⚠️ NEEDS WORK
Validation (deterministic): NOT RUN — REVIEW_AUTO_VALIDATE=0 and no head-matched report was supplied; the triaged target op_tests/test_flydsl_pa_mqa_logits_fp4_strided_pages.py is gfx950-only and this review box is gfx942. The PR's own CI stands in: on head commit 0030d05 the gfx950 leg of Aiter Test ran the new test green (all records "correctness": "passed", dense vs strided bit-identical, physical page offsets up to 4,613,299,200 bytes — past the 2^32 buffer-offset limit, plus 2 GiB / 4 GiB boundary cases), while the gfx942 leg fails on the new file's exit convention (finding 1).
Perf (advisory): NOT RUN — no gfx950 GPU on this review box (gfx942 only) and auto-validation disabled, so no validator base-vs-head timing exists; the PR supplies its own isolated before/after (only the three scorer files differ from baseline 246efe9, three alternating process repetitions, 20 warmup / 100 timed iterations): prefill 36.564 → 36.442 us (-0.33%), decode 37.704 → 37.779 us (+0.20%), overlapping observed ranges.

⚠️ [verified] op_tests/test_flydsl_pa_mqa_logits_fp4_strided_pages.py:493 ends main() with raise SystemExit(f"gfx950 required; ..."), which exits 1 — but aiter's CI executes every op_tests/test_*.py via plain python3 <file> on BOTH the MI35X (gfx950) and MI300X (gfx942) legs, and the sibling gfx950-only tests exit 0 with a printed "[skip]" (test_flydsl_pa_mqa_logits_fp4.py:699, test_flydsl_pa_mqa_logits_fp4_prefill.py:791). The PR's own CI run confirms the impact: the MI300X shard-5 job logs "gfx950 required; current architecture: gfx942" → "❌ Test failed: op_tests/test_flydsl_pa_mqa_logits_fp4_strided_pages.py", and the required Aiter Test Gate check stays red on every push until this is fixed (the other red MI300X shard is test_topk_select.py dying in hipErrorIllegalState — a file this PR does not touch, an unrelated flake). Author must replace the raise with the sibling skip convention (print "[skip]" and return, or exit 0) so the gfx942 leg goes green.

📝 [verified] CI invokes the new test with no flags, so the default run covers eager mode only — every record in the green gfx950 CI log reads "mode": "eager", "sensitivity_checked": false; the CUDAGraph capture/replay and the page-map / scale-permutation mutation checks (the parts guarding the production CUDA-graph decode path) ran only in the author's local --graph --sensitivity invocation. Reviewer should ask for those flags in the standard CI invocation, or for the local JSON logs to be attached to the PR.

Exit successfully on non-gfx950 runners and enable graph replay plus page-map/scale sensitivity checks in the default CI invocation.

Signed-off-by: LiuYinfeng01 <LiuYinfeng01@users.noreply.github.com>
Select a graph-stable CTA count from host-known batch and context buckets tuned on MI355X, while preserving explicit caller overrides. Cover the 8K/32K/64K/128K schedule table and MTP divisibility constraints.

Signed-off-by: LiuYinfeng01 <LiuYinfeng01@users.noreply.github.com>
Signed-off-by: LiuYinfeng01 <LiuYinfeng01@users.noreply.github.com>
@@ -0,0 +1,574 @@
# SPDX-License-Identifier: MIT

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

why another ut?

@github-actions

github-actions Bot commented Oct 7, 2026

Copy link
Copy Markdown
Contributor

No activity for 15 days, so this is now labelled stale. A push or a comment clears it; keep-open exempts it. Nothing is closed automatically.

@github-actions github-actions Bot added the stale The PR hasn't been updated in 2 weeks label Oct 7, 2026

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

FlyDSL stale The PR hasn't been updated in 2 weeks

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants