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