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

194 lines
6.1 KiB
Python

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