fanout-diffusion / scripts /train_consistency_aligned.py
dejanseo's picture
Add training pipelines, consistency distillation scripts, and interactive dashboard server
f08972a verified
Raw History Blame Contribute Delete
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
@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()