Files
K3/kda/ops/api.py
T
dela 584f7e9e73 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.
2026-08-25 14:43:17 +08:00

168 lines
5.5 KiB
Python

"""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),
)