"""Training-facing KDA operator with the same boundary as FLA's ``chunk_kda``.""" from __future__ import annotations import warnings from functools import lru_cache import torch import torch.nn.functional as F from .reference.chunkwise import DECAY_BLOCK, _EXP_LIMIT, naive_chunk_kda @lru_cache(maxsize=1) def _fla_chunk_kda(): try: from fla.ops.kda import chunk_kda except ImportError: return None return chunk_kda def _reference_chunk_size(T: int, requested: int) -> int: size = min(T, requested) while T % size: size -= 1 return size def chunk_kda( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, g: torch.Tensor, beta: torch.Tensor, *, A_log: torch.Tensor | None = None, dt_bias: torch.Tensor | None = None, scale: float | None = None, initial_state: torch.Tensor | None = None, output_final_state: bool = False, use_qk_l2norm_in_kernel: bool = False, use_gate_in_kernel: bool = False, use_beta_sigmoid_in_kernel: bool = False, safe_gate: bool = False, lower_bound: float | None = None, chunk_size: int = 64, backend: str = "reference", ): """Run an explicitly selected KDA implementation. ``reference`` and its legacy alias ``torch`` use this repository's differentiable PyTorch implementation. ``triton`` uses the vendored FLA NVIDIA Triton kernels in ``kda._fla`` (chunk_size 32 or 64, CUDA). ``fla`` is reserved for explicit upstream parity runs. """ supported = {"reference", "triton", "fla", "torch", "auto"} if backend not in supported: raise ValueError(f"backend must be one of {sorted(supported)}") if not use_qk_l2norm_in_kernel: # Backend-independent: this is a property of the recurrence, not of any # one implementation. warnings.warn( "use_qk_l2norm_in_kernel=False: KDA's chunkwise form assumes " "||k||=1 so that I + tril(A_kk*beta) has a convergent Neumann " "series. Unnormalised k makes the exact output grow like " "||k||^chunk_size and can reach inf on any backend.", RuntimeWarning, stacklevel=2, ) if backend == "auto": warnings.warn( "backend='auto' is deprecated and now selects the local reference backend; " "use backend='fla' explicitly for upstream FLA", DeprecationWarning, stacklevel=2, ) backend = "reference" if backend == "torch": backend = "reference" if backend == "triton": from .triton.chunk import chunk_kda as triton_chunk_kda fla_chunk = 32 if chunk_size <= 32 else 64 return triton_chunk_kda( q, k, v, g, beta, scale=scale, initial_state=initial_state, output_final_state=output_final_state, chunk_size=fla_chunk, use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, use_gate_in_kernel=use_gate_in_kernel, use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel, A_log=A_log, dt_bias=dt_bias, safe_gate=safe_gate, lower_bound=lower_bound, ) if backend == "fla": fused_op = _fla_chunk_kda() if fused_op is None: raise RuntimeError( "backend='fla' requires a complete flash-linear-attention installation" ) return fused_op( q, k, v, g, beta, A_log=A_log, dt_bias=dt_bias, scale=scale, initial_state=initial_state, output_final_state=output_final_state, use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, use_gate_in_kernel=use_gate_in_kernel, use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel, safe_gate=safe_gate, lower_bound=lower_bound, chunk_size=32 if chunk_size <= 32 else 64, ) if safe_gate and lower_bound is not None: # _decayed_dot exponentiates at most DECAY_BLOCK steps of gate decay, # and safe_gate bounds each step by |lower_bound|. budget = DECAY_BLOCK * abs(lower_bound) if budget > _EXP_LIMIT: raise ValueError( f"lower_bound={lower_bound} allows a gate span of {budget:.1f} " f"per {DECAY_BLOCK}-row block, which overflows exp() " f"(limit {_EXP_LIMIT:.1f}) and yields NaN. Use " f"|lower_bound| < {_EXP_LIMIT / DECAY_BLOCK:.2f} or " "backend='triton'." ) if use_qk_l2norm_in_kernel: q, k = F.normalize(q, dim=-1), F.normalize(k, dim=-1) if use_beta_sigmoid_in_kernel: beta = beta.sigmoid() if use_gate_in_kernel: if A_log is None: raise ValueError("A_log is required when use_gate_in_kernel=True") bias = 0 if dt_bias is None else dt_bias.view(g.shape[-2:]) gate_input = g + bias rate = A_log.exp().view(1, 1, -1, 1) if safe_gate: if lower_bound is None: raise ValueError("lower_bound is required when safe_gate=True") g = lower_bound * torch.sigmoid(rate * gate_input) else: g = -rate * F.softplus(gate_input) return naive_chunk_kda( q, k, v, g, beta, scale=scale, initial_state=initial_state, output_final_state=output_final_state, chunk_size=_reference_chunk_size(q.shape[1], chunk_size), )