fanout-diffusion / scripts /benchmark_f7.py
dejanseo's picture
Add training pipelines, consistency distillation scripts, and interactive dashboard server
69dcf19 verified
Raw History Blame Contribute Delete
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()