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
+155
View File
@@ -0,0 +1,155 @@
"""Pure-PyTorch chunked reference implementation of KDA."""
from __future__ import annotations
import math
import warnings
import torch
from einops import rearrange
#: ``exp`` overflows past this exponent in fp32 and bf16 (both top out at 3.4e38).
_EXP_LIMIT = math.log(torch.finfo(torch.float32).max)
#: Row-block size for :func:`_decayed_dot`.
#:
#: The g_ref GEMM exponentiates the gate span between the reference row and the
#: rows/columns it covers, so the block size caps that exponent at
#: ``DECAY_BLOCK * max|g|``. With the default ``lower_bound=-5`` gate that is
#: ``16 * 5 = 80 < ln(3.4e38) = 88.7``, i.e. fp32/bf16-safe for any chunk size.
#: Referencing a whole 64-row chunk instead would allow ``64 * 5 = 320`` and
#: overflow to NaN once the gate saturates.
DECAY_BLOCK = 16
def _decayed_dot(x: torch.Tensor, k: torch.Tensor, g: torch.Tensor) -> torch.Tensor:
"""Return ``A[..., i, j] = <x_i, exp(g_i-g_j) * k_j>`` (FLA g_ref GEMM).
Only the causal part (``j <= i``) is exact; callers mask the rest, which is
left at zero. Rows are processed in blocks of :data:`DECAY_BLOCK` against
the block's own first row, which is what bounds the exponent: for a row
block starting at ``r``, ``exp(g_i - g_ref)`` spans at most ``DECAY_BLOCK``
steps, and ``exp(g_ref - g_j)`` is ``<= 1`` for ``j < r`` and likewise spans
at most ``DECAY_BLOCK`` steps for ``j >= r``.
"""
C = g.shape[-2]
out = g.new_zeros(*g.shape[:-1], C)
for r in range(0, C, DECAY_BLOCK):
end = min(r + DECAY_BLOCK, C)
g_ref = g[..., r : r + 1, :]
rows = x[..., r:end, :] * (g[..., r:end, :] - g_ref).exp()
cols = k[..., :end, :] * (g_ref - g[..., :end, :]).exp()
out[..., r:end, :end] = rows @ cols.transpose(-1, -2)
return out
#: Whether :func:`naive_chunk_kda` checks the gate span against the ``exp``
#: budget. The check costs one device sync per call; set it to ``False`` if that
#: matters more than diagnosing a NaN.
CHECK_DECAY_SPAN = True
def _max_decay_span(g_cumsum: torch.Tensor) -> torch.Tensor:
"""Largest ``|g_ref - g_j|`` any row block will exponentiate."""
C = g_cumsum.shape[-2]
if C % DECAY_BLOCK == 0:
blocks = g_cumsum.unflatten(-2, (C // DECAY_BLOCK, DECAY_BLOCK))
return (blocks[..., :1, :] - blocks).abs().amax()
return torch.stack(
[
(g_cumsum[..., r : r + 1, :] - g_cumsum[..., r : r + DECAY_BLOCK, :])
.abs()
.amax()
for r in range(0, C, DECAY_BLOCK)
]
).amax()
def _warn_if_decay_span_overflows(g_cumsum: torch.Tensor) -> None:
"""Warn when a row block's gate span is about to overflow ``exp``.
``DECAY_BLOCK`` bounds this for the default ``safe_gate`` path, but an
unbounded gate (``-A.exp() * softplus(x)``) can still exceed it.
"""
span = _max_decay_span(g_cumsum).item()
if span > _EXP_LIMIT:
warnings.warn(
f"gate span within a {DECAY_BLOCK}-row block is {span:.1f} > "
f"{_EXP_LIMIT:.1f}; exp() will overflow to inf and the output will "
"be NaN. Reduce the gate magnitude (e.g. safe_gate with a smaller "
"|lower_bound|) or use backend='triton'.",
RuntimeWarning,
stacklevel=3,
)
def naive_chunk_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,
chunk_size: int = 64,
):
"""Chunk-parallel, inter-chunk recurrent KDA reference.
Shapes are ``q/k: [B,T,H,K]``, ``v: [B,T,HV,V]``,
``g: [B,T,HV,K]`` and ``beta: [B,T,HV]``.
"""
dtype = v.dtype
B, T, H, K = q.shape
HV, V = v.shape[2], v.shape[-1]
C = chunk_size
assert HV % H == 0, f"HV={HV} must be divisible by H={H}"
assert T % C == 0, f"T={T} must be divisible by chunk_size={C}"
scale = K**-0.5 if scale is None else scale
q, k = [
rearrange(x, "b (n c) h d -> b h n c d", c=C)
.repeat_interleave(HV // H, dim=1)
for x in (q, k)
]
v, g = [rearrange(x, "b (n c) h d -> b h n c d", c=C) for x in (v, g)]
beta = rearrange(beta, "b (n c) h -> b h n c", c=C)
q = q * scale
g = g.cumsum(dim=-2)
if CHECK_DECAY_SPAN:
_warn_if_decay_span_overflows(g)
# r_i + sum_{j<i} beta_j <k_i, exp(g_i-g_j)k_j> r_j
# = v_i - <exp(g_i)k_i, S_start>.
mask_upper = torch.triu(torch.ones(C, C, dtype=torch.bool, device=q.device))
mask_strict_upper = torch.triu(mask_upper, diagonal=1)
eye = torch.eye(C, dtype=q.dtype, device=q.device)
A_kk = _decayed_dot(k, k, g)
M = eye + (A_kk * beta[..., None, :]).masked_fill(mask_upper, 0)
W = torch.linalg.solve_triangular(M, g.exp() * k, upper=False)
U = torch.linalg.solve_triangular(M, v, upper=False)
# Output includes the current token, hence the diagonal is retained.
A_qk = (_decayed_dot(q, k, g) * beta[..., None, :]).masked_fill(mask_strict_upper, 0)
S = q.new_zeros(B, HV, K, V)
if initial_state is not None:
S = S + initial_state
o = v.new_empty(B, HV, T // C, C, V)
for n in range(T // C):
q_n, k_n, g_n = q[:, :, n], k[:, :, n], g[:, :, n]
r = U[:, :, n] - W[:, :, n] @ S
o[:, :, n] = (q_n * g_n.exp()) @ S + A_qk[:, :, n] @ r
decay = (g_n[:, :, -1:, :] - g_n).exp()
S = S * g_n[:, :, -1, :, None].exp()
S = S + (decay * k_n).transpose(-1, -2) @ (r * beta[:, :, n, :, None])
if not output_final_state:
S = None
return rearrange(o, "b h n c d -> b (n c) h d").to(dtype), S
# Backward-compatible name used by earlier notes/scripts.
naive_chunk_kda_fwd = naive_chunk_kda