AlexWortega commited on
Commit
954a618
·
verified ·
1 Parent(s): 0718aaa

initial commit: transformer x-ray (FastAPI + React, GIF export)

Browse files
.dockerignore ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ **/__pycache__
2
+ **/.venv
3
+ **/.env
4
+ **/node_modules
5
+ frontend/dist
6
+ static
7
+ .git
8
+ .DS_Store
9
+ *.log
10
+ hf_cache
.gitignore ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ __pycache__/
2
+ *.pyc
3
+ .venv/
4
+ venv/
5
+ .env
6
+ node_modules/
7
+ frontend/dist/
8
+ static/
9
+ .DS_Store
10
+ .idea/
11
+ .vscode/
12
+ *.log
13
+ .hf/
14
+ hf_cache/
Dockerfile ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ---- frontend build ----
2
+ FROM node:20-alpine AS frontend
3
+ WORKDIR /fe
4
+ COPY frontend/package.json frontend/package-lock.json* ./
5
+ RUN npm install --no-audit --no-fund
6
+ COPY frontend/ ./
7
+ RUN npm run build
8
+
9
+ # ---- runtime ----
10
+ FROM python:3.11-slim
11
+ ENV PYTHONUNBUFFERED=1 \
12
+ PIP_NO_CACHE_DIR=1 \
13
+ HF_HOME=/tmp/hf \
14
+ TRANSFORMERS_CACHE=/tmp/hf \
15
+ MAX_PARAMS=1500000000
16
+
17
+ RUN apt-get update && apt-get install -y --no-install-recommends \
18
+ git build-essential && rm -rf /var/lib/apt/lists/*
19
+
20
+ WORKDIR /app
21
+ COPY requirements.txt .
22
+ RUN pip install -r requirements.txt
23
+
24
+ COPY backend/ ./backend/
25
+ COPY --from=frontend /fe/dist ./static/
26
+
27
+ EXPOSE 7860
28
+ CMD ["uvicorn", "backend.main:app", "--host", "0.0.0.0", "--port", "7860"]
README.md CHANGED
@@ -1,10 +1,43 @@
1
  ---
2
- title: Transformer Xray
3
- emoji: 🏢
4
- colorFrom: pink
5
- colorTo: red
6
  sdk: docker
 
7
  pinned: false
 
 
8
  ---
9
 
10
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: Transformer X-Ray
3
+ emoji: 🔬
4
+ colorFrom: indigo
5
+ colorTo: purple
6
  sdk: docker
7
+ app_port: 7860
8
  pinned: false
9
+ license: mit
10
+ short_description: Visualize any HF model layer-by-layer, export GIF
11
  ---
12
 
13
+ # Transformer X-Ray
14
+
15
+ Interactive architecture visualizer for any model on the Hugging Face Hub.
16
+
17
+ - Paste a model ID (e.g. `google/vit-base-patch16-224`, `bert-base-uncased`, `prajjwal1/bert-tiny`)
18
+ - See every `nn.Module` from `transformers` as an interactive graph
19
+ - Hit **Play** — real tensors from a real forward pass animate through the graph: tokens → embeddings → attention → MLP → logits → probabilities
20
+ - Click any module to see its config, parameter count, and the actual tensor that flowed through it
21
+ - Export the playthrough as a **GIF**
22
+
23
+ ## Architecture
24
+
25
+ - **Backend** — FastAPI + `transformers` + `torch` (CPU). Loads the model from the Hub, traces it with `torch.fx` (with a `named_modules` fallback), registers forward hooks on every module, runs a sample input, and streams per-module tensor snapshots over a WebSocket.
26
+ - **Frontend** — Vite + React + React Flow. Renders the graph, drives the play timeline, and assembles the GIF in-browser with `gif.js`.
27
+
28
+ ## Local dev
29
+
30
+ ```bash
31
+ # backend
32
+ pip install -r requirements.txt
33
+ uvicorn backend.main:app --reload --port 7860
34
+
35
+ # frontend (separate terminal)
36
+ cd frontend && npm install && npm run dev
37
+ ```
38
+
39
+ The Vite dev server proxies `/api` and `/ws` to the backend.
40
+
41
+ ## Deploy
42
+
43
+ Push to a Hugging Face Space with `sdk: docker`. The included `Dockerfile` builds the frontend and serves it from the FastAPI app on port 7860.
backend/__init__.py ADDED
File without changes
backend/main.py ADDED
@@ -0,0 +1,145 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import asyncio
4
+ import json
5
+ import logging
6
+ import os
7
+ import time
8
+ import uuid
9
+ from dataclasses import asdict
10
+ from pathlib import Path
11
+ from typing import Any
12
+
13
+ from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect
14
+ from fastapi.middleware.cors import CORSMiddleware
15
+ from fastapi.responses import JSONResponse
16
+ from fastapi.staticfiles import StaticFiles
17
+ from pydantic import BaseModel
18
+
19
+ from .model_loader import CACHE, load_model
20
+ from .runner import get_last_run, run_model, store_run
21
+ from .tracer import build_graph
22
+
23
+ logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s %(message)s")
24
+ log = logging.getLogger("main")
25
+
26
+ app = FastAPI(title="Transformer X-Ray")
27
+
28
+ # Permissive CORS for local Vite dev. The deployed Space serves both apps from one origin.
29
+ app.add_middleware(
30
+ CORSMiddleware,
31
+ allow_origins=["*"],
32
+ allow_credentials=False,
33
+ allow_methods=["*"],
34
+ allow_headers=["*"],
35
+ )
36
+
37
+
38
+ class LoadRequest(BaseModel):
39
+ model_id: str
40
+
41
+
42
+ class RunRequest(BaseModel):
43
+ prompt: str | None = None
44
+
45
+
46
+ def _config_summary(cfg: Any) -> dict:
47
+ """Pull common config fields. Falls back to to_dict() for everything else."""
48
+ full = cfg.to_dict() if hasattr(cfg, "to_dict") else {}
49
+ keys = (
50
+ "model_type", "architectures", "hidden_size", "num_hidden_layers",
51
+ "num_attention_heads", "num_key_value_heads", "intermediate_size",
52
+ "vocab_size", "max_position_embeddings", "image_size", "patch_size",
53
+ "num_channels", "id2label", "label2id", "torch_dtype", "tie_word_embeddings",
54
+ )
55
+ return {k: full.get(k) for k in keys if k in full}
56
+
57
+
58
+ @app.post("/api/load")
59
+ def api_load(req: LoadRequest):
60
+ try:
61
+ loaded = load_model(req.model_id)
62
+ except ValueError as e:
63
+ raise HTTPException(status_code=413, detail=str(e))
64
+ except Exception as e:
65
+ log.exception("load failed")
66
+ raise HTTPException(status_code=400, detail=f"failed to load: {e}")
67
+ tok = loaded.tokenizer
68
+ tokenizer_info = None
69
+ if tok is not None:
70
+ tokenizer_info = {
71
+ "name_or_path": getattr(tok, "name_or_path", None),
72
+ "vocab_size": getattr(tok, "vocab_size", None),
73
+ "model_max_length": getattr(tok, "model_max_length", None),
74
+ "special_tokens": list(getattr(tok, "all_special_tokens", []) or []),
75
+ "is_fast": getattr(tok, "is_fast", None),
76
+ }
77
+ return {
78
+ "model_id": loaded.model_id,
79
+ "modality": loaded.modality,
80
+ "head_kind": loaded.head_kind,
81
+ "param_count": loaded.param_count,
82
+ "config": _config_summary(loaded.config),
83
+ "tokenizer": tokenizer_info,
84
+ "processor": type(loaded.processor).__name__ if loaded.processor else None,
85
+ }
86
+
87
+
88
+ @app.get("/api/graph")
89
+ def api_graph():
90
+ loaded = CACHE.current()
91
+ if loaded is None:
92
+ raise HTTPException(status_code=400, detail="no model loaded")
93
+ return build_graph(loaded.model)
94
+
95
+
96
+ @app.post("/api/run")
97
+ def api_run(req: RunRequest):
98
+ loaded = CACHE.current()
99
+ if loaded is None:
100
+ raise HTTPException(status_code=400, detail="no model loaded")
101
+ t0 = time.time()
102
+ result = run_model(loaded, req.prompt)
103
+ elapsed = time.time() - t0
104
+ store_run(result)
105
+ run_id = uuid.uuid4().hex[:12]
106
+ return {"run_id": run_id, "steps": len(result.events), "elapsed_s": round(elapsed, 3)}
107
+
108
+
109
+ @app.websocket("/ws/run/{run_id}")
110
+ async def ws_run(ws: WebSocket, run_id: str):
111
+ await ws.accept()
112
+ result = get_last_run()
113
+ if result is None:
114
+ await ws.send_json({"error": "no run available"})
115
+ await ws.close()
116
+ return
117
+ try:
118
+ for ev in result.events:
119
+ await ws.send_json({
120
+ "step": ev.step,
121
+ "path": ev.path,
122
+ "module_class": ev.module_class,
123
+ "kind": ev.kind,
124
+ "payload": ev.payload,
125
+ })
126
+ await ws.send_json({"kind": "stream_end"})
127
+ except WebSocketDisconnect:
128
+ return
129
+
130
+
131
+ @app.get("/api/health")
132
+ def health():
133
+ return {"ok": True}
134
+
135
+
136
+ # ---- static frontend ----
137
+ STATIC_DIR = Path(os.environ.get("STATIC_DIR", "static"))
138
+ if STATIC_DIR.exists():
139
+ app.mount("/", StaticFiles(directory=str(STATIC_DIR), html=True), name="static")
140
+ else:
141
+ @app.get("/")
142
+ def root():
143
+ return JSONResponse(
144
+ {"ok": True, "note": "frontend not built — run `npm run build` in frontend/ or use the dev server"}
145
+ )
backend/modality.py ADDED
@@ -0,0 +1,110 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import io
4
+ from typing import Any
5
+
6
+ import torch
7
+ from PIL import Image
8
+ import numpy as np
9
+
10
+
11
+ VISION_MODEL_TYPES = {
12
+ "vit", "deit", "beit", "swin", "swinv2", "convnext", "convnextv2",
13
+ "dinov2", "dinov2_with_registers", "siglip", "clip", "mobilevit", "regnet",
14
+ "resnet", "efficientnet",
15
+ }
16
+
17
+ AUDIO_MODEL_TYPES = {
18
+ "wav2vec2", "hubert", "whisper", "audio-spectrogram-transformer", "ast",
19
+ "speech_to_text", "unispeech", "unispeech-sat", "data2vec-audio",
20
+ }
21
+
22
+
23
+ def _placeholder_image(size: int = 224) -> Image.Image:
24
+ """Generate a small synthetic image — a soft gradient with a circle.
25
+
26
+ Avoids bundling binary assets and keeps the Space self-contained."""
27
+ arr = np.zeros((size, size, 3), dtype=np.uint8)
28
+ yy, xx = np.mgrid[0:size, 0:size]
29
+ arr[..., 0] = (xx * 255 / size).astype(np.uint8)
30
+ arr[..., 1] = (yy * 255 / size).astype(np.uint8)
31
+ arr[..., 2] = 128
32
+ cy, cx = size // 2, size // 2
33
+ mask = (yy - cy) ** 2 + (xx - cx) ** 2 < (size // 4) ** 2
34
+ arr[mask] = [240, 240, 240]
35
+ return Image.fromarray(arr)
36
+
37
+
38
+ def _placeholder_audio(seconds: float = 1.0, sr: int = 16000) -> np.ndarray:
39
+ t = np.linspace(0, seconds, int(seconds * sr), endpoint=False)
40
+ return (0.1 * np.sin(2 * np.pi * 440 * t)).astype(np.float32)
41
+
42
+
43
+ def detect_modality(model_type: str | None) -> str:
44
+ if not model_type:
45
+ return "text"
46
+ mt = model_type.lower()
47
+ if mt in VISION_MODEL_TYPES:
48
+ return "vision"
49
+ if mt in AUDIO_MODEL_TYPES:
50
+ return "audio"
51
+ return "text"
52
+
53
+
54
+ def build_inputs(
55
+ modality: str,
56
+ tokenizer: Any,
57
+ processor: Any,
58
+ config: Any,
59
+ prompt: str | None,
60
+ ) -> tuple[dict, dict]:
61
+ """Returns (model_inputs, ui_payload).
62
+
63
+ ui_payload is the human-readable description of the input shown in the TokenStrip
64
+ (token IDs + decoded pieces, or image preview metadata)."""
65
+ if modality == "vision" and processor is not None:
66
+ img = _placeholder_image()
67
+ inputs = processor(images=img, return_tensors="pt")
68
+ ui = {
69
+ "kind": "image",
70
+ "size": list(img.size),
71
+ "shape": list(inputs["pixel_values"].shape),
72
+ "patch_size": getattr(config, "patch_size", None),
73
+ }
74
+ return dict(inputs), ui
75
+
76
+ if modality == "audio" and processor is not None:
77
+ sr = getattr(processor, "sampling_rate", 16000) or 16000
78
+ wav = _placeholder_audio(sr=sr)
79
+ try:
80
+ inputs = processor(wav, sampling_rate=sr, return_tensors="pt")
81
+ except TypeError:
82
+ inputs = processor(wav, return_tensors="pt")
83
+ first = next(iter(inputs.values()))
84
+ ui = {
85
+ "kind": "audio",
86
+ "sample_rate": sr,
87
+ "duration_s": float(len(wav) / sr),
88
+ "shape": list(first.shape),
89
+ }
90
+ return dict(inputs), ui
91
+
92
+ # text
93
+ text = prompt or "The quick brown fox jumps over the lazy dog."
94
+ enc = tokenizer(text, return_tensors="pt", truncation=True, max_length=64)
95
+ ids = enc["input_ids"][0].tolist()
96
+ pieces = []
97
+ for i in ids:
98
+ try:
99
+ pieces.append(tokenizer.decode([i], skip_special_tokens=False))
100
+ except Exception:
101
+ pieces.append(f"<{i}>")
102
+ ui = {
103
+ "kind": "text",
104
+ "text": text,
105
+ "ids": ids,
106
+ "pieces": pieces,
107
+ "vocab_size": getattr(tokenizer, "vocab_size", None),
108
+ "shape": list(enc["input_ids"].shape),
109
+ }
110
+ return dict(enc), ui
backend/model_loader.py ADDED
@@ -0,0 +1,146 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import logging
4
+ import os
5
+ import threading
6
+ from dataclasses import dataclass
7
+ from typing import Any
8
+
9
+ import torch
10
+ from transformers import (
11
+ AutoConfig,
12
+ AutoModel,
13
+ AutoModelForCausalLM,
14
+ AutoModelForMaskedLM,
15
+ AutoModelForImageClassification,
16
+ AutoModelForAudioClassification,
17
+ AutoProcessor,
18
+ AutoTokenizer,
19
+ )
20
+
21
+ from .modality import detect_modality
22
+
23
+ log = logging.getLogger("model_loader")
24
+
25
+
26
+ MAX_PARAMS_DEFAULT = 1_500_000_000
27
+
28
+
29
+ @dataclass
30
+ class LoadedModel:
31
+ model_id: str
32
+ model: torch.nn.Module
33
+ config: Any
34
+ tokenizer: Any
35
+ processor: Any
36
+ modality: str
37
+ head_kind: str # "causal_lm" | "masked_lm" | "image_classification" | "audio_classification" | "encoder"
38
+ param_count: int
39
+
40
+
41
+ class _Cache:
42
+ """Single-slot cache. Loading another model evicts the previous one."""
43
+
44
+ def __init__(self) -> None:
45
+ self._lock = threading.Lock()
46
+ self._current: LoadedModel | None = None
47
+
48
+ def get(self, model_id: str) -> LoadedModel | None:
49
+ with self._lock:
50
+ if self._current and self._current.model_id == model_id:
51
+ return self._current
52
+ return None
53
+
54
+ def put(self, lm: LoadedModel) -> None:
55
+ with self._lock:
56
+ self._current = lm
57
+
58
+ def current(self) -> LoadedModel | None:
59
+ with self._lock:
60
+ return self._current
61
+
62
+
63
+ CACHE = _Cache()
64
+
65
+
66
+ def _pick_auto_class(config: Any, modality: str):
67
+ """Pick the most informative auto-class for this config so we get a real head."""
68
+ architectures = getattr(config, "architectures", None) or []
69
+ arch = architectures[0] if architectures else ""
70
+
71
+ if modality == "vision":
72
+ if "ImageClassification" in arch or hasattr(config, "id2label"):
73
+ return AutoModelForImageClassification, "image_classification"
74
+ return AutoModel, "encoder"
75
+ if modality == "audio":
76
+ if "AudioClassification" in arch:
77
+ return AutoModelForAudioClassification, "audio_classification"
78
+ return AutoModel, "encoder"
79
+
80
+ # text
81
+ if "ForCausalLM" in arch:
82
+ return AutoModelForCausalLM, "causal_lm"
83
+ if "ForMaskedLM" in arch:
84
+ return AutoModelForMaskedLM, "masked_lm"
85
+ # default — encoder backbone (BERT-like) without head
86
+ return AutoModel, "encoder"
87
+
88
+
89
+ def _count_params(model: torch.nn.Module) -> int:
90
+ return sum(p.numel() for p in model.parameters())
91
+
92
+
93
+ def load_model(model_id: str) -> LoadedModel:
94
+ cached = CACHE.get(model_id)
95
+ if cached is not None:
96
+ return cached
97
+
98
+ log.info("Loading config for %s", model_id)
99
+ config = AutoConfig.from_pretrained(model_id, trust_remote_code=False)
100
+ modality = detect_modality(getattr(config, "model_type", None))
101
+
102
+ auto_cls, head_kind = _pick_auto_class(config, modality)
103
+ log.info("Loading %s as %s (%s)", model_id, auto_cls.__name__, head_kind)
104
+
105
+ max_params = int(os.environ.get("MAX_PARAMS", MAX_PARAMS_DEFAULT))
106
+
107
+ # Load on CPU in float32. Keep it deterministic for the visualizer.
108
+ model = auto_cls.from_pretrained(
109
+ model_id,
110
+ torch_dtype=torch.float32,
111
+ low_cpu_mem_usage=True,
112
+ trust_remote_code=False,
113
+ )
114
+ n = _count_params(model)
115
+ if n > max_params:
116
+ del model
117
+ raise ValueError(
118
+ f"Model {model_id} has {n:,} params, exceeds MAX_PARAMS={max_params:,}."
119
+ )
120
+ model.eval()
121
+
122
+ tokenizer = None
123
+ processor = None
124
+ if modality == "text":
125
+ try:
126
+ tokenizer = AutoTokenizer.from_pretrained(model_id, use_fast=True)
127
+ except Exception:
128
+ tokenizer = AutoTokenizer.from_pretrained(model_id, use_fast=False)
129
+ else:
130
+ try:
131
+ processor = AutoProcessor.from_pretrained(model_id)
132
+ except Exception as e:
133
+ log.warning("AutoProcessor failed for %s: %s", model_id, e)
134
+
135
+ lm = LoadedModel(
136
+ model_id=model_id,
137
+ model=model,
138
+ config=config,
139
+ tokenizer=tokenizer,
140
+ processor=processor,
141
+ modality=modality,
142
+ head_kind=head_kind,
143
+ param_count=n,
144
+ )
145
+ CACHE.put(lm)
146
+ return lm
backend/runner.py ADDED
@@ -0,0 +1,174 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import logging
4
+ import threading
5
+ from dataclasses import dataclass, field
6
+ from typing import Any, Callable, Iterator
7
+
8
+ import torch
9
+ import torch.nn as nn
10
+
11
+ from .model_loader import LoadedModel
12
+ from .modality import build_inputs
13
+ from .snapshots import attention_snapshot, tensor_snapshot, topk_from_logits
14
+
15
+ log = logging.getLogger("runner")
16
+
17
+
18
+ @dataclass
19
+ class RunEvent:
20
+ step: int
21
+ path: str
22
+ module_class: str
23
+ kind: str # "module" | "attention" | "logits" | "input" | "done"
24
+ payload: dict
25
+
26
+
27
+ @dataclass
28
+ class RunResult:
29
+ events: list[RunEvent] = field(default_factory=list)
30
+
31
+
32
+ def _decoder_for(loaded: LoadedModel) -> Callable[[int], str]:
33
+ tok = loaded.tokenizer
34
+ if tok is not None:
35
+ def decode(i: int) -> str:
36
+ return tok.decode([i], skip_special_tokens=False)
37
+ return decode
38
+ id2label = getattr(loaded.config, "id2label", None) or {}
39
+ def decode(i: int) -> str:
40
+ return str(id2label.get(i) or id2label.get(str(i)) or f"<id:{i}>")
41
+ return decode
42
+
43
+
44
+ def _classify(cls_name: str) -> str:
45
+ from .tracer import classify
46
+ return classify(cls_name)
47
+
48
+
49
+ def run_model(loaded: LoadedModel, prompt: str | None) -> RunResult:
50
+ """Run a forward pass with hooks and return the ordered list of events."""
51
+ model = loaded.model
52
+ inputs, ui = build_inputs(
53
+ loaded.modality, loaded.tokenizer, loaded.processor, loaded.config, prompt,
54
+ )
55
+
56
+ result = RunResult()
57
+ result.events.append(
58
+ RunEvent(step=0, path="<input>", module_class="Input", kind="input", payload=ui)
59
+ )
60
+
61
+ counter = {"i": 1}
62
+ handles: list[Any] = []
63
+ seen: set[int] = set() # avoid double-recording the same module
64
+
65
+ def make_hook(path: str, cls_name: str):
66
+ def hook(module: nn.Module, inputs_, output):
67
+ mid = id(module)
68
+ if mid in seen:
69
+ return
70
+ seen.add(mid)
71
+ snap = tensor_snapshot(output)
72
+ if snap is None:
73
+ return
74
+ ev = RunEvent(
75
+ step=counter["i"],
76
+ path=path,
77
+ module_class=cls_name,
78
+ kind="module",
79
+ payload={"snapshot": snap, "kind": _classify(cls_name)},
80
+ )
81
+ counter["i"] += 1
82
+ result.events.append(ev)
83
+ return hook
84
+
85
+ for path, mod in model.named_modules():
86
+ if path == "":
87
+ continue
88
+ cls_name = type(mod).__name__
89
+ h = mod.register_forward_hook(make_hook(path, cls_name))
90
+ handles.append(h)
91
+
92
+ try:
93
+ forward_kwargs = dict(inputs)
94
+ # Only pass these flags if the model accepts them.
95
+ sig_params = getattr(model.forward, "__wrapped__", model.forward).__code__.co_varnames
96
+ if "output_attentions" in sig_params:
97
+ forward_kwargs["output_attentions"] = True
98
+ if "output_hidden_states" in sig_params:
99
+ forward_kwargs["output_hidden_states"] = True
100
+ with torch.no_grad():
101
+ outputs = model(**forward_kwargs)
102
+ finally:
103
+ for h in handles:
104
+ h.remove()
105
+
106
+ # Attention snapshots (cleaner than the raw module output).
107
+ attentions = getattr(outputs, "attentions", None)
108
+ if attentions is not None:
109
+ for layer_i, attn in enumerate(attentions):
110
+ snap = attention_snapshot(attn)
111
+ if snap is None:
112
+ continue
113
+ result.events.append(
114
+ RunEvent(
115
+ step=counter["i"],
116
+ path=f"<attention.layer.{layer_i}>",
117
+ module_class="Attention",
118
+ kind="attention",
119
+ payload={"snapshot": snap, "layer": layer_i},
120
+ )
121
+ )
122
+ counter["i"] += 1
123
+
124
+ # Final logits / probabilities.
125
+ logits = getattr(outputs, "logits", None)
126
+ if logits is not None:
127
+ decode = _decoder_for(loaded)
128
+ last_logits = logits
129
+ if logits.dim() >= 3:
130
+ # LM heads: take the last token's logits.
131
+ last_logits = logits[:, -1, :]
132
+ topk = topk_from_logits(last_logits, decode_fn=decode, k=10)
133
+ snap = tensor_snapshot(logits)
134
+ result.events.append(
135
+ RunEvent(
136
+ step=counter["i"],
137
+ path="<output.logits>",
138
+ module_class="Logits",
139
+ kind="logits",
140
+ payload={"snapshot": snap, "top_k": topk, "head_kind": loaded.head_kind},
141
+ )
142
+ )
143
+ counter["i"] += 1
144
+
145
+ result.events.append(
146
+ RunEvent(
147
+ step=counter["i"],
148
+ path="<done>",
149
+ module_class="Done",
150
+ kind="done",
151
+ payload={"total_steps": counter["i"]},
152
+ )
153
+ )
154
+ return result
155
+
156
+
157
+ # In-memory store of the most recent run so the WS endpoint can replay it.
158
+ _RUN_LOCK = threading.Lock()
159
+ _LAST_RUN: RunResult | None = None
160
+
161
+
162
+ def store_run(result: RunResult) -> None:
163
+ global _LAST_RUN
164
+ with _RUN_LOCK:
165
+ _LAST_RUN = result
166
+
167
+
168
+ def get_last_run() -> RunResult | None:
169
+ with _RUN_LOCK:
170
+ return _LAST_RUN
171
+
172
+
173
+ def iter_events(result: RunResult) -> Iterator[RunEvent]:
174
+ yield from result.events
backend/snapshots.py ADDED
@@ -0,0 +1,124 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import numpy as np
6
+ import torch
7
+
8
+
9
+ HEATMAP_MAX = 64
10
+ SAMPLE_FLAT = 64
11
+ TOPK = 10
12
+
13
+
14
+ def _to_numpy(t: torch.Tensor) -> np.ndarray:
15
+ return t.detach().to(torch.float32).cpu().numpy()
16
+
17
+
18
+ def _downsample_2d(arr: np.ndarray, max_side: int = HEATMAP_MAX) -> np.ndarray:
19
+ h, w = arr.shape
20
+ if h <= max_side and w <= max_side:
21
+ return arr
22
+ rh = max(1, h // max_side)
23
+ rw = max(1, w // max_side)
24
+ h2 = (h // rh) * rh
25
+ w2 = (w // rw) * rw
26
+ arr = arr[:h2, :w2]
27
+ return arr.reshape(h2 // rh, rh, w2 // rw, rw).mean(axis=(1, 3))
28
+
29
+
30
+ def _to_2d(arr: np.ndarray) -> np.ndarray | None:
31
+ if arr.ndim == 0:
32
+ return None
33
+ if arr.ndim == 1:
34
+ return arr.reshape(1, -1)
35
+ if arr.ndim == 2:
36
+ return arr
37
+ # Collapse all leading dims into rows; last dim is columns.
38
+ return arr.reshape(-1, arr.shape[-1])
39
+
40
+
41
+ def tensor_snapshot(t: Any) -> dict | None:
42
+ """Convert a tensor (or first tensor inside a tuple/list) into a JSON-safe summary."""
43
+ if isinstance(t, (tuple, list)):
44
+ for item in t:
45
+ snap = tensor_snapshot(item)
46
+ if snap is not None:
47
+ return snap
48
+ return None
49
+ if not isinstance(t, torch.Tensor):
50
+ return None
51
+ arr = _to_numpy(t)
52
+ flat = arr.reshape(-1)
53
+ if flat.size == 0:
54
+ return None
55
+ finite = flat[np.isfinite(flat)]
56
+ if finite.size == 0:
57
+ finite = np.zeros(1, dtype=np.float32)
58
+ stats = {
59
+ "mean": float(finite.mean()),
60
+ "std": float(finite.std()),
61
+ "min": float(finite.min()),
62
+ "max": float(finite.max()),
63
+ "abs_max": float(np.abs(finite).max()),
64
+ }
65
+ sample = flat[:SAMPLE_FLAT].tolist()
66
+ heatmap_payload: dict | None = None
67
+ arr2 = _to_2d(arr)
68
+ if arr2 is not None and arr2.size > 0:
69
+ ds = _downsample_2d(arr2)
70
+ heatmap_payload = {
71
+ "shape": list(ds.shape),
72
+ "values": ds.astype(np.float32).round(5).tolist(),
73
+ "vmin": float(ds.min()),
74
+ "vmax": float(ds.max()),
75
+ }
76
+ return {
77
+ "shape": list(arr.shape),
78
+ "dtype": str(t.dtype).replace("torch.", ""),
79
+ "stats": stats,
80
+ "sample": [round(float(x), 5) for x in sample],
81
+ "heatmap": heatmap_payload,
82
+ }
83
+
84
+
85
+ def topk_from_logits(logits: torch.Tensor, decode_fn, k: int = TOPK) -> list[dict]:
86
+ """Take logits over a vocab axis and return decoded top-k probabilities."""
87
+ if logits.dim() == 0:
88
+ return []
89
+ flat = logits
90
+ while flat.dim() > 1:
91
+ flat = flat[0] if flat.size(0) > 0 else flat.squeeze(0)
92
+ probs = torch.softmax(flat.float(), dim=-1)
93
+ top = torch.topk(probs, k=min(k, probs.numel()))
94
+ out = []
95
+ for prob, idx in zip(top.values.tolist(), top.indices.tolist()):
96
+ try:
97
+ decoded = decode_fn(int(idx))
98
+ except Exception:
99
+ decoded = f"<id:{idx}>"
100
+ out.append({"id": int(idx), "token": decoded, "prob": float(prob)})
101
+ return out
102
+
103
+
104
+ def attention_snapshot(attn: torch.Tensor) -> dict | None:
105
+ """Attention is [B, H, T, T]. Return [H, T, T] downsampled per head, plus a head-mean."""
106
+ if not isinstance(attn, torch.Tensor) or attn.dim() < 3:
107
+ return None
108
+ a = attn.detach().to(torch.float32).cpu()
109
+ if a.dim() == 4:
110
+ a = a[0] # batch 0
111
+ if a.dim() == 3:
112
+ head_mean = a.mean(dim=0).numpy()
113
+ else:
114
+ head_mean = a.numpy()
115
+ ds = _downsample_2d(head_mean, max_side=HEATMAP_MAX)
116
+ return {
117
+ "shape": list(attn.shape),
118
+ "heatmap": {
119
+ "shape": list(ds.shape),
120
+ "values": ds.astype(np.float32).round(5).tolist(),
121
+ "vmin": float(ds.min()),
122
+ "vmax": float(ds.max()),
123
+ },
124
+ }
backend/tracer.py ADDED
@@ -0,0 +1,130 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+
8
+
9
+ # Coarse module-class buckets used for graph styling on the frontend.
10
+ KIND_PATTERNS: list[tuple[str, tuple[str, ...]]] = [
11
+ ("Embedding", ("Embedding", "Embeddings", "PatchEmbed", "PositionalEmbedding")),
12
+ ("Attention", ("Attention", "SelfAttention", "MHA", "MultiHead")),
13
+ ("MLP", ("MLP", "FFN", "FeedForward", "Intermediate", "GLU", "MoE", "Expert")),
14
+ ("Norm", ("LayerNorm", "RMSNorm", "BatchNorm", "GroupNorm", "Norm")),
15
+ ("Linear", ("Linear",)),
16
+ ("Conv", ("Conv1d", "Conv2d", "Conv3d", "Conv")),
17
+ ("Activation", ("GELU", "ReLU", "SiLU", "Mish", "Swish", "Tanh", "Sigmoid")),
18
+ ("Dropout", ("Dropout",)),
19
+ ("Pooler", ("Pool", "Pooler")),
20
+ ("Block", ("Block", "Layer", "Encoder", "Decoder", "DecoderLayer", "EncoderLayer")),
21
+ ("Head", ("Head", "Classifier", "LMHead", "Predictor", "Projector")),
22
+ ]
23
+
24
+
25
+ def classify(module_class: str) -> str:
26
+ for kind, patterns in KIND_PATTERNS:
27
+ for p in patterns:
28
+ if p in module_class:
29
+ return kind
30
+ return "Other"
31
+
32
+
33
+ def _interesting_config(module: nn.Module) -> dict[str, Any]:
34
+ """Pull a few attributes off the module that are visualizer-friendly."""
35
+ out: dict[str, Any] = {}
36
+ for k in (
37
+ "num_heads",
38
+ "num_attention_heads",
39
+ "head_dim",
40
+ "hidden_size",
41
+ "intermediate_size",
42
+ "embed_dim",
43
+ "in_features",
44
+ "out_features",
45
+ "in_channels",
46
+ "out_channels",
47
+ "kernel_size",
48
+ "stride",
49
+ "padding",
50
+ "eps",
51
+ "normalized_shape",
52
+ "num_embeddings",
53
+ "embedding_dim",
54
+ "num_experts",
55
+ "top_k",
56
+ ):
57
+ v = getattr(module, k, None)
58
+ if v is None:
59
+ continue
60
+ if isinstance(v, (int, float, str, bool)):
61
+ out[k] = v
62
+ elif isinstance(v, (list, tuple)) and all(isinstance(x, (int, float)) for x in v):
63
+ out[k] = list(v)
64
+ return out
65
+
66
+
67
+ def build_graph(model: nn.Module) -> dict:
68
+ """Walk the module tree and return a JSON graph.
69
+
70
+ We build edges from the parent → each child in declaration order. This
71
+ captures the module hierarchy precisely; runtime dataflow is overlaid
72
+ later by the runner using the order of forward-hook firings.
73
+ """
74
+ nodes: list[dict] = []
75
+ edges: list[dict] = []
76
+
77
+ # Collect each module's parameter count cheaply.
78
+ param_counts: dict[str, int] = {}
79
+ for path, mod in model.named_modules():
80
+ param_counts[path] = sum(p.numel() for p in mod.parameters(recurse=False))
81
+
82
+ # Aggregate: every node also reports the *recursive* parameter count.
83
+ recursive: dict[str, int] = {}
84
+
85
+ def _walk(prefix: str, mod: nn.Module) -> int:
86
+ own = param_counts.get(prefix, 0)
87
+ total = own
88
+ for name, child in mod.named_children():
89
+ child_path = f"{prefix}.{name}" if prefix else name
90
+ total += _walk(child_path, child)
91
+ recursive[prefix] = total
92
+ return total
93
+
94
+ _walk("", model)
95
+
96
+ # Emit nodes + parent→child edges.
97
+ for path, mod in model.named_modules():
98
+ cls = type(mod).__name__
99
+ node_id = path or "<root>"
100
+ children = [
101
+ f"{path}.{name}" if path else name for name, _ in mod.named_children()
102
+ ]
103
+ nodes.append(
104
+ {
105
+ "id": node_id,
106
+ "path": path,
107
+ "module_class": cls,
108
+ "kind": classify(cls),
109
+ "params": recursive.get(path, 0),
110
+ "params_own": param_counts.get(path, 0),
111
+ "config": _interesting_config(mod),
112
+ "children": [c if c else "<root>" for c in children],
113
+ "depth": 0 if not path else path.count(".") + 1,
114
+ "is_leaf": len(children) == 0,
115
+ }
116
+ )
117
+ for c in children:
118
+ edges.append(
119
+ {
120
+ "from": node_id,
121
+ "to": c if c else "<root>",
122
+ "kind": "tree",
123
+ }
124
+ )
125
+
126
+ return {
127
+ "root": "<root>",
128
+ "nodes": nodes,
129
+ "edges": edges,
130
+ }
frontend/index.html ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <!doctype html>
2
+ <html lang="en">
3
+ <head>
4
+ <meta charset="UTF-8" />
5
+ <meta name="viewport" content="width=device-width, initial-scale=1.0" />
6
+ <title>Transformer X-Ray</title>
7
+ </head>
8
+ <body>
9
+ <div id="root"></div>
10
+ <script type="module" src="/src/main.tsx"></script>
11
+ </body>
12
+ </html>
frontend/package.json ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "name": "transformer-xray",
3
+ "private": true,
4
+ "version": "0.1.0",
5
+ "type": "module",
6
+ "scripts": {
7
+ "dev": "vite",
8
+ "build": "tsc -b --noEmit && vite build",
9
+ "preview": "vite preview"
10
+ },
11
+ "dependencies": {
12
+ "@dagrejs/dagre": "^1.1.4",
13
+ "gif.js": "^0.2.0",
14
+ "html-to-image": "^1.11.13",
15
+ "react": "^18.3.1",
16
+ "react-dom": "^18.3.1",
17
+ "reactflow": "^11.11.4",
18
+ "zustand": "^4.5.5"
19
+ },
20
+ "devDependencies": {
21
+ "@types/react": "^18.3.12",
22
+ "@types/react-dom": "^18.3.1",
23
+ "@vitejs/plugin-react": "^4.3.4",
24
+ "typescript": "^5.6.3",
25
+ "vite": "^5.4.11"
26
+ }
27
+ }
frontend/src/App.tsx ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import Graph from "./components/Graph";
2
+ import Header from "./components/Header";
3
+ import NodeCard from "./components/NodeCard";
4
+ import PlayBar from "./components/PlayBar";
5
+ import ProbsBar from "./components/ProbsBar";
6
+ import TokenStrip from "./components/TokenStrip";
7
+ import GifExport from "./components/GifExport";
8
+ import { useStore } from "./store";
9
+
10
+ export default function App() {
11
+ const error = useStore((s) => s.error);
12
+ return (
13
+ <div style={{ display: "grid", gridTemplateRows: "auto auto 1fr auto", height: "100vh", background: "#0b1220", color: "#e2e8f0" }}>
14
+ <Header />
15
+ <TokenStrip />
16
+ <main style={{ display: "grid", gridTemplateColumns: "1fr 360px", overflow: "hidden" }}>
17
+ <Graph />
18
+ <aside style={{ borderLeft: "1px solid #1f2937", background: "#0f172a", overflow: "auto", display: "flex", flexDirection: "column" }}>
19
+ <div style={{ borderBottom: "1px solid #1f2937" }}>
20
+ <NodeCard />
21
+ </div>
22
+ <div style={{ borderBottom: "1px solid #1f2937" }}>
23
+ <ProbsBar />
24
+ </div>
25
+ <GifExport />
26
+ </aside>
27
+ </main>
28
+ <PlayBar />
29
+ {error && (
30
+ <div style={{ position: "fixed", bottom: 80, right: 16, background: "#7f1d1d", color: "#fee2e2", padding: "8px 12px", borderRadius: 6, fontSize: 12, maxWidth: 360 }}>
31
+ {error}
32
+ </div>
33
+ )}
34
+ </div>
35
+ );
36
+ }
frontend/src/components/GifExport.tsx ADDED
@@ -0,0 +1,99 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useState } from "react";
2
+ // @ts-ignore — gif.js has no types
3
+ import GIF from "gif.js";
4
+ import { toPng } from "html-to-image";
5
+ import { useStore } from "../store";
6
+
7
+ export default function GifExport() {
8
+ const events = useStore((s) => s.events);
9
+ const setCurrentStep = useStore((s) => s.setCurrentStep);
10
+ const [busy, setBusy] = useState(false);
11
+ const [progress, setProgress] = useState(0);
12
+ const [fps, setFps] = useState(8);
13
+ const [maxFrames, setMaxFrames] = useState(60);
14
+
15
+ async function record() {
16
+ if (busy || events.length === 0) return;
17
+ setBusy(true);
18
+ setProgress(0);
19
+ try {
20
+ const target = document.getElementById("graph-canvas");
21
+ if (!target) throw new Error("graph not found");
22
+
23
+ const total = Math.min(maxFrames, events.length);
24
+ const stride = Math.max(1, Math.floor(events.length / total));
25
+
26
+ // gif.worker.js is bundled with the package; resolve at runtime so Vite serves it.
27
+ const workerUrl = new URL(
28
+ "gif.js/dist/gif.worker.js",
29
+ import.meta.url,
30
+ ).toString();
31
+ const gif = new GIF({
32
+ workers: 2,
33
+ quality: 8,
34
+ workerScript: workerUrl,
35
+ width: target.clientWidth,
36
+ height: target.clientHeight,
37
+ });
38
+
39
+ for (let i = 0; i < total; i++) {
40
+ const step = Math.min(events.length - 1, i * stride);
41
+ setCurrentStep(step);
42
+ // Allow React to repaint before snapshot.
43
+ await new Promise((r) => requestAnimationFrame(() => requestAnimationFrame(r)));
44
+ const dataUrl = await toPng(target, { cacheBust: true, pixelRatio: 1 });
45
+ const img = await new Promise<HTMLImageElement>((resolve, reject) => {
46
+ const im = new Image();
47
+ im.onload = () => resolve(im);
48
+ im.onerror = reject;
49
+ im.src = dataUrl;
50
+ });
51
+ gif.addFrame(img, { delay: Math.round(1000 / fps) });
52
+ setProgress(Math.round(((i + 1) / total) * 50));
53
+ }
54
+
55
+ gif.on("progress", (p: number) => setProgress(50 + Math.round(p * 50)));
56
+ const blob: Blob = await new Promise((resolve) => {
57
+ gif.on("finished", (b: Blob) => resolve(b));
58
+ gif.render();
59
+ });
60
+
61
+ const url = URL.createObjectURL(blob);
62
+ const a = document.createElement("a");
63
+ a.href = url;
64
+ a.download = `transformer-xray-${Date.now()}.gif`;
65
+ a.click();
66
+ URL.revokeObjectURL(url);
67
+ } catch (e) {
68
+ console.error(e);
69
+ alert(`gif export failed: ${(e as Error).message}`);
70
+ } finally {
71
+ setBusy(false);
72
+ setProgress(0);
73
+ }
74
+ }
75
+
76
+ return (
77
+ <div style={{ padding: 12, borderTop: "1px solid #1f2937" }}>
78
+ <div style={{ fontSize: 11, color: "#94a3b8", marginBottom: 6 }}>GIF export</div>
79
+ <div style={{ display: "flex", gap: 8, alignItems: "center", marginBottom: 6, fontSize: 11 }}>
80
+ <label style={{ color: "#94a3b8" }}>
81
+ fps
82
+ <input type="number" min={2} max={24} value={fps} onChange={(e) => setFps(Math.max(2, Math.min(24, Number(e.target.value))))} style={{ width: 50, marginLeft: 6 }} />
83
+ </label>
84
+ <label style={{ color: "#94a3b8" }}>
85
+ frames
86
+ <input type="number" min={10} max={200} value={maxFrames} onChange={(e) => setMaxFrames(Math.max(10, Math.min(200, Number(e.target.value))))} style={{ width: 60, marginLeft: 6 }} />
87
+ </label>
88
+ </div>
89
+ <button className="btn btn-primary" disabled={busy || events.length === 0} onClick={record}>
90
+ {busy ? `recording… ${progress}%` : "record GIF"}
91
+ </button>
92
+ {busy && (
93
+ <div style={{ marginTop: 6, height: 4, background: "#0b1220", borderRadius: 2, overflow: "hidden" }}>
94
+ <div style={{ width: `${progress}%`, height: "100%", background: "#38bdf8", transition: "width 120ms" }} />
95
+ </div>
96
+ )}
97
+ </div>
98
+ );
99
+ }
frontend/src/components/Graph.tsx ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useMemo } from "react";
2
+ import ReactFlow, { Background, Controls, MiniMap } from "reactflow";
3
+ import { useStore } from "../store";
4
+ import { layoutGraph } from "../lib/layout";
5
+ import ModuleNode from "./ModuleNode";
6
+
7
+ const nodeTypes = { module: ModuleNode };
8
+
9
+ export default function Graph() {
10
+ const graph = useStore((s) => s.graph);
11
+ const granularity = useStore((s) => s.granularity);
12
+ const setSelectedPath = useStore((s) => s.setSelectedPath);
13
+
14
+ const { nodes, edges } = useMemo(() => {
15
+ if (!graph) return { nodes: [], edges: [] };
16
+ return layoutGraph(graph, { granularity });
17
+ }, [graph, granularity]);
18
+
19
+ if (!graph) {
20
+ return (
21
+ <div style={{ display: "grid", placeItems: "center", height: "100%", color: "#94a3b8" }}>
22
+ Load a model to see its graph.
23
+ </div>
24
+ );
25
+ }
26
+
27
+ return (
28
+ <div id="graph-canvas" style={{ width: "100%", height: "100%", background: "#0b1220" }}>
29
+ <ReactFlow
30
+ nodes={nodes}
31
+ edges={edges}
32
+ nodeTypes={nodeTypes}
33
+ fitView
34
+ minZoom={0.05}
35
+ maxZoom={2}
36
+ onPaneClick={() => setSelectedPath(null)}
37
+ proOptions={{ hideAttribution: true }}
38
+ >
39
+ <Background color="#1f2937" gap={24} />
40
+ <Controls showInteractive={false} />
41
+ <MiniMap nodeColor={() => "#475569"} maskColor="rgba(11,18,32,0.7)" pannable zoomable />
42
+ </ReactFlow>
43
+ </div>
44
+ );
45
+ }
frontend/src/components/Header.tsx ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useState } from "react";
2
+ import { useStore } from "../store";
3
+ import { fetchGraph, loadModel } from "../lib/api";
4
+
5
+ const PRESETS = [
6
+ "prajjwal1/bert-tiny",
7
+ "bert-base-uncased",
8
+ "distilgpt2",
9
+ "google/vit-base-patch16-224",
10
+ "openai-community/gpt2",
11
+ ];
12
+
13
+ function fmt(n: number): string {
14
+ if (n >= 1e9) return (n / 1e9).toFixed(2) + "B";
15
+ if (n >= 1e6) return (n / 1e6).toFixed(2) + "M";
16
+ if (n >= 1e3) return (n / 1e3).toFixed(1) + "K";
17
+ return String(n);
18
+ }
19
+
20
+ export default function Header() {
21
+ const modelId = useStore((s) => s.modelId);
22
+ const setModelId = useStore((s) => s.setModelId);
23
+ const setLoadInfo = useStore((s) => s.setLoadInfo);
24
+ const setGraph = useStore((s) => s.setGraph);
25
+ const setLoading = useStore((s) => s.setLoading);
26
+ const setError = useStore((s) => s.setError);
27
+ const resetRun = useStore((s) => s.resetRun);
28
+ const loading = useStore((s) => s.loading);
29
+ const loadInfo = useStore((s) => s.loadInfo);
30
+ const granularity = useStore((s) => s.granularity);
31
+ const setGranularity = useStore((s) => s.setGranularity);
32
+ const [pending, setPending] = useState(modelId);
33
+
34
+ async function go(id: string) {
35
+ setLoading(true);
36
+ setError(null);
37
+ setLoadInfo(null);
38
+ setGraph(null);
39
+ resetRun();
40
+ try {
41
+ const info = await loadModel(id);
42
+ setLoadInfo(info);
43
+ const g = await fetchGraph();
44
+ setGraph(g);
45
+ setModelId(id);
46
+ } catch (e: any) {
47
+ setError(e?.message ?? String(e));
48
+ } finally {
49
+ setLoading(false);
50
+ }
51
+ }
52
+
53
+ return (
54
+ <header style={{ padding: "10px 14px", borderBottom: "1px solid #1f2937", background: "#0f172a", display: "flex", flexDirection: "column", gap: 8 }}>
55
+ <div style={{ display: "flex", alignItems: "center", gap: 8, flexWrap: "wrap" }}>
56
+ <span style={{ fontWeight: 800, color: "#f8fafc", letterSpacing: 0.3 }}>🔬 Transformer X-Ray</span>
57
+ <input
58
+ value={pending}
59
+ onChange={(e) => setPending(e.target.value)}
60
+ placeholder="huggingface-id e.g. bert-base-uncased"
61
+ onKeyDown={(e) => { if (e.key === "Enter") go(pending); }}
62
+ style={{
63
+ flex: 1,
64
+ minWidth: 280,
65
+ background: "#0b1220",
66
+ color: "#e2e8f0",
67
+ border: "1px solid #334155",
68
+ borderRadius: 6,
69
+ padding: "6px 10px",
70
+ fontFamily: "ui-monospace, monospace",
71
+ fontSize: 12,
72
+ }}
73
+ />
74
+ <button className="btn btn-primary" onClick={() => go(pending)} disabled={loading}>
75
+ {loading ? "loading…" : "load"}
76
+ </button>
77
+ <div style={{ display: "flex", gap: 4, marginLeft: 6 }}>
78
+ <span style={{ fontSize: 10, color: "#94a3b8", alignSelf: "center" }}>granularity</span>
79
+ {[0, 1, 2].map((g) => (
80
+ <button key={g} className={g === granularity ? "tab tab-active" : "tab"} onClick={() => setGranularity(g)}>
81
+ {["top", "blocks", "leaves"][g]}
82
+ </button>
83
+ ))}
84
+ </div>
85
+ </div>
86
+ <div style={{ display: "flex", gap: 4, flexWrap: "wrap", fontSize: 10 }}>
87
+ {PRESETS.map((p) => (
88
+ <button key={p} className="tab" onClick={() => { setPending(p); go(p); }}>{p}</button>
89
+ ))}
90
+ {loadInfo && (
91
+ <div style={{ marginLeft: "auto", color: "#94a3b8", fontFamily: "ui-monospace, monospace" }}>
92
+ {loadInfo.config.model_type} ·
93
+ {" "}{fmt(loadInfo.param_count)} params ·
94
+ {" "}{loadInfo.modality} ·
95
+ {" "}head <span style={{ color: "#fb923c" }}>{loadInfo.head_kind}</span>
96
+ {loadInfo.config.hidden_size && <> · hidden {loadInfo.config.hidden_size}</>}
97
+ {loadInfo.config.num_hidden_layers && <> · L={loadInfo.config.num_hidden_layers}</>}
98
+ {loadInfo.config.num_attention_heads && <> · H={loadInfo.config.num_attention_heads}</>}
99
+ </div>
100
+ )}
101
+ </div>
102
+ </header>
103
+ );
104
+ }
frontend/src/components/ModuleNode.tsx ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { Handle, Position, type NodeProps } from "reactflow";
2
+ import { useStore } from "../store";
3
+ import { colorFor } from "../lib/colors";
4
+ import type { GraphNode } from "../types";
5
+
6
+ function fmtParams(n: number): string {
7
+ if (n >= 1e9) return (n / 1e9).toFixed(2) + "B";
8
+ if (n >= 1e6) return (n / 1e6).toFixed(2) + "M";
9
+ if (n >= 1e3) return (n / 1e3).toFixed(1) + "K";
10
+ return String(n);
11
+ }
12
+
13
+ export default function ModuleNode({ data, id }: NodeProps<GraphNode>) {
14
+ const events = useStore((s) => s.events);
15
+ const currentStep = useStore((s) => s.currentStep);
16
+ const snapshotsByPath = useStore((s) => s.snapshotsByPath);
17
+ const selectedPath = useStore((s) => s.selectedPath);
18
+ const setSelectedPath = useStore((s) => s.setSelectedPath);
19
+
20
+ const c = colorFor(data.kind);
21
+ const ev = events[currentStep];
22
+ const isActive = ev && ev.path === data.path && ev.kind === "module";
23
+ const hasFlowed = !!snapshotsByPath[data.path];
24
+ const isSelected = selectedPath === id;
25
+
26
+ const snap = snapshotsByPath[data.path];
27
+ const shapeStr = snap ? `[${snap.shape.join(", ")}]` : "";
28
+
29
+ return (
30
+ <div
31
+ onClick={(e) => {
32
+ e.stopPropagation();
33
+ setSelectedPath(id);
34
+ }}
35
+ style={{
36
+ width: 220,
37
+ padding: "8px 10px",
38
+ borderRadius: 8,
39
+ background: c.bg,
40
+ border: `2px solid ${isSelected ? "#fff" : isActive ? c.accent : c.border}`,
41
+ color: c.text,
42
+ boxShadow: isActive ? `0 0 24px ${c.accent}` : hasFlowed ? `0 0 0 1px ${c.accent}` : "none",
43
+ opacity: hasFlowed || !ev ? 1 : 0.55,
44
+ transition: "box-shadow 120ms, opacity 120ms, border-color 120ms",
45
+ fontSize: 11,
46
+ lineHeight: 1.35,
47
+ fontFamily: "ui-monospace, Menlo, monospace",
48
+ }}
49
+ >
50
+ <Handle type="target" position={Position.Top} style={{ background: c.accent }} />
51
+ <div style={{ display: "flex", justifyContent: "space-between", gap: 6 }}>
52
+ <span style={{ fontWeight: 700, color: c.accent }}>{data.kind}</span>
53
+ <span style={{ opacity: 0.7 }}>{fmtParams(data.params)}</span>
54
+ </div>
55
+ <div style={{ fontWeight: 600, fontSize: 12, marginTop: 2 }}>{data.module_class}</div>
56
+ <div style={{ opacity: 0.7, fontSize: 10, marginTop: 2 }}>{data.path || "<root>"}</div>
57
+ {shapeStr && (
58
+ <div style={{ marginTop: 4, fontSize: 10, color: c.accent, fontWeight: 600 }}>{shapeStr}</div>
59
+ )}
60
+ <Handle type="source" position={Position.Bottom} style={{ background: c.accent }} />
61
+ </div>
62
+ );
63
+ }
frontend/src/components/NodeCard.tsx ADDED
@@ -0,0 +1,64 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useStore } from "../store";
2
+ import TensorView from "./TensorView";
3
+
4
+ function fmtParams(n: number): string {
5
+ if (n >= 1e9) return (n / 1e9).toFixed(2) + "B";
6
+ if (n >= 1e6) return (n / 1e6).toFixed(2) + "M";
7
+ if (n >= 1e3) return (n / 1e3).toFixed(1) + "K";
8
+ return String(n);
9
+ }
10
+
11
+ export default function NodeCard() {
12
+ const graph = useStore((s) => s.graph);
13
+ const selectedPath = useStore((s) => s.selectedPath);
14
+ const snapshotsByPath = useStore((s) => s.snapshotsByPath);
15
+
16
+ if (!graph) return null;
17
+ const node = selectedPath ? graph.nodes.find((n) => n.id === selectedPath) : null;
18
+ const snap = node ? snapshotsByPath[node.path] : undefined;
19
+
20
+ if (!node) {
21
+ return (
22
+ <div style={{ padding: 12, color: "#64748b", fontSize: 12 }}>
23
+ Click a node to see its config and the tensor that flowed through it.
24
+ </div>
25
+ );
26
+ }
27
+
28
+ return (
29
+ <div style={{ padding: 12, fontSize: 12, color: "#e2e8f0" }}>
30
+ <div style={{ fontSize: 13, fontWeight: 700, color: "#f8fafc" }}>{node.module_class}</div>
31
+ <div style={{ fontSize: 10, opacity: 0.6, fontFamily: "ui-monospace, monospace", marginTop: 2, wordBreak: "break-all" }}>
32
+ {node.path || "<root>"}
33
+ </div>
34
+ <div style={{ display: "flex", gap: 8, marginTop: 8, fontSize: 11 }}>
35
+ <span className="pill">{node.kind}</span>
36
+ <span className="pill">{fmtParams(node.params)} params</span>
37
+ {node.is_leaf && <span className="pill">leaf</span>}
38
+ </div>
39
+
40
+ {Object.keys(node.config).length > 0 && (
41
+ <div style={{ marginTop: 10 }}>
42
+ <div style={{ color: "#94a3b8", fontSize: 11, marginBottom: 4 }}>module config</div>
43
+ <div style={{ display: "grid", gridTemplateColumns: "auto 1fr", gap: "2px 8px", fontFamily: "ui-monospace, monospace", fontSize: 11 }}>
44
+ {Object.entries(node.config).map(([k, v]) => (
45
+ <div key={k} style={{ display: "contents" }}>
46
+ <span style={{ color: "#94a3b8" }}>{k}</span>
47
+ <span style={{ color: "#e2e8f0" }}>{JSON.stringify(v)}</span>
48
+ </div>
49
+ ))}
50
+ </div>
51
+ </div>
52
+ )}
53
+
54
+ <div style={{ marginTop: 12, color: "#94a3b8", fontSize: 11 }}>output tensor</div>
55
+ {snap ? (
56
+ <TensorView snap={snap} />
57
+ ) : (
58
+ <div style={{ marginTop: 6, color: "#64748b", fontSize: 11 }}>
59
+ run the model to see this module's output
60
+ </div>
61
+ )}
62
+ </div>
63
+ );
64
+ }
frontend/src/components/PlayBar.tsx ADDED
@@ -0,0 +1,118 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useStore } from "../store";
2
+ import { usePlayback, useRunController } from "../hooks/useRun";
3
+ import { colorFor } from "../lib/colors";
4
+
5
+ const SPEEDS = [0.5, 1, 2, 4];
6
+
7
+ export default function PlayBar() {
8
+ usePlayback();
9
+ const { start } = useRunController();
10
+ const events = useStore((s) => s.events);
11
+ const currentStep = useStore((s) => s.currentStep);
12
+ const playing = useStore((s) => s.playing);
13
+ const speed = useStore((s) => s.speed);
14
+ const setPlaying = useStore((s) => s.setPlaying);
15
+ const setSpeed = useStore((s) => s.setSpeed);
16
+ const setCurrentStep = useStore((s) => s.setCurrentStep);
17
+ const running = useStore((s) => s.running);
18
+ const loaded = useStore((s) => s.loadInfo);
19
+ const promptDefault = loaded?.modality === "text" ? "Paris is the capital of [MASK]." : "";
20
+
21
+ const ev = events[currentStep];
22
+ const accent = ev?.kind === "module"
23
+ ? colorFor(ev.payload?.kind ?? "Other").accent
24
+ : ev?.kind === "attention"
25
+ ? "#a78bfa"
26
+ : ev?.kind === "logits"
27
+ ? "#fb923c"
28
+ : "#38bdf8";
29
+
30
+ // Build milestone markers — first appearance of each kind + first/middle/last "Block" snapshot.
31
+ const milestones: { step: number; label: string }[] = [];
32
+ if (events.length) {
33
+ const seen = new Set<string>();
34
+ events.forEach((e, i) => {
35
+ if (e.kind === "input" && !seen.has("input")) { milestones.push({ step: i, label: "Tokenize" }); seen.add("input"); }
36
+ if (e.kind === "module" && e.payload?.kind === "Embedding" && !seen.has("emb")) { milestones.push({ step: i, label: "Embed" }); seen.add("emb"); }
37
+ if (e.kind === "module" && e.payload?.kind === "Block" && !seen.has("block0")) { milestones.push({ step: i, label: "Block 0" }); seen.add("block0"); }
38
+ if (e.kind === "module" && e.payload?.kind === "Head" && !seen.has("head")) { milestones.push({ step: i, label: "Head" }); seen.add("head"); }
39
+ if (e.kind === "logits" && !seen.has("logits")) { milestones.push({ step: i, label: "Logits → Softmax" }); seen.add("logits"); }
40
+ });
41
+ }
42
+
43
+ return (
44
+ <div style={{ padding: "10px 14px", borderTop: "1px solid #1f2937", background: "#0f172a" }}>
45
+ <div style={{ display: "flex", alignItems: "center", gap: 10, flexWrap: "wrap" }}>
46
+ <button
47
+ className="btn btn-primary"
48
+ disabled={!loaded || running}
49
+ onClick={() => start(promptDefault || null)}
50
+ >
51
+ {running ? "running…" : "▶ Play"}
52
+ </button>
53
+ <button
54
+ className="btn"
55
+ disabled={events.length === 0}
56
+ onClick={() => setPlaying(!playing)}
57
+ >
58
+ {playing ? "❚❚ pause" : "▶ resume"}
59
+ </button>
60
+ <div style={{ display: "flex", gap: 2 }}>
61
+ {SPEEDS.map((s) => (
62
+ <button key={s} className={s === speed ? "tab tab-active" : "tab"} onClick={() => setSpeed(s)}>
63
+ {s}×
64
+ </button>
65
+ ))}
66
+ </div>
67
+ <div style={{ flex: 1, minWidth: 200, position: "relative" }}>
68
+ <input
69
+ type="range"
70
+ min={0}
71
+ max={Math.max(0, events.length - 1)}
72
+ value={currentStep}
73
+ onChange={(e) => setCurrentStep(Number(e.target.value))}
74
+ style={{ width: "100%", accentColor: accent }}
75
+ />
76
+ <div style={{ position: "relative", height: 14, marginTop: -4 }}>
77
+ {milestones.map((m) => (
78
+ <button
79
+ key={m.label}
80
+ onClick={() => setCurrentStep(m.step)}
81
+ title={m.label}
82
+ style={{
83
+ position: "absolute",
84
+ left: `${(m.step / Math.max(1, events.length - 1)) * 100}%`,
85
+ transform: "translateX(-50%)",
86
+ fontSize: 9,
87
+ padding: "1px 6px",
88
+ borderRadius: 8,
89
+ background: "#1e293b",
90
+ color: "#cbd5e1",
91
+ border: "1px solid #334155",
92
+ cursor: "pointer",
93
+ whiteSpace: "nowrap",
94
+ }}
95
+ >
96
+ {m.label}
97
+ </button>
98
+ ))}
99
+ </div>
100
+ </div>
101
+ <div style={{ fontFamily: "ui-monospace, monospace", fontSize: 11, color: "#94a3b8", minWidth: 70, textAlign: "right" }}>
102
+ {currentStep} / {Math.max(0, events.length - 1)}
103
+ </div>
104
+ </div>
105
+ <div style={{ fontSize: 11, color: "#94a3b8", marginTop: 6, fontFamily: "ui-monospace, monospace" }}>
106
+ {ev ? (
107
+ <>
108
+ <span style={{ color: accent, fontWeight: 600 }}>[{ev.kind}]</span>{" "}
109
+ <span style={{ color: "#e2e8f0" }}>{ev.module_class}</span>{" "}
110
+ <span style={{ opacity: 0.7 }}>{ev.path}</span>
111
+ </>
112
+ ) : (
113
+ "no run yet — press Play to capture a forward pass"
114
+ )}
115
+ </div>
116
+ </div>
117
+ );
118
+ }
frontend/src/components/ProbsBar.tsx ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useStore } from "../store";
2
+
3
+ export default function ProbsBar() {
4
+ const topK = useStore((s) => s.topK);
5
+ const loadInfo = useStore((s) => s.loadInfo);
6
+
7
+ if (!topK || topK.length === 0) {
8
+ return (
9
+ <div style={{ padding: 12, color: "#64748b", fontSize: 11 }}>
10
+ top-k probabilities will appear here after the run completes.
11
+ </div>
12
+ );
13
+ }
14
+
15
+ const max = Math.max(...topK.map((t) => t.prob));
16
+ return (
17
+ <div style={{ padding: 12 }}>
18
+ <div style={{ fontSize: 11, color: "#94a3b8", marginBottom: 6 }}>
19
+ head: <span style={{ color: "#fb923c" }}>{loadInfo?.head_kind}</span>
20
+ </div>
21
+ <div style={{ display: "flex", flexDirection: "column", gap: 4 }}>
22
+ {topK.map((t, i) => (
23
+ <div key={i} style={{ display: "grid", gridTemplateColumns: "60px 1fr 60px", alignItems: "center", gap: 8 }}>
24
+ <span
25
+ style={{
26
+ fontFamily: "ui-monospace, monospace",
27
+ fontSize: 11,
28
+ color: "#e2e8f0",
29
+ whiteSpace: "nowrap",
30
+ overflow: "hidden",
31
+ textOverflow: "ellipsis",
32
+ }}
33
+ title={t.token}
34
+ >
35
+ {t.token}
36
+ </span>
37
+ <div style={{ background: "#0b1220", borderRadius: 3, height: 14, overflow: "hidden", border: "1px solid #1f2937" }}>
38
+ <div
39
+ style={{
40
+ width: `${(t.prob / max) * 100}%`,
41
+ height: "100%",
42
+ background: i === 0 ? "#fb923c" : "#60a5fa",
43
+ transition: "width 240ms",
44
+ }}
45
+ />
46
+ </div>
47
+ <span style={{ fontFamily: "ui-monospace, monospace", fontSize: 10, color: "#94a3b8", textAlign: "right" }}>
48
+ {(t.prob * 100).toFixed(2)}%
49
+ </span>
50
+ </div>
51
+ ))}
52
+ </div>
53
+ </div>
54
+ );
55
+ }
frontend/src/components/TensorView.tsx ADDED
@@ -0,0 +1,122 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useEffect, useRef, useState } from "react";
2
+ import { viridis } from "../lib/colors";
3
+ import type { Heatmap, Snapshot } from "../types";
4
+
5
+ function HeatmapCanvas({ hm }: { hm: Heatmap }) {
6
+ const ref = useRef<HTMLCanvasElement>(null);
7
+ useEffect(() => {
8
+ const canvas = ref.current;
9
+ if (!canvas) return;
10
+ const [h, w] = hm.shape;
11
+ canvas.width = w;
12
+ canvas.height = h;
13
+ const ctx = canvas.getContext("2d");
14
+ if (!ctx) return;
15
+ const img = ctx.createImageData(w, h);
16
+ const span = hm.vmax - hm.vmin || 1;
17
+ for (let i = 0; i < h; i++) {
18
+ for (let j = 0; j < w; j++) {
19
+ const v = hm.values[i][j];
20
+ const t = (v - hm.vmin) / span;
21
+ const c = viridis(t);
22
+ const m = c.match(/\d+/g)!;
23
+ const idx = (i * w + j) * 4;
24
+ img.data[idx] = +m[0];
25
+ img.data[idx + 1] = +m[1];
26
+ img.data[idx + 2] = +m[2];
27
+ img.data[idx + 3] = 255;
28
+ }
29
+ }
30
+ ctx.putImageData(img, 0, 0);
31
+ }, [hm]);
32
+ return (
33
+ <canvas
34
+ ref={ref}
35
+ style={{
36
+ width: "100%",
37
+ maxWidth: 320,
38
+ imageRendering: "pixelated",
39
+ borderRadius: 4,
40
+ border: "1px solid #1f2937",
41
+ background: "#0b1220",
42
+ }}
43
+ />
44
+ );
45
+ }
46
+
47
+ export default function TensorView({ snap }: { snap: Snapshot }) {
48
+ const [mode, setMode] = useState<"heatmap" | "numbers" | "histogram">("heatmap");
49
+
50
+ const fmt = (x: number) => x.toFixed(Math.abs(x) < 0.01 ? 4 : 3);
51
+ const buckets = (() => {
52
+ const vals = snap.sample;
53
+ if (vals.length === 0) return [];
54
+ const min = Math.min(...vals);
55
+ const max = Math.max(...vals);
56
+ const bins = 20;
57
+ const span = max - min || 1;
58
+ const counts = new Array(bins).fill(0);
59
+ for (const v of vals) {
60
+ const b = Math.min(bins - 1, Math.floor(((v - min) / span) * bins));
61
+ counts[b]++;
62
+ }
63
+ const peak = Math.max(...counts) || 1;
64
+ return counts.map((c, i) => ({ pct: c / peak, label: min + (i / bins) * span }));
65
+ })();
66
+
67
+ return (
68
+ <div style={{ marginTop: 8 }}>
69
+ <div style={{ display: "flex", gap: 4, marginBottom: 6 }}>
70
+ {(["heatmap", "numbers", "histogram"] as const).map((m) => (
71
+ <button
72
+ key={m}
73
+ onClick={() => setMode(m)}
74
+ className={mode === m ? "tab tab-active" : "tab"}
75
+ >
76
+ {m}
77
+ </button>
78
+ ))}
79
+ </div>
80
+ <div style={{ fontSize: 11, color: "#94a3b8", marginBottom: 6, fontFamily: "ui-monospace, monospace" }}>
81
+ shape [{snap.shape.join(", ")}] · {snap.dtype} · μ={fmt(snap.stats.mean)} σ={fmt(snap.stats.std)} ·
82
+ min={fmt(snap.stats.min)} max={fmt(snap.stats.max)}
83
+ </div>
84
+ {mode === "heatmap" && (
85
+ snap.heatmap ? <HeatmapCanvas hm={snap.heatmap} /> : <div style={{ color: "#64748b" }}>no 2D projection</div>
86
+ )}
87
+ {mode === "numbers" && (
88
+ <div
89
+ style={{
90
+ display: "grid",
91
+ gridTemplateColumns: "repeat(8, 1fr)",
92
+ gap: 2,
93
+ fontFamily: "ui-monospace, monospace",
94
+ fontSize: 10,
95
+ }}
96
+ >
97
+ {snap.sample.map((v, i) => (
98
+ <span key={i} style={{ padding: "2px 4px", background: "#0b1220", borderRadius: 2, textAlign: "right", color: "#cbd5e1" }}>
99
+ {fmt(v)}
100
+ </span>
101
+ ))}
102
+ </div>
103
+ )}
104
+ {mode === "histogram" && (
105
+ <div style={{ display: "flex", alignItems: "flex-end", gap: 2, height: 80 }}>
106
+ {buckets.map((b, i) => (
107
+ <div
108
+ key={i}
109
+ title={fmt(b.label)}
110
+ style={{
111
+ flex: 1,
112
+ background: viridis(0.3 + b.pct * 0.5),
113
+ height: `${Math.max(2, b.pct * 100)}%`,
114
+ borderRadius: 1,
115
+ }}
116
+ />
117
+ ))}
118
+ </div>
119
+ )}
120
+ </div>
121
+ );
122
+ }
frontend/src/components/TokenStrip.tsx ADDED
@@ -0,0 +1,68 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useStore } from "../store";
2
+
3
+ export default function TokenStrip() {
4
+ const inputUI = useStore((s) => s.inputUI);
5
+ const loadInfo = useStore((s) => s.loadInfo);
6
+
7
+ if (!inputUI) {
8
+ return (
9
+ <div style={{ padding: "8px 14px", borderBottom: "1px solid #1f2937", background: "#0f172a", color: "#64748b", fontSize: 11 }}>
10
+ input will appear here once you press Play
11
+ </div>
12
+ );
13
+ }
14
+
15
+ if (inputUI.kind === "image") {
16
+ return (
17
+ <div style={{ padding: "8px 14px", borderBottom: "1px solid #1f2937", background: "#0f172a", color: "#cbd5e1", fontSize: 11, fontFamily: "ui-monospace, monospace", display: "flex", gap: 16 }}>
18
+ <span><span style={{ color: "#94a3b8" }}>image</span> {inputUI.size?.[0]}×{inputUI.size?.[1]}</span>
19
+ <span><span style={{ color: "#94a3b8" }}>pixel_values</span> [{inputUI.shape.join(", ")}]</span>
20
+ {inputUI.patch_size && <span><span style={{ color: "#94a3b8" }}>patch</span> {inputUI.patch_size}</span>}
21
+ </div>
22
+ );
23
+ }
24
+ if (inputUI.kind === "audio") {
25
+ return (
26
+ <div style={{ padding: "8px 14px", borderBottom: "1px solid #1f2937", background: "#0f172a", color: "#cbd5e1", fontSize: 11, fontFamily: "ui-monospace, monospace", display: "flex", gap: 16 }}>
27
+ <span><span style={{ color: "#94a3b8" }}>audio</span> {inputUI.duration_s?.toFixed(2)}s @ {inputUI.sample_rate}Hz</span>
28
+ <span><span style={{ color: "#94a3b8" }}>shape</span> [{inputUI.shape.join(", ")}]</span>
29
+ </div>
30
+ );
31
+ }
32
+
33
+ const ids: number[] = inputUI.ids ?? [];
34
+ const pieces: string[] = inputUI.pieces ?? [];
35
+ const vocab = inputUI.vocab_size ?? loadInfo?.tokenizer?.vocab_size;
36
+
37
+ return (
38
+ <div style={{ padding: "8px 14px", borderBottom: "1px solid #1f2937", background: "#0f172a", overflowX: "auto" }}>
39
+ <div style={{ display: "flex", gap: 6, alignItems: "center" }}>
40
+ <div style={{ fontSize: 10, color: "#94a3b8", whiteSpace: "nowrap" }}>
41
+ tokens [{ids.length}] · vocab {vocab?.toLocaleString()}
42
+ </div>
43
+ <div style={{ display: "flex", gap: 4 }}>
44
+ {ids.map((id, i) => (
45
+ <div
46
+ key={i}
47
+ style={{
48
+ background: "#1e293b",
49
+ border: "1px solid #334155",
50
+ borderRadius: 6,
51
+ padding: "3px 6px",
52
+ fontSize: 11,
53
+ fontFamily: "ui-monospace, monospace",
54
+ color: "#e2e8f0",
55
+ whiteSpace: "nowrap",
56
+ lineHeight: 1.2,
57
+ }}
58
+ title={`id ${id}`}
59
+ >
60
+ <div style={{ color: "#38bdf8", fontSize: 9 }}>{id}</div>
61
+ <div>{pieces[i]?.replace(/^▁/, "_") ?? ""}</div>
62
+ </div>
63
+ ))}
64
+ </div>
65
+ </div>
66
+ </div>
67
+ );
68
+ }
frontend/src/hooks/useRun.ts ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { useEffect, useRef } from "react";
2
+ import { useStore } from "../store";
3
+ import { runModel, streamRun } from "../lib/api";
4
+
5
+ export function useRunController() {
6
+ const wsRef = useRef<WebSocket | null>(null);
7
+ const setRunning = useStore((s) => s.setRunning);
8
+ const resetRun = useStore((s) => s.resetRun);
9
+ const appendEvent = useStore((s) => s.appendEvent);
10
+ const setError = useStore((s) => s.setError);
11
+
12
+ async function start(prompt: string | null) {
13
+ resetRun();
14
+ setRunning(true);
15
+ setError(null);
16
+ try {
17
+ const { run_id } = await runModel(prompt);
18
+ wsRef.current = streamRun(
19
+ run_id,
20
+ (ev) => appendEvent(ev),
21
+ () => setRunning(false),
22
+ (e) => {
23
+ console.error("ws error", e);
24
+ setError("websocket error");
25
+ setRunning(false);
26
+ },
27
+ );
28
+ } catch (e: any) {
29
+ setError(e?.message ?? String(e));
30
+ setRunning(false);
31
+ }
32
+ }
33
+
34
+ useEffect(() => {
35
+ return () => {
36
+ wsRef.current?.close();
37
+ };
38
+ }, []);
39
+
40
+ return { start };
41
+ }
42
+
43
+ /** Drives playback: advances `currentStep` over time when `playing` is true. */
44
+ export function usePlayback() {
45
+ const playing = useStore((s) => s.playing);
46
+ const speed = useStore((s) => s.speed);
47
+ const events = useStore((s) => s.events);
48
+ const setCurrentStep = useStore((s) => s.setCurrentStep);
49
+ const setPlaying = useStore((s) => s.setPlaying);
50
+ const stepRef = useRef(0);
51
+ const currentStep = useStore((s) => s.currentStep);
52
+
53
+ useEffect(() => {
54
+ stepRef.current = currentStep;
55
+ }, [currentStep]);
56
+
57
+ useEffect(() => {
58
+ if (!playing) return;
59
+ const interval = Math.max(40, 250 / speed);
60
+ const id = window.setInterval(() => {
61
+ const next = stepRef.current + 1;
62
+ if (next >= events.length) {
63
+ setPlaying(false);
64
+ return;
65
+ }
66
+ stepRef.current = next;
67
+ setCurrentStep(next);
68
+ }, interval);
69
+ return () => window.clearInterval(id);
70
+ }, [playing, speed, events.length, setCurrentStep, setPlaying]);
71
+ }
frontend/src/lib/api.ts ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import type { Graph, LoadResponse, RunEvent } from "../types";
2
+
3
+ async function jsonOrThrow<T>(res: Response): Promise<T> {
4
+ if (!res.ok) {
5
+ const t = await res.text();
6
+ throw new Error(t || `${res.status} ${res.statusText}`);
7
+ }
8
+ return res.json();
9
+ }
10
+
11
+ export async function loadModel(model_id: string): Promise<LoadResponse> {
12
+ const res = await fetch("/api/load", {
13
+ method: "POST",
14
+ headers: { "content-type": "application/json" },
15
+ body: JSON.stringify({ model_id }),
16
+ });
17
+ return jsonOrThrow<LoadResponse>(res);
18
+ }
19
+
20
+ export async function fetchGraph(): Promise<Graph> {
21
+ const res = await fetch("/api/graph");
22
+ return jsonOrThrow<Graph>(res);
23
+ }
24
+
25
+ export async function runModel(prompt: string | null): Promise<{ run_id: string; steps: number; elapsed_s: number }> {
26
+ const res = await fetch("/api/run", {
27
+ method: "POST",
28
+ headers: { "content-type": "application/json" },
29
+ body: JSON.stringify({ prompt }),
30
+ });
31
+ return jsonOrThrow(res);
32
+ }
33
+
34
+ export function streamRun(run_id: string, onEvent: (e: RunEvent) => void, onEnd: () => void, onError: (e: Event) => void): WebSocket {
35
+ const proto = location.protocol === "https:" ? "wss:" : "ws:";
36
+ const ws = new WebSocket(`${proto}//${location.host}/ws/run/${run_id}`);
37
+ ws.onmessage = (m) => {
38
+ try {
39
+ const data = JSON.parse(m.data);
40
+ if (data.kind === "stream_end") {
41
+ onEnd();
42
+ ws.close();
43
+ return;
44
+ }
45
+ onEvent(data as RunEvent);
46
+ } catch {
47
+ // ignore malformed
48
+ }
49
+ };
50
+ ws.onerror = onError;
51
+ ws.onclose = onEnd;
52
+ return ws;
53
+ }
frontend/src/lib/colors.ts ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ export const KIND_COLORS: Record<string, { bg: string; border: string; text: string; accent: string }> = {
2
+ Embedding: { bg: "#1e293b", border: "#38bdf8", text: "#e0f2fe", accent: "#38bdf8" },
3
+ Attention: { bg: "#1e1b4b", border: "#a78bfa", text: "#ede9fe", accent: "#a78bfa" },
4
+ MLP: { bg: "#172554", border: "#60a5fa", text: "#dbeafe", accent: "#60a5fa" },
5
+ Norm: { bg: "#1f2937", border: "#9ca3af", text: "#f3f4f6", accent: "#9ca3af" },
6
+ Linear: { bg: "#0f172a", border: "#22d3ee", text: "#cffafe", accent: "#22d3ee" },
7
+ Conv: { bg: "#0c4a6e", border: "#fbbf24", text: "#fef3c7", accent: "#fbbf24" },
8
+ Activation: { bg: "#3f3f46", border: "#facc15", text: "#fef9c3", accent: "#facc15" },
9
+ Dropout: { bg: "#27272a", border: "#a1a1aa", text: "#fafafa", accent: "#a1a1aa" },
10
+ Pooler: { bg: "#1e3a8a", border: "#34d399", text: "#d1fae5", accent: "#34d399" },
11
+ Block: { bg: "#3b0764", border: "#f472b6", text: "#fce7f3", accent: "#f472b6" },
12
+ Head: { bg: "#7c2d12", border: "#fb923c", text: "#ffedd5", accent: "#fb923c" },
13
+ Other: { bg: "#262626", border: "#737373", text: "#e5e5e5", accent: "#737373" },
14
+ };
15
+
16
+ export function colorFor(kind: string) {
17
+ return KIND_COLORS[kind] ?? KIND_COLORS.Other;
18
+ }
19
+
20
+ export function viridis(t: number): string {
21
+ // Approximate viridis colormap. t in [0,1].
22
+ const stops: [number, number, number][] = [
23
+ [68, 1, 84],
24
+ [59, 82, 139],
25
+ [33, 145, 140],
26
+ [94, 201, 98],
27
+ [253, 231, 37],
28
+ ];
29
+ const x = Math.max(0, Math.min(1, t)) * (stops.length - 1);
30
+ const i = Math.floor(x);
31
+ const f = x - i;
32
+ const a = stops[i];
33
+ const b = stops[Math.min(i + 1, stops.length - 1)];
34
+ const r = Math.round(a[0] + (b[0] - a[0]) * f);
35
+ const g = Math.round(a[1] + (b[1] - a[1]) * f);
36
+ const bl = Math.round(a[2] + (b[2] - a[2]) * f);
37
+ return `rgb(${r},${g},${bl})`;
38
+ }
frontend/src/lib/layout.ts ADDED
@@ -0,0 +1,64 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import dagre from "@dagrejs/dagre";
2
+ import type { Edge, Node } from "reactflow";
3
+ import type { Graph } from "../types";
4
+
5
+ const NODE_W = 220;
6
+ const NODE_H = 64;
7
+
8
+ export type LayoutOptions = {
9
+ granularity: number; // 0 = roots only, 1 = blocks, 2 = leaves
10
+ };
11
+
12
+ const DEPTH_FOR_GRAN: Record<number, number> = { 0: 1, 1: 3, 2: 99 };
13
+
14
+ export function layoutGraph(graph: Graph, opts: LayoutOptions): { nodes: Node[]; edges: Edge[] } {
15
+ const maxDepth = DEPTH_FOR_GRAN[opts.granularity] ?? 3;
16
+ const visible = new Set<string>();
17
+ for (const n of graph.nodes) {
18
+ if (n.id === "<root>") continue;
19
+ if (n.depth <= maxDepth) visible.add(n.id);
20
+ }
21
+ // If no leaves at this depth, surface their nearest visible ancestors instead.
22
+ if (visible.size === 0) {
23
+ for (const n of graph.nodes) if (n.id !== "<root>") visible.add(n.id);
24
+ }
25
+
26
+ const idToNode = new Map(graph.nodes.map((n) => [n.id, n]));
27
+ const g = new dagre.graphlib.Graph();
28
+ g.setDefaultEdgeLabel(() => ({}));
29
+ g.setGraph({ rankdir: "TB", nodesep: 30, ranksep: 50 });
30
+ visible.forEach((id) => g.setNode(id, { width: NODE_W, height: NODE_H }));
31
+
32
+ // Build parent→child edges only between visible nodes.
33
+ const edges: Edge[] = [];
34
+ for (const e of graph.edges) {
35
+ if (e.from === "<root>") continue;
36
+ if (visible.has(e.from) && visible.has(e.to)) {
37
+ g.setEdge(e.from, e.to);
38
+ edges.push({
39
+ id: `${e.from}->${e.to}`,
40
+ source: e.from,
41
+ target: e.to,
42
+ type: "smoothstep",
43
+ animated: false,
44
+ style: { stroke: "#6b7280", strokeWidth: 1.5 },
45
+ });
46
+ }
47
+ }
48
+
49
+ dagre.layout(g);
50
+
51
+ const nodes: Node[] = [];
52
+ visible.forEach((id) => {
53
+ const layout = g.node(id);
54
+ const data = idToNode.get(id);
55
+ if (!data || !layout) return;
56
+ nodes.push({
57
+ id,
58
+ type: "module",
59
+ position: { x: layout.x - NODE_W / 2, y: layout.y - NODE_H / 2 },
60
+ data: { ...data },
61
+ });
62
+ });
63
+ return { nodes, edges };
64
+ }
frontend/src/main.tsx ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import React from "react";
2
+ import ReactDOM from "react-dom/client";
3
+ import App from "./App";
4
+ import "reactflow/dist/style.css";
5
+ import "./styles.css";
6
+
7
+ ReactDOM.createRoot(document.getElementById("root")!).render(
8
+ <React.StrictMode>
9
+ <App />
10
+ </React.StrictMode>,
11
+ );
frontend/src/store.ts ADDED
@@ -0,0 +1,93 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { create } from "zustand";
2
+ import type { Graph, LoadResponse, RunEvent, Snapshot } from "./types";
3
+
4
+ type SnapshotMap = Record<string, Snapshot>;
5
+
6
+ type State = {
7
+ modelId: string;
8
+ loadInfo: LoadResponse | null;
9
+ graph: Graph | null;
10
+ events: RunEvent[];
11
+ snapshotsByPath: SnapshotMap;
12
+ attentionByLayer: Record<number, any>;
13
+ inputUI: any | null;
14
+ topK: any[] | null;
15
+ currentStep: number;
16
+ playing: boolean;
17
+ speed: number;
18
+ selectedPath: string | null;
19
+ loading: boolean;
20
+ running: boolean;
21
+ error: string | null;
22
+ granularity: number; // 0 top-level, 1 blocks, 2 leaves
23
+
24
+ setModelId: (s: string) => void;
25
+ setLoading: (b: boolean) => void;
26
+ setRunning: (b: boolean) => void;
27
+ setError: (s: string | null) => void;
28
+ setLoadInfo: (l: LoadResponse | null) => void;
29
+ setGraph: (g: Graph | null) => void;
30
+ resetRun: () => void;
31
+ appendEvent: (ev: RunEvent) => void;
32
+ setCurrentStep: (n: number) => void;
33
+ setPlaying: (b: boolean) => void;
34
+ setSpeed: (n: number) => void;
35
+ setSelectedPath: (p: string | null) => void;
36
+ setGranularity: (n: number) => void;
37
+ };
38
+
39
+ export const useStore = create<State>((set) => ({
40
+ modelId: "prajjwal1/bert-tiny",
41
+ loadInfo: null,
42
+ graph: null,
43
+ events: [],
44
+ snapshotsByPath: {},
45
+ attentionByLayer: {},
46
+ inputUI: null,
47
+ topK: null,
48
+ currentStep: 0,
49
+ playing: false,
50
+ speed: 1,
51
+ selectedPath: null,
52
+ loading: false,
53
+ running: false,
54
+ error: null,
55
+ granularity: 1,
56
+
57
+ setModelId: (s) => set({ modelId: s }),
58
+ setLoading: (b) => set({ loading: b }),
59
+ setRunning: (b) => set({ running: b }),
60
+ setError: (s) => set({ error: s }),
61
+ setLoadInfo: (l) => set({ loadInfo: l }),
62
+ setGraph: (g) => set({ graph: g }),
63
+ resetRun: () =>
64
+ set({
65
+ events: [],
66
+ snapshotsByPath: {},
67
+ attentionByLayer: {},
68
+ inputUI: null,
69
+ topK: null,
70
+ currentStep: 0,
71
+ playing: false,
72
+ }),
73
+ appendEvent: (ev) =>
74
+ set((s) => {
75
+ const next: Partial<State> = { events: [...s.events, ev] };
76
+ if (ev.kind === "module" && ev.payload?.snapshot) {
77
+ next.snapshotsByPath = { ...s.snapshotsByPath, [ev.path]: ev.payload.snapshot };
78
+ } else if (ev.kind === "attention") {
79
+ next.attentionByLayer = { ...s.attentionByLayer, [ev.payload.layer]: ev.payload.snapshot };
80
+ } else if (ev.kind === "input") {
81
+ next.inputUI = ev.payload;
82
+ } else if (ev.kind === "logits") {
83
+ next.topK = ev.payload.top_k;
84
+ next.snapshotsByPath = { ...s.snapshotsByPath, [ev.path]: ev.payload.snapshot };
85
+ }
86
+ return next as State;
87
+ }),
88
+ setCurrentStep: (n) => set({ currentStep: n }),
89
+ setPlaying: (b) => set({ playing: b }),
90
+ setSpeed: (n) => set({ speed: n }),
91
+ setSelectedPath: (p) => set({ selectedPath: p }),
92
+ setGranularity: (n) => set({ granularity: n }),
93
+ }));
frontend/src/styles.css ADDED
@@ -0,0 +1,68 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ * { box-sizing: border-box; }
2
+ html, body, #root { height: 100%; margin: 0; }
3
+ body {
4
+ font-family: ui-sans-serif, system-ui, -apple-system, Segoe UI, Roboto, sans-serif;
5
+ background: #0b1220;
6
+ color: #e2e8f0;
7
+ }
8
+
9
+ button { font-family: inherit; }
10
+
11
+ .btn {
12
+ background: #1e293b;
13
+ color: #e2e8f0;
14
+ border: 1px solid #334155;
15
+ border-radius: 6px;
16
+ padding: 5px 12px;
17
+ font-size: 12px;
18
+ cursor: pointer;
19
+ transition: background 120ms, border-color 120ms;
20
+ }
21
+ .btn:hover:not(:disabled) { background: #334155; border-color: #475569; }
22
+ .btn:disabled { opacity: 0.5; cursor: not-allowed; }
23
+
24
+ .btn-primary {
25
+ background: #38bdf8;
26
+ color: #0b1220;
27
+ border-color: #0ea5e9;
28
+ font-weight: 600;
29
+ }
30
+ .btn-primary:hover:not(:disabled) { background: #0ea5e9; }
31
+
32
+ .tab {
33
+ background: transparent;
34
+ color: #94a3b8;
35
+ border: 1px solid #334155;
36
+ border-radius: 4px;
37
+ padding: 3px 8px;
38
+ font-size: 11px;
39
+ cursor: pointer;
40
+ transition: background 120ms, color 120ms;
41
+ }
42
+ .tab:hover { color: #e2e8f0; }
43
+ .tab-active { background: #38bdf8; color: #0b1220; border-color: #0ea5e9; font-weight: 600; }
44
+
45
+ .pill {
46
+ background: #1e293b;
47
+ color: #cbd5e1;
48
+ border: 1px solid #334155;
49
+ border-radius: 999px;
50
+ padding: 2px 8px;
51
+ font-size: 10px;
52
+ font-family: ui-monospace, monospace;
53
+ }
54
+
55
+ input[type="number"], input[type="text"] {
56
+ background: #0b1220;
57
+ color: #e2e8f0;
58
+ border: 1px solid #334155;
59
+ border-radius: 4px;
60
+ padding: 3px 6px;
61
+ font-family: ui-monospace, monospace;
62
+ font-size: 11px;
63
+ }
64
+
65
+ ::-webkit-scrollbar { width: 8px; height: 8px; }
66
+ ::-webkit-scrollbar-track { background: #0b1220; }
67
+ ::-webkit-scrollbar-thumb { background: #334155; border-radius: 4px; }
68
+ ::-webkit-scrollbar-thumb:hover { background: #475569; }
frontend/src/types.ts ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ export type Heatmap = {
2
+ shape: [number, number];
3
+ values: number[][];
4
+ vmin: number;
5
+ vmax: number;
6
+ };
7
+
8
+ export type Snapshot = {
9
+ shape: number[];
10
+ dtype: string;
11
+ stats: { mean: number; std: number; min: number; max: number; abs_max: number };
12
+ sample: number[];
13
+ heatmap: Heatmap | null;
14
+ };
15
+
16
+ export type GraphNode = {
17
+ id: string;
18
+ path: string;
19
+ module_class: string;
20
+ kind: string;
21
+ params: number;
22
+ params_own: number;
23
+ config: Record<string, unknown>;
24
+ children: string[];
25
+ depth: number;
26
+ is_leaf: boolean;
27
+ };
28
+
29
+ export type GraphEdge = { from: string; to: string; kind: string };
30
+
31
+ export type Graph = { root: string; nodes: GraphNode[]; edges: GraphEdge[] };
32
+
33
+ export type RunEventKind = "input" | "module" | "attention" | "logits" | "done" | "stream_end";
34
+
35
+ export type RunEvent = {
36
+ step: number;
37
+ path: string;
38
+ module_class: string;
39
+ kind: RunEventKind;
40
+ payload: any;
41
+ };
42
+
43
+ export type LoadResponse = {
44
+ model_id: string;
45
+ modality: "text" | "vision" | "audio";
46
+ head_kind: string;
47
+ param_count: number;
48
+ config: Record<string, any>;
49
+ tokenizer: null | {
50
+ name_or_path: string;
51
+ vocab_size: number;
52
+ model_max_length: number;
53
+ special_tokens: string[];
54
+ is_fast: boolean;
55
+ };
56
+ processor: string | null;
57
+ };
frontend/src/vite-env.d.ts ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ /// <reference types="vite/client" />
2
+ declare module "gif.js";
frontend/tsconfig.json ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "compilerOptions": {
3
+ "target": "ES2020",
4
+ "useDefineForClassFields": true,
5
+ "lib": ["ES2020", "DOM", "DOM.Iterable"],
6
+ "module": "ESNext",
7
+ "skipLibCheck": true,
8
+ "moduleResolution": "bundler",
9
+ "allowImportingTsExtensions": true,
10
+ "resolveJsonModule": true,
11
+ "isolatedModules": true,
12
+ "noEmit": true,
13
+ "jsx": "react-jsx",
14
+ "strict": true,
15
+ "noUnusedLocals": false,
16
+ "noUnusedParameters": false,
17
+ "noFallthroughCasesInSwitch": true
18
+ },
19
+ "include": ["src"]
20
+ }
frontend/vite.config.ts ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { defineConfig } from "vite";
2
+ import react from "@vitejs/plugin-react";
3
+
4
+ export default defineConfig({
5
+ plugins: [react()],
6
+ server: {
7
+ port: 5173,
8
+ proxy: {
9
+ "/api": "http://localhost:7860",
10
+ "/ws": { target: "ws://localhost:7860", ws: true },
11
+ },
12
+ },
13
+ build: {
14
+ outDir: "dist",
15
+ sourcemap: false,
16
+ },
17
+ });
requirements.txt ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ --extra-index-url https://download.pytorch.org/whl/cpu
2
+ torch==2.4.1
3
+ transformers==4.46.3
4
+ tokenizers>=0.20
5
+ safetensors>=0.4
6
+ huggingface_hub>=0.26
7
+ fastapi==0.115.5
8
+ uvicorn[standard]==0.32.1
9
+ pydantic>=2.8
10
+ numpy>=1.26
11
+ pillow>=10.4
12
+ sentencepiece>=0.2
13
+ protobuf>=4.25