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