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 "