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.
This commit is contained in:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+5
View File
@@ -0,0 +1,5 @@
"""Incremental recurrent KDA implementations and state containers."""
from .fused import KDAState, fused_recurrent_kda, fused_recurrent_kda_step
__all__ = ["KDAState", "fused_recurrent_kda", "fused_recurrent_kda_step"]
+73
View File
@@ -0,0 +1,73 @@
"""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"]