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