Download train.py from solintellegence/SolPix: direct link, hf CLI and curl.
- Browser
- Download file 21.4 kB
-
https://huggingface.co/solintellegence/SolPix/resolve/main/train.py
- Command line
-
hf download hf://solintellegence/SolPix/train.py
-
curl -L -o train.py https://huggingface.co/solintellegence/SolPix/resolve/main/train.py
21.4 kB
| 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() | |
| } | |
| 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() | |