Files
K3/inspect_tensors.py
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

125 lines
3.6 KiB
Python

"""Inspect KDA activations: print stats, TorchLens extract, or TensorLens web UI.
uv run python inspect_tensors.py # shape / min / max / nan
uv run python inspect_tensors.py --lens torch # named activations
uv run python inspect_tensors.py --lens web # http://127.0.0.1:8000
"""
from __future__ import annotations
import argparse
import torch
from kda.models.causal_lm import CausalLM
from kda.models.config import KDAConfig
MODULES = (
"embedding",
"blocks.0.attn.q_proj",
"blocks.0.attn.k_proj",
"blocks.0.attn.v_proj",
"blocks.0.attn",
"blocks.0.ffn",
"blocks.0",
"norm",
)
def _model() -> CausalLM:
torch.manual_seed(51)
return CausalLM(
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",
)
).eval()
def _tokens() -> torch.Tensor:
return torch.tensor([[1, 2, 3, 4]])
def _named_modules(model: torch.nn.Module) -> dict[str, torch.nn.Module]:
return dict(model.named_modules())
def capture_activations(model: CausalLM, x: torch.Tensor) -> dict[str, torch.Tensor]:
captured: dict[str, torch.Tensor] = {}
hooks = []
modules = _named_modules(model)
for name in MODULES:
module = modules[name]
def _hook(_module, _inp, out, key=name):
captured[key] = out.detach()
hooks.append(module.register_forward_hook(_hook))
with torch.no_grad():
captured["logits"] = model(x).detach()
for hook in hooks:
hook.remove()
return captured
def print_stats(tensors: dict[str, torch.Tensor]) -> None:
print(f"{'name':28} {'shape':18} {'dtype':10} {'min':>10} {'max':>10} {'mean':>10} nan/inf")
for name, tensor in tensors.items():
finite = torch.isfinite(tensor)
n_bad = int((~finite).sum())
stats = tensor.float() if tensor.is_floating_point() else tensor
print(
f"{name:28} {str(tuple(tensor.shape)):18} {str(tensor.dtype):10} "
f"{stats.min().item():10.4f} {stats.max().item():10.4f} "
f"{stats.float().mean().item():10.4f} {n_bad}"
)
def inspect_torchlens(model: CausalLM, x: torch.Tensor) -> None:
import torchlens as tl
names = [*MODULES, "output"]
with torch.no_grad():
acts = tl.extract(model, x, names)
print_stats(acts)
def inspect_web(model: CausalLM, x: torch.Tensor, host: str, port: int) -> None:
from tensorlens.tensorlens import trace
from tensorlens.web.server import app
acts = capture_activations(model, x)
for name, tensor in acts.items():
trace(name, tensor.cpu().float().numpy(), normalization="minmax")
trace("lm_head.weight", model.lm_head.weight.detach().cpu().float().numpy(), normalization="minmax")
print(f"TensorLens: http://{host}:{port} (Ctrl-C to stop)")
app.run(host=host, port=port, debug=False, use_reloader=False)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--lens", choices=("print", "torch", "web"), default="print")
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=8000)
args = parser.parse_args()
model = _model()
x = _tokens()
if args.lens == "torch":
inspect_torchlens(model, x)
return
if args.lens == "web":
inspect_web(model, x, args.host, args.port)
return
print_stats(capture_activations(model, x))
if __name__ == "__main__":
main()