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:
@@ -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"]
|
||||
Reference in New Issue
Block a user