Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
194 lines
6.1 KiB
Python
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"]))
|