Files
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

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"]