Download scripts/journal_server.py from dejanseo/fanout-diffusion: direct link, hf CLI and curl.
- Browser
- Download file 16.1 kB
-
https://huggingface.co/dejanseo/fanout-diffusion/resolve/main/scripts/journal_server.py
- Command line
-
hf download hf://dejanseo/fanout-diffusion/scripts/journal_server.py
-
curl -L -o journal_server.py https://huggingface.co/dejanseo/fanout-diffusion/resolve/main/scripts/journal_server.py
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 | |
| 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, | |
| } | |
| 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 | |
| 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 | |
| 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 | |
| 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) | |