from __future__ import annotations import argparse from contextlib import ExitStack from datetime import datetime import math import os import random import time from pathlib import Path import torch import torch.distributed as dist import torch.nn.functional as F from torch.nn.parallel import DistributedDataParallel from torch.utils.data import DataLoader from solpix import SolPix, flow_matching_loss from solpix.data import ( LatentTextShardDataset, ResolutionBucketBatchSampler, apply_classifier_free_dropout, collate_latent_text, ) from solpix.training import REPAHead, representation_alignment_loss, sample_logit_normal_time def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Train the SolPix-50M latent flow transformer") parser.add_argument("--data-dir", required=True, help="Directory containing prepared .pt tensor shards") parser.add_argument("--output-dir", required=True, help="Directory for resumeable checkpoints") parser.add_argument("--max-steps", type=int, required=True, help="Number of optimizer updates") parser.add_argument("--stop-after", type=int, default=None, help="Stop after this many updates in this stage") parser.add_argument("--resume", type=str, default=None, help="Checkpoint path, or 'latest' in output-dir") parser.add_argument("--warm-start", type=str, default=None, help="Load generator weights but start a fresh optimizer run") parser.add_argument("--empty-prompt", type=str, required=True, help=".pt file with pre-encoded empty prompt embeddings") parser.add_argument("--batch-size", type=int, default=8, help="Per-device batch size") parser.add_argument("--gradient-accumulation", type=int, default=1) parser.add_argument("--learning-rate", type=float, default=2e-4) parser.add_argument("--weight-decay", type=float, default=0.01) parser.add_argument("--warmup-steps", type=int, default=4000) parser.add_argument("--min-lr-ratio", type=float, default=0.1) parser.add_argument("--cfg-dropout", type=float, default=0.1) parser.add_argument("--repa-weight", type=float, default=0.0, help="Set >0 to train with precomputed teacher_features") parser.add_argument("--repa-decay-start", type=float, default=0.75, help="Fraction of training after which REPA linearly decays") parser.add_argument("--ema-decay", type=float, default=0.9999) parser.add_argument("--grad-clip", type=float, default=1.0) parser.add_argument("--checkpoint-every", type=int, default=2000) parser.add_argument("--log-every", type=int, default=20) parser.add_argument("--num-workers", type=int, default=4) parser.add_argument("--seed", type=int, default=1234) parser.add_argument("--precision", choices=("auto", "bf16", "fp8", "fp16", "fp32"), default="auto") parser.add_argument( "--stop-at", type=str, default=None, help="Stop at an ISO-8601 timestamp with a timezone; save a final checkpoint before exiting", ) parser.add_argument( "--stop-file", type=str, default=None, help="Stop at the next logging boundary when this file appears; save a final checkpoint", ) parser.add_argument("--device", choices=("auto", "cuda", "mps", "cpu"), default="auto") return parser.parse_args() def init_distributed(device_choice: str) -> tuple[bool, int, int, int, torch.device]: world_size = int(os.environ.get("WORLD_SIZE", "1")) rank = int(os.environ.get("RANK", "0")) local_rank = int(os.environ.get("LOCAL_RANK", "0")) distributed = world_size > 1 if device_choice == "auto": if torch.cuda.is_available(): device_choice = "cuda" elif torch.backends.mps.is_available(): device_choice = "mps" else: device_choice = "cpu" if distributed and device_choice != "cuda": raise ValueError("multi-process distributed training currently requires CUDA") if device_choice == "cuda": device = torch.device("cuda", local_rank if distributed else 0) torch.cuda.set_device(device) else: device = torch.device(device_choice) if distributed: dist.init_process_group(backend="nccl", init_method="env://") return distributed, rank, world_size, local_rank, device def choose_dtype(precision: str, device: torch.device) -> torch.dtype | None: if precision == "fp32": return None if precision == "fp8": if device.type != "cuda": raise ValueError("FP8 linear training requires CUDA") if not torch.cuda.is_bf16_supported(): raise ValueError("FP8 linear training uses bfloat16 for non-linear operations, but this GPU lacks bf16") return torch.bfloat16 if precision == "auto": if device.type == "cuda": return torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16 if device.type == "mps": return None return None if precision == "bf16": if device.type == "mps": raise ValueError("bf16 autocast is not enabled for MPS in this launcher; use auto or fp16") return torch.bfloat16 return torch.float16 class ModelEMA: def __init__(self, model: torch.nn.Module, decay: float): self.decay = decay self.shadow = { name: value.detach().clone() for name, value in model.state_dict().items() } @torch.no_grad() def update(self, model: torch.nn.Module) -> None: for name, value in model.state_dict().items(): value = value.detach() shadow = self.shadow[name] if shadow.is_floating_point(): shadow.lerp_(value, 1.0 - self.decay) else: shadow.copy_(value) def state_dict(self) -> dict[str, torch.Tensor]: return self.shadow def load_state_dict(self, state: dict[str, torch.Tensor]) -> None: self.shadow = {name: value.clone() for name, value in state.items()} def load_empty_prompt(path: str, device: torch.device) -> tuple[torch.Tensor, torch.Tensor | None]: payload = torch.load(path, map_location="cpu", weights_only=True) if isinstance(payload, torch.Tensor): return payload.to(device), None if not isinstance(payload, dict) or "text_embeddings" not in payload: raise ValueError("empty-prompt file must be a tensor or a dict with text_embeddings and optional text_mask") embeddings = payload["text_embeddings"] mask = payload.get("text_mask") if not isinstance(embeddings, torch.Tensor): raise ValueError("empty-prompt text_embeddings must be a tensor") return embeddings.to(device), mask.to(device).bool() if isinstance(mask, torch.Tensor) else None def learning_rate_at(step: int, args: argparse.Namespace) -> float: if args.warmup_steps > 0 and step < args.warmup_steps: scale = (step + 1) / args.warmup_steps else: span = max(args.max_steps - args.warmup_steps, 1) progress = min(max((step - args.warmup_steps) / span, 0.0), 1.0) cosine = 0.5 * (1.0 + math.cos(math.pi * progress)) scale = args.min_lr_ratio + (1 - args.min_lr_ratio) * cosine return args.learning_rate * scale def _cpu_state(state: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: return {key: value.detach().cpu() for key, value in state.items()} def save_checkpoint( path: Path, model: torch.nn.Module, ema: ModelEMA, optimizer: torch.optim.Optimizer, step: int, args: argparse.Namespace, repa_head: torch.nn.Module | None, ) -> None: bare_model = model.module if isinstance(model, DistributedDataParallel) else model bare_repa_head = ( repa_head.module if isinstance(repa_head, DistributedDataParallel) else repa_head ) payload = { "step": step, "model": _cpu_state(bare_model.state_dict()), "ema": _cpu_state(ema.state_dict()), "optimizer": optimizer.state_dict(), "model_config": bare_model.config.__dict__, "train_args": vars(args), "repa_head": _cpu_state(bare_repa_head.state_dict()) if bare_repa_head is not None else None, } path.parent.mkdir(parents=True, exist_ok=True) temporary = path.with_suffix(path.suffix + ".tmp") torch.save(payload, temporary) temporary.replace(path) def main() -> None: args = parse_args() if args.max_steps < 1 or args.gradient_accumulation < 1: raise ValueError("max-steps and gradient-accumulation must be positive") if args.stop_after is not None and args.stop_after < 1: raise ValueError("stop-after must be positive when provided") if args.resume and args.warm_start: raise ValueError("choose either --resume or --warm-start") stop_at_ts: float | None = None if args.stop_at: stop_at = datetime.fromisoformat(args.stop_at) if stop_at.tzinfo is None: raise ValueError("--stop-at must include a timezone offset") stop_at_ts = stop_at.timestamp() if args.log_every < 1 or args.num_workers < 0 or args.checkpoint_every < 0: raise ValueError("log-every must be positive; num-workers and checkpoint-every cannot be negative") if args.learning_rate <= 0 or args.weight_decay < 0 or args.grad_clip <= 0: raise ValueError("learning-rate and grad-clip must be positive; weight-decay cannot be negative") if not 0 <= args.cfg_dropout <= 1 or not 0 <= args.ema_decay < 1: raise ValueError("cfg-dropout must be in [0,1] and ema-decay in [0,1)") distributed, rank, world_size, _, device = init_distributed(args.device) seed = args.seed + rank random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) dataset = LatentTextShardDataset(args.data_dir) if dataset.latent_channels != 32 or dataset.text_dim != 768: raise ValueError( f"default SolPix expects 32 latent channels and 768 text features; " f"found {dataset.latent_channels} and {dataset.text_dim}" ) if args.repa_weight < 0 or not 0 <= args.repa_decay_start < 1: raise ValueError("repa-weight must be nonnegative and repa-decay-start must be in [0,1)") if args.repa_weight > 0 and not dataset.has_teacher_features: raise ValueError("REPA is enabled but the data shards have no teacher_features") batch_sampler = ResolutionBucketBatchSampler( dataset, args.batch_size, seed=args.seed, drop_last=True, rank=rank, world_size=world_size, ) loader = DataLoader( dataset, batch_sampler=batch_sampler, collate_fn=collate_latent_text, num_workers=args.num_workers, pin_memory=device.type == "cuda", persistent_workers=args.num_workers > 0, ) if len(loader) < args.gradient_accumulation: raise ValueError("dataset yields fewer microbatches per epoch than gradient-accumulation requires") model: torch.nn.Module = SolPix().to(device) if args.precision == "fp8": try: from torchao.float8 import Float8LinearConfig, convert_to_float8_training except ImportError as exc: raise RuntimeError("FP8 requested, but torchao is not installed") from exc fp8_config = Float8LinearConfig.from_recipe_name("tensorwise") all_linears = [module for module in model.modules() if isinstance(module, torch.nn.Linear)] eligible_linears = sum( module.in_features % 16 == 0 and module.out_features % 16 == 0 for module in all_linears ) def _fp8_dim_supported(module: torch.nn.Module, _fqn: str) -> bool: return ( isinstance(module, torch.nn.Linear) and module.in_features % 16 == 0 and module.out_features % 16 == 0 ) model = convert_to_float8_training( model, config=fp8_config, module_filter_fn=_fp8_dim_supported ) if rank == 0: print( f"FP8: torchao tensorwise float8 training for {eligible_linears}/" f"{len(all_linears)} eligible Linear layers; bf16 for small/incompatible layers" ) if distributed: model = DistributedDataParallel(model, device_ids=[device.index], output_device=device.index) bare_model = model.module if isinstance(model, DistributedDataParallel) else model repa_head: torch.nn.Module | None = None if args.repa_weight > 0: repa_head = REPAHead(teacher_width=dataset.repa_dim).to(device) if distributed: repa_head = DistributedDataParallel( repa_head, device_ids=[device.index], output_device=device.index ) parameters = list(model.parameters()) if repa_head is not None: parameters.extend(repa_head.parameters()) optimizer = torch.optim.AdamW( parameters, lr=args.learning_rate, betas=(0.9, 0.95), weight_decay=args.weight_decay, ) ema = ModelEMA(bare_model, args.ema_decay) start_step = 0 load_path = args.resume or args.warm_start if load_path: resume_path = Path(load_path) if load_path == "latest": resume_path = Path(args.output_dir) / "latest.pt" checkpoint = torch.load(resume_path, map_location=device, weights_only=False) if checkpoint.get("model_config") != bare_model.config.__dict__: raise ValueError("checkpoint model configuration does not match this SolPix build") if args.warm_start: # Warm starts initialize a new experiment from EMA generator weights. bare_model.load_state_dict(checkpoint.get("ema", checkpoint["model"])) ema.load_state_dict(checkpoint.get("ema", checkpoint["model"])) else: bare_model.load_state_dict(checkpoint["model"]) ema.load_state_dict(checkpoint["ema"]) if repa_head is not None: if checkpoint.get("repa_head") is None: raise ValueError("checkpoint has no REPA head; use --warm-start to begin a new REPA run") bare_repa_head = repa_head.module if isinstance(repa_head, DistributedDataParallel) else repa_head bare_repa_head.load_state_dict(checkpoint["repa_head"]) elif checkpoint.get("repa_head") is not None: raise ValueError("checkpoint has a REPA head; resume with --repa-weight > 0") optimizer.load_state_dict(checkpoint["optimizer"]) start_step = int(checkpoint["step"]) precision_dtype = choose_dtype(args.precision, device) scaler = torch.cuda.amp.GradScaler( enabled=device.type == "cuda" and precision_dtype == torch.float16 ) empty_embeddings, empty_mask = load_empty_prompt(args.empty_prompt, device) output_dir = Path(args.output_dir) if rank == 0: output_dir.mkdir(parents=True, exist_ok=True) print( f"SolPix: {bare_model.parameter_count():,} parameters; {dataset.total_samples:,} samples; " f"{len(dataset.paths)} shards; device={device}; world_size={world_size}" ) step = start_step epoch = 0 running_loss = 0.0 running_count = 0 last_log = time.time() stage_end = min(args.max_steps, start_step + args.stop_after) if args.stop_after is not None else args.max_steps while step < stage_end: batch_sampler.set_epoch(epoch) model.train() optimizer.zero_grad(set_to_none=True) for micro_step, batch in enumerate(loader): batch_latents = batch["latents"].to(device, non_blocking=True) text_embeddings = batch["text_embeddings"].to(device, non_blocking=True) text_mask = batch["text_mask"].to(device, non_blocking=True) text_embeddings, text_mask = apply_classifier_free_dropout( text_embeddings, text_mask, empty_embeddings, empty_mask, args.cfg_dropout, ) autocast_enabled = precision_dtype is not None finish_window = ( (micro_step + 1) % args.gradient_accumulation == 0 or micro_step + 1 == len(loader) ) with ExitStack() as contexts: if not finish_window and isinstance(model, DistributedDataParallel): contexts.enter_context(model.no_sync()) if not finish_window and isinstance(repa_head, DistributedDataParallel): contexts.enter_context(repa_head.no_sync()) with torch.autocast(device_type=device.type, dtype=precision_dtype, enabled=autocast_enabled): if repa_head is None: loss, _, _ = flow_matching_loss( model, batch_latents, text_embeddings, text_mask, ) else: time_batch = sample_logit_normal_time(batch_latents.shape[0], device) noise = torch.randn_like(batch_latents) time_view = time_batch.to(batch_latents.dtype).reshape(-1, 1, 1, 1) noisy_latents = (1 - time_view) * batch_latents + time_view * noise target_velocity = noise - batch_latents prediction = model( noisy_latents, time_batch, text_embeddings, text_mask, return_features=True, ) flow_loss = F.mse_loss(prediction.velocity.float(), target_velocity.float()) alignment = representation_alignment_loss( prediction, batch["teacher_features"].to(device, non_blocking=True), repa_head, ) progress = step / max(args.max_steps, 1) if progress <= args.repa_decay_start: repa_scale = 1.0 else: repa_scale = (1.0 - progress) / (1.0 - args.repa_decay_start) loss = flow_loss + args.repa_weight * max(repa_scale, 0.0) * alignment window_start = (micro_step // args.gradient_accumulation) * args.gradient_accumulation window_size = min(args.gradient_accumulation, len(loader) - window_start) scaled_loss = loss / window_size scaler.scale(scaled_loss).backward() running_loss += float(loss.detach()) running_count += 1 if not finish_window: continue lr = learning_rate_at(step, args) for group in optimizer.param_groups: group["lr"] = lr scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(parameters, args.grad_clip) scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True) ema.update(bare_model) step += 1 if rank == 0 and step % args.log_every == 0: elapsed = max(time.time() - last_log, 1e-6) mean_loss = running_loss / max(running_count, 1) print(f"step={step} loss={mean_loss:.5f} lr={lr:.3e} steps_per_sec={args.log_every / elapsed:.2f}") running_loss = 0.0 running_count = 0 last_log = time.time() if rank == 0 and args.checkpoint_every > 0 and step % args.checkpoint_every == 0: save_checkpoint(output_dir / f"step_{step:08d}.pt", model, ema, optimizer, step, args, repa_head) save_checkpoint(output_dir / "latest.pt", model, ema, optimizer, step, args, repa_head) if stop_at_ts is not None and time.time() >= stop_at_ts: stage_end = step if rank == 0: print(f"Reached requested stop time at step {step}; saving final checkpoint") break if ( args.stop_file and step % args.log_every == 0 and Path(args.stop_file).is_file() ): stage_end = step if rank == 0: print(f"Received stop request at step {step}; saving final checkpoint") break if step >= stage_end: break epoch += 1 if distributed: dist.barrier() if rank == 0: save_checkpoint(output_dir / f"step_{step:08d}.pt", model, ema, optimizer, step, args, repa_head) save_checkpoint(output_dir / "latest.pt", model, ema, optimizer, step, args, repa_head) state = "Training complete" if step >= args.max_steps else "Stage complete" print(f"{state} at step {step}; checkpoint: {output_dir / 'latest.pt'}") if distributed: dist.destroy_process_group() if __name__ == "__main__": main()