Skip to content

[FlyDSL] perf(pa_decode): tune MTP work budgets and planned reduction - #5628

Closed
zhiding512 wants to merge 2 commits into
flydsl-pa-decode-tilefrom
zhiding512/flydsl-pa-mtp-plan-tuning
Closed

zhiding512 wants to merge 2 commits into
flydsl-pa-decode-tilefrom
zhiding512/flydsl-pa-mtp-plan-tuning

Conversation

@zhiding512

Copy link
Copy Markdown
Contributor

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

  • Let planned calls infer max_context_partition_num from work_plan, expose actual GPU counts through plan.num_partitions, and inherit the partition limit when refreshing a plan. Existing explicit matching arguments remain supported.
  • Add an optional host total_context_length hint. Together with query_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.
  • Use a 128-thread register reducer for sufficiently large D128 planned grids and iterate only each request's actual partitions, while retaining a maximum of 256. Preserve the target branch's sink normalization, masked fallback, and window consistency checks.
  • Add static/planned benchmark entry points with time, effective bandwidth, partition metadata, and configurable timer settings. Keep new correctness coverage in the existing PA CI test module.

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.

Contexts Page V layout Base (µs) PR (µs) PR effective TB/s Time reduction
B200, all 200000 16 plain 3000.13 2295.66 4.607 23.5%
B200, all 200000 16 transposed 2999.26 2332.80 4.534 22.2%
B200, all 200000 128 plain 2999.39 2325.20 4.545 22.5%
B200, all 200000 128 transposed 3227.13 2459.28 4.297 23.8%
B16, one 100000 + fifteen 1024 128 transposed 45.11 48.93 0.633 -8.5%
B16, one 200000 + fifteen 1024 128 transposed 55.60 57.50 0.998 -3.4%

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.
  • Added coverage for work-hint sizing with windows, compact reduction with sinks across partition-chunk boundaries, and compact graph replay after changing lengths and sink logits.
  • All 36 case/layout executions in the alternating performance comparison pass the causal FP32 reference and static/planned checks at atol=rtol=0.005.
  • The MTP/TP benchmark smoke sweep passes all four static/planned and uniform/one-long results.
  • Black, Ruff, and git diff --check pass.
HIP_VISIBLE_DEVICES=7 AITER_LOG_MORE=1 python3 -m op_tests.benchmark_flydsl_pa_decode_plan \
  --cases b200_uniform --block-size 16 128 --trans-v 0 1 \
  --num-iters 101 --num-warmup 5 --num-rotate-args 3 \
  --output b200_results.json

API usage and scratch layouts: docs/flydsl_pa_decode_plan.md.

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.
@zufayu
zufayu requested a review from coderfeli September 18, 2026 01:16
@zhiding512 zhiding512 closed this Sep 22, 2026
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