"""L6: FLA fused recurrent KDA decode with optional step cache.""" from __future__ import annotations from dataclasses import dataclass import torch from kda._fla.ops.kda.fused_recurrent import fused_recurrent_kda as _fused_recurrent_kda @dataclass class KDAState: """Mutable recurrent state cache: ``S`` is ``[B, HV, K, V]``.""" S: torch.Tensor pos: int = 0 def reset(self): self.S.zero_() self.pos = 0 def fused_recurrent_kda_step( state: KDAState, q_t: torch.Tensor, k_t: torch.Tensor, v_t: torch.Tensor, g_t: torch.Tensor, beta_t: torch.Tensor, scale: float | None = None, ): """Single-token step. Inputs are ``[B, H|HV, ...]`` (no time dim).""" o, ht = _fused_recurrent_kda( q_t.unsqueeze(1), k_t.unsqueeze(1), v_t.unsqueeze(1), g_t.unsqueeze(1), beta_t.unsqueeze(1), scale=scale, initial_state=state.S, output_final_state=True, ) state.S = ht state.pos += 1 return o.squeeze(1) def fused_recurrent_kda( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, g: torch.Tensor, beta: torch.Tensor, scale: float | None = None, initial_state: torch.Tensor | None = None, output_final_state: bool = False, **kwargs, ): return _fused_recurrent_kda( q, k, v, g, beta, scale=scale, initial_state=initial_state, output_final_state=output_final_state, **kwargs, ) __all__ = ["KDAState", "fused_recurrent_kda", "fused_recurrent_kda_step"]