Skip to content

[MiniMax-M3] Route the paged decode to aiter's FlyDSL kernel with work planner - #2366

Merged
valarLip merged 13 commits into
mainfrom
flydsl-pa-decode-routing
Sep 23, 2026
Merged

valarLip merged 13 commits into
mainfrom
flydsl-pa-decode-routing

Conversation

@yitingw1

@yitingw1 yitingw1 commented Sep 23, 2026 •

Copy link
Copy Markdown
Contributor

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/min context 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:

conc p90 interactivity throughput/chip
1 −0.0% −0.7%
10 +4.0% +3.6%
15 +11.8% +1.0%
20 +20.5% +2.7%

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 as 512/batch and 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

Argument Required Description
ATOM_PA_FLYDSL=1 No (default 0) Route the paged decode to FlyDSL where its domain covers the call. Off by default.
ATOM_PA_FLYDSL_PLAN=1 No (default 1) Use #5546's work planner on the dense decode. Requires 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 -> gluon lines, and flydsl work plan: num_seqs=... max_partitions=256 appearing once per cudagraph capture size per rank and never during the run.

3. Design

Aspect gluon (default) ATOM_PA_FLYDSL=1 + ATOM_PA_FLYDSL_PLAN=1
Split count dense_decode_splits(), capped at 32 same not used — plan.max_partitions instead
Partition allocation uniform across the batch uniform per request, by real context length, under a workgroup budget
Ceiling 32 (C++ PS reduce limit) 32 256, aiter's own default, never overridden
Where decided per pa_decode call per call once per forward, in the metadata builder
When allocated — — cudagraph capture only; runtime refreshes in place
Call sites all all in FlyDSL's domain whoever passes a work_plan — dense only

Two 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 own num_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 contiguous block_tables/context_lens, and a row count that divides evenly and stays inside the batch. The and short-circuits, so with the env off the capability check never runs.

The planner is kept off the sparse call sites. M3 issues 63 pa_decode calls 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_plan defaults to None, and only the dense site in attention_mha.py hands one in, matching the boundary dense_decode_splits already 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_lens alone, 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.

 build_for_cudagraph_capture(bs)        ← plan MINTED here, outside the graph
        │  one entry per (num_seqs, kv_heads, device), never replaced:
        │  a single slot would free tensors an earlier graph already baked in
        ▼
 warmup forward (eager)                 op reads plan → planned path recorded
        ▼
 torch.cuda.graph(...)  capture         only the refresh kernel is captured
        │                               (no allocation, no D2H → capturable)
        ▼
 replay, every step                     prepare_decode / prepare_mtp_decode
                                        refresh in place from live context_lens

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_bs long, not scheduled_bs: the op derives its own batch as q.shape[0] // max_seqlen_q, and aiter validates reduce_info.shape == (num_seqs, 2) against exactly that. Every draft pass republishes its plan through workinfos — each pass advances context_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_partitions is 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_utils and the mooncake connector now read ROCR_VISIBLE_DEVICES before HIP_VISIBLE_DEVICES. ROCR filters and renumbers the device table before HIP indexes the result, so a recipe setting both cuts the visible count twice; dropping HIP_VISIBLE_DEVICES then left the NUMA binding reading nothing and put all four ranks on one node, which deadlocks at allocate_kv_cache.

Integration: needs aiter 94dca7bc6 or later — that commit (#4332, merged 2026-09-22) brings both the FlyDSL paged decode and #5546's planner into main. On an older aiter the import guard keeps the default path working; with ATOM_PA_FLYDSL=1 it would fail at startup.

4. Test Plan

  • tests/test_paged_attention_dispatch.py — dispatch, capability gate and plan wiring, CPU only.
  • gsm8k, full 1319 questions, splits 32 vs 128, order 32/128/128/32 — the accuracy precondition for splitting past the C++ reduce cap.
  • Planner drift against an fp32 reference, on the planned path, swept over the ceiling the planner actually reaches.
  • End-to-end agentic trace, TP4 + EAGLE3, MXFP4, 3600 s per point, against the same machine's gluon curve.

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=True moved 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 capacity is min(batch × max_partitions, max(batch, 2 × CU / kv_heads)) = 512 here, so each request's share is about 512/batch and at batch 1 the realized split count equals max_partitions exactly. 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):

max relative deviation ‖Δ‖/‖ref‖
across all 36 arms 2.84 – 4.76% 3.62 – 3.80%
paired 8 → 256, per cell ±0.5pp, both signs flat

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):

conc gluon this PR Δ tok/s/chip Δ ITL p90
1 297.5 297.4 −0.0% −0.7% 3.36 → 3.36
2 274.9 291.6 +6.1% −2.6% 3.64 → 3.43
10 206.6 214.9 +4.0% +3.6% 4.84 → 4.65
15 172.8 193.3 +11.8% +1.0% 5.79 → 5.17
20 133.0 160.3 +20.5% +2.7% 7.52 → 6.24

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 tok/s/chip vs published p90 interactivity vs published
24 42,661 +3.7% 138.9 +17.9%
28 46,608 +3.4% 118.2 +16.1%
32 47,782 +3.4% 85.6 +9.7%
40 58,624 +12.2% 87.7 +30.2%

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

  1. No aiter version pin. The code assumes aiter.ops.flydsl.pa_decode exists whenever ATOM_PA_FLYDSL=1. It landed in 94dca7bc6 (2026-09-22); nothing in this repo enforces that floor.
  2. The planner is inert under --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.
  3. The conc 2 point is noisy. +6.1% interactivity with −2.6% throughput does not match the shape of the other points; two independent runs at that concurrency agree with each other and disagree with this one. Treat it as run-to-run variance.
  4. The drift sweep is nine cells, not a tail. The bound it replaces was set by the worst of 370 shapes. Nine cells covering the production context and batch range show no growth, which is not the same as showing the tail is gone.
  5. Not attributed at kernel level. Per-kernel traces at bs=1 and bs=4 show the planner as a wash: the attention tile gets 11–38% faster, but PS reduce gets ~14% dearer because planned partials are packed per task. The end-to-end gain must come from larger, more uneven batches, which those two traces do not cover.
  6. Widens beyond MiniMax-M3. With ATOM_PA_FLYDSL=1, any fp8 page-16 MHA decode takes FlyDSL, including the vLLM bridge. The SGLang bridge passes compute_type=torch.bfloat16, which the gate rejects, so it stays on gluon.

Submission Checklist

yitingw1 and others added 8 commits September 21, 2026 10:32
`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>
@github-actions

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every eligible PR before approval:

  • ✅ Pre Checkin: Black, Ruff, catalog schema validation, non-GPU unit tests

Heavy model tests:

  • ✅ Run after the PR is approved and Pre Checkin passes
  • ✅ Run immediately when an approval review is submitted
  • ✅ Can be requested before approval with labels
Label Tests
ci:full Run all heavy PR model tests: native ATOM, vLLM, and SGLang
ci:atom Run native ATOM model accuracy tests
ci:vllm Run ATOM vLLM OOT model accuracy tests
ci:sglang Run ATOM SGLang model accuracy tests

Heavy jobs are skipped when the PR is not approved and no matching ci:* label is present.
Add labels via the sidebar or gh pr edit 2366 --add-label <label>

@yitingw1 yitingw1 changed the title # [MiniMax-M3] Route the paged decode to aiter's FlyDSL kernel with work planner [MiniMax-M3] Route the paged decode to aiter's FlyDSL kernel with work planner Sep 23, 2026
@valarLip

Copy link
Copy Markdown
Collaborator

Review — FlyDSL pa_decode routing

Reviewed 0e1374302 in a detached worktree. The new tests pass here (113 passed, 0 skipped).

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:

pa_decode() is called with a ps= keyword that aiter's FlyDSL kernel does not accept and does not absorb via **kwargs. Every routed call raises TypeError. No test exercises the real call, which is why the suite is green.

The second is structural:

The gate's docstring promises that everything aiter hard-rejects is rejected here first, but about eight of aiter's rejects are missing; and per-forward state that every sibling artefact keeps as an AttentionMetaData field now lives in a process-global dict and a thread-local context lookup instead. Together these turn several would-be fallbacks into crashes, and several would-be crashes into silence.

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]

base_attention.py:433 passes ps=ps. inspect.signature(aiter.ops.flydsl.pa_decode.pa_decode) ends ... alibi_slopes, sinks, sliding_window=0, work_plan=None — no ps, no **kwargs — and the same holds at 94dca7bc67, the aiter floor this PR names.

With ATOM_PA_FLYDSL=1 (which the recipe now sets), the first dense decode layer that passes the gate dies during cudagraph warmup with TypeError: pa_decode() got an unexpected keyword argument 'ps': worker dead, server still answering /metrics. That is precisely the failure mode the gate exists to prevent.

The argument is also redundant — the gate already rejects not ps before reaching this call.

The reason this shipped green is worth fixing alongside it: every wiring test monkeypatches plan_pa_decode with raising=False, which creates the attribute rather than asserting it exists, and nothing pins aiter's pa_decode signature. An inspect.signature assertion against both symbols would have caught this and would catch the next aiter rename.

2. MTP draft attends over a stale plan, and flydsl_plan_matches accepts it [verified]

GDNAttentionMetadataBuilder.prepare_mtp_decode (gdn_attn.py:1518) and Qwen4ExpMetadataBuilder.prepare_mtp_decode (qwen4_exp_attn.py:627) override without calling super() and return dicts with no flydsl_work_plan key. Both classes do inherit prepare_decode, which sets it — so the draft pass runs on the target's plan.

eagle_proposer.py:746 bumps attn_metadata.context_lens[:running_bs] += 1 in place before each draft step, and aiter's planner bakes per-task begin/end tiles plus seq_ctx into the plan (pa_decode_plan.py, tl.store(work + slot*4 + 3, seq_ctx, active)). flydsl_plan_matches compares only reduce_info.shape[0] and num_kv_heads — neither changes — so the stale plan is accepted, not refused.

On a GDN/Qwen4-Exp hybrid with MTP, fp8 page-16 KV, ATOM_PA_FLYDSL=1 and an eager or piecewise-eager draft step, the draft attends over the pre-bump context and never sees the KV of the token it just drafted. Silent acceptance-rate and accuracy loss, no exception.

The AST guard at tests/test_paged_attention_dispatch.py:629 parses aiter_attention only, so these two overrides are invisible to it.

3. Per-forward state left the metadata dataclass

flydsl_work_plan is the first per-forward attention artefact to skip AttentionMetaData entirely — work_meta_data, work_indptr and reduce_indptr are all declared fields, and gdn_attn.py:106 declares flydsl_prefill_metadata the same way. Two consequences, both invisible today and both load-bearing under TBO:

The plan is fetched from the thread-local forward context rather than from the fwd_ctx the caller already holds (base_attention.py:364) [verified]. attention_mha.paged_attention_triton takes fwd_ctx: ForwardContext, reads attn_metadata = fwd_ctx.attn_metadata at line 560 and threads context_lens, block_tables and max_seqlen_q off it explicitly; the new code calls get_forward_context() and re-reads. atom/utils/forward_context.py resolves thread-local first and falls back to the module global, so a TBO worker thread that never installed its own context reads whatever the other ubatch thread wrote last. Since flydsl_plan_matches compares only batch and kv-head count, two equal-sized ubatches pass and the kernel is handed the other ubatch's work_info.

_FLYDSL_PLAN_SCRATCH is one process-global triple shared by every layer, model, builder and stream (base_attention.py:245) [verified]. The static path it replaces allocates exp_sums / max_logits / temporary_output fresh per call (attention_mha.py:587-596), so two in-flight pa_decode calls could never alias. The new key (num_kv_heads, capacity, rows, head_dim, out_dtype, device.index) is identical across every rung of the capture ladder, because capacity = min(batch*max_parts, max(batch, ceil(2*CU/kv_heads))) is a constant 512 for every batch under the workgroup budget at kv_heads=1. Under capture_tbo_graph ("threads + multi-stream captured in graph") two equal-shaped microbatches write the same buffers concurrently — wrong logits, no error, vanishing under HIP_LAUNCH_BLOCKING.

This is only accidentally safe today, because ubatch metadata carries no plan and falls back to per-call scratch. Nothing states the single-stream assumption and no test pins it. Putting the plan on AttentionMetaData like its siblings resolves both facets at once.

4. The dummy-run short-circuit was dropped 30 lines from where it is applied [verified]

aiter_attention.py:892-908 passes running_bs == 0 or get_forward_context().context.is_dummy_run to _mtp_prepare_decode_metadata_kernel as a no-op flag, with the comment: "Dummy runs skip the draft attention, so keep this launch as a no-op: their synthetic context_lens can point past block_tables."

refresh_flydsl_plan(context_lens[:running_bs]) at line 932 then runs unconditionally over that same buffer. Every DP-padding or dummy draft step pays a planner launch it can never use, and because _flydsl_plans hands back the one plan object whose tensors the real target's captured graph baked pointers into, the dummy step rewrites the live plan's work_info from synthetic lengths. Correctness rests entirely on the next real step refreshing first; nothing states or tests that ordering.

5. The gate is narrower than its docstring claims [verified]

_flydsl_pa_decode_num_seqs (base_attention.py:188) says "Every clause below is a hard reject there." Missing versus aiter's pa_decode: num_hgroups != head_dim // 16 (pa_decode.py:337), block_tables.shape[0] != num_seqs (:313), output.shape != query.shape (:308), block_tables.dtype != int32 (:400), context_lengths.dtype != int32 (:403), query.dtype not in (bf16, f16) (:411), output.dtype != query.dtype (:414), num_q_heads % num_kv_heads != 0 (:379), query.stride(2) != 1 (:429), value_cache.dtype != expected_fp8 (:426).

The int32 one is worth calling out on its own: refresh_flydsl_plan does guard context_lens.dtype is not torch.int32. A model handing int64 lengths therefore gets work_plan=None and still routes into FlyDSL via the static path, where aiter raises TypeError: context_lengths must be int32 on the first decode step. The one dtype the author knew about is guarded on the path that survives it and unguarded on the path that cannot.

6. The measured bf16 split clamp is bypassed, and the comment retiring it cites no measurement [verified]

The original comment at base_attention.py:64 records two separate bounds behind PA_DENSE_SPLIT_MAX = 32: the C++ PS-reduce build limit (gluon-only, legitimately inapplicable to the planner) and a numerics bound — "temporary_output is bf16 ... the worst shape measured drifts 20pp further from an fp32 reference at 64 than at 8."

_flydsl_plan_scratch allocates temporary_output in output.dtype, i.e. bf16, so the round-trip is identical. plan_pa_decode defaults max_partitions to the device CU count (256 on MI355X) and the planner clamps per-sequence counts only at MAX_PARTS, so a long request in an uneven batch gets up to 256 splits — 8x past the measured-safe ceiling. Known Limitation #4 concedes the gsm8k arm tested 128, not 256.

Either re-measure the drift at the planner's actual ceiling, or pass max_partitions — which §7 wants anyway.

7. Taking aiter's max_partitions default disables its bounded and vectorized reduce, and a test pins that choice [verified]

int(work_plan.max_partitions) reaches aiter as max_context_partition_num → context_partition_num in launch_pa_decode_ps_reduce. There, bounded_plan_logits requires context_partition_num <= 64, and vectorize_plan_logits is gated on it (pa_decode.py:116-137). At the CU-count default of 256 both are permanently off, so every planned call compiles the unbounded, unvectorized reducer.

That directly explains this PR's own Known Limitation #5 ("PS reduce gets ~14% dearer"). And aiter's comment on the flag is a safety statement — "so out-of-bounds reads cannot wrap back into valid data" — not only a speed one.

test_ceiling_is_left_at_the_aiter_default asserts "max_partitions" not in seen, so capping at 64 goes red and reads as a regression. The test conflates "do not feed the static split count in" with "never set the ceiling"; those are different claims and only the first one is intended.

8. The plan cache is unbounded and the planner specializes on batch [verified, measured here]

_flydsl_plans (aiter_attention.py:821) is never evicted, and aiter's planner kernel takes batch as a tl.constexpr. Measured on this box: a first plan_pa_decode at a new batch value costs 65-72 ms cold (1.2-1.7 ms with a warm triton disk cache); steady-state refresh is 16 µs.

Under --enforce-eager, piecewise skips, above-ladder batches or non-unified DP steps, n tracks the real batch, so each new value mints a plan (never freed: GPU int32 capacity x 4 plus n x 2), one logger.info, and one kernel specialization, up to the 4096 cap. This is the per-shape-JIT-leaking-into-serving failure that base_attention.py:89-94 already guards against in this same file ("the PS reduce compiles one variant per distinct count"); the planner reintroduces it. Startup also pays it per capture rung — roughly 2.6 s over a 20-rung ladder x 2 q-buckets on a cold cache.

9. The feature-off path is no longer free [verified, measured here]

base_attention.py:341 sits outside if flydsl_seqs:. envs.<bool> costs 0.69 µs per access because atom/utils/envs.py:824 re-runs the lambda — and os.getenv — on every attribute access with no cache; the route_sig tuple plus set lookup costs another 0.12 µs.

At 63 pa_decode calls per step that is ~51 µs of CPU per decode step on every deployment that never enables this feature — about 1% of a 5 ms interactivity-bound step, and not elided by graph replay since attention is a piecewise split op. With the env on it is ~92 µs/step. Hoisting the env read to a module constant and moving the route log inside the guard fixes both.

Separately, paged_attention_triton still allocates its three torch.empty scratch tensors unconditionally and the planned branch then discards them — ~3 MB/step of pure allocator churn.

10. Three builders bake the static path into their graphs [reported]

build_for_cudagraph_capture in gdn_attn.py:1551, qwen4_exp_attn.py:673 and _build_ubatch_decode_metadata never attach a plan. The route decision is Python executed at capture time, so with flydsl_work_plan absent, getattr(md, "flydsl_work_plan", None) returns None and the static branch is what gets recorded. Every later refresh_flydsl_plan is then a 16 µs GPU kernel feeding a graph that never reads it.

The "plan refused" warning cannot fire here either, because work_plan is None skips it. So on GDN/Qwen4-Exp/TBO an A/B of ATOM_PA_FLYDSL_PLAN=0/1 measures pure overhead and reads as "the planner does not help" — exactly the ambiguity the routing log was added to remove.

11. Deleting HIP_VISIBLE_DEVICES from the recipe is correct, and it silently breaks NUMA binding [verified, measured here]

The deletion is a real fix: ROCR_VISIBLE_DEVICES=0,1,4,5 HIP_VISIBLE_DEVICES=0,1,4,5 yields torch.cuda.device_count() == 2 on this box (ROCR filters first, then HIP indexes into the filtered set and 4/5 are out of range), versus 4 with ROCR alone. But it arrives unexplained and unrelated to FlyDSL, and it has a cost.

numa_utils.py:54 reads os.environ.get("HIP_VISIBLE_DEVICES") or os.environ.get("CUDA_VISIBLE_DEVICES") and never consults ROCR_VISIBLE_DEVICES. With HIP removed, visible is empty, _physical_index falls back to identity, and logical ranks 2/3 map to physical GPUs 2/3 instead of 4/5. Physical cards sorted by BDF split [0,0,0,0,1,1,1,1] across nodes, ATOM_AUTO_NUMA_BIND defaults to 1, and the recipe sets no ATOM_NUMA_NODE — so all four ranks land on node 0, the condition this same recipe documents as fatal ("timed out on the allocate_kv_cache barrier after 600 s").

Fix-then-sweep: mooncake_connector.py:641 has the identical HIP-or-CUDA chain.

12. A test asserts the gate accepts a call aiter rejects [verified]

test_head_dim_stays_a_subset_of_aiters[256-True] (tests/test_paged_attention_dispatch.py:845) asserts self._call(q=torch.empty(16, 16, 256, ...)) == 4 while _call's baseline k_cache stays (4, 1, 8, 128, 16) — num_hgroups=8, which encodes head_dim 128. aiter rejects exactly that at pa_decode.py:337.

The docstring presents this acceptance as the contract ("Narrower than aiter on purpose; widening is what breaks"), so the missing k_cache.shape[2] * 16 == q.shape[-1] clause will not be added, and any deployment where those disagree routes into FlyDSL and dies on a shape error inside aiter.

13. Conventions

Smaller, but several are repo rules rather than preferences.

run_pa_decode_gluon now dispatches two kernels under a gluon-specific name. ATOM/CLAUDE.md Name-matches-function is mandatory — "When behavior changes, rename immediately". The PR edited the docstring from "Run the AITER paged-attention Gluon decode kernel" to drop the word rather than renaming, leaving four call sites — including the vLLM and SGLang bridges — importing a symbol named gluon with no signal that ATOM_PA_FLYDSL reroutes them.

Neither new env var is documented. grep FLYDSL docs/ finds neither ATOM_PA_FLYDSL nor ATOM_PA_FLYDSL_PLAN among the 96 documented ATOM_* vars, and docs/model_ops_guide.md:152 still states the decode kernel is torch.ops.aiter.pa_decode_gluon.

The comment at line 160 says the mirrored limits are "pinned against its source by a test", and they are not. TestFlyDSLConstantsMatchAiter pins only the block-size whitelist and the arch tuple; _FLYDSL_PA_MAX_PARTITIONS, _FLYDSL_PA_TILE and _FLYDSL_PLAN_MAX_BATCH are unpinned.

Five new assertions test source text rather than behaviour (inspect.getsource / ast.parse at tests lines 494, 508, 600, 686, 611-668). A test that greps the code it tests passes for reasons unrelated to what the code does. test_scratch_is_keyed_by_capacity also leaves ~25 MB of live CUDA buffers in the never-evicted module global with no teardown.

envs.py:525 says "+18.6% interactivity at conc 20"; the recipe says "+20.5%" for the same measurement.


Claims checked and not upheld

Recorded so they are not re-raised:

  • The planner is not inert. running_tokens == running_bs * max_seqlen_q holds by construction (forward_context.py:370-378), so flydsl_plan_matches passes on every decode step in the shipped TP4/no-DP recipe. The gate docstring's stated reason for the num_seqs design is factually wrong, but the design works.
  • monkeypatch.setattr(envs, ...) does not leak — the pre-existing autouse keep_envs_lazy fixture cleans it.
  • Zero-length padded rows come back as 0.0, not uninitialized memory.

…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>
@valarLip
valarLip merged commit 94cde4b into main Sep 23, 2026
74 of 80 checks passed
@valarLip
valarLip deleted the flydsl-pa-decode-routing branch September 23, 2026 12:01
yhl-amd added a commit to SemiAnalysisAI/InferenceX that referenced this pull request Sep 26, 2026
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>
yhl-amd added a commit to SemiAnalysisAI/InferenceX that referenced this pull request Sep 28, 2026
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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants