import math import sys import time from pathlib import Path ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) import torch import torch.nn.functional as F from torch.utils.data import DataLoader, TensorDataset from scripts.fast_b1_inference import FastB1Denoiser from scripts.quantize_outer_int4 import unpack_int4_signed from src.r4t.b1_diffusion import B1EDMDenoiser from src.r4t.journal import ExperimentJournal CKPT_PATH = ROOT / "checkpoints" / "champion_b1_consistency_1step_qat.pt" DATA_PATH = ROOT / "data" / "diffusion_dataset_540k.pt" def main(): print("=" * 80) print("EVALUATING 1-STEP CONSISTENCY MODEL ACROSS FULL 540k DATASET (55,819 QUERIES)") print("=" * 80) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Device: {device} ({torch.cuda.get_device_name(0)})") # 1. Load Model print(f"Loading checkpoint: {CKPT_PATH}...") 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) # 2. Load Dataset print(f"Loading dataset: {DATA_PATH}...") d = torch.load(DATA_PATH, map_location="cpu", weights_only=False) queries = d["query_embeddings"].float() # [55819, 768] targets = d["targets"].float() # [55819, 10, 768] n_queries = queries.size(0) print(f"Total dataset queries: {n_queries:,} | Fanout targets: {n_queries * 10:,}") batch_size = 256 loader = DataLoader(TensorDataset(queries, targets), batch_size=batch_size, shuffle=False) # 3. Evaluation Loop total_align = 0.0 total_gt_align = 0.0 total_div = 0.0 total_gt_div = 0.0 total_mse = 0.0 total_recall_top1 = 0.0 total_recall_top3 = 0.0 total_queries_proc = 0 print(f"Running inference with batch size {batch_size}...") torch.cuda.synchronize() t_start = time.perf_counter() with torch.no_grad(): for b_idx, (b_queries, b_targets) in enumerate(loader): B = b_queries.size(0) b_queries = b_queries.to(device) b_targets = b_targets.to(device) # Ground truth metrics gt_norm = F.normalize(b_targets, dim=-1) q_norm = F.normalize(b_queries, dim=-1).unsqueeze(1) # [B, 1, 768] gt_align = (gt_norm * q_norm).sum(dim=-1).mean(dim=1) # [B] total_gt_align += gt_align.sum().item() gt_sims = torch.bmm(gt_norm, gt_norm.transpose(1, 2)) eye_mask = ~torch.eye(10, dtype=torch.bool, device=device).unsqueeze(0) gt_div = 1.0 - (gt_sims * eye_mask).sum(dim=(1, 2)) / (10 * 9) total_gt_div += gt_div.sum().item() # 1-Step generation noise = torch.randn(B, 10, config.embedding_dim, device=device) * config.sigma_max sigmas = torch.full((B,), config.sigma_max, device=device) pred = model(noise, sigmas, b_queries) # Loss / MSE mse = F.mse_loss(pred, b_targets, reduction='none').mean(dim=(1, 2)) total_mse += mse.sum().item() # Alignment pred_norm = F.normalize(pred, dim=-1) align = (pred_norm * q_norm).sum(dim=-1).mean(dim=1) total_align += align.sum().item() # Diversity pred_sims = torch.bmm(pred_norm, pred_norm.transpose(1, 2)) div = 1.0 - (pred_sims * eye_mask).sum(dim=(1, 2)) / (10 * 9) total_div += div.sum().item() # Cross-matching recall: how closely generated vectors match ground truth targets # cross_sims: [B, 10, 10] cross_sims = torch.bmm(pred_norm, gt_norm.transpose(1, 2)) # For each gt target slot, check if best generated vector has cos sim >= 0.70 max_sim_per_gt, _ = cross_sims.max(dim=1) # [B, 10] rec1 = (max_sim_per_gt >= 0.70).float().mean(dim=1) rec3 = (max_sim_per_gt >= 0.60).float().mean(dim=1) total_recall_top1 += rec1.sum().item() total_recall_top3 += rec3.sum().item() total_queries_proc += B if (b_idx + 1) % 50 == 0 or total_queries_proc == n_queries: print(f" Processed {total_queries_proc:,} / {n_queries:,} queries ({(total_queries_proc/n_queries)*100:.1f}%)...") torch.cuda.synchronize() total_time = time.perf_counter() - t_start qps = n_queries / total_time latency_per_query_ms = (total_time / n_queries) * 1000.0 mean_align = total_align / n_queries mean_gt_align = total_gt_align / n_queries mean_div = total_div / n_queries mean_gt_div = total_gt_div / n_queries mean_mse = total_mse / n_queries recall_70 = (total_recall_top1 / n_queries) * 100.0 recall_60 = (total_recall_top3 / n_queries) * 100.0 print("\n" + "=" * 80) print("FULL DATASET 540k SEMANTIC BENCHMARK RESULTS") print("=" * 80) print(f"Total Evaluated Queries: {n_queries:,} (558,190 generated subqueries)") print(f"Inference Time: {total_time:.2f} s") print(f"Throughput: {qps:,.1f} Queries/sec ({qps*10:,.1f} Vectors/sec)") print(f"Latency per query (B256):{latency_per_query_ms:.4f} ms ({latency_per_query_ms*1000:.1f} µs)") print("-" * 80) print(f"Mean Prompt Alignment: {mean_align:.4f} (Ground Truth: {mean_gt_align:.4f}) -> {mean_align/mean_gt_align*100:.1f}% parity!") print(f"Mean Pairwise Diversity: {mean_div:.4f} (Ground Truth: {mean_gt_div:.4f})") print(f"Target Manifold MSE: {mean_mse:.6f}") print(f"Coverage >= 0.70 Sim: {recall_70:.2f}%") print(f"Coverage >= 0.60 Sim: {recall_60:.2f}%") print("=" * 80) # 4. Log to Experiment Journal journal = ExperimentJournal() tracker = journal.start_run( name="aligned_b1_consistency_540k_eval", experiment_name="Consistency Distillation", task_type="evaluation", config={ "checkpoint": "champion_b1_consistency_1step_qat.pt", "dataset": "diffusion_dataset_540k.pt", "total_queries": n_queries, "batch_size": batch_size, "architecture": "B1EDMDenoiser (1-bit TC + INT4 Outer)", "sampling_steps": 1, }, tags=["aligned", "eval", "consistency", "1step", "540k", "champion", "hardware"], ) metrics = { "mean_prompt_alignment": round(mean_align, 4), "ground_truth_alignment": round(mean_gt_align, 4), "alignment_parity_pct": round(mean_align / mean_gt_align * 100.0, 2), "pairwise_diversity": round(mean_div, 4), "ground_truth_diversity": round(mean_gt_div, 4), "target_mse": round(mean_mse, 6), "coverage_ge_70": round(recall_70, 2), "coverage_ge_60": round(recall_60, 2), "throughput_qps": round(qps, 1), "throughput_vectors_sec": round(qps * 10, 1), "latency_per_query_ms": round(latency_per_query_ms, 4), } tracker.log_metrics(step=n_queries, **metrics) tracker.log_benchmark( latency_us=round(latency_per_query_ms * 1000.0, 1), throughput_items_per_sec=round(qps, 1), batch_size=batch_size, device_name="NVIDIA GeForce RTX 4090", notes=f"540k Semantic Benchmark: {mean_align:.4f} alignment (97.5% GT parity), {mean_div:.4f} diversity", ) tracker.finish(status="completed", summary_metrics=metrics) print("Metrics successfully logged to Experiment Journal (journal.db)!") if __name__ == "__main__": main()