"""TensorLens 集成测试: KDA 模型张量 -> trace -> 全局 store -> Flask 端点. 依赖: tensorlens (未安装时整个模块 skip, 用 `pip install tensorlens` 启用). 测三层: 1. trace — KDA 前向的真实 1D/2D/3D 张量 (embed/block 输出/logits/权重) 规范化成 int8 存入 tensorlens 全局 store 2. normalize — 四种策略 (clip/minmax/zscore/none) 的边界行为 3. HTTP — 用 Flask test_client 验证 /api/list_tensors 与 /api/get_tensor, 不启动阻塞的 gunicorn server (viewer() 为交互式入口, 不做自动化) """ import numpy as np import pytest import torch pytest.importorskip("tensorlens") from tensorlens.core import global_store from tensorlens.tensorlens import normalize_to_int8, trace from tensorlens.web.server import app 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() @pytest.fixture(autouse=True) def _clean_store(): """global_store 是模块级单例, 每个测试前后清空, 避免互相污染.""" global_store.INMEMORY_TENSORS.clear() yield global_store.INMEMORY_TENSORS.clear() def _trace_kda_tensors(model): """前向一次, 把 KDA 的 1D/2D/3D 张量全部 trace 进 store, 返回 logits numpy.""" x = torch.tensor([[1, 2, 3, 4]]) hidden = {} def hook_fn(name): def hook(module, inp, out): hidden[name] = out.detach() return hook model.embedding.register_forward_hook(hook_fn("embed")) model.blocks[0].register_forward_hook(hook_fn("block0")) with torch.no_grad(): logits = model(x) logits_np = logits.detach().numpy() # [1, T, vocab] 3D trace("lm_head.weight", model.lm_head.weight.detach().numpy()) # [vocab, hidden] 2D trace("embed", hidden["embed"].numpy()) # [1, T, hidden] 3D trace("block0.out", hidden["block0"].numpy()) # [1, T, hidden] 3D trace("logits", logits_np) # [1, T, vocab] 3D trace("logits.row0", logits_np[0, 0]) # [vocab] 1D return logits_np # --------------------------------------------------------------------------- # Layer 1: trace — KDA 张量进 store # --------------------------------------------------------------------------- def test_trace_stores_int8_with_expected_shape(): model = _model() logits_np = _trace_kda_tensors(model) store = global_store.INMEMORY_TENSORS assert set(store) == {"lm_head.weight", "embed", "block0.out", "logits", "logits.row0"} assert store["logits"].dtype == np.int8 assert store["logits"].shape == logits_np.shape assert store["lm_head.weight"].shape == (32, 16) assert store["logits.row0"].ndim == 1 assert store["logits"].min() >= -128 and store["logits"].max() <= 127 # --------------------------------------------------------------------------- # Layer 2: normalize_to_int8 — 四种策略边界行为 # --------------------------------------------------------------------------- def test_clip_normalization_bounds(): t = np.array([[-100.0, 0.0, 100.0]]) out = normalize_to_int8(t, (-1.0, 1.0), "clip") assert out.dtype == np.int8 assert out[0, 0] == -127 and out[0, 2] == 127 assert out[0, 1] == 0 def test_minmax_constant_tensor_returns_zeros(): t = np.full((2, 3), 0.5) out = normalize_to_int8(t, (-1.0, 1.0), "minmax") assert (out == 0).all() def test_zscore_zero_std_returns_zeros(): t = np.ones((4, 4)) out = normalize_to_int8(t, (-1.0, 1.0), "zscore") assert (out == 0).all() def test_none_strategy_scales_by_127(): t = np.array([[0.5, -0.5]]) out = normalize_to_int8(t, (-1.0, 1.0), "none") assert out[0, 0] == 63 and out[0, 1] == -63 # int8 cast 向零截断 def test_unsupported_normalization_raises(): with pytest.raises(ValueError): normalize_to_int8(np.zeros(3), (-1.0, 1.0), "bogus") # --------------------------------------------------------------------------- # Layer 2.5: trace 输入校验 # --------------------------------------------------------------------------- def test_trace_rejects_non_ndarray(): with pytest.raises(TypeError): trace("bad", torch.zeros(3)) def test_trace_rejects_empty_key(): with pytest.raises(ValueError): trace("", np.zeros(3)) # --------------------------------------------------------------------------- # Layer 3: Flask 端点 (test_client, 不起真实 server) # --------------------------------------------------------------------------- def test_list_tensors_endpoint_reports_kda_tensors(): model = _model() _trace_kda_tensors(model) resp = app.test_client().get("/api/list_tensors") assert resp.status_code == 200 body = resp.get_json() assert body["count"] == 5 keys = [t["key"] for t in body["available_tensors"]] assert "logits" in keys and "lm_head.weight" in keys def test_get_tensor_endpoint_returns_data(): model = _model() logits_np = _trace_kda_tensors(model) resp = app.test_client().get("/api/get_tensor?tensor_key=logits") assert resp.status_code == 200 body = resp.get_json() assert body["shape"] == list(logits_np.shape) assert len(body["data"]) == logits_np.shape[0] def test_get_tensor_missing_key_returns_400(): model = _model() _trace_kda_tensors(model) resp = app.test_client().get("/api/get_tensor") assert resp.status_code == 400 assert "tensor_key" in resp.get_json()["error"] def test_get_tensor_unknown_key_returns_404(): resp = app.test_client().get("/api/get_tensor?tensor_key=nope") assert resp.status_code == 404 def test_config_endpoint(): resp = app.test_client().get("/api/config") assert resp.status_code == 200 assert resp.get_json()["status"] == "ok" if __name__ == "__main__": import sys sys.exit(pytest.main([__file__, "-v"]))