Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
156 lines
5.7 KiB
Python
156 lines
5.7 KiB
Python
"""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
|