[Triton/Gluon] chunk_kimi_delta_attn: accept a non-fp32 KDA state - #5249
Conversation
`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
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
PR title tags & labels: |
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>
There was a problem hiding this comment.
🟡 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’sinitial_statecontract by removing the fp32-only check and updating the docstring. - In FlashKDA, explicitly upcasts
h0loads to fp32 inside_flash_kda_seg_scan_kernel. - Allocates FlashKDA
final_stateusingh0.dtype(instead of always fp32) whenoutput_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.
| 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) | ||
|
|
There was a problem hiding this comment.
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.
| `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 |
There was a problem hiding this comment.
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.
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>
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>
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>
…dtype doc Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
148dec1 to
3696211
Compare
There was a problem hiding this comment.
🟡 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
There was a problem hiding this comment.
🟢 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
There was a problem hiding this comment.
🟡 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
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>
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>
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>
What
chunk_kimi_delta_attnrequired an fp32initial_stateand always returned an fp32final_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_kernelup-casts theh0tile explicitly, the way_flash_kda_segment_kernelalready did.flash_kda_fwdno longer castsh0, and allocatesfinal_stateath0's dtype (fp32 when there is noh0, as before). Thetl.storeinto it was already an implicit cast to the pointer's element type.chunk_kimi_delta_attnis 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_statefinal_state.dtypematches 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.