Download scripts/train_diffusion.py from dejanseo/fanout-diffusion: direct link, hf CLI and curl.
- Browser
- Download file 13 kB
-
https://huggingface.co/dejanseo/fanout-diffusion/resolve/main/scripts/train_diffusion.py
- Command line
-
hf download hf://dejanseo/fanout-diffusion/scripts/train_diffusion.py
-
curl -L -o train_diffusion.py https://huggingface.co/dejanseo/fanout-diffusion/resolve/main/scripts/train_diffusion.py
13 kB
| import argparse | |
| 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 torch.utils.data import DataLoader, TensorDataset | |
| from src.r4t.config import DiffusionConfig | |
| from src.r4t.diffusion import ( | |
| EDMDenoiser, | |
| ExponentialMovingAverage, | |
| diffusion_loss, | |
| sample_edm, | |
| ) | |
| ROOT = Path(__file__).resolve().parents[1] | |
| DEFAULT_540K = ROOT / "data" / "diffusion_dataset_540k.pt" | |
| FALLBACK_DATA = ROOT / "data" / "diffusion_dataset.pt" | |
| CHECKPOINT_DIR = ROOT / "checkpoints" | |
| def train(args, tracker=None): | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| print(f"Using device: {device}") | |
| data_path = Path(args.data_path) if args.data_path else (DEFAULT_540K if DEFAULT_540K.exists() else FALLBACK_DATA) | |
| if not data_path.exists(): | |
| raise FileNotFoundError(f"Dataset not found at {data_path}. Run scripts/prepare_diffusion_dataset_540k.py first.") | |
| print(f"Loading dataset from {data_path}...") | |
| data = torch.load(data_path, map_location="cpu", weights_only=False) | |
| queries = data["query_embeddings"].float() | |
| targets = data["targets"].float() | |
| sigma_data = float(data.get("sigma_data", 0.0361)) | |
| dim = int(data.get("embedding_dim", 768)) | |
| N, L, D = targets.shape | |
| print(f"Loaded {N:,} query-fanout pairs: sequence length L={L}, embedding dimension D={D}") | |
| print(f"Empirical sigma_data: {sigma_data:.4f}") | |
| # Train / Val Split (90/10) | |
| perm = torch.randperm(N) | |
| val_size = max(1, int(N * 0.1)) | |
| train_indices = perm[val_size:] | |
| val_indices = perm[:val_size] | |
| train_queries, train_targets = queries[train_indices], targets[train_indices] | |
| val_queries, val_targets = queries[val_indices], targets[val_indices] | |
| print(f"Split: {len(train_queries):,} training samples, {len(val_queries):,} validation samples.") | |
| train_dataset = TensorDataset(train_queries, train_targets) | |
| val_dataset = TensorDataset(val_queries, val_targets) | |
| train_loader = DataLoader( | |
| train_dataset, | |
| batch_size=args.batch_size, | |
| shuffle=True, | |
| drop_last=len(train_dataset) > args.batch_size, | |
| pin_memory=torch.cuda.is_available(), | |
| ) | |
| val_loader = DataLoader( | |
| val_dataset, | |
| batch_size=args.batch_size, | |
| shuffle=False, | |
| pin_memory=torch.cuda.is_available(), | |
| ) | |
| # Initialize Model | |
| config = DiffusionConfig( | |
| sequence_length=L, | |
| embedding_dim=D, | |
| hidden_dim=args.hidden_dim, | |
| mlp_dim=args.mlp_dim, | |
| heads=args.heads, | |
| layers=args.layers, | |
| dropout=args.dropout, | |
| sigma_min=args.sigma_min, | |
| sigma_max=args.sigma_max, | |
| sigma_data=sigma_data, | |
| condition_drop_probability=0.1, | |
| cfg_strength=0.1, | |
| sampling_steps=args.sampling_steps, | |
| ) | |
| model = EDMDenoiser(config).to(device) | |
| ema = ExponentialMovingAverage(model, decay=0.999) | |
| param_count = sum(p.numel() for p in model.parameters()) | |
| print(f"Initialized EDMDenoiser ({param_count / 1e6:.2f}M parameters).") | |
| journal_tracker = tracker | |
| own_tracker = False | |
| if journal_tracker is None and getattr(args, "journal", False): | |
| from src.r4t.journal import ExperimentJournal | |
| journal = ExperimentJournal() | |
| r_name = args.run_name or f"{args.layers}L-{args.target_sorting}-lr{args.lr}" | |
| journal_tracker = journal.start_run( | |
| name=r_name, | |
| experiment_name=args.experiment_name, | |
| task_type="diffusion", | |
| config={ | |
| "layers": args.layers, | |
| "hidden_dim": args.hidden_dim, | |
| "mlp_dim": args.mlp_dim, | |
| "heads": args.heads, | |
| "lr": args.lr, | |
| "epochs": args.epochs, | |
| "batch_size": args.batch_size, | |
| "target_ordering": args.target_sorting, | |
| "sigma_max": args.sigma_max, | |
| "sigma_min": args.sigma_min, | |
| "params_m": param_count / 1e6, | |
| }, | |
| tags=[f"{args.layers}l", args.target_sorting, f"lr_{args.lr}"], | |
| ) | |
| own_tracker = True | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-4) | |
| total_steps = len(train_loader) * args.epochs | |
| warmup_steps = int(len(train_loader) * args.warmup_epochs) | |
| if warmup_steps > 0: | |
| warmup_sched = torch.optim.lr_scheduler.LinearLR(optimizer, start_factor=0.05, total_iters=warmup_steps) | |
| cosine_sched = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max(1, total_steps - warmup_steps), eta_min=args.lr * 0.05) | |
| scheduler = torch.optim.lr_scheduler.SequentialLR(optimizer, schedulers=[warmup_sched, cosine_sched], milestones=[warmup_steps]) | |
| else: | |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=total_steps, eta_min=args.lr * 0.05) | |
| CHECKPOINT_DIR.mkdir(parents=True, exist_ok=True) | |
| best_val_loss = float("inf") | |
| best_epoch = 1 | |
| best_checkpoint_path = CHECKPOINT_DIR / args.checkpoint_name | |
| latest_checkpoint_path = CHECKPOINT_DIR / "latest_diffusion_model.pt" | |
| use_amp = torch.cuda.is_available() | |
| amp_dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16 | |
| print(f"Mixed precision AMP: {use_amp} ({amp_dtype})") | |
| print(f"Target slot ordering strategy: {args.target_sorting}") | |
| print("\nStarting Diffusion Training:") | |
| print(f" Epochs: {args.epochs}") | |
| print(f" Batch size: {args.batch_size}") | |
| print(f" Batches per epoch: {len(train_loader):,}") | |
| print(f" Peak learning rate: {args.lr}") | |
| print(f" Target checkpoint: {best_checkpoint_path}") | |
| print("=" * 60) | |
| t0 = time.time() | |
| for epoch in range(1, args.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) | |
| if args.target_sorting == "random": | |
| perms = torch.argsort(torch.rand(b_targets.shape[0], L, device=device), dim=1) | |
| b_targets = torch.gather(b_targets, 1, perms.unsqueeze(-1).expand(-1, -1, D)) | |
| elif args.target_sorting == "cosine": | |
| # Sort descending by cosine similarity with prompt query | |
| 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)) | |
| optimizer.zero_grad() | |
| with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=use_amp): | |
| loss = diffusion_loss(model, b_targets, b_queries) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) | |
| optimizer.step() | |
| scheduler.step() | |
| ema.update(model) | |
| train_loss_total += loss.item() * len(b_queries) | |
| train_loss = train_loss_total / len(train_dataset) | |
| # Validation with deterministic seed for true comparability | |
| model.eval() | |
| 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) | |
| if args.target_sorting == "cosine": | |
| 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)) | |
| with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=use_amp): | |
| 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_dataset) | |
| checkpoint = { | |
| "epoch": epoch, | |
| "config": config, | |
| "model_state_dict": model.state_dict(), | |
| "ema_state_dict": ema.state_dict(), | |
| "val_loss": val_loss, | |
| "sigma_data": sigma_data, | |
| "dim": D, | |
| "L": L, | |
| "target_sorting": args.target_sorting, | |
| } | |
| torch.save(checkpoint, latest_checkpoint_path) | |
| is_best = val_loss < best_val_loss | |
| if is_best: | |
| best_val_loss = val_loss | |
| best_epoch = epoch | |
| torch.save(checkpoint, best_checkpoint_path) | |
| # Log metrics to journal if active | |
| if journal_tracker is not None: | |
| journal_tracker.log_metrics( | |
| step=epoch * len(train_loader), | |
| epoch=epoch, | |
| train_loss=train_loss, | |
| val_loss=val_loss, | |
| lr=scheduler.get_last_lr()[0], | |
| ) | |
| ep_time = time.time() - ep_t0 | |
| elapsed = time.time() - t0 | |
| best_marker = " [BEST SAVED]" if is_best else "" | |
| print( | |
| f"Epoch [{epoch:3d}/{args.epochs:3d}] | " | |
| f"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}{best_marker} | " | |
| f"LR: {scheduler.get_last_lr()[0]:.2e} | Ep Time: {ep_time:.1f}s | Elapsed: {elapsed:.1f}s", | |
| flush=True, | |
| ) | |
| print("\n" + "=" * 60) | |
| print(f"Training Complete! Best Val Loss: {best_val_loss:.4f}") | |
| print(f"Saved best model checkpoint to: {best_checkpoint_path}") | |
| # Benchmark and finish journal tracking if active and owned by this process | |
| if journal_tracker is not None and own_tracker: | |
| # Quick benchmark of 10-vector ODE generation | |
| sample_q = val_dataset[0][0].unsqueeze(0).to(device) | |
| t_bench0 = time.time() | |
| with torch.no_grad(): | |
| sample_edm(model, sample_q, sampling_steps=16, cfg_strength=0.1) | |
| torch.cuda.synchronize() | |
| bench_lat_ms = (time.time() - t_bench0) * 1000.0 | |
| journal_tracker.log_benchmark( | |
| latency_us=bench_lat_ms * 1000.0, | |
| throughput_items_per_sec=1000.0 / max(1.0, bench_lat_ms), | |
| batch_size=1, | |
| device_name=torch.cuda.get_device_name(0) if torch.cuda.is_available() else "CPU", | |
| notes=f"16-step Heun ODE sampling with {args.layers} layers", | |
| ) | |
| journal_tracker.finish( | |
| status="completed", | |
| summary_metrics={ | |
| "best_val_loss": best_val_loss, | |
| "best_epoch": best_epoch, | |
| "sampling_latency_ms": bench_lat_ms, | |
| }, | |
| ) | |
| if __name__ == "__main__": | |
| parser = argparse.ArgumentParser(description="Train Continuous Diffusion Retriever on 540k query pairs") | |
| parser.add_argument("--data-path", type=str, default=None, help="Path to .pt dataset") | |
| parser.add_argument("--epochs", type=int, default=50, help="Number of training epochs") | |
| parser.add_argument("--warmup-epochs", type=int, default=2, help="Number of linear warmup epochs") | |
| parser.add_argument("--batch-size", type=int, default=128, help="Batch size (e.g. 128 for RTX 4090)") | |
| parser.add_argument("--lr", type=float, default=3e-4, help="Peak learning rate") | |
| parser.add_argument("--hidden-dim", type=int, default=512, help="Transformer hidden dim") | |
| parser.add_argument("--mlp-dim", type=int, default=1024, help="Transformer feedforward dim") | |
| parser.add_argument("--heads", type=int, default=8, help="Number of attention heads") | |
| parser.add_argument("--layers", type=int, default=4, help="Number of decoder layers") | |
| parser.add_argument("--dropout", type=float, default=0.1, help="Dropout probability") | |
| parser.add_argument("--sigma-min", type=float, default=0.0001, help="EDM minimum noise level") | |
| parser.add_argument("--sigma-max", type=float, default=80.0, help="EDM maximum noise level") | |
| parser.add_argument("--sampling-steps", type=int, default=16, help="Sampling ODE steps") | |
| parser.add_argument("--checkpoint-name", type=str, default="best_diffusion_model.pt", help="Checkpoint filename") | |
| parser.add_argument("--target-sorting", type=str, default="random", choices=["random", "cosine", "none"], help="Target slot ordering: random, cosine, or none") | |
| parser.add_argument("--journal", action="store_true", help="Log experiment runs and metrics to journal.db") | |
| parser.add_argument("--experiment-name", type=str, default="Diffusion Sweep", help="Experiment group name for journal") | |
| parser.add_argument("--run-name", type=str, default=None, help="Custom run name for journal") | |
| args = parser.parse_args() | |
| train(args) | |