Download scripts/benchmark_f7.py from dejanseo/fanout-diffusion: direct link, hf CLI and curl.
- Browser
- Download file 5.59 kB
-
https://huggingface.co/dejanseo/fanout-diffusion/resolve/main/scripts/benchmark_f7.py
- Command line
-
hf download hf://dejanseo/fanout-diffusion/scripts/benchmark_f7.py
-
curl -L -o benchmark_f7.py https://huggingface.co/dejanseo/fanout-diffusion/resolve/main/scripts/benchmark_f7.py
5.59 kB
| import time | |
| import sys | |
| from pathlib import Path | |
| ROOT = Path(__file__).resolve().parent.parent | |
| if str(ROOT) not in sys.path: | |
| sys.path.insert(0, str(ROOT)) | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| import onnxruntime as ort | |
| from sentence_transformers import SentenceTransformer | |
| from src.r4t.b1_diffusion import B1EDMDenoiser | |
| QUERIES = [ | |
| "wireless noise cancelling headphones", | |
| "espresso coffee machines", | |
| "running shoes for marathon training", | |
| "smart home security cameras", | |
| "mountain bike repair and maintenance", | |
| "vintage mechanical wristwatches", | |
| "keto diet meal prep recipes", | |
| "dslr camera lenses for portrait photography", | |
| "electric guitar amplifiers and pedals", | |
| "ergonomic office chairs for lower back pain", | |
| ] | |
| def main(): | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| print(f"Running benchmark on {device} ({torch.cuda.get_device_name(0)})...\n") | |
| # 1. Load Embedder | |
| embedder = SentenceTransformer("google/embeddinggemma-300m", model_kwargs={"torch_dtype": torch.bfloat16}, device=device) | |
| # 2. Load Taxonomy Bank | |
| tax_data = torch.load("data/taxonomy_embeddings.pt", map_location="cpu", weights_only=False) | |
| tax_emb = F.normalize(tax_data["embeddings"].float(), dim=-1).to(device) | |
| tax_names = tax_data["names"] | |
| # 3. Load Natural Subquery Index | |
| sub_data = torch.load("data/subqueries_index.pt", map_location="cpu", weights_only=False) | |
| sub_texts = sub_data["subqueries"] | |
| sub_emb = F.normalize(sub_data["embeddings"].float(), dim=-1).to(device) | |
| # 4. Load Champion B1 (1-Step Baseline ONNX) | |
| providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] | |
| b1_session = ort.InferenceSession("checkpoints/champion_b1_consistency_1step.onnx", providers=providers) | |
| # 5. Load Frontier F7 (Calibrated Champion: 768d, 2048 MLP, Gram 0.3) | |
| ckpt_f7 = torch.load("checkpoints/frontier_f7_calibrated_champion.pt", map_location=device, weights_only=False) | |
| m_f7 = B1EDMDenoiser(ckpt_f7["config"], backend="tc", pure_1bit=False).to(device) | |
| m_f7.load_state_dict(ckpt_f7["model_state_dict"]) | |
| m_f7.eval() | |
| print("=" * 105) | |
| print(f"{'#':<3} | {'Query':<36} | {'Model':<22} | {'Align':<6} | {'Div':<6} | {'Latency':<7} | {'Decoded Taxonomy / Subqueries'}") | |
| print("=" * 105) | |
| b1_aligns, b1_divs = [], [] | |
| f7_aligns, f7_divs = [], [] | |
| mask = ~torch.eye(10, dtype=torch.bool, device=device) | |
| for idx, q_text in enumerate(QUERIES, 1): | |
| with torch.no_grad(): | |
| z_q_tensor = embedder.encode([q_text], convert_to_tensor=True, device=device, normalize_embeddings=True).float() | |
| z_q_np = z_q_tensor.cpu().numpy() | |
| # --- Champion B1 (Baseline) --- | |
| b1_noise = (np.random.randn(1, 10, 768) * 80.0).astype(np.float32) | |
| b1_sigma = np.full((1,), 80.0, dtype=np.float32) | |
| t0 = time.perf_counter() | |
| b1_out_np = b1_session.run(["denoised_fanout"], {"noisy_slots": b1_noise, "sigma": b1_sigma, "query_emb": z_q_np})[0][0] | |
| b1_ms = (time.perf_counter() - t0) * 1000.0 | |
| b1_out = torch.from_numpy(b1_out_np).to(device) | |
| b1_norm = F.normalize(b1_out, dim=-1) | |
| b1_align = (b1_norm @ z_q_tensor.T).mean().item() | |
| b1_sim = b1_norm @ b1_norm.T | |
| b1_div = 1.0 - b1_sim[mask].mean().item() | |
| b1_tax_idx = (b1_norm @ tax_emb.T).argmax(dim=-1).tolist() | |
| b1_tax_terms = list(dict.fromkeys([tax_names[i] for i in b1_tax_idx]))[:3] | |
| b1_sub_idx = (b1_norm @ sub_emb.T).argmax(dim=-1).tolist() | |
| b1_sub_terms = list(dict.fromkeys([sub_texts[i] for i in b1_sub_idx]))[:2] | |
| # --- Frontier F7 (Calibrated Champion) --- | |
| f7_noise = torch.randn(1, 10, 768, device=device) * 80.0 | |
| f7_sigma = torch.full((1,), 80.0, device=device) | |
| t0 = time.perf_counter() | |
| f7_out = m_f7(f7_noise, f7_sigma, z_q_tensor)[0] | |
| torch.cuda.synchronize() | |
| f7_ms = (time.perf_counter() - t0) * 1000.0 | |
| f7_norm = F.normalize(f7_out, dim=-1) | |
| f7_align = (f7_norm @ z_q_tensor.T).mean().item() | |
| f7_sim = f7_norm @ f7_norm.T | |
| f7_div = 1.0 - f7_sim[mask].mean().item() | |
| f7_tax_idx = (f7_norm @ tax_emb.T).argmax(dim=-1).tolist() | |
| f7_tax_terms = list(dict.fromkeys([tax_names[i] for i in f7_tax_idx]))[:3] | |
| f7_sub_idx = (f7_norm @ sub_emb.T).argmax(dim=-1).tolist() | |
| f7_sub_terms = list(dict.fromkeys([sub_texts[i] for i in f7_sub_idx]))[:2] | |
| b1_aligns.append(b1_align) | |
| b1_divs.append(b1_div) | |
| f7_aligns.append(f7_align) | |
| f7_divs.append(f7_div) | |
| print(f"{idx:<3} | {q_text[:36]:<36} | Champion B1 (Base) | {b1_align:.3f} | {b1_div:.3f} | {b1_ms:.2f}ms | {', '.join(b1_tax_terms)}") | |
| print(f"{'':<3} | {'':<36} | Frontier F7 (Calib) | {f7_align:.3f} | {f7_div:.3f} | {f7_ms:.2f}ms | {', '.join(f7_tax_terms)}") | |
| print(f"{'':<3} | {'Subqueries B1':<36} | -> {'; '.join(b1_sub_terms)}") | |
| print(f"{'':<3} | {'Subqueries F7':<36} | -> {'; '.join(f7_sub_terms)}") | |
| print("-" * 105) | |
| print("\n" + "=" * 105) | |
| print("BENCHMARK SUMMARY (10 DIVERSE QUERIES)") | |
| print("=" * 105) | |
| print(f"Champion B1 (Baseline): Mean Align: {np.mean(b1_aligns):.3f} | Mean Div: {np.mean(b1_divs):.3f}") | |
| print(f"Frontier F7 (Calib): Mean Align: {np.mean(f7_aligns):.3f} | Mean Div: {np.mean(f7_divs):.3f}") | |
| if __name__ == "__main__": | |
| main() | |