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