[MiniMax-M3] Route the paged decode to aiter's FlyDSL kernel with work planner - #2366
Conversation
`PA_DENSE_SPLIT_MAX = 32` is not a tuned value. It is the smaller of two
unrelated bounds that happen to agree near there: production aiter does not
build the C++ PS reduce past 64 and has no working fallback under it, and
`temporary_output` is bf16, so each extra split adds a round trip through the
PS combine.
aiter #4332 removes the first bound. Deciding whether the second one still
binds needs the cap measured against the shipping value on the same tree,
which a constant cannot express. `dense_decode_splits` computes
`min(PA_DENSE_SPLIT_MAX, cdiv(128, running_bs))`, so the cap only binds at
`running_bs <= 2` -- the difference is invisible without an A/B at that shape.
Default stays 32 and behaviour is unchanged. Raising it on production aiter
still fails at engine init rather than at runtime.
`int(os.getenv(...) or 32)` rather than a getenv default: an exported-but-empty
variable would otherwise raise `int('')` at import and take the engine down
before it starts.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
…n env
ATOM_PA_FLYDSL=1 sends the paged decode to aiter's FlyDSL implementation
(aiter PR #4332) where FlyDSL's domain covers the call, and to gluon
otherwise. A measurement switch, not a shipping default: the two take the
same arguments and compute the same thing, and the env exists to A/B them
without maintaining two ATOM trees. Default 0, so nothing changes until an
operator opts in on an aiter that carries #4332.
_flydsl_pa_decode_num_seqs mirrors the kernel's own validation -- alibi,
sinks, externally quantized FP8 queries, ps=False, a positive sliding window,
a non-256 tile, a partition count outside [1, 256], an unsupported head dim,
and any cache that is not [nb, Hkv, D//16, page, 16] at 1 byte per element are
all hard rejects in aiter, not preferences. Enabling the env therefore cannot
turn a working configuration into an exception.
The returned count is the substance. ATOM pads two axes independently:
context_lens to running_bs for CUDA-graph identity, with
[scheduled_bs:running_bs] zeroed, and q to running_tokens for the MoE
all_gather. forward_context.Context stores both "because the ratio is not
always max_seqlen_q" and asserts no rectangle between them. gluon absorbs the
mismatch by deriving its batch as q.shape[0] // query_length and letting the
surplus zero-length rows do no work; FlyDSL demands the rectangle and raises
ValueError: query.shape[0] (12) must equal
context_lengths.shape[0] * query_length (4 * 4)
on an ordinary step. Recovering the count the way gluon does and slicing the
five per-sequence arguments to it hands FlyDSL the same rectangle, dropping
exactly the zeroed tail. All slices are views on dim 0.
The gate is deliberately not restricted to max_seqlen_q == 1. The dense path
is where FlyDSL's headroom over gluon lives, it runs max_seqlen_q ==
num_spec + 1, and FlyDSL carries a query_length == 4 MTP4 grid split tuned for
exactly that shape.
ATOM_PA_FLYDSL_PLAN=1 additionally opts into aiter #5546's GPU work planner.
The plan and its scratch are allocated once and only refreshed: allocation is
illegal inside graph capture and a captured graph bakes in the pointers, and a
planned call does not use the caller's static buffers -- planned partials are
packed as [kv_heads, capacity, query_rows(, D)] against the static API's
[num_seqs, kv_heads, partitions, query_rows(, D)]. The planner refuses batches
past 4096 while M3's sparse call site folds query tokens into num_seqs and
reaches 32768 on a prefill-as-decode step, so those fall back to the static
path rather than raising.
Measured negative on this workload and defaulted off: against the same tree
with the planner off it costs 3.86% at conc 1, 11.50% at conc 10 and 18.28% at
conc 20. The cause is not batch uniformity -- 31% of decode steps have max/min
KV length >= 4 -- but the split count, which plan_pa_decode and pa_decode
share: dense_decode_splits gives 16 at running_bs 5 and 8 at running_bs 24,
and at those counts a batch skewed 2.5x and one skewed 120x both measure
1.00x. #5546's own conclusion is that the data "support using dynamic plans
for uneven KV work, not enabling them universally".
Both routes are logged once per shape signature. A run where the env is set
but every call still lands on gluon is otherwise indistinguishable from one
where FlyDSL simply did not help.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
…and keep it off the sparse call sites `_flydsl_work_plan` handed `plan_pa_decode` the static split count from `dense_decode_splits`, overwriting that argument's own default. The two are different quantities: `get_recommended_splits` gives every request the same count and documents itself as "not a variable-work scheduler", while the plan's `max_partitions` is a per-request ceiling it then divides under a workgroup budget. Clamping the ceiling to the static count leaves a long request with the same share as a short one, which is the planner's entire mechanism. Read back from `plan.reduce_info`: a 5-request batch at ceiling 16 splits [16,4,4,4,4] with the long request pinned at the ceiling, and at the default 256 it splits [256,20,21,20,21]. `ATOM_PA_FLYDSL_PLAN_MAX` now carries that ceiling, 0 meaning "omit the argument and take aiter's default". Do not read the ceiling off aiter's unit test: it is the only caller of `plan_pa_decode` in that tree and it binds the two deliberately, building the plan from the same count it hands `pa_decode` so the paths must agree numerically. That makes it a correctness test which cannot see the conflation. The planner was also a global switch inside `run_pa_decode_gluon`, so enabling it reached the two sparse call sites as well -- 57 of the 63 `pa_decode` calls per step. Their context is a fixed topk window, so there is no unevenness to rebalance and only the refresh and the task packing remain. On uniform batches the planner is a loss that the ceiling cannot rescue: 0.56x at B8/257 and 0.84x at B200/200000, flat from ceiling 4 through 256. `allow_work_plan` moves that choice to the caller, where the split policy already lives, and only the dense site opts in -- the same boundary `dense_decode_splits` already draws. A plan build now logs its shape once per cached key, which is the only evidence that the planner went where it was meant to rather than being assumed to. Measured at conc 20 on the agentic trace (one operating point), single-user output throughput 157.19 -> 178.49 tok/s and ITL 7.97 -> 6.96 ms, against 139.20 for the same arm before these changes. Prefix hit 96.56 vs 96.63, input length +1.6%, request count +1.0%. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
…ur env switches into one
The plan is a function of `context_lens` alone, so every layer of one forward
wants the same one. Building it inside the attention op re-ran the planner
kernel for each of M3's three dense layers, and the module-level cache that
existed to paper over that had to key itself on shape to stay correct. It now
lives where the rest of the per-forward metadata is built: `prepare_decode`
attaches it, `prepare_mtp_decode` refreshes it for each draft pass (their
context lengths differ by a token, so reusing the target's plan would point the
kernel at the wrong KV ranges -- wrong output, not merely slower), and the op
reads it off `get_forward_context().attn_metadata`.
Measured before the move: 6 planner kernels per step, and `envs.__getattr__` on
all 63 `pa_decode` calls rather than the 6 that could use a plan. After: 4 and
1. The env read is the larger of the two -- one `__getattr__` costs ~0.97 us and
63 of them are ~0.5% of a 12.8 ms step.
Scratch allocation stays in the op. It is sized from the query tensor the op is
holding and happens once per shape, so only the per-step work moved out.
The env surface collapses to `ATOM_PA_FLYDSL_PLAN`, now on by default:
* ATOM_PA_FLYDSL is gone. It was a measurement switch for A/B-ing FlyDSL
against gluon without two trees; FlyDSL is the decode path. The capability
check that falls back to gluon for shapes outside FlyDSL's domain stays --
it mirrors the kernel's own validation and is not a switch.
* ATOM_PA_FLYDSL_PLAN_MAX is gone, and the call now omits `max_partitions`
entirely so `plan_pa_decode` takes its own default. The one time this tree
chose the ceiling it chose the static split count, which clamps every
request to the same share and removes the planner's mechanism.
* ATOM_PA_DENSE_SPLIT_MAX is gone and PA_DENSE_SPLIT_MAX is a constant again.
With the planner on, a planned call is told `plan.max_partitions`, so the
cap no longer sets the production partition count; it survives for the
gluon fallback, where both halves of its bound still hold. Raising it would
only enlarge the static scratch a planned call never reads.
* ATOM_PA_FLYDSL_PLAN defaults to 1. On the agentic trace, SA-convention
interactivity (1/itl_p90) is +24.6% at conc 20 and +8.4% at conc 10, and
the gain runs monotonically from the slow tail to the fast one (p10 +24.6%,
p50 +19.3%, p90 +5.0%) -- it lifts the floor, which is the end the metric
watches.
The routing log no longer hangs off the deleted env, and is unconditional: with
FlyDSL as the default path, which route a call took is worth always recording.
tests/test_paged_attention_dispatch.py: 74 passed.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Prose cut, facts kept. What stays is the four things a reader would otherwise re-derive the hard way: the ceiling must not come from the static split count, each draft pass needs its own refresh (reusing the target's plan points the kernel at the wrong KV ranges -- wrong output, not merely slower), scratch is allocated in the op because its shapes come from the query tensor, and the FlyDSL/gluon split is a capability check rather than a switch. The measured numbers move to the one place they inform a decision: why the planner defaults on. Also fixes a contradiction left by the previous commit. `run_pa_decode_gluon` still described itself as a "MEASUREMENT SWITCH, not a shipping default" after FlyDSL had become the path. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Review found the previous commit's premise broken: decode runs from captured
graphs (cudagraph_mode=FULL), and `build_for_cudagraph_capture` never attached
a plan, so the op saw None at capture time and the STATIC path is what got
recorded. Every later refresh then worked on a graph that does not read it --
no error, no slowdown, just a planner that quietly does nothing and an
end-to-end number that looks fine.
The capture flow (model_runner.py) is what makes the fix possible:
3857 build_capture(bs) outside the graph -- first plan allocated
3913 self.model(...) warmup, eager; the op sees the plan
3969 with torch.cuda.graph(...) capture; the plan is already there
So the plan is attached in `build_for_cudagraph_capture` too, and the builder
keeps one per (batch, kv_heads, device) instead of a single slot. The single
slot was the more dangerous half: `build_capture` runs once per capture-ladder
size, so the second rung replaced the first rung's plan while the first rung's
graph had already baked in its pointers.
Two more from the same review:
* The plan was built from `context_lens[:scheduled_bs]` while the op derives
its own batch as `q.shape[0] // max_seqlen_q`, which is `running_bs`.
aiter validates `reduce_info.shape == (num_seqs, 2)` and raises, so any
step whose batch was not exactly on the ladder would have killed the
worker. `context_lens` is already `running_bs` long with the tail zeroed.
* `prepare_mtp_decode` dropped the refresh result and relied on the returned
object being the same one. It now publishes through `workinfos`, which the
caller splats into the metadata, so a None result clears a stale plan
rather than leaving it in place.
`flydsl_plan_matches` is extracted rather than inlined so the guard is
testable: a plan built for another batch now falls back to the static path
instead of reaching aiter's validate.
tests/test_paged_attention_dispatch.py: 74 -> 83. Eight of the nine are
falsifiable by reverting the behaviour they describe, each reddening exactly
one test; the ninth is a source check on the capture attach, labelled as such,
because driving a real capture needs a model runner and what it guards fails
silently.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
…y default aiter #4332 is unmerged and the kernel is not fully tested, so gluon goes back to being the shipping decode path and FlyDSL becomes the opt-in: ATOM_PA_FLYDSL=1, off by default. The two gates run in series -- the env, then the capability check that mirrors the kernel's own validation -- and the `and` short-circuits, so with the env off the check does not run. Shapes outside FlyDSL's domain still fall back to gluon either way. The work planner (aiter #5546) only feeds the FlyDSL kernel, so it now requires ATOM_PA_FLYDSL=1 as well. Without that, a plan would be built and refreshed every step for a kernel that never reads it, and the "flydsl work plan" log line -- what an A/B arm is checked with -- would show up on a run that is entirely gluon. Tests 84 -> 86: the default must be off, the dispatch must consult the env before the capability check, and the builder must not plan with FlyDSL off. Each of the three goes red on its own reversion. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
The parameter now matches the env that controls it, ATOM_PA_FLYDSL_PLAN. "work plan" is aiter's own term for the object, but on its own it does not say the plan belongs to the FlyDSL path -- which matters now that gluon is the default and FlyDSL is the opt-in. Comments trimmed throughout: the _flydsl_pa_decode_num_seqs docstring keeps why it returns a count rather than a bool (ATOM pads its sequence and row axes independently, so the two disagree by a padding slot on an ordinary step; gluon absorbs that, FlyDSL raises) and drops the shape algebra and the verbatim traceback. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
🏷️ CI GuideRuns automatically on every eligible PR before approval:
Heavy model tests:
|
Review — FlyDSL pa_decode routingReviewed The routing design is the right shape — a narrow gate in front of an opt-in kernel, with a static fallback behind it — but two things need to be resolved before this can land, and they are of different kinds. The first is blunt:
The second is structural:
Provenance: [verified] means the deciding code was read in the PR-head worktree or the command was run there. [reported] means the shape matches the code but was not traced end to end. 1. The routed path cannot run [verified]
With The argument is also redundant — the gate already rejects The reason this shipped green is worth fixing alongside it: every wiring test monkeypatches 2. MTP draft attends over a stale plan, and
|
…binding Two gates in series guard the FlyDSL paged decode: ATOM_PA_FLYDSL is the opt-in, and _flydsl_pa_decode_num_seqs mirrors aiter's own rejects so an unsupported call falls back to gluon instead of raising inside the kernel and killing the worker. The gate was missing several of those rejects -- q/output dtype and shape agreement, both head_dim stride checks, the q-heads-over-kv-heads divisor, the v_cache dtype, and contiguity of the two caches and the two per-sequence arrays. Plans are now only ever minted during cudagraph capture. aiter's planner takes batch as a tl.constexpr, so a runtime batch that no graph replays would cost a cold kernel specialization and a plan nothing can free; the static path is the right answer for it. Runtime only refreshes in place. The one-shot "no plan for this batch" notice waits until a capture has happened, since before that every call legitimately lands there and the notice would be spent on a step that says nothing. max_partitions stays unset, at plan_pa_decode's own default: every number measured on this branch was measured there, and the review's suggestion to clamp it to 64 is a perf claim nothing here backs. Measuring instead, on the planned path, across nine cells spanning the trace's p50/p90/p99 contexts at batch 1/8/24: drift from an fp32 reference does not grow with the split count anywhere, and the ceiling only ever binds at batch 1. The plan's capacity is min(batch * max_partitions, max(batch, 2 * CU / kv_heads)) = 512 here, so each sequence's share is about 512/batch and a 64 ceiling would have changed nothing above batch 2 while capping the one place the planner has the most to rebalance. numa_utils and the mooncake connector now read ROCR_VISIBLE_DEVICES first. ROCR filters and renumbers before HIP indexes the result, so setting both cut the visible device count twice; dropping HIP from the recipe left the NUMA binding reading nothing and put all four ranks on one node. Tests go from 86 to 127. The call-site contract is now checked by walking the AST of all three builders and comparing the value of `create`, not its presence in one file -- a sed that moved `create=True` onto a runtime path had left every test green. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Main migrated single-node AgentX to native srt-slurm (#3428) and deleted benchmarks/single_node/agentic (#3461). Accept the script deletion and move this PR's change into the declarative recipe: bump the engine image to rocm/atom-dev:nightly_202609231248 (ROCm/ATOM#2366) and set ATOM_PA_FLYDSL=1 and ATOM_PA_FLYDSL_PLAN=1. The bash-only GPU-mask fix no longer applies. The changelog entry follows the default eval policy. 中文:main 已将单节点 AgentX 迁移到 srt-slurm 并删除旧 bash 脚本。本 PR 改为在 YAML 配方中更新镜像至 nightly_202609231248 并启用 ATOM_PA_FLYDSL / ATOM_PA_FLYDSL_PLAN;GPU mask 修复随旧脚本一起移除,changelog 采用默认 eval 策略。 Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Move the MiniMax-M3 MI355X ATOM AgentX srt-slurm recipe to rocm/atom-dev:nightly_202609231248, which contains ROCm/ATOM#2366, and set ATOM_PA_FLYDSL=1 and ATOM_PA_FLYDSL_PLAN=1. Drop TP4 C32 to bound the sweep. 中文:将 MiniMax-M3 MI355X ATOM AgentX srt-slurm 配方镜像更新为包含 ROCm/ATOM#2366 的 rocm/atom-dev:nightly_202609231248,并设置 ATOM_PA_FLYDSL=1 与 ATOM_PA_FLYDSL_PLAN=1;移除 TP4 C32 以控制 sweep 规模。 Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
1. Motivation
gluon's paged decode gives every request in a batch the same number of KV splits. On an agentic trace a decode batch is wildly uneven — measured over 30,400 steps at conc 32, the step-internal
max/mincontext ratio is p50 2.49, p90 22.2, max 1610, and 31.2% of steps sit above 4× — so one long request owns the critical path while the short ones idle.aiter #4332 (FlyDSL paged attention) plus #5546 (a Triton GPU planner that hands each request a partition count proportional to its real context length) address exactly that shape. Measured on MiniMax-M3 MXFP4 + EAGLE3, TP4, same machine as the current published curve:
The gain grows with concurrency because that is when the batch has enough unevenness to rebalance; at conc 1 there is nothing to rebalance, and the curve confirms it. Past conc 20 it narrows again — the planner's task budget is
2 × CU, so each request's share falls as512/batchand there is less left to redistribute.Both paths are opt-in and gluon stays the default: the FlyDSL kernel is not fully tested on our shapes yet, so it is something you turn on, not something you inherit.
2. Usage
ATOM_PA_FLYDSL=10)ATOM_PA_FLYDSL_PLAN=11)ATOM_PA_FLYDSL=1; with FlyDSL off no plan is built at all.ATOM_PA_FLYDSL=1 python3 -m atom.entrypoints.openai_server \ --model amd/MiniMax-M3-MXFP4 --tensor-parallel-size 4 ...Verify from the server log:
pa_decode -> flydsl[N seqs]with no-> gluonlines, andflydsl work plan: num_seqs=... max_partitions=256appearing once per cudagraph capture size per rank and never during the run.3. Design
ATOM_PA_FLYDSL=1+ ATOM_PA_FLYDSL_PLAN=1dense_decode_splits(), capped at 32plan.max_partitionsinsteadpa_decodecallwork_plan— dense onlyTwo gates in series. The env is the opt-in;
_flydsl_pa_decode_num_seqs()then mirrors the kernel's own validation so an unsupported shape falls back to gluon instead of raising from inside aiter. Every clause is a hard reject on aiter's side, not a preference: fp8 compute and fp8 page-16 cache,block_size ∈ {16, 64, 128}taken from the cache's dim −2, gfx942/gfx950, partition size and count bounds,head_dim ∈ {64} ∪ {128…1024 step 128}and agreeing with the cache's ownnum_hgroups × 16, q dtype, output dtype and shape matching q, a contiguous head_dim axis on both q and output, q heads divisible by kv heads, matching k/v cache dtypes, int32 and contiguousblock_tables/context_lens, and a row count that divides evenly and stays inside the batch. Theandshort-circuits, so with the env off the capability check never runs.The planner is kept off the sparse call sites. M3 issues 63
pa_decodecalls per step, 57 of them sparse. Sparse reads a fixed top-k window, so its contexts are near-uniform — the one shape where the planner is a measured net loss (0.56× on an all-equal batch, and a flat curve across every ceiling, so it is not a tuning problem). There is no flag:work_plandefaults toNone, and only the dense site inattention_mha.pyhands one in, matching the boundarydense_decode_splitsalready draws. The vLLM and SGLang bridges run under someone else's forward context and pass none.Plan lifetime vs cudagraph. The plan depends on
context_lensalone, which every layer of one forward shares, so it is built once per forward rather than per call. Decode replays captured graphs, so it must exist at capture time or the static path is what gets recorded — silently, with a benchmark that still looks reasonable.Minting happens only on that first line. aiter's planner takes batch as a
tl.constexpr, so an unseen value costs a kernel specialization — 65–72 ms cold — and a plan that can never be freed, because some captured graph may have baked its pointers in. A runtime batch with no plan is one no graph will replay, so the static path is the right answer for it rather than a stall. The cost of that rule is that the planner is inert under--enforce-eager.Rows are
running_bslong, notscheduled_bs: the op derives its own batch asq.shape[0] // max_seqlen_q, and aiter validatesreduce_info.shape == (num_seqs, 2)against exactly that. Every draft pass republishes its plan throughworkinfos— each pass advancescontext_lens, so reusing the target's plan would point the kernel at the wrong KV ranges, which is a wrong answer rather than a slow one.max_partitionsis left unset, at aiter's own default. Setting it from the static split count clamps every request alike and removes the planner's whole mechanism; this tree did that once. Clamping it to 64 was considered and rejected — see §5.Device visibility.
numa_utilsand the mooncake connector now readROCR_VISIBLE_DEVICESbeforeHIP_VISIBLE_DEVICES. ROCR filters and renumbers the device table before HIP indexes the result, so a recipe setting both cuts the visible count twice; droppingHIP_VISIBLE_DEVICESthen left the NUMA binding reading nothing and put all four ranks on one node, which deadlocks atallocate_kv_cache.Integration: needs aiter
94dca7bc6or later — that commit (#4332, merged 2026-09-22) brings both the FlyDSL paged decode and #5546's planner intomain. On an older aiter the import guard keeps the default path working; withATOM_PA_FLYDSL=1it would fail at startup.4. Test Plan
tests/test_paged_attention_dispatch.py— dispatch, capability gate and plan wiring, CPU only.5. Test Result
Unit: 127 passed (74 → 127). Each new guard was verified to go red on its own reversion, by injecting the defect and re-running — including the two the previous round's tests did not catch: a
create=Truemoved onto a runtime path, and the head_dim whitelist deleted outright.Accuracy (gsm8k, 5-shot / 20-shot): 32 → 0.9674 / 0.9659, 128 → 0.9636 / 0.9644, all far above the 0.93 gate. Paired McNemar p=0.441 / 0.831; 98% of questions agree question-by-question and the 2% that disagree split evenly both ways.
Planner drift at the real ceiling. The plan's
capacityismin(batch × max_partitions, max(batch, 2 × CU / kv_heads))= 512 here, so each request's share is about512/batchand at batch 1 the realized split count equalsmax_partitionsexactly. The trace's single-request contexts are p50 120k / p90 334k / p99 764k tokens — 470 to 2986 tiles — so low-concurrency steps really do run at the 256 ceiling, four times past anything the old bf16 round-trip argument measured. Swept over nine cells (those three contexts × batch 1/8/24, one shared fp32 reference per cell):No growth with the split count anywhere, so the ceiling stays at aiter's default. The sweep also shows the ceiling only ever binds at batch 1: at batch 8 and 24 the 64 and 256 arms receive identical per-request counts, so clamping to 64 would change nothing above batch 2 while capping the one case the planner has the most to rebalance.
End-to-end, same machine, paired against the pre-FlyDSL gluon curve (p90 interactivity =
1/itl_p90, SemiAnalysis convention):Completed-request counts match within 1.3% at every point, so the deltas are rates rather than a different request mix. Routing was asserted per point:
gluon=0, and the plan built only during capture.High concurrency, measured on a second node, so stated against the published curve rather than added to the paired table above:
conc 40 runs with LMCache CPU offload on, as does the published curve's own conc-40 point; the other three are GPU-resident on both sides.
6. Known Limitations
aiter.ops.flydsl.pa_decodeexists wheneverATOM_PA_FLYDSL=1. It landed in94dca7bc6(2026-09-22); nothing in this repo enforces that floor.--enforce-eager. Plans are minted only during cudagraph capture, which never happens there. This is the deliberate cost of not specializing a planner kernel mid-serving.ATOM_PA_FLYDSL=1, any fp8 page-16 MHA decode takes FlyDSL, including the vLLM bridge. The SGLang bridge passescompute_type=torch.bfloat16, which the gate rejects, so it stays on gluon.Submission Checklist