[FlyDSL] perf(pa_decode): tune MTP work budgets and planned reduction - #5628
Closed
zhiding512 wants to merge 2 commits into
Closed
zhiding512 wants to merge 2 commits into
zhiding512 wants to merge 2 commits into
Conversation
Infer partition bounds from reusable plans, size long-MTP plans from host-known work, and reduce only actual partitions on large D128 grids. Preserve sliding-window and attention-sink support from the current tile branch, cap windowed work hints, and keep regression coverage in the existing PA test module. Validation: 489 PA tests; MTP benchmark smoke checks; alternating base/PR performance runs on gfx950; Black and Ruff.
Retain the single parametrized PA test and shared input/reference path. Port plan inference, work-hint and compact-reduction checks to DecodeCase and adapt benchmark imports. Validation: 88 parametrized PA cases passed with eager execution, graph replay and contract checks.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Large MTP batches can receive only two or three partitions per request under the default two-CTA-per-CU plan budget. Size long-MTP plans from host-known KV work and use a compact reducer for large output grids. For B200/C200000/MTP4, measured plan + attention + reduction latency drops from 3.00–3.23 ms to 2.30–2.46 ms on gfx950.
Targets
flydsl-pa-decode-tile(#4332), building on #5546 and the tile branch's sliding-window and attention-sink support.Changes
max_context_partition_numfromwork_plan, expose actual GPU counts throughplan.num_partitions, and inherit the partition limit when refreshing a plan. Existing explicit matching arguments remain supported.total_context_lengthhint. Together withquery_length, it supplies a work-based partition floor for long MTP workloads. B200/C200k/MTP4 selects 16 partitions per request and capacity 3200. Explicit budgets and partition caps retain precedence; windowed hints are capped by the causal-window union and tile alignment.The B200 planned scratch allocation grows from 8.25 MiB to 51.56 MiB per rank. Work hints size buffers at creation; length changes still refresh the existing GPU metadata, with fixed capacity for graph replay.
Performance
Compared against target commit
fd825332f, alternating base and PR runs on GPU 7 (gfx950, 256 CUs). BF16 Q/O, FP8 KV with per-token FP32 scales, MTP4, Hq16/Hkv1/D128 (one TP=4 rank). Each value is the median of three runs with 101 iterations, five warmups, and three allocation sets, using the same sparse-page helper. Time includes a plan refresh on every call, attention, and reduction.The B16 mixed-length regression remains a tradeoff of the compact-reduction dispatch. The work-budget increase primarily targets large, long-context MTP batches.
Bandwidth counts logical Q/O, valid KV/scales, referenced page-table entries, and context lengths once per rank. It excludes padding, sparse holes, repeated loads, scratch, and plan metadata; it is not hardware-counter HBM traffic and is not multiplied by MTP or TP. Timings exclude CPU dispatch, allocation, reference calculation, and JIT compilation.
Validation
python3 -m pytest -q op_tests/test_flydsl_pa_decode.py: 88 passed on gfx950, using the target branch’s consolidated parameterized test, including sliding windows, sinks, inferred partition arguments, work hints, non-power-of-two counts, FP16, poisoned scratch, and graph replay.atol=rtol=0.005.git diff --checkpass.API usage and scratch layouts:
docs/flydsl_pa_decode_plan.md.