perf(kimi-k3): make the KDA temporal state dtype configurable - #2130
Conversation
🏷️ CI GuideRuns automatically on every eligible PR before approval:
Heavy model tests:
|
41ef086 to
8c9e0b8
Compare
There was a problem hiding this comment.
🟡 Changes recommended
There’s a concrete correctness issue in fused_sigmoid_gating_delta_rule_update where initial_state is treated as optional but dereferenced unconditionally.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Pull request overview
Adds an environment-variable switch to control the storage dtype of the KDA temporal (recurrent) state pool for KDA models, enabling a lower-bandwidth decode path while keeping the default behavior unchanged.
Changes:
- Introduce
ATOM_KDA_SSM_DTYPE(defaultfp32) and plumb it intoGDNStateMixin._state_dtypes()forkimi_linear/glm5_next_text. - Adjust the fused sigmoid gating kernel’s
BVcap based on the state element size. - Add/extend tests to validate accepted dtype values, default behavior, pool sizing, and layout-id differentiation.
File summaries
| File | Description |
|---|---|
| tests/test_kda_ssm_dtype.py | New unit tests covering accepted env values, default, rejection of unknown values, and pool sizing impact. |
| tests/test_kda_layout_id.py | Extends layout-id assertions to distinguish fp16 vs bf16 temporal dtype (same size, different semantics). |
| atom/utils/envs.py | Adds ATOM_KDA_SSM_DTYPE env var definition and documentation. |
| atom/model_ops/fla_ops/fused_sigmoid_gating.py | Makes BV tuning depend on the state’s element size for bandwidth-bound KDA decode. |
| atom/model_ops/attentions/kimi_mla_gdn_attn.py | Updates paged-checkpoint transfer docstring to reflect configurable temporal dtype. |
| atom/model_ops/attentions/gdn_attn.py | Adds dtype mapping + env parsing/validation and uses it for KDA temporal-state dtype selection. |
Review details
- Files reviewed: 6/6 changed files
- Comments generated: 2
- Review effort level: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
8c9e0b8 to
290a47f
Compare
There was a problem hiding this comment.
🔵 Needs a closer look
It changes a performance-critical recurrent-state dtype and kernel tiling behavior that can affect numerical stability and runtime characteristics and warrants final human review.
Review details
Suppressed comments (1)
atom/model_ops/attentions/gdn_attn.py:446
- The comment says kimi_linear's temporal side “breaks it at any setting”, but with ATOM_KDA_SSM_DTYPE set to match config.torch_dtype (e.g. bf16), the two dtypes can still agree. This is documentation-only, but it’s misleading about when the midstep-exactness argument fails.
Exact, not approximate, when it is turned back on: `h` is `k.new_empty`
and `_state_dtypes` returns `config.torch_dtype`, so slicing `h` rounds
exactly where a shortened forward would. That rests on the two dtypes
agreeing; kimi_linear's temporal side is dtype-configurable
(ATOM_KDA_SSM_DTYPE) and breaks it at any setting, so it overrides
(`_KimiMLAGDNCommon.state_transfer`).
- Files reviewed: 5/5 changed files
- Comments generated: 0 new
- Review effort level: Lite
290a47f to
a60b2d3
Compare
a60b2d3 to
808dbcc
Compare
There was a problem hiding this comment.
🟡 Changes recommended
There are a few correctness/robustness issues to address (notably an API contract mismatch around initial_state being effectively required despite a None default, plus missing test coverage for the new env var).
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Review details
Suppressed comments (2)
atom/model_ops/fla_ops/fused_sigmoid_gating.py:237
initial_stateis annotated with a default ofNone, but this function unconditionally dereferences it (element_size(),dtype,stride). If a caller relies on the default it will raise an AttributeError; consider validating early and raising a clear ValueError (or make the parameter required).
# State-bandwidth bound: a 2-byte state wants twice the V per block to keep
# the same bytes in flight (HV=12, K=V=128, N=64: 18.2 -> 16.1 us, while
# fp32 goes 24.6 -> 25.2).
bv_cap = 64 if initial_state.element_size() <= 2 else 32
BK, BV = triton.next_power_of_2(K), min(triton.next_power_of_2(V), bv_cap)
atom/model_ops/attentions/gdn_attn.py:444
- Docstring grammar: “and need not” is missing the implied verb/object, which makes the sentence hard to parse. Suggest clarifying that the temporal dtype need not agree with the model dtype.
agreeing; kimi_linear's temporal side is `ATOM_KDA_SSM_DTYPE` and need
not, so it overrides unconditionally
(`_KimiMLAGDNCommon.state_transfer`).
- Files reviewed: 5/5 changed files
- Comments generated: 1
- Review effort level: Lite
| # "fp32" | "fp16" | "bf16". Storage dtype of the KDA temporal state pool, | ||
| # whose per-token traffic dominates KDA decode; the recurrence itself | ||
| # always accumulates in fp32. fp16 over bf16 when narrowing: the state is | ||
| # O(1), so bf16's range buys nothing and its short mantissa costs accuracy. | ||
| "ATOM_KDA_SSM_DTYPE": lambda: os.getenv("ATOM_KDA_SSM_DTYPE", "fp32").lower(), |
808dbcc to
9c0256c
Compare
There was a problem hiding this comment.
🟢 Approval recommended
The functional changes are localized and default-preserving, with validation and tests updated; only a minor docstring grammar nit was found.
Review details
- Files reviewed: 5/5 changed files
- Comments generated: 1
- Review effort level: Lite
KDA decode is state-bandwidth bound: 69 layers each stream a
[12, 128, 128] fp32 recurrent state per token, and the fused gating
kernel already runs at 71-90% of achievable bandwidth. Halving the
state's element size is the lever that is left.
ATOM_KDA_SSM_DTYPE ("fp32" | "fp16" | "bf16", default "fp32") picks the
storage dtype of the temporal state pool for the KDA families
(kimi_linear and glm5_next_text). The default is unchanged, so main
behaves identically.
The dtype is decided in one place, GDNStateMixin._state_dtypes. Pool
sizing, per-request allocation, the checkpoint plane shapes and the
checkpoint layout id all derive their bytes from it. No cast is
introduced anywhere: the recurrence accumulates in fp32 whatever the
pool stores, and every path -- prefill, decode, spec-decode, ReplaySSM
-- reads and writes the state through the destination pointer's element
type. Prefill needs ROCm/aiter#5249.
The gating kernel's BV cap now follows the state's element size. It was
tuned for a 4-byte state; a 2-byte one wants twice the V per block to
keep the same bytes in flight (HV=12, K=V=128, N=64: 18.2 -> 16.1 us at
BV=64), while fp32 is slightly worse there (24.6 -> 25.2 us).
fp16 rather than bf16 for the narrow setting: q and k are L2-normalized
in-kernel so the state is O(1) and bf16's range buys nothing, while its
three fewer mantissa bits cost roughly 8x the error.
Measured on MI355X, Kimi-K3 TP8, per recipes/Kimi-K3.md:
GSM8K, 1319 questions, 5-shot, greedy (strict = flexible):
fp32 0.9659 fp16 0.9591 +- 0.0055
Serving, 256 in / 1024 out, --ignore-eos:
fp32 fp16
c=32 tok/s 1144.71 1164.95 +1.8%
TPOT ms 26.91 26.55 -1.3%
ITL ms 33.86 32.13 -5.1%
c=64 tok/s 1902.21 1913.54 +0.6%
TPOT ms 32.34 31.75 -1.8%
ITL ms 40.83 39.48 -3.3%
State pool: 56.17 -> 29.04 MB per slot, 3.35 -> 1.73 GB for 64 slots,
which the KV pool takes over (53.58 -> 55.61 GB, +1236 blocks).
Flipping the default is the accuracy owner's call: GSM8K's short
generations do not exercise long-context state accumulation, and the
evidence that the error does not grow with sequence length is offline
numerics rather than an end-to-end run.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
9c0256c to
b82994e
Compare
There was a problem hiding this comment.
🟢 Approval recommended
The functional changes are localized, validated with an updated layout-id test, and the only remaining feedback is a minor docstring grammar nit.
Review details
Suppressed comments (1)
atom/model_ops/attentions/gdn_attn.py:435
- Docstring reads "and need not" without an object; adding "agree" makes the sentence grammatical and clearer about why the override is unconditional.
agreeing; kimi_linear's temporal side is `ATOM_GDN_SSM_DTYPE` and need
not, so it overrides unconditionally
(`_KimiMLAGDNCommon.state_transfer`).
- Files reviewed: 5/5 changed files
- Comments generated: 0 new
- Review effort level: Lite
Apply the KDA temporal state dtype optimization from ATOM #2130. Co-authored-by: Cursor <cursoragent@cursor.com>
* (recipe) add Kimi-K3 AgentX recipe Document the per-concurrency ATOM serving bands for the Weka AgentX workload. * (recipe) tune Kimi-K3 AgentX ReplaySSM and AITER comm reuse Turn ReplaySSM off at C1/C2/C4 and on at C16, and keep AITER comm-group reuse only on C56/C64. * (recipe) document container prerequisites and add the C14 point Two container-level facts decide whether this recipe reproduces its own numbers, and both of them fail quietly, so add a prerequisites section covering each: - triton must be 3.7.x. The kimi_k3_agentic_0903 image shipped 3.8.0, which costs ~2.1x on prefill TTFT here (C1 p50 0.79s -> 1.5s) while leaving decode untouched. Measured across eleven runs: the gap survives swapping aiter's .so between revisions, swapping the ATOM checkout, and switching between a full CI prebuild and a lean JIT build, so triton is the only variable that moves it. - DRAFT_MODEL_PATH must name a local directory. The default is a Hugging Face repo id the image has no cache entry for, so leaving it unset makes every rank fetch the same 7 GB checkpoint with no error printed, the log stopped after "Loading drafter model...", and every GPU at 0%. Document the symptom next to the fix, because the hang is otherwise indistinguishable from a wedged server. C14 was in the benchmark scripts but never in the recipe. It reuses the C16 server recipe verbatim, including a CUDA-graph width pinned to 32; deriving it from 2*CONC gives 28 and the server fails during CUDA-graph warmup. Routing GRAPH_MAX through CUDAGRAPH_MAX_NUM_SEQS lets the C14 arm pin it while every other concurrency keeps the 2*CONC default unchanged. Verified against run_kimik3_atom_dspark_server.sh: all eleven concurrencies now agree item by item across DCP, max-num-seqs, batched tokens, GPU utilization, LMCache, ReplaySSM, spec tokens, acceptance length, AITER comm reuse and the recomputed graph_max. * (recipe) set ATOM_GDN_SSM_DTYPE=fp16 for Kimi-K3 AgentX Apply the KDA temporal state dtype optimization from ATOM #2130. Co-authored-by: Cursor <cursoragent@cursor.com> * (recipe) Kimi-K3 AgentX: per-concurrency params from the 12-point MI355X sweep (#2208) * (recipe) Kimi-K3 AgentX: band params from the 12-point MI355X sweep Replaces the per-concurrency values with the ones a full 12-point AgentX sweep actually ran on MI355X (8x, ATOM 97c15a7 + aiter 6b3f14b5f, aiperf 0.12.0, 3600 s profiling window per point). * GPU util 0.88/0.86 -> 0.90 at every band. Measured on this box as available_for_kv 39.42 -> 51.11 GB, a +40.6% KV sub-pool. * max-num-batched-tokens 4096 -> 8192 at C8 and C12, so it is now 8192 everywhere. Cross-round TTFT p50 deltas move monotonically through zero across C8..C56, consistent with the larger prefill chunk. * AITER_REUSE_IDENTICAL_COMM_GROUPS is now a DCP rule rather than two hardcoded branches: 1 at every DCP=8 band, 0 at DCP=1. C32/C40/C48 measured faster with it on; C1 measured 3.3-3.9% WORSE ITL with it on, which is why the DCP=1 bands keep it off. C8-C16 are still under A/B and may move in a follow-up. * Adds the missing C48 band (max-num-seqs 96, graph_max 96). It was in neither the table nor the case, so CONC=48 hit the unsupported exit. * Pins ATOM_USE_FLYDSL_GATHER_KV_B_PROJ=1. It already defaults to 1 (atom/utils/envs.py), so this changes nothing today -- it is set explicitly so the recipe keeps reproducing the measured numbers if the default flips. DCP, spec tokens, acceptance length, LMCache sizes, ReplaySSM, max-num-seqs and graph_max are unchanged; every row was re-verified against the sweep's own launcher. * (recipe) Kimi-K3 AgentX: hoist band-invariant values out of the case The per-concurrency case repeated nine values that are now identical across every band, so a single-band change had twelve places to go wrong. Hoist the shared defaults above the case and leave only the genuine per-band deltas inside it: 119 lines -> 51. Behaviour is unchanged. All twelve concurrencies were evaluated before and after and every field -- DCP, max-num-seqs, batched tokens, GPU util, ENABLE_LMCACHE, LMCache size, ReplaySSM, reuse, spec tokens, acceptance length, CUDAGRAPH_MAX_NUM_SEQS and the derived graph_max -- is byte-identical. max-num-seqs is now derived as 2*CONC on the throughput bands rather than written out five times; it was already exactly 2*CONC there. The two bands that genuinely deviate keep explicit overrides with the reason attached: C12 at 24, and C14's pinned CUDAGRAPH_MAX_NUM_SEQS=32. LMCACHE_MAX_LOCAL_CPU_SIZE now carries its 128 default on the DCP=1 bands too. That is inert: it is only exported inside the ENABLE_LMCACHE==1 branch, which those bands do not take. * (recipe) trim the ATOM_USE_FLYDSL_GATHER_KV_B_PROJ comment * (recipe) Kimi-K3 AgentX: comm-group reuse off below C56 A seven-point A/B on MI355X (C8/12/14/16/32/40/48, one reuse=1 leg and one reuse=0 leg each, identical builds) measured reuse=0 faster at every point: ITL p50 -7.3% at C8 narrowing to -1.9% at C48, total tok/chip +3.8% to +1.2%. Split into 10-minute buckets, 39 of 42 buckets favour reuse=0 (one-sided p ~ 3e-9), which is well clear of the ~7% run-to-run noise floor that makes any single full-window comparison on this workload unreliable. The advantage shrinks monotonically as the KV pool tightens -- reuse=0 gives up roughly 13k blocks, and its peak KV usage climbs from 77.6% at C8 to 97.6% at C48, where it is the first point to spend any time above 95% and where maxRunning starts being pushed back. C56/C64 were not measured and use a larger LMCache tier and max-num-seqs, so they keep the previous setting. This also removes the claim that the 12-point sweep measured enabling reuse at C32/C40/C48 as a gain. That reading came from single-run full-window aggregates and does not survive the paired A/B. * (recipe) Kimi-K3 AgentX: C72 + C80 bands, drop the C14 CUDA-graph pin (#2213) * (recipe) Kimi-K3 AgentX: add the C72 throughput band Same shape as C64 -- DCP=8, no spec, 192 GiB LMCache, ReplaySSM off -- with max-num-seqs following the 2*CONC rule every band from C32 up obeys, so 144, and graph_max 144 with spec disabled. Measured on MI355X against the C64 band it extends: total tok/chip 12794 vs 12221, so C64 is not the throughput plateau. ITL p50 moves 77.11 -> 88.60 and P90 interactivity 9.9 -> 8.4, which is the trade the band exists to expose. maxRunning reached 119, so the wider capture range is genuinely reached rather than captured and left unused. Reuse stays enabled here: the CONC >= 56 gate already covers C72, and a paired A/B measured reuse=1 ahead at C72 by a wider margin than at C64 (ITL bucket mean +4.5% and total tok/chip +3.1% in its favour, 6 of 6 ten-minute buckets agreeing), consistent with C72 sitting further past the C56 crossover. * (recipe) Kimi-K3 AgentX: drop the C14 CUDA-graph pin C14 pinned CUDAGRAPH_MAX_NUM_SEQS=32 with the justification that deriving graph_max from 2*CONC "would give 28 and the server fails during CUDA-graph warmup". That is false on ATOM 97c15a7, and the pin has no effect. Capture sizes are SEQUENCES, not tokens. ModelRunner truncates the declared ladder to max_schedulable_decode_bs() = min(max_num_seqs, max_num_batched_tokens // full_q_len) (atom/model_engine/model_runner.py:142) -- it divides the speculative width out, while the launcher multiplies it in. For C14 that is min(32, 8192//4) = 32, independent of CUDAGRAPH_MAX_NUM_SEQS. Sizes above the bound are dropped with a warning, not asserted on; a batch past the ladder just falls back to eager. So both settings capture exactly bs=2..32: pinned CGMS=32 -> declared [2..128] -> "= 2048) = 32; dropping" x8 ranks default CGMS=28 -> declared [2..112] -> "= 2048) = 32; dropping" x8 ranks Measured on 8xMI355X, AgentX 3600 s, C14, AITER_REUSE_IDENTICAL_COMM_GROUPS=0 both legs, ATOM 97c15a7 / aiter 6b3f14b5f / triton 3.7.0: TTFT p50 TTFT p90 ITL p50 ITL p90 Inter p90 tot tok/s CGMS=32 804.3 2140.0 12.42 19.18 52.1 57148 CGMS=28 775.8 2028.4 12.36 19.18 52.1 57060 ITL p90 and Inter p90 are equal to the digit; throughput differs by -0.15%. Ten-minute buckets have identical n per bucket and ITL deltas of -5.4/+1.2/+0.2/-0.9/-1.0/+3.6 %, three each way, all inside the ~7% run-to-run noise floor. The server reached READY in ~330 s unpinned. The table's graph_max for C14 becomes 2*14*(1+3) = 112. With the pin gone, every row now satisfies the stated rule graph_max = 2 * CONC * (1 + spec); C14 was the only exception. Co-Authored-By: Claude <noreply@anthropic.com> * (recipe) Kimi-K3 AgentX: add the C80 throughput band C80 is the same shape as C64/C72 -- DCP=8, no spec, 192 GiB LMCache, max-num-seqs following the 2*CONC rule (160). It needs no new branch: the existing `CONC -ge 56` guards already pick up the larger CPU tier and AITER_REUSE_IDENTICAL_COMM_GROUPS=1. Measured on 8xMI355X, AgentX 3600 s, ATOM 97c15a7 / aiter 6b3f14b5f / triton 3.7.0, reuse=1: CONC TTFT p50 TTFT p90 ITL p50 ITL p90 Inter p90 tot tok/chip out tok/chip 64 1356.1 3802.3 77.11 100.66 9.9 12221 88.56 72 1489.8 3963.2 88.60 119.50 8.4 12794 89.42 80 1549.8 4405.3 103.07 130.43 7.7 12839 95.75 **Total throughput plateaus at C72.** C64 -> C72 is +4.7% on tokens/chip; C72 -> C80 is +0.35%, well inside the ~7% run-to-run noise floor, while ITL p90 rises 9% and Inter p90 falls from 8.4 to 7.7. Read `tot` with care at the top of the ladder. C80's ISL drops to 107777 (vs C72's 116894) because the trace serves more, shorter requests (3443 vs 3165), so output throughput still climbs +7.1% (715 -> 766 tok/s) while total flattens. The two agree at every band from C8 to C72 and diverge for the first time here. Saturation is unambiguous across three independent signals: KV usage samples >=95% KV max maxRunning C64 2/528 96.5% 84 C72 17/545 96.2% 119 C80 48/556 99.7% 114 maxRunning stops rising at C80 while KV hits 99.7% -- the scheduler is bounded by the KV pool, not by the concurrency parameter. This is the same ordering variable that decides the AITER_REUSE_IDENTICAL_COMM_GROUPS crossover at C56. C80 is included for completeness of the band table. C72 is the knee. Co-Authored-By: Claude <noreply@anthropic.com> --------- Co-authored-by: Claude <noreply@anthropic.com> --------- Co-authored-by: xytpai <xytpai@foxmail.com> Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: gbyu-amd <Guanbao.Yu@amd.com> Co-authored-by: Claude <noreply@anthropic.com>
What
ATOM_GDN_SSM_DTYPE("fp32"|"fp16"|"bf16", default"fp32") picks the storage dtype of the KDA temporal (recurrent) state pool, forkimi_linearandglm5_next_text. The default is unchanged, so main behaves identically.KDA decode is state-bandwidth bound — 69 layers each stream a
[12, 128, 128]fp32 state per token, and the fused gating kernel already runs at 71-90% of achievable bandwidth. Halving the element size is the lever that is left.The dtype is decided in one place,
GDNStateMixin._state_dtypes(). Pool sizing, per-request allocation, the checkpoint plane shapes and the checkpoint layout id all derive their bytes from it, so nothing else needed to change; the layout id already names both dtypes, so a build that flips the variable cannot read another's checkpoint images.No cast is introduced anywhere. The recurrence accumulates in fp32 whatever the pool stores, and every path — prefill, decode, spec-decode, ReplaySSM — reads and writes the state through the destination pointer's element type.
One other change: the gating kernel's
BVcap now follows the state's element size. It was tuned for a 4-byte state; a 2-byte one wants twice the V per block to keep the same bytes in flight (HV=12, K=V=128, N=64: 18.2 → 16.1 us at BV=64), while fp32 is slightly worse there (24.6 → 25.2 us).fp16 rather than bf16 for the narrow setting: q/k are L2-normalized in-kernel so the state is O(1) and bf16's range buys nothing, while its 3 fewer mantissa bits cost roughly 8× the error.
Depends on
ROCm/aiter#5249 —
chunk_kimi_delta_attn(the prefill kernel) previously required an fp32initial_state.Accuracy
GSM8K, 1319 questions, 5-shot, greedy.
strict-matchandflexible-extractagree in both arms.Both are inside the range
recipes/Kimi-K3.mdrecords as verified (0.9538–0.9591).Performance
MI355X, Kimi-K3 TP8, per
recipes/Kimi-K3.md. Serving benchmark, 256 in / 1024 out,--ignore-eos.Memory
Note on the default
This PR only adds the switch and the data. Flipping the default is the accuracy owner's call: GSM8K's short generations do not exercise long-context state accumulation, and the evidence that the error does not grow with sequence length is offline numerics rather than an end-to-end run.