Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
74 lines
1.6 KiB
Python
74 lines
1.6 KiB
Python
"""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"]
|