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)