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 as nn import torch.nn.functional as F from sentence_transformers import SentenceTransformer from torch.utils.data import DataLoader, TensorDataset from b1_tensor_core import B1Linear from scripts.fast_b1_inference import FastB1Denoiser from scripts.quantize_outer_int4 import pack_int4_signed, unpack_int4_signed from src.r4t.b1_diffusion import B1EDMDenoiser from src.r4t.config import DiffusionConfig from src.r4t.diffusion import ExponentialMovingAverage from src.r4t.journal import ExperimentJournal DATA_PATH = ROOT / "data" / "diffusion_dataset_540k.pt" TAXONOMY_PATH = ROOT / "data" / "taxonomy_embeddings.pt" BASE_CHECKPOINT = ROOT / "checkpoints" / "b1_tc_10ep_champion.pt" OUTPUT_CHECKPOINT = ROOT / "checkpoints" / "champion_b1_consistency_1step_qat.pt" OUTPUT_FULL_CHECKPOINT = ROOT / "checkpoints" / "champion_b1_consistency_1step.pt" TEST_QUERIES = [ "running shoes and athletic sneakers", "espresso coffee machines and barista accessories", "wireless noise cancelling headphones and audio gear", "organic gardening tools and indoor plant care", "python machine learning algorithms and deep neural networks", ] def fake_quantize_int4(w: torch.Tensor) -> torch.Tensor: """Straight-Through Estimator (STE) for symmetric INT4 quantization.""" max_val = w.abs().max(dim=-1, keepdim=True).values.clamp_min(1e-8) scale = max_val / 7.0 q = torch.clamp(torch.round(w / scale), -8, 7) w_q = (q * scale - w).detach() + w return w_q def load_taxonomy_bank(device): if not TAXONOMY_PATH.exists(): return None, None tax_data = torch.load(TAXONOMY_PATH, map_location="cpu", weights_only=False) emb = F.normalize(tax_data["embeddings"].float(), dim=-1).to(device) names = tax_data["names"] return emb, names @torch.no_grad() def evaluate_qualitative(fast_model, embedder, tax_emb, tax_names, device): all_alignments = [] all_diversities = [] decoded_results = {} for q_text in TEST_QUERIES: z_q = embedder.encode([q_text], convert_to_tensor=True, normalize_embeddings=True, device=device).float() cached_mem = fast_model.precompute_cross_memory(z_q) shape = (1, fast_model.config.sequence_length, fast_model.config.embedding_dim) noise = torch.randn(shape, device=device) * fast_model.config.sigma_max sigma = torch.full((1,), fast_model.config.sigma_max, device=device) fanout = fast_model.forward_with_cached_memory(noise, sigma, cached_mem)[0] fanout = F.normalize(fanout, dim=-1) align_scores = (fanout @ z_q.T).squeeze(-1) mean_align = align_scores.mean().item() all_alignments.append(mean_align) sim_mat = fanout @ fanout.T mask = ~torch.eye(10, dtype=torch.bool, device=device) pairwise_sim = sim_mat[mask].mean().item() all_diversities.append(1.0 - pairwise_sim) top_matches = fanout @ tax_emb.T best_indices = top_matches.argmax(dim=-1).tolist() terms = [] for slot_idx, tax_idx in enumerate(best_indices): cat_name = tax_names[tax_idx] slot_sim = top_matches[slot_idx, tax_idx].item() terms.append(f"{cat_name} (sim: {slot_sim:.3f}, align: {align_scores[slot_idx].item():.3f})") decoded_results[q_text] = { "alignment": mean_align, "diversity": 1.0 - pairwise_sim, "terms": terms, } mean_align = sum(all_alignments) / len(all_alignments) mean_div = sum(all_diversities) / len(all_diversities) return mean_align, mean_div, decoded_results def benchmark_cuda_graph_1step(fast_model, device, batch_sizes=[1, 16, 64]): results = {} D = fast_model.config.embedding_dim L = fast_model.config.sequence_length for B in batch_sizes: dummy_q = torch.randn(B, D, device=device) dummy_q = F.normalize(dummy_q, dim=-1) cached_mem = fast_model.precompute_cross_memory(dummy_q) shape = (B, L, D) static_noise = torch.randn(shape, device=device) * fast_model.config.sigma_max static_sigma = torch.full((B,), fast_model.config.sigma_max, device=device) # Warmup graph stream g_stream = torch.cuda.Stream() g_stream.wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(g_stream): for _ in range(5): _ = fast_model.forward_with_cached_memory(static_noise, static_sigma, cached_mem) torch.cuda.current_stream().wait_stream(g_stream) # Capture graph g = torch.cuda.CUDAGraph() with torch.cuda.graph(g, stream=g_stream): _ = fast_model.forward_with_cached_memory(static_noise, static_sigma, cached_mem) # Benchmark Replay torch.cuda.synchronize() times = [] for _ in range(100): t0 = time.perf_counter() g.replay() torch.cuda.synchronize() times.append((time.perf_counter() - t0) * 1000.0) avg_ms = sum(times) / len(times) p50_ms = sorted(times)[int(len(times) * 0.5)] p95_ms = sorted(times)[int(len(times) * 0.95)] qps = (B * 1000.0) / avg_ms results[B] = {"avg_ms": avg_ms, "p50_ms": p50_ms, "p95_ms": p95_ms, "qps": qps} return results def main(): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Compute device: {device} ({torch.cuda.get_device_name(0)})") print(f"Loading 540k dataset from {DATA_PATH}...") dataset_dict = torch.load(DATA_PATH, map_location="cpu", weights_only=False) queries = dataset_dict["query_embeddings"].float() targets = dataset_dict["targets"].float() N, L, D = targets.shape n_train = int(0.9 * N) train_queries, val_queries = queries[:n_train], queries[n_train:] train_targets, val_targets = targets[:n_train], targets[n_train:] batch_size = 128 train_loader = DataLoader(TensorDataset(train_queries, train_targets), batch_size=batch_size, shuffle=True, pin_memory=True) val_loader = DataLoader(TensorDataset(val_queries, val_targets), batch_size=batch_size, shuffle=False, pin_memory=True) print(f"Loading Base Champion Checkpoint: {BASE_CHECKPOINT}...") ckpt = torch.load(BASE_CHECKPOINT, map_location=device, weights_only=False) config = ckpt["config"] model = B1EDMDenoiser(config, backend="tc", pure_1bit=False).to(device) if "ema_state_dict" in ckpt and "shadow" in ckpt["ema_state_dict"]: shadow = ckpt["ema_state_dict"]["shadow"] model.load_state_dict({k: shadow[k].to(device) for k in shadow}) else: model.load_state_dict(ckpt["model_state_dict"]) ema = ExponentialMovingAverage(model, decay=0.999) print("Loading SentenceTransformer and taxonomy bank...") embedder = SentenceTransformer("google/embeddinggemma-300m", model_kwargs={"torch_dtype": torch.bfloat16}, device=device) tax_emb, tax_names = load_taxonomy_bank(device) epochs = 10 lr = 3.0e-4 optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs * len(train_loader), eta_min=1e-5) journal = ExperimentJournal() tracker = journal.start_run( name="Champion B1: Aligned 1-Step Consistency QAT (True Manifold)", experiment_name="1-Bit Tensor Core Innovation", task_type="diffusion", config={ "architecture": "consistency_distillation_aligned_qat", "backend": "b1_tc", "precision": "1bit_core_int4_qat_outer", "layers": config.layers, "hidden_dim": config.hidden_dim, "epochs": epochs, "lr": lr, "batch_size": batch_size, "sampling_steps": 1, "target_ordering": "cosine", "loss_weights": "recon_1.0_align_1.0_gram_0.5", }, tags=["1bit", "tensor_core", "qat", "int4", "consistency_distillation", "1step", "aligned", "champion"], ) print("\n" + "=" * 95) print("STARTING 10-EPOCH ALIGNED CONSISTENCY QAT TRAINING") print("Alignment-Centric Loss: Direct Target MSE + Query Cosine Alignment + Target Gram Matching") print("No destructive <0.30 orthogonality penalty!") print(f"Epochs: {epochs} | Batch Size: {batch_size} | Base Learning Rate: {lr}") print("=" * 95) sigma_max = config.sigma_max outer_modules = [ model.backbone.input_projection, model.backbone.output_projection, model.backbone.query_projection, model.backbone.time_mlp[0], model.backbone.time_mlp[2], ] for epoch in range(1, epochs + 1): ep_t0 = time.time() model.train() train_loss_total = 0.0 for b_queries, b_targets in train_loader: b_queries = b_queries.to(device, non_blocking=True) b_targets = b_targets.to(device, non_blocking=True) B = len(b_queries) # Cosine target ordering (sort slots by descending alignment with query) sims = torch.einsum("bd,bld->bl", b_queries, b_targets) sorted_idx = torch.argsort(sims, dim=1, descending=True) b_targets = torch.gather(b_targets, 1, sorted_idx.unsqueeze(-1).expand(-1, -1, D)) optimizer.zero_grad() # Apply INT4 STE fake quantization to outer weights saved_weights = [] for m in outer_modules: saved_weights.append(m.weight.data.clone()) m.weight.data = fake_quantize_int4(m.weight) # 1. 1-Step Prediction from pure noise pure_noise = torch.randn_like(b_targets) * sigma_max sigmas_max = torch.full((B,), sigma_max, device=device) pred_1step = model(pure_noise, sigmas_max, b_queries) # 2. Intermediate noise level prediction rnd_normal = torch.randn((B,), device=device) sigmas_mid = (rnd_normal * 1.2 - 1.2).exp() noisy_mid = b_targets + sigmas_mid[:, None, None] * torch.randn_like(b_targets) pred_mid = model(noisy_mid, sigmas_mid, b_queries) # Restore unquantized FP32 weights for backward pass (STE) for m, saved_w in zip(outer_modules, saved_weights): m.weight.data = saved_w # ========================================================================= # Alignment-Centric Loss Formulation: # 1. Target Reconstruction MSE (1-step and mid-step) # ========================================================================= loss_recon = F.mse_loss(pred_1step, b_targets) + 0.5 * F.mse_loss(pred_mid, b_targets) # ========================================================================= # 2. Query Cosine Alignment Matching # Match the exact per-slot query cosine alignment distribution of the ground truth # ========================================================================= pred_norm = F.normalize(pred_1step, dim=-1) target_norm = F.normalize(b_targets, dim=-1) q_norm = F.normalize(b_queries, dim=-1) pred_align = torch.einsum("bd,bld->bl", q_norm, pred_norm) target_align = torch.einsum("bd,bld->bl", q_norm, target_norm) loss_align = F.mse_loss(pred_align, target_align) + 0.2 * F.relu(0.50 - pred_align).mean() # ========================================================================= # 3. Target Gram Matrix Matching (True Manifold Covariance) # Matches the exact pairwise geometric relations of ground truth fanouts (~0.63) # ========================================================================= pred_gram = torch.bmm(pred_norm, pred_norm.transpose(1, 2)) target_gram = torch.bmm(target_norm, target_norm.transpose(1, 2)) loss_gram = F.mse_loss(pred_gram, target_gram) total_loss = loss_recon + 1.0 * loss_align + 0.5 * loss_gram total_loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() # STE weight clamp for B1Linear with torch.no_grad(): for m in model.modules(): if isinstance(m, B1Linear): m.weight.clamp_(-1.0, 1.0) scheduler.step() ema.update(model) train_loss_total += total_loss.item() * B train_loss = train_loss_total / len(train_queries) # Validation model.eval() val_loss_total = 0.0 val_align_total = 0.0 val_div_total = 0.0 with torch.no_grad(): for b_queries, b_targets in val_loader: b_queries = b_queries.to(device, non_blocking=True) b_targets = b_targets.to(device, non_blocking=True) B = len(b_queries) sims = torch.einsum("bd,bld->bl", b_queries, b_targets) sorted_idx = torch.argsort(sims, dim=1, descending=True) b_targets = torch.gather(b_targets, 1, sorted_idx.unsqueeze(-1).expand(-1, -1, D)) pure_noise = torch.randn_like(b_targets) * sigma_max sigmas_max = torch.full((B,), sigma_max, device=device) pred_1step = model(pure_noise, sigmas_max, b_queries) v_loss = F.mse_loss(pred_1step, b_targets) val_loss_total += v_loss.item() * B p_norm = F.normalize(pred_1step, dim=-1) align = (p_norm * b_queries[:, None, :]).sum(dim=-1).mean().item() val_align_total += align * B sim_mat = torch.bmm(p_norm, p_norm.transpose(1, 2)) mask = ~torch.eye(L, dtype=torch.bool, device=device).unsqueeze(0).expand(B, -1, -1) div = (1.0 - sim_mat[mask].mean().item()) * B val_div_total += div val_loss = val_loss_total / len(val_queries) val_align = val_align_total / len(val_queries) val_div = val_div_total / len(val_queries) ep_time = time.time() - ep_t0 print(f"Epoch [{epoch:2d}/{epochs}] | Train: {train_loss:.4f} | Val MSE: {val_loss:.4f} | 1-Step Align: {val_align:.3f} | Div: {val_div:.3f} | Time: {ep_time:.1f}s") tracker.log_metrics( step=epoch * len(train_loader), epoch=epoch, train_loss=train_loss, val_loss=val_loss, val_alignment=val_align, val_diversity=val_div, lr=scheduler.get_last_lr()[0], ) # Freeze & Quantize Outer Projections to INT4 print("\nFreezing Model & Packing INT4 Outer Adapters...") model.eval() model.freeze_for_inference() outer_keys = [ "backbone.input_projection.weight", "backbone.output_projection.weight", "backbone.query_projection.weight", "backbone.time_mlp.0.weight", "backbone.time_mlp.2.weight", ] export_dict = { "config": config, "weights": {}, "int4_outer": {}, } state = model.state_dict() for k, v in state.items(): if k in outer_keys: packed, scale = pack_int4_signed(v.float()) export_dict["int4_outer"][k] = { "packed": packed.cpu(), "scale": scale.cpu(), } elif "packed_weight" in k: export_dict["weights"][k] = v.cpu() elif "weight" in k and any(proj in k for proj in ["self_attn", "cross_attn", "mlp"]): continue else: export_dict["weights"][k] = v.to(torch.float16).cpu() torch.save(export_dict, OUTPUT_CHECKPOINT) file_size_bytes = OUTPUT_CHECKPOINT.stat().st_size print(f"Saved Aligned QAT INT4 Checkpoint: {OUTPUT_CHECKPOINT} ({file_size_bytes / (1024*1024):.2f} MB)") # Unpack quantized weights for benchmark and evaluation for k in outer_keys: p = export_dict["int4_outer"][k]["packed"].to(device) s = export_dict["int4_outer"][k]["scale"].to(device) state[k].copy_(unpack_int4_signed(p, s)) # FastB1Denoiser & CUDA Graph Benchmark print("\nCompiling FastB1Denoiser and Capturing CUDA Graph...") fast_model = FastB1Denoiser(model) bench_results = benchmark_cuda_graph_1step(fast_model, device, [1, 16, 64]) print("\nCUDA Graph 1-Step Latency Benchmark (Aligned QAT INT4):") for B, r in bench_results.items(): print(f" Batch {B:2d}: {r['avg_ms']:.3f} ms (P50: {r['p50_ms']:.3f} ms, P95: {r['p95_ms']:.3f} ms) -> {r['qps']:9.1f} QPS") # Qualitative Evaluation print("\nEvaluating Decoded Fanout Taxonomy Terms...") mean_align, mean_div, decoded = evaluate_qualitative(fast_model, embedder, tax_emb, tax_names, device) for q, d in decoded.items(): print(f"\nQuery: \"{q}\"") print(f" Alignment: {d['alignment']:.3f} | Diversity: {d['diversity']:.3f}") for s_idx, t in enumerate(d["terms"][:5]): print(f" #{s_idx+1}: {t}") tracker.log_benchmark( latency_us=int(bench_results[1]["avg_ms"] * 1000), throughput_items_per_sec=bench_results[1]["qps"], device_name=torch.cuda.get_device_name(0), notes=f"Aligned 1-Step INT4: {bench_results[1]['avg_ms']:.3f} ms (B=1), Align: {mean_align:.3f}", ) tracker.finish( status="completed", summary_metrics={ "disk_size_mb": file_size_bytes / (1024 * 1024), "latency_b1_ms": bench_results[1]["avg_ms"], "latency_b64_ms": bench_results[64]["avg_ms"], "throughput_b1_qps": bench_results[1]["qps"], "throughput_b64_qps": bench_results[64]["qps"], "prompt_alignment": mean_align, "pairwise_diversity": mean_div, "val_loss": val_loss, "best_val_loss": val_loss, "sampling_steps": 1, "mode": "aligned_consistency_distillation_qat_int4", }, ) print("\n" + "=" * 95) print(f"All done! Aligned QAT Checkpoint: {OUTPUT_CHECKPOINT.name} ({file_size_bytes / (1024*1024):.2f} MB)") print(f"Latency B=1: {bench_results[1]['avg_ms']:.3f} ms | Alignment: {mean_align:.3f} | Diversity: {mean_div:.3f}") print("=" * 95) if __name__ == "__main__": main()