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,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
|
||||
Reference in New Issue
Block a user