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:
+167
@@ -0,0 +1,167 @@
|
||||
"""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),
|
||||
)
|
||||
Reference in New Issue
Block a user