File size: 5,588 Bytes
69dcf19
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
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()