Spaces:
Running
Running
initial commit: transformer x-ray (FastAPI + React, GIF export)
Browse files- .dockerignore +10 -0
- .gitignore +14 -0
- Dockerfile +28 -0
- README.md +38 -5
- backend/__init__.py +0 -0
- backend/main.py +145 -0
- backend/modality.py +110 -0
- backend/model_loader.py +146 -0
- backend/runner.py +174 -0
- backend/snapshots.py +124 -0
- backend/tracer.py +130 -0
- frontend/index.html +12 -0
- frontend/package.json +27 -0
- frontend/src/App.tsx +36 -0
- frontend/src/components/GifExport.tsx +99 -0
- frontend/src/components/Graph.tsx +45 -0
- frontend/src/components/Header.tsx +104 -0
- frontend/src/components/ModuleNode.tsx +63 -0
- frontend/src/components/NodeCard.tsx +64 -0
- frontend/src/components/PlayBar.tsx +118 -0
- frontend/src/components/ProbsBar.tsx +55 -0
- frontend/src/components/TensorView.tsx +122 -0
- frontend/src/components/TokenStrip.tsx +68 -0
- frontend/src/hooks/useRun.ts +71 -0
- frontend/src/lib/api.ts +53 -0
- frontend/src/lib/colors.ts +38 -0
- frontend/src/lib/layout.ts +64 -0
- frontend/src/main.tsx +11 -0
- frontend/src/store.ts +93 -0
- frontend/src/styles.css +68 -0
- frontend/src/types.ts +57 -0
- frontend/src/vite-env.d.ts +2 -0
- frontend/tsconfig.json +20 -0
- frontend/vite.config.ts +17 -0
- requirements.txt +13 -0
.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
|
| 3 |
-
emoji:
|
| 4 |
-
colorFrom:
|
| 5 |
-
colorTo:
|
| 6 |
sdk: docker
|
|
|
|
| 7 |
pinned: false
|
|
|
|
|
|
|
| 8 |
---
|
| 9 |
|
| 10 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|