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