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:
@@ -0,0 +1,107 @@
|
||||
"""Public backend selection must be explicit and reproducible."""
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from kda.models.config import KDAConfig
|
||||
from kda.ops.api import chunk_kda
|
||||
|
||||
|
||||
def _inputs():
|
||||
torch.manual_seed(41)
|
||||
B, T, H, HV, K, V = 1, 4, 1, 2, 2, 2
|
||||
return (
|
||||
torch.randn(B, T, H, K),
|
||||
torch.randn(B, T, H, K),
|
||||
torch.randn(B, T, HV, V),
|
||||
-torch.rand(B, T, HV, K),
|
||||
torch.rand(B, T, HV),
|
||||
)
|
||||
|
||||
|
||||
def test_default_config_enables_kernel_side_transforms():
|
||||
config = KDAConfig()
|
||||
assert config.use_gate_in_kernel
|
||||
assert config.use_qk_l2norm_in_kernel
|
||||
assert config.use_beta_sigmoid_in_kernel
|
||||
assert config.lower_bound == -5.0
|
||||
|
||||
|
||||
def test_default_backend_is_local_reference():
|
||||
inputs = _inputs()
|
||||
default_output, _ = chunk_kda(*inputs, chunk_size=4, use_qk_l2norm_in_kernel=True)
|
||||
reference_output, _ = chunk_kda(*inputs, chunk_size=4, backend="reference", use_qk_l2norm_in_kernel=True)
|
||||
torch.testing.assert_close(default_output, reference_output)
|
||||
assert KDAConfig().kda_backend == "reference"
|
||||
|
||||
|
||||
def test_auto_is_deprecated_reference_alias():
|
||||
inputs = _inputs()
|
||||
with pytest.warns(DeprecationWarning, match="selects the local reference backend"):
|
||||
auto_output, _ = chunk_kda(*inputs, chunk_size=4, backend="auto", use_qk_l2norm_in_kernel=True)
|
||||
reference_output, _ = chunk_kda(*inputs, chunk_size=4, backend="reference", use_qk_l2norm_in_kernel=True)
|
||||
torch.testing.assert_close(auto_output, reference_output)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
|
||||
def test_triton_backend_matches_reference():
|
||||
torch.manual_seed(41)
|
||||
B, T, H, HV, K, V = 2, 64, 4, 8, 32, 32
|
||||
device = "cuda"
|
||||
q = torch.randn(B, T, H, K, device=device)
|
||||
k = torch.randn(B, T, H, K, device=device)
|
||||
v = torch.randn(B, T, HV, V, device=device)
|
||||
g = -torch.rand(B, T, HV, K, device=device)
|
||||
beta = torch.rand(B, T, HV, device=device)
|
||||
triton_out, _ = chunk_kda(
|
||||
q, k, v, g, beta, chunk_size=64, backend="triton", use_qk_l2norm_in_kernel=True
|
||||
)
|
||||
reference_out, _ = chunk_kda(
|
||||
q, k, v, g, beta, chunk_size=64, backend="reference", use_qk_l2norm_in_kernel=True
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
triton_out.float(), reference_out.float(), atol=2e-3, rtol=2e-3
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
|
||||
def test_unnormalized_parity_is_not_meaningful():
|
||||
"""对拍必须用与训练一致的 q/k L2-norm (守卫测试).
|
||||
|
||||
M = I + tril(A_kk*beta, -1) 是单位下三角, 行列式恒为 1、特征值全为 1,
|
||||
不存在"失去对角占优"一说. 真正的机制是 M = I + N 中 N 严格下三角 (幂零),
|
||||
故 M^-1 = sum_{m<C} (-N)^m, ‖M^-1‖ ~ ‖N‖^C —— 关于 chunk 长度指数增长.
|
||||
L2-norm 的作用是把 ‖N‖ 压到 1 以下让这个 Neumann 级数收敛.
|
||||
|
||||
不归一化时 ‖k‖>1, 精确解本身就以 ‖k‖^chunk_size 增长 (实测输出 1e31 量级
|
||||
乃至 inf). 这不是某个后端的缺陷: fp64 逐步递推 (完全不含三角求解) 与分块
|
||||
形式吻合到 8e-15, 说明爆炸就是该递推的真值. 两个后端的相对误差其实很接近
|
||||
(1.3e-3 → 2.6e-2), 但绝对差随输出量级一起走, 任何一致性容差都必然失败.
|
||||
"""
|
||||
torch.manual_seed(41)
|
||||
B, T, H, HV, K, V = 2, 64, 4, 8, 32, 32
|
||||
device = "cuda"
|
||||
q = torch.randn(B, T, H, K, device=device) # 故意不归一化
|
||||
k = torch.randn(B, T, H, K, device=device)
|
||||
v = torch.randn(B, T, HV, V, device=device)
|
||||
g = -torch.rand(B, T, HV, K, device=device)
|
||||
beta = torch.rand(B, T, HV, device=device)
|
||||
|
||||
with pytest.warns(RuntimeWarning, match="use_qk_l2norm_in_kernel=False"):
|
||||
triton_out, _ = chunk_kda(q, k, v, g, beta, chunk_size=64, backend="triton")
|
||||
reference_out, _ = chunk_kda(q, k, v, g, beta, chunk_size=64, backend="reference")
|
||||
|
||||
diff = (triton_out.float() - reference_out.float()).abs().max()
|
||||
assert not torch.isfinite(diff).item() or diff.item() > 1e-2, (
|
||||
f"unnormalised parity diff {diff.item():.3e} should be pathological"
|
||||
)
|
||||
|
||||
|
||||
def test_triton_backend_does_not_import_upstream_fla():
|
||||
import sys
|
||||
|
||||
for name in list(sys.modules):
|
||||
if name == "fla" or name.startswith("fla."):
|
||||
sys.modules.pop(name)
|
||||
from kda.ops.triton.chunk import chunk_kda as _triton_chunk_kda # noqa: F401
|
||||
|
||||
assert not any(n == "fla" or n.startswith("fla.") for n in sys.modules)
|
||||
Reference in New Issue
Block a user