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,83 @@
|
||||
"""TorchLens 集成测试: 计算图展开 KDA 模型并提取逐层激活.
|
||||
|
||||
依赖: torchlens (未安装时整个模块 skip, 用 `pip install torchlens` 启用).
|
||||
|
||||
验证三点:
|
||||
1. trace 能展开 KDA 模型 —— 关键子模块 (embedding / attention 各投影 /
|
||||
block 输出 / norm / lm_head 输出) 的激活被捕获且形状正确
|
||||
2. 展开不改变模型行为 —— trace 记录的输出与直接 forward 严格一致
|
||||
3. extract 便捷接口 —— 按模块名批量取激活
|
||||
|
||||
torchlens 2.34 的 trace[key] 返回 Op 对象, 取原始 tensor 用 `.tensor`.
|
||||
"""
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
pytest.importorskip("torchlens")
|
||||
|
||||
import torchlens as tl
|
||||
|
||||
from kda.models.causal_lm import CausalLM
|
||||
from kda.models.config import KDAConfig
|
||||
|
||||
|
||||
def _model():
|
||||
torch.manual_seed(51)
|
||||
config = KDAConfig(
|
||||
hidden_size=16,
|
||||
num_hidden_layers=1,
|
||||
num_heads=2,
|
||||
num_value_heads=2,
|
||||
head_dim=4,
|
||||
chunk_size=4,
|
||||
vocab_size=32,
|
||||
intermediate_size=32,
|
||||
kda_backend="reference",
|
||||
)
|
||||
return CausalLM(config).eval()
|
||||
|
||||
|
||||
def test_trace_captures_kda_submodule_activations():
|
||||
model = _model()
|
||||
x = torch.tensor([[1, 2, 3, 4]])
|
||||
with torch.no_grad():
|
||||
trace = tl.trace(model, x, capture=tl.options.CaptureOptions(verbose=False))
|
||||
|
||||
# 模块激活被捕获, 形状正确
|
||||
assert tuple(trace["embedding"].tensor.shape) == (1, 4, 16)
|
||||
assert tuple(trace["blocks.0.attn.q_proj"].tensor.shape) == (1, 4, 8) # H*K = 2*4
|
||||
assert tuple(trace["blocks.0.attn.v_proj"].tensor.shape) == (1, 4, 8) # HV*V = 2*4
|
||||
assert tuple(trace["blocks.0.attn"].tensor.shape) == (1, 4, 16)
|
||||
assert tuple(trace["blocks.0.ffn"].tensor.shape) == (1, 4, 16)
|
||||
assert tuple(trace["norm"].tensor.shape) == (1, 4, 16)
|
||||
assert tuple(trace["output"].tensor.shape) == (1, 4, 32) # vocab_size
|
||||
|
||||
|
||||
def test_trace_does_not_change_model_behavior():
|
||||
model = _model()
|
||||
x = torch.tensor([[1, 2, 3, 4]])
|
||||
with torch.no_grad():
|
||||
trace = tl.trace(model, x, capture=tl.options.CaptureOptions(verbose=False))
|
||||
traced_logits = trace["output"].tensor
|
||||
direct_logits = model(x)
|
||||
torch.testing.assert_close(traced_logits, direct_logits)
|
||||
|
||||
|
||||
def test_extract_returns_activations_by_module_name():
|
||||
model = _model()
|
||||
x = torch.tensor([[1, 2, 3, 4]])
|
||||
with torch.no_grad():
|
||||
acts = tl.extract(
|
||||
model, x, ["embedding", "blocks.0.attn.q_proj", "blocks.0", "output"]
|
||||
)
|
||||
assert set(acts) == {"embedding", "blocks.0.attn.q_proj", "blocks.0", "output"}
|
||||
assert tuple(acts["embedding"].shape) == (1, 4, 16)
|
||||
assert tuple(acts["blocks.0.attn.q_proj"].shape) == (1, 4, 8)
|
||||
assert tuple(acts["blocks.0"].shape) == (1, 4, 16)
|
||||
assert tuple(acts["output"].shape) == (1, 4, 32)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import sys
|
||||
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
Reference in New Issue
Block a user