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,34 @@
|
||||
"""L1: gradcheck for naive_recurrent_kda.
|
||||
|
||||
验证策略: torch.autograd.gradcheck 走 forward+backward 五个梯度.
|
||||
强制 dtype=float64; eps=1e-6, atol=1e-4.
|
||||
|
||||
shape (small):
|
||||
B=2, T=8, H=2, HV=4, K=4, V=4
|
||||
"""
|
||||
import torch
|
||||
|
||||
from kda.ops.reference.recurrent import naive_kda
|
||||
|
||||
|
||||
def test_gradcheck():
|
||||
B, T, H, HV, K, V = 2, 8, 2, 4, 4, 4
|
||||
torch.manual_seed(0)
|
||||
|
||||
# 所有输入都需要 requires_grad=True
|
||||
q = torch.randn(B, T, H, K, dtype=torch.float64, requires_grad=True)
|
||||
k = torch.randn(B, T, H, K, dtype=torch.float64, requires_grad=True)
|
||||
v = torch.randn(B, T, HV, V, dtype=torch.float64, requires_grad=True)
|
||||
g = torch.randn(B, T, HV, K, dtype=torch.float64, requires_grad=True) * 0.1
|
||||
beta = torch.rand(B, T, HV, dtype=torch.float64, requires_grad=True)
|
||||
|
||||
assert torch.autograd.gradcheck(
|
||||
lambda q, k, v, g, b: naive_kda(q, k, v, g, b, output_final_state=True),
|
||||
(q, k, v, g, beta),
|
||||
eps=1e-6, atol=1e-4, rtol=1e-3,
|
||||
), "L1 gradcheck 失败"
|
||||
print("L1 gradcheck: PASSED")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_gradcheck()
|
||||
Reference in New Issue
Block a user