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
+7
View File
@@ -0,0 +1,7 @@
"""Local Triton KDA kernels vendored from FLA chunk_{fwd,intra,bwd,wy,gate}."""
from .chunk import ChunkKDAFunction, chunk_kda
from .chunk_fwd import chunk_kda_fwd
from .gate import kda_gate_fwd
__all__ = ["ChunkKDAFunction", "chunk_kda", "chunk_kda_fwd", "kda_gate_fwd"]
+5
View File
@@ -0,0 +1,5 @@
"""FLA ``chunk_kda`` surface used by ``ops.api`` backend='triton'."""
from kda._fla.ops.kda.chunk import ChunkKDAFunction, chunk_kda
__all__ = ["ChunkKDAFunction", "chunk_kda"]
+5
View File
@@ -0,0 +1,5 @@
"""Vendored FLA chunk KDA backward."""
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
__all__ = ["chunk_kda_bwd"]
+37
View File
@@ -0,0 +1,37 @@
"""Vendored FLA chunk KDA forward, returning ``(o, ht)`` like the public op."""
from __future__ import annotations
import torch
from kda._fla.ops.kda.chunk import chunk_kda
from kda._fla.ops.kda.chunk_fwd import chunk_kda_fwd as fla_chunk_kda_fwd
__all__ = ["chunk_kda_fwd", "fla_chunk_kda_fwd"]
def chunk_kda_fwd(
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,
chunk_size: int = 64,
**kwargs,
):
"""Chunked KDA forward with FLA kernels. Returns ``(o, ht)``."""
return chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=chunk_size,
**kwargs,
)
+36
View File
@@ -0,0 +1,36 @@
"""Vendored FLA KDA gate fusion (standard + safe gate + chunk cumsum)."""
from __future__ import annotations
import torch
from kda._fla.ops.kda.gate import (
kda_gate_bwd,
kda_gate_chunk_cumsum,
kda_gate_fwd as _kda_gate_fwd,
)
DEFAULT_LOWER_BOUND = -5.0
def kda_gate_fwd(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = DEFAULT_LOWER_BOUND,
):
return _kda_gate_fwd(
g,
A_log=A_log,
dt_bias=dt_bias,
lower_bound=lower_bound,
output_dtype=g.dtype,
)
__all__ = [
"DEFAULT_LOWER_BOUND",
"kda_gate_bwd",
"kda_gate_chunk_cumsum",
"kda_gate_fwd",
]
+5
View File
@@ -0,0 +1,5 @@
"""Vendored FLA WY recompute used by the chunk KDA backward."""
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
__all__ = ["recompute_w_u_fwd"]