Repository navigation
[FlyDSL] [gfx950] Add explicit page strides and 64-bit rebase for FP4 MQA logits - #5518
LiuYinfeng01 wants to merge 5 commits into
Conversation
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
PR title tags & labels: |
95fe436 to
563488a
Compare
…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>
563488a to
0030d05
Compare
|
Advisory review (static + hand-run; not a merge gate). Validation/Perf ran no GPU stage — reasons are on their lines below. Findings tagged 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):
📝 [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>
25129cc to
7109ff1
Compare
Signed-off-by: LiuYinfeng01 <LiuYinfeng01@users.noreply.github.com>
7109ff1 to
7db41be
Compare
| @@ -0,0 +1,574 @@ | |||
| # SPDX-License-Identifier: MIT | |||
|
No activity for 15 days, so this is now labelled |
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), andblock_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 Torch2.12.0+rocm10.0.0and Triton3.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. Eachshape 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.
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.000000against the exactFP4-dequant reference, separating kernel correctness from the expected
quantization difference.
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.
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_numremains an override.Correctness / regression
Local MI355X, gfx950, GPU 7. Four-commit branch: functional kernel commit
d6807379, consolidated test/benchmark commit0030d058, CI coverage fix063f3eb5, and auto-CTA tuning commit6ab04df3; all are authored by LiuYinfeng01. Torch2.12.0+git6bbd260, HIP7.2.53211; original validation used image Triton3.7.1; the complete correctness suite and standard shape sweep were rerun with the upstream CI Triton3.8.0+amd.rocm7.2.0.git111ff227pin. 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.-infmasks retained[skip]and exits 0Negative control: running the new small strided case against upstream scorer files from
246efe9105baproduced 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
246efe9105baand 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.
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.jsonArtifacts 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