fanout-diffusion / scripts /journal_server.py
dejanseo's picture
Add training pipelines, consistency distillation scripts, and interactive dashboard server
69dcf19 verified
Raw History Blame Contribute Delete
16.1 kB
import json
import os
import sqlite3
from pathlib import Path
from typing import Any, Dict, List, Optional
from fastapi import FastAPI, HTTPException, Query
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import HTMLResponse, JSONResponse
from fastapi.staticfiles import StaticFiles
import sys
ROOT_DIR = Path(__file__).resolve().parents[1]
if str(ROOT_DIR) not in sys.path:
sys.path.insert(0, str(ROOT_DIR))
DB_PATH = os.path.join(ROOT_DIR, "journal.db")
STATIC_DIR = os.path.join(ROOT_DIR, "static")
app = FastAPI(title="Diffusion Fan-Out Experiment Journal", version="1.0.0")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
def get_db():
conn = sqlite3.connect(DB_PATH)
conn.row_factory = sqlite3.Row
return conn
def try_parse_json(payload: Optional[str]) -> Any:
if not payload:
return None
try:
return json.loads(payload)
except Exception:
return payload
@app.get("/api/overview")
def get_overview():
with get_db() as conn:
cur = conn.cursor()
cur.execute("PRAGMA journal_mode")
journal_mode = cur.fetchone()[0]
cur.execute("SELECT COUNT(*) FROM runs")
total_runs = cur.fetchone()[0]
cur.execute("SELECT COUNT(*) FROM experiments")
total_experiments = cur.fetchone()[0]
cur.execute("SELECT COUNT(*) FROM metrics")
total_metrics = cur.fetchone()[0]
cur.execute("SELECT COUNT(*) FROM benchmarks")
total_benchmarks = cur.fetchone()[0]
cur.execute(
"""
SELECT r.id, r.name, r.task_type, r.status, r.tags_json, r.summary_metrics_json, r.config_json,
b.latency_us, b.throughput_items_per_sec, b.ns_per_neuron
FROM runs r
LEFT JOIN benchmarks b ON r.id = b.run_id
ORDER BY r.started_at DESC
"""
)
runs_data = []
for row in cur.fetchall():
runs_data.append(
{
"id": row["id"],
"name": row["name"],
"task_type": row["task_type"],
"status": row["status"],
"tags": try_parse_json(row["tags_json"]) or [],
"summary": try_parse_json(row["summary_metrics_json"]) or {},
"config": try_parse_json(row["config_json"]) or {},
"latency_us": row["latency_us"],
"throughput": row["throughput_items_per_sec"],
"ns_per_neuron": row["ns_per_neuron"],
}
)
# Collect loss series across all completed/running runs
cur.execute(
"""
SELECT r.id, r.name, m.epoch, m.step, m.metric_name, m.value
FROM runs r
JOIN metrics m ON r.id = m.run_id
WHERE m.metric_name IN ('val_loss', 'loss', 'train_loss')
ORDER BY r.id, m.step ASC
"""
)
loss_series = {}
for row in cur.fetchall():
rid = row["id"]
mname = row["metric_name"]
if rid not in loss_series:
loss_series[rid] = {
"id": rid,
"name": row["name"],
"val_loss": {"epochs": [], "steps": [], "values": []},
"train_loss": {"epochs": [], "steps": [], "values": []},
}
target_key = "val_loss" if "val" in mname else "train_loss"
loss_series[rid][target_key]["epochs"].append(row["epoch"])
loss_series[rid][target_key]["steps"].append(row["step"])
loss_series[rid][target_key]["values"].append(row["value"])
return {
"journal_mode": journal_mode,
"total_runs": total_runs,
"total_experiments": total_experiments,
"total_metrics": total_metrics,
"total_benchmarks": total_benchmarks,
"loss_series": list(loss_series.values()),
"runs": runs_data,
}
@app.get("/api/runs")
def get_runs(task_type: Optional[str] = None):
with get_db() as conn:
cur = conn.cursor()
query = """
SELECT r.*, e.name as experiment_name,
b.latency_us, b.throughput_items_per_sec, b.ns_per_neuron
FROM runs r
LEFT JOIN experiments e ON r.experiment_id = e.id
LEFT JOIN benchmarks b ON r.id = b.run_id
"""
params = []
if task_type and task_type != "all":
query += " WHERE r.task_type = ?"
params.append(task_type)
query += " ORDER BY r.started_at DESC"
cur.execute(query, params)
rows = cur.fetchall()
results = []
for row in rows:
results.append(
{
"id": row["id"],
"experiment_id": row["experiment_id"],
"experiment_name": row["experiment_name"],
"name": row["name"],
"task_type": row["task_type"],
"status": row["status"],
"config": try_parse_json(row["config_json"]) or {},
"tags": try_parse_json(row["tags_json"]) or [],
"summary": try_parse_json(row["summary_metrics_json"]) or {},
"latency_us": row["latency_us"],
"throughput": row["throughput_items_per_sec"],
"started_at": row["started_at"],
"completed_at": row["completed_at"],
}
)
return results
@app.get("/api/runs/{run_id}")
def get_run_detail(run_id: str):
with get_db() as conn:
cur = conn.cursor()
cur.execute(
"""
SELECT r.*, e.name as experiment_name, e.description as experiment_description,
b.latency_us, b.throughput_items_per_sec, b.ns_per_neuron, b.notes as benchmark_notes, b.device_name
FROM runs r
LEFT JOIN experiments e ON r.experiment_id = e.id
LEFT JOIN benchmarks b ON r.id = b.run_id
WHERE r.id = ?
""",
(run_id,),
)
row = cur.fetchone()
if not row:
raise HTTPException(status_code=404, detail="Run not found")
cur.execute(
"""
SELECT step, epoch, metric_name, value, timestamp
FROM metrics
WHERE run_id = ?
ORDER BY step ASC, id ASC
""",
(run_id,),
)
metric_rows = cur.fetchall()
metrics = {}
for mr in metric_rows:
m_name = mr["metric_name"]
if m_name not in metrics:
metrics[m_name] = {"steps": [], "epochs": [], "values": []}
metrics[m_name]["steps"].append(mr["step"])
metrics[m_name]["epochs"].append(mr["epoch"])
metrics[m_name]["values"].append(mr["value"])
cur.execute(
"""
SELECT id, step, sample_type, input_payload, output_payload, target_payload, metadata_json, timestamp
FROM samples
WHERE run_id = ?
ORDER BY step ASC
""",
(run_id,),
)
sample_rows = cur.fetchall()
samples = []
for sr in sample_rows:
samples.append(
{
"id": sr["id"],
"step": sr["step"],
"sample_type": sr["sample_type"],
"input": try_parse_json(sr["input_payload"]),
"output": try_parse_json(sr["output_payload"]),
"target": try_parse_json(sr["target_payload"]),
"metadata": try_parse_json(sr["metadata_json"]),
"timestamp": sr["timestamp"],
}
)
return {
"id": row["id"],
"name": row["name"],
"task_type": row["task_type"],
"status": row["status"],
"experiment_name": row["experiment_name"],
"experiment_description": row["experiment_description"],
"config": try_parse_json(row["config_json"]) or {},
"tags": try_parse_json(row["tags_json"]) or [],
"summary": try_parse_json(row["summary_metrics_json"]) or {},
"benchmark": (
{
"latency_us": row["latency_us"],
"throughput": row["throughput_items_per_sec"],
"ns_per_neuron": row["ns_per_neuron"],
"device_name": row["device_name"],
"notes": row["benchmark_notes"],
}
if row["latency_us"] is not None
else None
),
"started_at": row["started_at"],
"completed_at": row["completed_at"],
"metrics": metrics,
"samples": samples,
}
# ---------------------------------------------------------------------------
# Live Playground Engine & Endpoint (1.64 MB INT4 1-Step Model)
# ---------------------------------------------------------------------------
from pydantic import BaseModel
class PlaygroundRequest(BaseModel):
query: str
_playground_engine = None
def get_playground_engine():
global _playground_engine
if _playground_engine is not None:
return _playground_engine
import torch
import torch.nn.functional as F
from sentence_transformers import SentenceTransformer
from scripts.fast_b1_inference import FastB1Denoiser
from scripts.quantize_outer_int4 import unpack_int4_signed
from src.r4t.b1_diffusion import B1EDMDenoiser
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
ckpt_path = ROOT_DIR / "checkpoints" / "champion_b1_consistency_1step_qat.pt"
if not ckpt_path.exists():
ckpt_path = ROOT_DIR / "checkpoints" / "champion_b1_consistency_1step_int4.pt"
if not ckpt_path.exists():
ckpt_path = ROOT_DIR / "checkpoints" / "champion_b1_consistency_1step.pt"
ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
config = ckpt["config"]
model = B1EDMDenoiser(config, backend="tc", pure_1bit=False).to(device)
model.freeze_for_inference()
state = model.state_dict()
if "weights" in ckpt:
for k, v in ckpt["weights"].items():
if k in state:
state[k].copy_(v.to(device))
if "int4_outer" in ckpt:
for k, d in ckpt["int4_outer"].items():
state[k].copy_(unpack_int4_signed(d["packed"].to(device), d["scale"].to(device)))
elif "model_state_dict" in ckpt:
model.load_state_dict(ckpt["model_state_dict"], strict=False)
model.eval()
fast_model = FastB1Denoiser(model)
embed_kwargs = {"torch_dtype": torch.bfloat16} if device.type == "cuda" else {}
embedder = SentenceTransformer("google/embeddinggemma-300m", model_kwargs=embed_kwargs, device=device)
tax_path = ROOT_DIR / "data" / "taxonomy_embeddings.pt"
tax_emb = None
tax_names = None
if tax_path.exists():
tax_data = torch.load(tax_path, map_location="cpu", weights_only=False)
tax_emb = F.normalize(tax_data["embeddings"].float(), dim=-1).to(device)
tax_names = tax_data["names"]
subq_path = ROOT_DIR / "data" / "subqueries_index.pt"
subq_emb = None
subqueries = None
if subq_path.exists():
subq_data = torch.load(subq_path, map_location="cpu", weights_only=False)
subq_emb = subq_data["embeddings"].to(device)
subqueries = subq_data["subqueries"]
_playground_engine = {
"model": fast_model,
"embedder": embedder,
"tax_emb": tax_emb,
"tax_names": tax_names,
"subq_emb": subq_emb,
"subqueries": subqueries,
"device": device,
}
return _playground_engine
@app.post("/api/playground/fanout")
def run_fanout(req: PlaygroundRequest):
import time
import torch
import torch.nn.functional as F
query_text = req.query.strip()
if not query_text:
raise HTTPException(status_code=400, detail="Query cannot be empty")
engine = get_playground_engine()
device = engine["device"]
fast_model = engine["model"]
embedder = engine["embedder"]
tax_emb = engine["tax_emb"]
tax_names = engine["tax_names"]
subq_emb = engine.get("subq_emb")
subqueries = engine.get("subqueries")
# 1. Embed query
t_emb_start = time.perf_counter()
z_q = embedder.encode([query_text], convert_to_tensor=True, normalize_embeddings=True, device=device).float()
embed_ms = (time.perf_counter() - t_emb_start) * 1000.0
# 2. 1-Step Fanout Generation
cached_mem = fast_model.precompute_cross_memory(z_q)
shape = (1, fast_model.config.sequence_length, fast_model.config.embedding_dim)
noise = torch.randn(shape, device=device) * fast_model.config.sigma_max
sigma = torch.full((1,), fast_model.config.sigma_max, device=device)
if device.type == "cuda":
torch.cuda.synchronize()
t_gen_start = time.perf_counter()
fanout = fast_model.forward_with_cached_memory(noise, sigma, cached_mem)[0]
if device.type == "cuda":
torch.cuda.synchronize()
gen_ms = (time.perf_counter() - t_gen_start) * 1000.0
fanout_norm = F.normalize(fanout, dim=-1)
# 3. Prompt alignment & pairwise diversity
align_scores = (fanout_norm @ z_q.T).squeeze(-1).tolist()
mean_align = sum(align_scores) / len(align_scores)
sim_mat = fanout_norm @ fanout_norm.T
mask = ~torch.eye(10, dtype=torch.bool, device=device)
diversity = 1.0 - sim_mat[mask].mean().item()
# 4. Taxonomy decoding
slots = []
if tax_emb is not None:
top_matches = fanout_norm @ tax_emb.T
top_indices = top_matches.argmax(dim=-1).tolist()
for idx, t_idx in enumerate(top_indices):
sim_score = top_matches[idx, t_idx].item()
slots.append({
"slot": idx + 1,
"category": tax_names[t_idx],
"similarity": round(sim_score, 3),
"alignment": round(align_scores[idx], 3),
})
else:
for idx in range(10):
slots.append({
"slot": idx + 1,
"category": f"Subquery Cluster #{idx+1}",
"similarity": 1.0,
"alignment": round(align_scores[idx], 3),
})
# 5. Natural subquery decoding (540k index)
if subq_emb is not None and subqueries is not None:
subq_matches = fanout_norm.half() @ subq_emb.T
top_subq_idx = subq_matches.argmax(dim=-1).tolist()
for idx, sq_idx in enumerate(top_subq_idx):
slots[idx]["subquery"] = subqueries[sq_idx]
slots[idx]["subquery_similarity"] = round(subq_matches[idx, sq_idx].item(), 3)
return {
"query": query_text,
"latency_ms": round(gen_ms, 3),
"embed_latency_ms": round(embed_ms, 2),
"total_latency_ms": round(gen_ms + embed_ms, 2),
"mean_alignment": round(mean_align, 3),
"pairwise_diversity": round(diversity, 3),
"slots": slots,
"model_info": {
"checkpoint": "champion_b1_consistency_1step_int4.pt",
"size_mb": 1.64,
"sampling_steps": 1,
"hardware": "RTX 4090 Tensor Cores (PTX MMA)",
}
}
# Frontend routes
@app.get("/", response_class=HTMLResponse)
@app.get("/overview", response_class=HTMLResponse)
@app.get("/runs", response_class=HTMLResponse)
@app.get("/drilldown", response_class=HTMLResponse)
@app.get("/playground", response_class=HTMLResponse)
@app.get("/run/{run_id}", response_class=HTMLResponse)
def serve_index(run_id: Optional[str] = None):
index_file = os.path.join(STATIC_DIR, "index.html")
if os.path.exists(index_file):
with open(index_file, "r", encoding="utf-8") as f:
return f.read()
return "<h1>Journal static/index.html not found</h1>"
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="127.0.0.1", port=8000)