Skip to content

[Triton/Gluon] chunk_kimi_delta_attn: accept a non-fp32 KDA state - #5249

Merged
valarLip merged 2 commits into
mainfrom
kda-state-dtype
Sep 8, 2026
Merged

valarLip merged 2 commits into
mainfrom
kda-state-dtype

Conversation

@XiaobingSuper

Copy link
Copy Markdown
Contributor

What

chunk_kimi_delta_attn required an fp32 initial_state and always returned an fp32 final_state. A caller whose KDA state pool is 2-byte had to cast in and back out — two extra passes over the state per prefill chunk.

The recurrence already accumulates in fp32 whatever the state is stored as, so the constraint was only on the load and the allocation:

  • _flash_kda_seg_scan_kernel up-casts the h0 tile explicitly, the way _flash_kda_segment_kernel already did.
  • flash_kda_fwd no longer casts h0, and allocates final_state at h0's dtype (fp32 when there is no h0, as before). The tl.store into it was already an implicit cast to the pointer's element type.
  • The fp32-only check in chunk_kimi_delta_attn is dropped.

fp32 callers are unaffected — same dtypes in, same dtypes out.

Accuracy

Against an fp32 reference, B=1, T=2048, H=12, K=V=128, 2 varlen sequences:

initial_state output rel err final_state rel err
fp32 (reference) (reference)
fp16 1.272e-04 2.079e-04
bf16 1.409e-04 1.663e-03

final_state.dtype matches the input dtype in every case.

Why

Consumed by ROCm/ATOM's ATOM_KDA_SSM_DTYPE, which makes the Kimi-K3 KDA temporal state pool's storage dtype configurable. KDA decode is state-bandwidth bound, so halving the element size is worth ~2 ms/step over 69 layers.

`chunk_kimi_delta_attn` required an fp32 `initial_state` and always
returned an fp32 `final_state`. A caller whose state pool is 2-byte had
to cast on the way in and back out on the way out, which is two extra
passes over the state per prefill chunk.

The recurrence already accumulates in fp32 whatever the state is stored
as, so the constraint was only on the load and the allocation:

- `_flash_kda_seg_scan_kernel` up-casts the `h0` tile explicitly, the way
  `_flash_kda_segment_kernel` already did.
- `flash_kda_fwd` no longer casts `h0`, and allocates `final_state` at
  `h0`'s dtype (fp32 when there is no `h0`, as before). The `tl.store`
  into it was already an implicit cast to the pointer's element type.
- The fp32-only check in `chunk_kimi_delta_attn` is dropped.

fp32 callers are unaffected. Measured against an fp32 reference
(B=1, T=2048, H=12, K=V=128, 2 varlen sequences):

  initial_state   out rel err   state rel err
  fp16            1.272e-04     2.079e-04
  bf16            1.409e-04     1.663e-03
@XiaobingSuper
XiaobingSuper requested review from a team and a lite review from Copilot September 3, 2026 11:58
@github-actions github-actions Bot changed the title [Triton] chunk_kimi_delta_attn: accept a non-fp32 KDA state [Triton/Gluon] chunk_kimi_delta_attn: accept a non-fp32 KDA state Sep 3, 2026
@github-actions

github-actions Bot commented Sep 3, 2026

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every PR:

  • ✅ Pre-checks (submodule verification, code formatting)
  • ✅ Aiter op tests (gfx942 + gfx950)
  • ✅ Triton tests on MI35X (only when aiter/ops/triton/** or related paths are changed)

Extended tests (opt-in via labels):

Label Tests
ci:gfx1250-ffm-triton Run the five-shard gfx1250 FFM Triton test suite
ci:triton-300x Run an additional Triton test job on MI300X in PRs; main branch always runs both MI35X and MI300X
multigpu Aiter multi-GPU tests on the 8-GPU runner
ci:sglang SGLang integration tests: DeepSeek-R1-MXFP4 accuracy, Qwen 3.5 accuracy
ci:atom ATOM benchmark: DeepSeek-R1-0528, GPT-OSS-120B
ci:atom_full ATOM accuracy suite for PR and main models from ATOM models_accuracy.json
ci:vllm vLLM benchmark: GPT-OSS-120B, DeepSeek-R1-0528, Kimi-K2.5
ci:all All standard extended tests (excludes ci:atom_full)

Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 5249 --add-label <label>

PR title tags & labels:
Component tags ([Triton/Gluon], [HIP], [CK], [ASM], ...) are added to the PR title and as PR labels automatically from the changed files and re-synced on every push — change-type tags like [fix]/[Perf], op tags like [MLA], and human labels (ci:*) are left untouched. Add the no-auto-title label to opt this PR out.

XiaobingSuper added a commit to ROCm/ATOM that referenced this pull request Sep 3, 2026
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>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Allowing non-fp32 final_state requires corresponding dtype-safe stores in the FlashKDA kernel (and the public API docs/validation need to match actual per-path behavior).

Once you've addressed the issues Copilot identified, you can request another Copilot review.

Pull request overview

Updates the Kimi Delta Attention (chunk_kimi_delta_attn) Triton path to accept a non-fp32 recurrent state tensor, reducing unnecessary casts and enabling smaller state-storage dtypes for bandwidth-bound workloads.

Changes:

  • Loosens chunk_kimi_delta_attn’s initial_state contract by removing the fp32-only check and updating the docstring.
  • In FlashKDA, explicitly upcasts h0 loads to fp32 inside _flash_kda_seg_scan_kernel.
  • Allocates FlashKDA final_state using h0.dtype (instead of always fp32) when output_final_state=True.
File summaries
File Description
aiter/ops/triton/kimi_delta_attn/chunk_delta_attn.py Updates public API docs and removes the fp32-only initial_state validation.
aiter/ops/triton/_triton_kernels/chunk_delta_attn/flash_kda.py Adjusts FlashKDA state load/upcast behavior and final-state allocation dtype.
Review details
  • Files reviewed: 2/2 changed files
  • Comments generated: 3
  • Review effort level: Lite

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment on lines 842 to 846
if output_final_state:
shape = (N, H, V, K) if state_v_first else (N, H, K, V)
final_state = torch.empty(shape, dtype=torch.float32, device=dev)
state_dtype = h0.dtype if h0 is not None else torch.float32
final_state = torch.empty(shape, dtype=state_dtype, device=dev)

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The store was already correct (Triton implicitly converts to the pointer element type, and fp16/bf16 both round-trip in the numerics table above), but the surrounding out/h_out stores cast explicitly, so I made these match.

Comment thread aiter/ops/triton/kimi_delta_attn/chunk_delta_attn.py
Comment on lines 99 to 103
`state_v_first=True`), for `N` input sequences. Any float dtype; the
recurrence accumulates in fp32 whatever the state is stored as. For
equal-length inputs `N` equals the batch size `B`. Default: `None`.
output_final_state (bool):
Whether to return the final state, same shape and dtype as

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Correct, and I missed it — the default pipeline allocates final_state as fp32 in chunk_delta_h.py. Docstring now states the dtype per path rather than promising one.

XiaobingSuper added a commit to ROCm/ATOM that referenced this pull request Sep 3, 2026
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>
Copilot AI review requested due to automatic review settings September 3, 2026 12:14
XiaobingSuper added a commit to ROCm/ATOM that referenced this pull request Sep 3, 2026
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>
XiaobingSuper added a commit to ROCm/ATOM that referenced this pull request Sep 3, 2026
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>
@XiaobingSuper
XiaobingSuper requested a review from zufayu September 3, 2026 12:18
…dtype doc

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

The two newly added GEMM tuning JSONs are placed under the legacy flat configs/gemm/ path (not discoverable by the current nested resolve_config_dir() layout), and the new state-dtype contract needs targeted test coverage to prevent regressions.

Once you've addressed the issues Copilot identified, you can request another Copilot review.

Review details
  • Files reviewed: 2/2 changed files
  • Comments generated: 1
  • Review effort level: Lite

Comment thread aiter/ops/triton/kimi_delta_attn/chunk_delta_attn.py
Copilot AI review requested due to automatic review settings September 3, 2026 12:22

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟢 Approval recommended

The functional changes are narrowly scoped and consistent across validation, kernel load/store behavior, and allocation, with only a minor doc clarification suggested.

Review details
  • Files reviewed: 2/2 changed files
  • Comments generated: 1
  • Review effort level: Lite

Comment thread aiter/ops/triton/kimi_delta_attn/chunk_delta_attn.py

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Low-precision state handling lacks regression tests, and the FlashKDA docstring remains outdated.

Once you've addressed the issues Copilot identified, you can request another Copilot review.

Review details
  • Files reviewed: 2/2 changed files
  • Comments generated: 2
  • Review effort level: Balanced

Comment thread aiter/ops/triton/_triton_kernels/chunk_delta_attn/flash_kda.py
Comment thread aiter/ops/triton/_triton_kernels/chunk_delta_attn/flash_kda.py
XiaobingSuper added a commit to ROCm/ATOM that referenced this pull request Sep 8, 2026
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>
@valarLip
valarLip merged commit 93316d2 into main Sep 8, 2026
73 checks passed
@valarLip
valarLip deleted the kda-state-dtype branch September 8, 2026 05:12
XiaobingSuper added a commit to ROCm/ATOM that referenced this pull request Sep 8, 2026
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>
valarLip pushed a commit to ROCm/ATOM that referenced this pull request Sep 8, 2026
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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants