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