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 sentence_transformers import SentenceTransformer from torch.utils.data import DataLoader, TensorDataset from scripts.fast_b1_inference import FastB1Denoiser, fast_sample_edm_8step from src.r4t.b1_diffusion import B1EDMDenoiser from src.r4t.config import DiffusionConfig from src.r4t.diffusion import diffusion_loss, sample_edm from src.r4t.journal import ExperimentJournal CHAMPION_CKPT = Path("checkpoints/b1_tc_10ep_champion.pt") INT4_EXPORT_PATH = Path("checkpoints/champion_b1_tc_int4_outer.pt") DATA_PATH = Path("data/diffusion_dataset_540k.pt") TAXONOMY_PATH = Path("data/taxonomy_embeddings.pt") def pack_int4_signed(tensor: torch.Tensor): """ Symmetric per-channel INT4 quantization: Quantizes float tensor in range [-8, 7] and packs pairs of 4-bit nibbles into uint8. """ # Per-row scale: [M, 1] max_val = tensor.abs().max(dim=-1, keepdim=True).values.clamp_min(1e-8) scale = max_val / 7.0 # Range -7 to +7 (or -8 to 7) q = torch.clamp(torch.round(tensor / scale), -8, 7).to(torch.int8) # Convert signed 4-bit to unsigned 4-bit [0..15] q_u = (q & 0x0F).to(torch.uint8) # Pack adjacent elements (dim=-1 must be even) # low nibble = even, high nibble = odd q_even = q_u[..., 0::2] q_odd = q_u[..., 1::2] packed = (q_odd << 4) | q_even return packed, scale.to(torch.float16) def unpack_int4_signed(packed: torch.Tensor, scale: torch.Tensor): """Unpacks pairs of 4-bit nibbles from uint8 back to float16 tensor.""" q_even = (packed & 0x0F).to(torch.int8) q_odd = ((packed >> 4) & 0x0F).to(torch.int8) # Sign extend from 4-bit to 8-bit q_even = torch.where(q_even >= 8, q_even - 16, q_even) q_odd = torch.where(q_odd >= 8, q_odd - 16, q_odd) M = packed.shape[0] K = packed.shape[1] * 2 unpacked = torch.empty((M, K), dtype=torch.float16, device=packed.device) unpacked[:, 0::2] = q_even.to(torch.float16) unpacked[:, 1::2] = q_odd.to(torch.float16) return unpacked * scale.to(unpacked.device) def evaluate_qualitative(model, embedder, tax_emb, tax_names, device): test_queries = [ "quantum computing algorithms for cryptography", "renewable energy storage systems and solar cells", "deep neural networks for medical image diagnostics", ] total_alignment = 0.0 total_diversity = 0.0 model.eval() with torch.no_grad(): for q_text in test_queries: q_emb = embedder.encode([q_text], convert_to_tensor=True, device=device).float() q_emb = F.normalize(q_emb, dim=-1) subq_traj = sample_edm(model, q_emb, sampling_steps=8, cfg_strength=0.1) subq_emb = F.normalize(subq_traj[0], dim=-1) # Alignment sim_prompt = (subq_emb @ q_emb.squeeze(0)).mean().item() total_alignment += sim_prompt # Diversity sim_matrix = subq_emb @ subq_emb.T L = sim_matrix.shape[0] mask = ~torch.eye(L, dtype=torch.bool, device=device) pairwise_div = 1.0 - sim_matrix[mask].mean().item() total_diversity += pairwise_div avg_align = total_alignment / len(test_queries) avg_div = total_diversity / len(test_queries) return avg_align, avg_div def main(): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Device: {device} ({torch.cuda.get_device_name(0)})") print(f"Loading champion checkpoint from {CHAMPION_CKPT}...") ckpt = torch.load(CHAMPION_CKPT, 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"]) model.eval() model.freeze_for_inference() # Outer adapters to quantize to INT4 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() total_unpacked_bytes = 0 total_int4_bytes = 0 print("\nQuantizing Outer Adapters to Symmetric INT4:") print("--------------------------------------------------------------------------------") print(f"{'Layer':<40} | {'Orig FP16':<12} | {'INT4 Packed':<12} | {'MSE Error':<10}") print("--------------------------------------------------------------------------------") for k, v in state.items(): if k in outer_keys: orig_bytes = v.numel() * 2 total_unpacked_bytes += orig_bytes packed, scale = pack_int4_signed(v.float()) recon = unpack_int4_signed(packed, scale) mse = F.mse_loss(recon.float(), v.float()).item() export_dict["int4_outer"][k] = { "packed": packed.cpu(), "scale": scale.cpu(), } int4_bytes = packed.numel() + scale.numel() * 2 total_int4_bytes += int4_bytes print(f"{k:<40} | {orig_bytes/1024:>8.1f} KB | {int4_bytes/1024:>8.1f} KB | {mse:.2e}") elif "packed_weight" in k: export_dict["weights"][k] = v.cpu() int4_bytes = v.numel() * v.element_size() total_int4_bytes += int4_bytes total_unpacked_bytes += int4_bytes elif "weight" in k and any(proj in k for proj in ["self_attn", "cross_attn", "mlp"]): # Skip uncompressed latent FP32 weights of B1Linear layers continue else: # Biases, LayerNorms, positional embeddings (FP16) v_fp16 = v.to(torch.float16).cpu() export_dict["weights"][k] = v_fp16 int4_bytes = v_fp16.numel() * 2 total_int4_bytes += int4_bytes total_unpacked_bytes += int4_bytes print("--------------------------------------------------------------------------------") print(f"Original Hybrid Model Size: {total_unpacked_bytes / (1024*1024):.2f} MB") print(f"INT4 Outer Quantized Model: {total_int4_bytes / (1024*1024):.2f} MB (raw tensors)") # Save to disk torch.save(export_dict, INT4_EXPORT_PATH) file_size_bytes = INT4_EXPORT_PATH.stat().st_size print(f"\nSaved INT4 Deployment Checkpoint: {INT4_EXPORT_PATH}") print(f"File Size on Disk: {file_size_bytes / 1024:.1f} KB ({file_size_bytes / (1024*1024):.2f} MB)") # Apply reconstructed weights back to model to test validation loss & fidelity for k in outer_keys: p = export_dict["int4_outer"][k]["packed"].to(device) s = export_dict["int4_outer"][k]["scale"].to(device) recon = unpack_int4_signed(p, s) state[k].copy_(recon) print("\nValidating Fidelity of INT4-Quantized Model on Dataset...") 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) val_queries, val_targets = queries[n_train:], targets[n_train:] val_loader = DataLoader(TensorDataset(val_queries, val_targets), batch_size=128, shuffle=False) val_loss_total = 0.0 val_gen = torch.Generator(device=device).manual_seed(1337) 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) sims = torch.einsum("bd,bld->bl", F.normalize(b_queries, dim=-1), F.normalize(b_targets, dim=-1)) sorted_idx = torch.argsort(sims, dim=1, descending=True) b_targets = torch.gather(b_targets, 1, sorted_idx.unsqueeze(-1).expand(-1, -1, D)) v_loss = diffusion_loss(model, b_targets, b_queries, generator=val_gen) val_loss_total += v_loss.item() * len(b_queries) val_loss = val_loss_total / len(val_queries) print(f"Validation Loss after INT4 Outer Quantization: {val_loss:.4f} (Baseline FP16: 0.6885)") print("\nEvaluating Qualitative Decoding (EmbeddingGemma)...") embedder = SentenceTransformer("google/embeddinggemma-300m", model_kwargs={"torch_dtype": torch.bfloat16}, device=device) tax_dict = torch.load(TAXONOMY_PATH, map_location=device, weights_only=False) tax_emb = tax_dict["embeddings"].to(device).float() tax_names = tax_dict["names"] align, div = evaluate_qualitative(model, embedder, tax_emb, tax_names, device) print(f"INT4 Outer Model: Prompt Alignment = {align:.3f} (FP16: 0.324) | Diversity = {div:.3f} (FP16: 0.844)") # Benchmark latency fast_model = FastB1Denoiser(model) dummy_q = torch.randn(1, 768, device=device) dummy_q = F.normalize(dummy_q, dim=-1) # CUDA Graph g_stream = torch.cuda.Stream() g_stream.wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(g_stream): for _ in range(3): _ = fast_sample_edm_8step(fast_model, dummy_q) torch.cuda.current_stream().wait_stream(g_stream) g = torch.cuda.CUDAGraph() with torch.cuda.graph(g, stream=g_stream): _ = fast_sample_edm_8step(fast_model, dummy_q) 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) p95_ms = sorted(times)[int(len(times) * 0.95)] qps = 1000.0 / avg_ms print(f"\nCUDA Graph 8-Step Heun Latency: {avg_ms:.2f} ms (P95: {p95_ms:.2f} ms, {qps:.1f} QPS)") # Log to journal journal = ExperimentJournal() tracker = journal.start_run( name="Champion B1-TC: INT4 Outer Quantization (1.5 MB)", experiment_name="1-Bit Tensor Core Innovation", task_type="diffusion", config={ "core": "1-bit_ptx_mma", "outer": "symmetric_int4_packed", "layers": 2, "sampling_steps": 8, "checkpoint_size_mb": file_size_bytes / (1024 * 1024), }, tags=["1bit", "tensor_core", "int4", "compression", "quantization"], ) tracker.log_benchmark( latency_us=int(avg_ms * 1000), throughput_items_per_sec=qps, device_name=torch.cuda.get_device_name(0), notes=f"INT4 outer quantized model: {file_size_bytes / (1024*1024):.2f} MB, {avg_ms:.2f} ms latency", ) tracker.finish( status="completed", summary_metrics={ "val_loss": val_loss, "prompt_alignment": align, "pairwise_diversity": div, "latency_ms": avg_ms, "file_size_mb": file_size_bytes / (1024 * 1024), "file_size_kb": file_size_bytes / 1024, }, ) print("\nLogged INT4 champion experiment to journal.db!") if __name__ == "__main__": main()