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
+47
View File
@@ -0,0 +1,47 @@
"""PyTorch references for the two KDA gate activations."""
from __future__ import annotations
import torch
import torch.nn.functional as F
def kda_gate_reference(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
*,
safe_gate: bool = False,
lower_bound: float | None = None,
) -> torch.Tensor:
"""Compute the official KDA gate semantics in PyTorch.
``A_log`` is head-wise with shape ``[HV]`` and ``dt_bias`` is
per-dimension with shape ``[HV, K]`` (or flattened to ``[HV*K]``).
"""
HV, K = g.shape[-2:]
gate_input = g if dt_bias is None else g + dt_bias.view(HV, K)
rate = A_log.view(HV, 1).exp()
if safe_gate:
if lower_bound is None:
raise ValueError("lower_bound is required when safe_gate=True")
return lower_bound * torch.sigmoid(rate * gate_input)
return -rate * F.softplus(gate_input)
def kda_gate_naive(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
) -> torch.Tensor:
"""Compatibility name matching FLA's reference gate convention."""
return kda_gate_reference(
g,
A_log,
dt_bias,
safe_gate=lower_bound is not None,
lower_bound=lower_bound,
)
__all__ = ["kda_gate_naive", "kda_gate_reference"]