Files
K3/tests/integration/test_torchlens.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

84 lines
2.8 KiB
Python

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