"""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] = `` (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 r_j # = v_i - . 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