diff --git a/README.md b/README.md index 00362d2..edcb40b 100644 --- a/README.md +++ b/README.md @@ -19,6 +19,14 @@ This repository contains: * ⚡️ Pre-trained class-conditional SiT models trained on ImageNet 256x256 * 🛸 A SiT [training script](train.py) using PyTorch DDP +## Experiment backup and resume + +The project-specific recovery procedure for the base, rotation-layer, and +convolution-layer experiments is documented in [docs/RESUME_GUIDE.md](docs/RESUME_GUIDE.md). +It starts from downloading gated ImageNet-1K, recreates the ImageFolder layout, +downloads the backed-up checkpoints/results from Hugging Face, verifies a +checkpoint by sampling, and resumes the matching model implementation. + ## Setup First, download and set up the repo: @@ -166,4 +174,3 @@ versus 2.06 in the paper). ## License This project is under the MIT license. See [LICENSE](LICENSE.txt) for details. - diff --git a/run_train.sh b/run_train.sh index 1d30f42..0f1ac15 100644 --- a/run_train.sh +++ b/run_train.sh @@ -3,9 +3,11 @@ torchrun \ --nproc_per_node=8 \ train.py \ --model SiT-S/2 \ ---epochs=400 \ ---data-path /home/jiayou.zhang/hom/personal/imagenet_dataset/images/train \ +--epochs=800 \ +--data-path /home/nvidia/datasets/imagenet-1k/train \ --wandb \ ---global-batch-size=1024 +--global-batch-size=1024 \ +--run-name 005-SiT-S-2-Linear-velocity-None \ +--ckpt /home/nvidia/SiT-Complementary/results/005-SiT-S-2-Linear-velocity-None/checkpoints/0550000.pt # batch_size x 4, lr x 2 diff --git a/run_train_conv.sh b/run_train_conv.sh old mode 100644 new mode 100755 index a3fa493..98b68ee --- a/run_train_conv.sh +++ b/run_train_conv.sh @@ -1,11 +1,46 @@ -torchrun \ ---nnodes=1 \ ---nproc_per_node=8 \ -train_conv.py \ ---model SiT-S/2 \ ---epochs=200 \ ---data-path /home/jiayou.zhang/hom/personal/imagenet_dataset/images/train \ ---wandb \ ---global-batch-size=1024 - -# batch_size x 4, lr x 2 +#!/usr/bin/env bash +set -euo pipefail + +cd /home/nvidia/SiT-Complementary + +export WANDB_KEY +WANDB_KEY="$(python -c 'import netrc; print(netrc.netrc().authenticators("api.wandb.ai")[2])')" + +# Match the base and rotation-layer runs while keeping a separate W&B run. +export WANDB_MODE=offline +export WANDB_DIR=/data/nvidia/SiT-conv-layer-bs256/wandb +export SIT_FID_COMPARISON_OUTPUT_DIR=/home/nvidia/SiT-comparisons/bs256-lr1e-4-800ep +mkdir -p "$WANDB_DIR" + +exec torchrun \ + --nnodes=1 \ + --nproc_per_node=8 \ + train_conv.py \ + --model SiT-S/2 \ + --epochs 800 \ + --data-path /home/nvidia/datasets/imagenet-1k/train \ + --results-dir /data/nvidia/SiT-conv-layer-bs256/results-800ep \ + --global-batch-size 256 \ + --learning-rate 0.0001 \ + --global-seed 0 \ + --vae ema \ + --num-workers 4 \ + --log-every 100 \ + --ckpt-every 50000 \ + --sample-every 10000 \ + --cfg-scale 4.0 \ + --run-name SiT-S-2-ConvLayer-bs256-lr1e-4-800ep \ + --fid-every-checkpoint \ + --fid-every 250000 \ + --fid-num-samples 50000 \ + --fid-reference /home/nvidia/evaluation/reference/discon-download/VIRTUAL_imagenet256_labeled.npz \ + --fid-history /data/nvidia/SiT-conv-layer-bs256/results-800ep/SiT-S-2-ConvLayer-bs256-lr1e-4-800ep/fid_cfg1_50k.tsv \ + --fid-per-proc-batch-size 64 \ + --fid-inception-batch-size 128 \ + --fid-num-workers 8 \ + --fid-sampling-steps 250 \ + --fid-seed 0 \ + --fid-stop-consecutive-increases 3 \ + --fid-stop-min-absolute-rise 0.25 \ + --fid-stop-min-relative-rise 0.005 \ + --wandb diff --git a/run_train_rot_layer.sh b/run_train_rot_layer.sh old mode 100644 new mode 100755 index e5bc80a..7396223 --- a/run_train_rot_layer.sh +++ b/run_train_rot_layer.sh @@ -1,11 +1,47 @@ -torchrun \ ---nnodes=1 \ ---nproc_per_node=8 \ -train_rot_layer.py \ ---model SiT-S/2 \ ---epochs=200 \ ---data-path /home/jiayou.zhang/hom/personal/imagenet_dataset/images/train \ ---wandb \ ---global-batch-size=1024 - -# batch_size x 4, lr x 2 +#!/usr/bin/env bash +set -euo pipefail + +cd /home/nvidia/SiT-Complementary + +export WANDB_KEY +WANDB_KEY="$(python -c 'import netrc; print(netrc.netrc().authenticators("api.wandb.ai")[2])')" + +# Match the base run's low-overhead W&B recording setup. This is a new run +# because its name (and therefore deterministic W&B run ID) is unique. +export WANDB_MODE=offline +export WANDB_DIR=/home/nvidia/SiT-rot-layer-bs256/wandb +mkdir -p "$WANDB_DIR" + +exec torchrun \ + --nnodes=1 \ + --nproc_per_node=8 \ + train_rot_layer.py \ + --model SiT-S/2 \ + --epochs 800 \ + --data-path /home/nvidia/datasets/imagenet-1k/train \ + --results-dir /home/nvidia/SiT-rot-layer-bs256/results-200ep \ + --global-batch-size 256 \ + --learning-rate 0.0001 \ + --global-seed 0 \ + --vae ema \ + --num-workers 4 \ + --log-every 100 \ + --ckpt-every 50000 \ + --sample-every 10000 \ + --cfg-scale 4.0 \ + --run-name SiT-S-2-RotLayer-bs256-lr1e-4-200ep \ + --ckpt /home/nvidia/SiT-rot-layer-bs256/results-200ep/SiT-S-2-RotLayer-bs256-lr1e-4-200ep/checkpoints/1000800.pt \ + --fid-every-checkpoint \ + --fid-every 250000 \ + --fid-num-samples 50000 \ + --fid-reference /home/nvidia/evaluation/reference/discon-download/VIRTUAL_imagenet256_labeled.npz \ + --fid-history /home/nvidia/SiT-rot-layer-bs256/results-200ep/SiT-S-2-RotLayer-bs256-lr1e-4-200ep/fid_cfg1_50k.tsv \ + --fid-per-proc-batch-size 64 \ + --fid-inception-batch-size 128 \ + --fid-num-workers 8 \ + --fid-sampling-steps 250 \ + --fid-seed 0 \ + --fid-stop-consecutive-increases 3 \ + --fid-stop-min-absolute-rise 0.25 \ + --fid-stop-min-relative-rise 0.005 \ + --wandb diff --git a/sample.py b/sample.py index 8bd86b5..2abb502 100644 --- a/sample.py +++ b/sample.py @@ -31,7 +31,8 @@ def main(mode, args): assert args.image_size == 256, "512x512 models are not yet available for auto-download." # remove this line when 512x512 models are available learn_sigma = args.image_size == 256 else: - learn_sigma = False + # train.py uses the model default learn_sigma=True. + learn_sigma = True # Load model: latent_size = args.image_size // 8 diff --git a/sample_ddp.py b/sample_ddp.py index 346b846..20d8202 100644 --- a/sample_ddp.py +++ b/sample_ddp.py @@ -8,15 +8,21 @@ evaluation metrics via the ADM repo: https://github.com/openai/guided-diffusion/ For a simple single-GPU/CPU sampling script, see sample.py. """ +import importlib +import os + import torch import torch.distributed as dist -from models_rot_head import SiT_models + +MODEL_MODULE_NAME = os.environ.get("SIT_MODEL_MODULE", "models") +model_module = importlib.import_module(MODEL_MODULE_NAME) +SiT_models = model_module.SiT_models +MODEL_IMPLEMENTATION_PATH = os.path.realpath(model_module.__file__) from download import find_model from transport import create_transport, Sampler from diffusers.models import AutoencoderKL from train_utils import parse_ode_args, parse_sde_args, parse_transport_args from tqdm import tqdm -import os from PIL import Image import numpy as np import math @@ -47,6 +53,12 @@ def main(mode, args): """ torch.backends.cuda.matmul.allow_tf32 = args.tf32 # True: fast but may lead to some small numerical differences assert torch.cuda.is_available(), "Sampling with DDP requires at least one GPU. sample.py supports CPU-only usage" + expected_model_module = os.environ.get("SIT_EXPECTED_MODEL_MODULE") + if expected_model_module and MODEL_MODULE_NAME != expected_model_module: + raise RuntimeError( + f"Expected model module {expected_model_module!r}, but loaded " + f"{MODEL_MODULE_NAME!r} from {MODEL_IMPLEMENTATION_PATH}" + ) torch.set_grad_enabled(False) # Setup DDP: @@ -57,6 +69,8 @@ def main(mode, args): torch.manual_seed(seed) torch.cuda.set_device(device) print(f"Starting rank={rank}, seed={seed}, world_size={dist.get_world_size()}.") + if rank == 0: + print(f"Model implementation: {MODEL_MODULE_NAME} ({MODEL_IMPLEMENTATION_PATH})") if args.ckpt is None: assert args.model == "SiT-XL/2", "Only SiT-XL/2 models are available for auto-download." @@ -65,6 +79,8 @@ def main(mode, args): assert args.image_size == 256, "512x512 models are not yet available for auto-download." # remove this line when 512x512 models are available learn_sigma = args.image_size == 256 else: + # train.py constructs custom checkpoints with the model default + # learn_sigma=True, so preserve that architecture for strict loading. learn_sigma = True # Load model: @@ -123,7 +139,7 @@ def main(mode, args): model_string_name = args.model.replace("/", "-") ckpt_string_name = os.path.basename(args.ckpt).replace(".pt", "") if args.ckpt else "pretrained" if mode == "ODE": - folder_name = f"{model_string_name}-rot-head-{ckpt_string_name}-" \ + folder_name = f"{model_string_name}-{ckpt_string_name}-" \ f"cfg-{args.cfg_scale}-{args.per_proc_batch_size}-"\ f"{mode}-{args.num_sampling_steps}-{args.sampling_method}" elif mode == "SDE": diff --git a/train.py b/train.py index ece1e65..a24028f 100644 --- a/train.py +++ b/train.py @@ -21,10 +21,19 @@ from copy import deepcopy from glob import glob from time import time import argparse +import csv +import importlib import logging +import math import os - -from models import SiT_models +import re +import shutil +from itertools import islice + +MODEL_MODULE_NAME = os.environ.get("SIT_MODEL_MODULE", "models") +model_module = importlib.import_module(MODEL_MODULE_NAME) +SiT_models = model_module.SiT_models +MODEL_IMPLEMENTATION_PATH = os.path.realpath(model_module.__file__) from download import find_model from transport import create_transport, Sampler from diffusers.models import AutoencoderKL @@ -82,6 +91,20 @@ def create_logger(logging_dir): return logger +class SkipBatchSampler: + """Skip already-consumed batches without loading or transforming their images.""" + + def __init__(self, batch_sampler, skip): + self.batch_sampler = batch_sampler + self.skip = skip + + def __iter__(self): + return islice(iter(self.batch_sampler), self.skip, None) + + def __len__(self): + return max(0, len(self.batch_sampler) - self.skip) + + def center_crop_arr(pil_image, image_size): """ Center cropping implementation from ADM. @@ -103,6 +126,177 @@ def center_crop_arr(pil_image, image_size): return Image.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size]) +@torch.no_grad() +def evaluate_checkpoint_fid(ema, vae, transport_sampler, args, train_steps, + device, rank, logger, experiment_dir): + """Evaluate EMA with CFG=1 while preserving the training RNG trajectory.""" + world_size = dist.get_world_size() + local_batch = args.fid_per_proc_batch_size + global_batch = local_batch * world_size + total_samples = math.ceil(args.fid_num_samples / global_batch) * global_batch + history_path = args.fid_history or os.path.join(experiment_dir, "fid_cfg1_50k.tsv") + sample_dir = os.path.join( + experiment_dir, "fid_cfg1_work", f"{train_steps:07d}" + ) + + # A completed record is reusable after a restart. Rank 0 decides and tells + # every worker, so all ranks take the same collective path. + already_done = False + if rank == 0 and os.path.isfile(history_path): + with open(history_path, newline="") as f: + for row in csv.DictReader(f, delimiter="\t"): + if int(row["step"]) == train_steps and row["status"] == "ok": + already_done = True + break + done_tensor = torch.tensor(int(already_done), device=device) + dist.broadcast(done_tensor, src=0) + if done_tensor.item(): + logger.info(f"Reusing recorded CFG=1 FID for checkpoint {train_steps:07d}") + return False + + cpu_rng_state = torch.get_rng_state() + cuda_rng_state = torch.cuda.get_rng_state(device) + torch.manual_seed(args.fid_seed * world_size + rank) + torch.cuda.manual_seed(args.fid_seed * world_size + rank) + + if rank == 0: + os.makedirs(sample_dir, exist_ok=True) + # A prior interrupted attempt may contain a partial sample set. + for name in os.listdir(sample_dir): + if name.endswith(".png"): + os.remove(os.path.join(sample_dir, name)) + logger.info( + f"Evaluating checkpoint {train_steps:07d}: CFG=1, " + f"requested={args.fid_num_samples:,}, actual={total_samples:,}" + ) + dist.barrier() + + sample_fn = transport_sampler.sample_ode(num_steps=args.fid_sampling_steps) + latent_size = args.image_size // 8 + iterations = total_samples // global_batch + for batch_index in range(iterations): + z = torch.randn(local_batch, 4, latent_size, latent_size, device=device) + y = torch.randint(0, args.num_classes, (local_batch,), device=device) + samples = sample_fn(z, ema.forward, y=y)[-1] + samples = vae.decode(samples / 0.18215).sample + samples = torch.clamp(127.5 * samples + 128.0, 0, 255) + samples = samples.permute(0, 2, 3, 1).to("cpu", dtype=torch.uint8).numpy() + for local_index, sample in enumerate(samples): + image_index = batch_index * global_batch + local_index * world_size + rank + Image.fromarray(sample).save(os.path.join(sample_dir, f"{image_index:06d}.png")) + if batch_index % 10 == 0: + dist.barrier() + dist.barrier() + + fid_value = 0.0 + stop_requested = False + previous_step = None + previous_fid = None + if rank == 0: + from pytorch_fid.fid_score import calculate_fid_given_paths + + fid_value = float(calculate_fid_given_paths( + [args.fid_reference, sample_dir], + batch_size=args.fid_inception_batch_size, + device="cuda:0", + dims=2048, + num_workers=args.fid_num_workers, + )) + + prior_rows = [] + if os.path.isfile(history_path): + with open(history_path, newline="") as f: + prior_rows = [ + row for row in csv.DictReader(f, delimiter="\t") + if row["status"] == "ok" and int(row["step"]) < train_steps + ] + trend_rows = sorted( + ( + (int(row["step"]), float(row["fid"])) + for row in prior_rows + ), + key=lambda item: item[0], + ) + trend_rows.append((train_steps, fid_value)) + required_points = args.fid_stop_consecutive_increases + 1 + recent_trend = trend_rows[-required_points:] + if prior_rows: + previous_step, previous_fid = trend_rows[-2] + if len(recent_trend) == required_points: + consecutive_increases = all( + right_fid > left_fid + for (_, left_fid), (_, right_fid) + in zip(recent_trend, recent_trend[1:]) + ) + cumulative_rise = recent_trend[-1][1] - recent_trend[0][1] + required_rise = max( + args.fid_stop_min_absolute_rise, + recent_trend[0][1] * args.fid_stop_min_relative_rise, + ) + stop_requested = consecutive_increases and cumulative_rise >= required_rise + + os.makedirs(os.path.dirname(history_path), exist_ok=True) + needs_header = not os.path.isfile(history_path) or os.path.getsize(history_path) == 0 + with open(history_path, "a", newline="") as f: + writer = csv.writer(f, delimiter="\t", lineterminator="\n") + if needs_header: + writer.writerow([ + "step", "checkpoint", "status", "fid", "cfg", + "num_requested", "num_png", "seed", "timestamp_utc" + ]) + from datetime import datetime, timezone + writer.writerow([ + train_steps, + os.path.join(experiment_dir, "checkpoints", f"{train_steps:07d}.pt"), + "ok", repr(fid_value), "1.0", args.fid_num_samples, + total_samples, args.fid_seed, + datetime.now(timezone.utc).isoformat(), + ]) + + logger.info(f"Checkpoint {train_steps:07d} CFG=1 PyTorch FID: {fid_value:.9f}") + comparison_output_dir = os.environ.get("SIT_FID_COMPARISON_OUTPUT_DIR") + if comparison_output_dir: + try: + from tools.plot_fid_training_curves import generate_plot + generated = generate_plot( + comparison_output_dir, + conv_history=history_path, + ) + logger.info( + f"Updated FID comparison plot: {generated['png']}" + ) + except Exception: + # A reporting artifact must never interrupt model training. + logger.exception("Could not update the FID comparison plot") + if args.wandb: + wandb_utils.log({"eval/fid_cfg1_50k": fid_value}, step=train_steps) + if stop_requested: + marker = os.path.join(experiment_dir, "FID_REGRESSION_STOPPED") + with open(marker, "w") as f: + f.write( + f"sustained FID regression over {args.fid_stop_consecutive_increases} " + f"consecutive checkpoints: step {recent_trend[0][0]} " + f"FID {recent_trend[0][1]:.9f} -> step {train_steps} " + f"FID {fid_value:.9f}\n" + ) + logger.error( + f"FID increased for {args.fid_stop_consecutive_increases} consecutive " + f"checkpoints, from {recent_trend[0][1]:.9f} at step " + f"{recent_trend[0][0]} to {fid_value:.9f}; stopping after " + f"checkpoint {train_steps:07d}." + ) + shutil.rmtree(sample_dir) + + result = torch.tensor([fid_value, float(stop_requested)], device=device) + dist.broadcast(result, src=0) + dist.barrier() + + # Evaluation must not perturb the random stream used by resumed training. + torch.set_rng_state(cpu_rng_state) + torch.cuda.set_rng_state(cuda_rng_state, device) + return bool(result[1].item()) + + ################################################################################# # Training Loop # ################################################################################# @@ -112,6 +306,37 @@ def main(args): Trains a new SiT model. """ assert torch.cuda.is_available(), "Training currently requires at least one GPU." + expected_model_module = os.environ.get("SIT_EXPECTED_MODEL_MODULE") + if expected_model_module and MODEL_MODULE_NAME != expected_model_module: + raise RuntimeError( + f"Expected model module {expected_model_module!r}, but loaded " + f"{MODEL_MODULE_NAME!r} from {MODEL_IMPLEMENTATION_PATH}" + ) + + # Load resume metadata before creating the output directory or WandB run. The + # checkpoint hyperparameters remain authoritative; only runtime location, + # target epoch, and logging options may be overridden by the command line. + resume_checkpoint = None + resume_step = 0 + if args.ckpt is not None: + runtime_args = args + resume_checkpoint = torch.load(args.ckpt, map_location="cpu", weights_only=False) + checkpoint_args = resume_checkpoint["args"] + runtime_names = ( + "data_path", "results_dir", "epochs", "wandb", "ckpt", "run_name", + "fid_every_checkpoint", "fid_every", "fid_num_samples", "fid_reference", + "fid_history", "fid_per_proc_batch_size", "fid_inception_batch_size", + "fid_num_workers", "fid_sampling_steps", "fid_seed", + "fid_stop_consecutive_increases", "fid_stop_min_absolute_rise", + "fid_stop_min_relative_rise", + ) + for name in runtime_names: + setattr(checkpoint_args, name, getattr(runtime_args, name)) + args = checkpoint_args + match = re.fullmatch(r"(\d+)\.pt", os.path.basename(args.ckpt)) + if match is None: + raise ValueError("Cannot infer the training step from checkpoint filename; expected NNNNNNN.pt") + resume_step = int(match.group(1)) # Setup DDP: dist.init_process_group("nccl") @@ -129,13 +354,16 @@ def main(args): os.makedirs(args.results_dir, exist_ok=True) # Make results folder (holds all experiment subfolders) experiment_index = len(glob(f"{args.results_dir}/*")) model_string_name = args.model.replace("/", "-") # e.g., SiT-XL/2 --> SiT-XL-2 (for naming folders) - experiment_name = f"{experiment_index:03d}-{model_string_name}-" \ - f"{args.path_type}-{args.prediction}-{args.loss_weight}" + experiment_name = args.run_name or (f"{experiment_index:03d}-{model_string_name}-" \ + f"{args.path_type}-{args.prediction}-{args.loss_weight}") experiment_dir = f"{args.results_dir}/{experiment_name}" # Create an experiment folder checkpoint_dir = f"{experiment_dir}/checkpoints" # Stores saved model checkpoints os.makedirs(checkpoint_dir, exist_ok=True) logger = create_logger(experiment_dir) logger.info(f"Experiment directory created at {experiment_dir}") + logger.info( + f"Model implementation: {MODEL_MODULE_NAME} ({MODEL_IMPLEMENTATION_PATH})" + ) entity = os.environ["ENTITY"] project = os.environ["PROJECT"] @@ -143,6 +371,10 @@ def main(args): wandb_utils.initialize(args, entity, experiment_name, project) else: logger = create_logger(None) + experiment_dir = None + path_objects = [experiment_dir] + dist.broadcast_object_list(path_objects, src=0) + experiment_dir = path_objects[0] # Create model: assert args.image_size % 8 == 0, "Image size must be divisible by 8 (for the VAE encoder)." @@ -155,14 +387,6 @@ def main(args): # Note that parameter initialization is done within the SiT constructor ema = deepcopy(model).to(device) # Create an EMA of the model for use after training - if args.ckpt is not None: - ckpt_path = args.ckpt - state_dict = find_model(ckpt_path) - model.load_state_dict(state_dict["model"]) - ema.load_state_dict(state_dict["ema"]) - opt.load_state_dict(state_dict["opt"]) - args = state_dict["args"] - requires_grad(ema, False) model = DDP(model.to(device), device_ids=[device]) @@ -178,7 +402,16 @@ def main(args): logger.info(f"SiT Parameters: {sum(p.numel() for p in model.parameters()):,}") # Setup optimizer (we used default Adam betas=(0.9, 0.999) and a constant learning rate of 1e-4 in our paper): - opt = torch.optim.AdamW(model.parameters(), lr=2e-4, weight_decay=0) + opt = torch.optim.AdamW(model.parameters(), lr=args.learning_rate, weight_decay=0) + logger.info( + f"Optimizer: AdamW(lr={args.learning_rate:g}, weight_decay=0, " + "betas=(0.9, 0.999))" + ) + if resume_checkpoint is not None: + model.module.load_state_dict(resume_checkpoint["model"]) + ema.load_state_dict(resume_checkpoint["ema"]) + opt.load_state_dict(resume_checkpoint["opt"]) + logger.info(f"Resumed model, EMA, and optimizer from {args.ckpt} at step {resume_step:,}") # Setup data: transform = transforms.Compose([ @@ -207,12 +440,13 @@ def main(args): logger.info(f"Dataset contains {len(dataset):,} images ({args.data_path})") # Prepare models for training: - update_ema(ema, model.module, decay=0) # Ensure EMA is initialized with synced weights + if resume_checkpoint is None: + update_ema(ema, model.module, decay=0) # Initialize EMA only for a new run. model.train() # important! This enables embedding dropout for classifier-free guidance ema.eval() # EMA model should always be in eval mode # Variables for monitoring/logging purposes: - train_steps = 0 + train_steps = resume_step log_steps = 0 running_loss = 0 start_time = time() @@ -235,11 +469,38 @@ def main(args): sample_model_kwargs = dict(y=ys) model_fn = ema.forward - logger.info(f"Training for {args.epochs} epochs...") - for epoch in range(args.epochs): + steps_per_epoch = len(loader) + target_steps = args.epochs * steps_per_epoch + start_epoch, batches_to_skip = divmod(train_steps, steps_per_epoch) + logger.info( + f"Training to {args.epochs} total epochs ({target_steps:,} steps); " + f"starting at epoch {start_epoch}, batch {batches_to_skip}, step {train_steps:,}." + ) + # Establish a same-protocol 50k baseline for the resume checkpoint before + # comparing later checkpoints against it. A recorded baseline is reused. + if args.fid_every_checkpoint and resume_checkpoint is not None: + if evaluate_checkpoint_fid( + ema, vae, transport_sampler, args, train_steps, + device, rank, logger, experiment_dir, + ): + logger.error("Resume checkpoint already violates the recorded FID trend; exiting.") + cleanup() + return + start_time = time() + + stop_requested = False + for epoch in range(start_epoch, args.epochs): sampler.set_epoch(epoch) logger.info(f"Beginning epoch {epoch}...") - for x, y in loader: + epoch_loader = loader + if epoch == start_epoch and batches_to_skip: + epoch_loader = DataLoader( + dataset, + batch_sampler=SkipBatchSampler(loader.batch_sampler, batches_to_skip), + num_workers=args.num_workers, + pin_memory=True, + ) + for x, y in epoch_loader: x = x.to(device) y = y.to(device) with torch.no_grad(): @@ -290,6 +551,14 @@ def main(args): torch.save(checkpoint, checkpoint_path) logger.info(f"Saved checkpoint to {checkpoint_path}") dist.barrier() + if args.fid_every_checkpoint and train_steps % args.fid_every == 0: + stop_requested = evaluate_checkpoint_fid( + ema, vae, transport_sampler, args, train_steps, + device, rank, logger, experiment_dir, + ) + start_time = time() + if stop_requested: + break if train_steps % args.sample_every == 0 and train_steps > 0: logger.info("Generating EMA samples...") @@ -308,6 +577,21 @@ def main(args): wandb_utils.log_image(out_samples, train_steps) logging.info("Generating EMA samples done.") + if stop_requested: + break + + if rank == 0 and not stop_requested and train_steps % args.ckpt_every != 0: + checkpoint = { + "model": model.module.state_dict(), + "ema": ema.state_dict(), + "opt": opt.state_dict(), + "args": args, + } + checkpoint_path = f"{checkpoint_dir}/{train_steps:07d}.pt" + torch.save(checkpoint, checkpoint_path) + logger.info(f"Saved final checkpoint to {checkpoint_path}") + dist.barrier() + model.eval() # important! This disables randomized embedding dropout # do any sampling/FID calculation/etc. with ema (or model) in eval mode ... @@ -325,6 +609,7 @@ if __name__ == "__main__": parser.add_argument("--num-classes", type=int, default=1000) parser.add_argument("--epochs", type=int, default=1400) parser.add_argument("--global-batch-size", type=int, default=256) + parser.add_argument("--learning-rate", type=float, default=1e-4) parser.add_argument("--global-seed", type=int, default=0) parser.add_argument("--vae", type=str, choices=["ema", "mse"], default="ema") # Choice doesn't affect training parser.add_argument("--num-workers", type=int, default=4) @@ -335,6 +620,27 @@ if __name__ == "__main__": parser.add_argument("--wandb", action="store_true") parser.add_argument("--ckpt", type=str, default=None, help="Optional path to a custom SiT checkpoint") + parser.add_argument("--run-name", type=str, default=None, + help="Experiment directory and WandB run name (useful when resuming)") + parser.add_argument("--fid-every-checkpoint", action="store_true", + help="Run periodic CFG=1 PyTorch FID checks and stop on sustained regression") + parser.add_argument("--fid-every", type=int, default=50_000, + help="Training-step interval between FID checks") + parser.add_argument("--fid-num-samples", type=int, default=50_000) + parser.add_argument("--fid-reference", type=str, + default="/home/nvidia/evaluation/reference/discon-download/VIRTUAL_imagenet256_labeled.npz") + parser.add_argument("--fid-history", type=str, default=None) + parser.add_argument("--fid-per-proc-batch-size", type=int, default=64) + parser.add_argument("--fid-inception-batch-size", type=int, default=128) + parser.add_argument("--fid-num-workers", type=int, default=8) + parser.add_argument("--fid-sampling-steps", type=int, default=250) + parser.add_argument("--fid-seed", type=int, default=0) + parser.add_argument("--fid-stop-consecutive-increases", type=int, default=3, + help="Stop only after this many consecutive checkpoint FID increases") + parser.add_argument("--fid-stop-min-absolute-rise", type=float, default=0.25, + help="Minimum cumulative absolute FID rise required to stop") + parser.add_argument("--fid-stop-min-relative-rise", type=float, default=0.005, + help="Minimum cumulative relative FID rise required to stop") parse_transport_args(parser) args = parser.parse_args() diff --git a/train_conv.py b/train_conv.py old mode 100644 new mode 100755 index ee1640b..9989900 --- a/train_conv.py +++ b/train_conv.py @@ -1,341 +1,13 @@ -# This source code is licensed under the license found in the -# LICENSE file in the root directory of this source tree. +#!/usr/bin/env python3 +"""Run the shared SiT trainer with the convolutional-layer model implementation.""" -""" -A minimal training script for SiT using PyTorch DDP. -""" -import torch -# the first flag below was False when we tested this script but True makes A100 training a lot faster: -torch.backends.cuda.matmul.allow_tf32 = True -torch.backends.cudnn.allow_tf32 = True -import torch.distributed as dist -from torch.nn.parallel import DistributedDataParallel as DDP -from torch.utils.data import DataLoader -from torch.utils.data.distributed import DistributedSampler -from torchvision.datasets import ImageFolder -from torchvision import transforms -import numpy as np -from collections import OrderedDict -from PIL import Image -from copy import deepcopy -from glob import glob -from time import time -import argparse -import logging import os +import runpy -from models_conv import SiT_models -from download import find_model -from transport import create_transport, Sampler -from diffusers.models import AutoencoderKL -from train_utils import parse_transport_args -import wandb_utils +# Select models_conv.py before train.py imports the model registry, and require +# the shared trainer to fail closed if another model module is selected. +os.environ["SIT_MODEL_MODULE"] = "models_conv" +os.environ["SIT_EXPECTED_MODEL_MODULE"] = "models_conv" -################################################################################# -# Training Helper Functions # -################################################################################# - -@torch.no_grad() -def update_ema(ema_model, model, decay=0.9999): - """ - Step the EMA model towards the current model. - """ - ema_params = OrderedDict(ema_model.named_parameters()) - model_params = OrderedDict(model.named_parameters()) - - for name, param in model_params.items(): - # TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed - ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay) - - -def requires_grad(model, flag=True): - """ - Set requires_grad flag for all parameters in a model. - """ - for p in model.parameters(): - p.requires_grad = flag - - -def cleanup(): - """ - End DDP training. - """ - dist.destroy_process_group() - - -def create_logger(logging_dir): - """ - Create a logger that writes to a log file and stdout. - """ - if dist.get_rank() == 0: # real logger - logging.basicConfig( - level=logging.INFO, - format='[\033[34m%(asctime)s\033[0m] %(message)s', - datefmt='%Y-%m-%d %H:%M:%S', - handlers=[logging.StreamHandler(), logging.FileHandler(f"{logging_dir}/log.txt")] - ) - logger = logging.getLogger(__name__) - else: # dummy logger (does nothing) - logger = logging.getLogger(__name__) - logger.addHandler(logging.NullHandler()) - return logger - - -def center_crop_arr(pil_image, image_size): - """ - Center cropping implementation from ADM. - https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126 - """ - while min(*pil_image.size) >= 2 * image_size: - pil_image = pil_image.resize( - tuple(x // 2 for x in pil_image.size), resample=Image.BOX - ) - - scale = image_size / min(*pil_image.size) - pil_image = pil_image.resize( - tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC - ) - - arr = np.array(pil_image) - crop_y = (arr.shape[0] - image_size) // 2 - crop_x = (arr.shape[1] - image_size) // 2 - return Image.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size]) - - -################################################################################# -# Training Loop # -################################################################################# - -def main(args): - """ - Trains a new SiT model. - """ - assert torch.cuda.is_available(), "Training currently requires at least one GPU." - - # Setup DDP: - dist.init_process_group("nccl") - assert args.global_batch_size % dist.get_world_size() == 0, f"Batch size must be divisible by world size." - rank = dist.get_rank() - device = rank % torch.cuda.device_count() - seed = args.global_seed * dist.get_world_size() + rank - torch.manual_seed(seed) - torch.cuda.set_device(device) - print(f"Starting rank={rank}, seed={seed}, world_size={dist.get_world_size()}.") - local_batch_size = int(args.global_batch_size // dist.get_world_size()) - - # Setup an experiment folder: - if rank == 0: - os.makedirs(args.results_dir, exist_ok=True) # Make results folder (holds all experiment subfolders) - experiment_index = len(glob(f"{args.results_dir}/*")) - model_string_name = args.model.replace("/", "-") # e.g., SiT-XL/2 --> SiT-XL-2 (for naming folders) - experiment_name = f"{experiment_index:03d}-{model_string_name}-conv-" \ - f"{args.path_type}-{args.prediction}-{args.loss_weight}" - experiment_dir = f"{args.results_dir}/{experiment_name}" # Create an experiment folder - checkpoint_dir = f"{experiment_dir}/checkpoints" # Stores saved model checkpoints - os.makedirs(checkpoint_dir, exist_ok=True) - logger = create_logger(experiment_dir) - logger.info(f"Experiment directory created at {experiment_dir}") - - entity = os.environ["ENTITY"] - project = os.environ["PROJECT"] - if args.wandb: - wandb_utils.initialize(args, entity, experiment_name, project) - else: - logger = create_logger(None) - - # Create model: - assert args.image_size % 8 == 0, "Image size must be divisible by 8 (for the VAE encoder)." - latent_size = args.image_size // 8 - model = SiT_models[args.model]( - input_size=latent_size, - num_classes=args.num_classes - ) - - # Note that parameter initialization is done within the SiT constructor - ema = deepcopy(model).to(device) # Create an EMA of the model for use after training - - if args.ckpt is not None: - ckpt_path = args.ckpt - state_dict = find_model(ckpt_path) - model.load_state_dict(state_dict["model"]) - ema.load_state_dict(state_dict["ema"]) - opt.load_state_dict(state_dict["opt"]) - args = state_dict["args"] - - requires_grad(ema, False) - - model = DDP(model.to(device), device_ids=[device]) - transport = create_transport( - args.path_type, - args.prediction, - args.loss_weight, - args.train_eps, - args.sample_eps - ) # default: velocity; - transport_sampler = Sampler(transport) - vae = AutoencoderKL.from_pretrained(f"stabilityai/sd-vae-ft-{args.vae}").to(device) - logger.info(f"SiT Parameters: {sum(p.numel() for p in model.parameters()):,}") - - # Setup optimizer (we used default Adam betas=(0.9, 0.999) and a constant learning rate of 1e-4 in our paper): - opt = torch.optim.AdamW(model.parameters(), lr=2e-4, weight_decay=0) - - # Setup data: - transform = transforms.Compose([ - transforms.Lambda(lambda pil_image: center_crop_arr(pil_image, args.image_size)), - transforms.RandomHorizontalFlip(), - transforms.ToTensor(), - transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True) - ]) - dataset = ImageFolder(args.data_path, transform=transform) - sampler = DistributedSampler( - dataset, - num_replicas=dist.get_world_size(), - rank=rank, - shuffle=True, - seed=args.global_seed - ) - loader = DataLoader( - dataset, - batch_size=local_batch_size, - shuffle=False, - sampler=sampler, - num_workers=args.num_workers, - pin_memory=True, - drop_last=True - ) - logger.info(f"Dataset contains {len(dataset):,} images ({args.data_path})") - - # Prepare models for training: - update_ema(ema, model.module, decay=0) # Ensure EMA is initialized with synced weights - model.train() # important! This enables embedding dropout for classifier-free guidance - ema.eval() # EMA model should always be in eval mode - - # Variables for monitoring/logging purposes: - train_steps = 0 - log_steps = 0 - running_loss = 0 - start_time = time() - - # Labels to condition the model with (feel free to change): - ys = torch.randint(1000, size=(local_batch_size,), device=device) - use_cfg = args.cfg_scale > 1.0 - # Create sampling noise: - n = ys.size(0) - zs = torch.randn(n, 4, latent_size, latent_size, device=device) - - # Setup classifier-free guidance: - if use_cfg: - zs = torch.cat([zs, zs], 0) - y_null = torch.tensor([1000] * n, device=device) - ys = torch.cat([ys, y_null], 0) - sample_model_kwargs = dict(y=ys, cfg_scale=args.cfg_scale) - model_fn = ema.forward_with_cfg - else: - sample_model_kwargs = dict(y=ys) - model_fn = ema.forward - - logger.info(f"Training for {args.epochs} epochs...") - for epoch in range(args.epochs): - sampler.set_epoch(epoch) - logger.info(f"Beginning epoch {epoch}...") - for x, y in loader: - x = x.to(device) - y = y.to(device) - with torch.no_grad(): - # Map input images to latent space + normalize latents: - x = vae.encode(x).latent_dist.sample().mul_(0.18215) - model_kwargs = dict(y=y) - loss_dict = transport.training_losses(model, x, model_kwargs) - loss = loss_dict["loss"].mean() - opt.zero_grad() - loss.backward() - opt.step() - update_ema(ema, model.module) - - # Log loss values: - running_loss += loss.item() - log_steps += 1 - train_steps += 1 - if train_steps % args.log_every == 0: - # Measure training speed: - torch.cuda.synchronize() - end_time = time() - steps_per_sec = log_steps / (end_time - start_time) - # Reduce loss history over all processes: - avg_loss = torch.tensor(running_loss / log_steps, device=device) - dist.all_reduce(avg_loss, op=dist.ReduceOp.SUM) - avg_loss = avg_loss.item() / dist.get_world_size() - logger.info(f"(step={train_steps:07d}) Train Loss: {avg_loss:.4f}, Train Steps/Sec: {steps_per_sec:.2f}") - if args.wandb: - wandb_utils.log( - { "train loss": avg_loss, "train steps/sec": steps_per_sec }, - step=train_steps - ) - # Reset monitoring variables: - running_loss = 0 - log_steps = 0 - start_time = time() - - # Save SiT checkpoint: - if train_steps % args.ckpt_every == 0 and train_steps > 0: - if rank == 0: - checkpoint = { - "model": model.module.state_dict(), - "ema": ema.state_dict(), - "opt": opt.state_dict(), - "args": args - } - checkpoint_path = f"{checkpoint_dir}/{train_steps:07d}.pt" - torch.save(checkpoint, checkpoint_path) - logger.info(f"Saved checkpoint to {checkpoint_path}") - dist.barrier() - - if train_steps % args.sample_every == 0 and train_steps > 0: - logger.info("Generating EMA samples...") - with torch.no_grad(): - sample_fn = transport_sampler.sample_ode() # default to ode sampling - samples = sample_fn(zs, model_fn, **sample_model_kwargs)[-1] - dist.barrier() - - if use_cfg: #remove null samples - samples, _ = samples.chunk(2, dim=0) - samples = vae.decode(samples / 0.18215).sample - out_samples = torch.zeros((args.global_batch_size, 3, args.image_size, args.image_size), device=device) - dist.all_gather_into_tensor(out_samples, samples) - - if args.wandb: - wandb_utils.log_image(out_samples, train_steps) - logging.info("Generating EMA samples done.") - - model.eval() # important! This disables randomized embedding dropout - # do any sampling/FID calculation/etc. with ema (or model) in eval mode ... - - logger.info("Done!") - cleanup() - - -if __name__ == "__main__": - # Default args here will train SiT-XL/2 with the hyperparameters we used in our paper (except training iters). - parser = argparse.ArgumentParser() - parser.add_argument("--data-path", type=str, required=True) - parser.add_argument("--results-dir", type=str, default="results") - parser.add_argument("--model", type=str, choices=list(SiT_models.keys()), default="SiT-XL/2") - parser.add_argument("--image-size", type=int, choices=[256, 512], default=256) - parser.add_argument("--num-classes", type=int, default=1000) - parser.add_argument("--epochs", type=int, default=1400) - parser.add_argument("--global-batch-size", type=int, default=256) - parser.add_argument("--global-seed", type=int, default=0) - parser.add_argument("--vae", type=str, choices=["ema", "mse"], default="ema") # Choice doesn't affect training - parser.add_argument("--num-workers", type=int, default=4) - parser.add_argument("--log-every", type=int, default=100) - parser.add_argument("--ckpt-every", type=int, default=50_000) - parser.add_argument("--sample-every", type=int, default=10_000) - parser.add_argument("--cfg-scale", type=float, default=4.0) - parser.add_argument("--wandb", action="store_true") - parser.add_argument("--ckpt", type=str, default=None, - help="Optional path to a custom SiT checkpoint") - - parse_transport_args(parser) - args = parser.parse_args() - main(args) +runpy.run_module("train", run_name="__main__") diff --git a/train_rot_layer.py b/train_rot_layer.py old mode 100644 new mode 100755 index 3b40586..8fdc85a --- a/train_rot_layer.py +++ b/train_rot_layer.py @@ -1,341 +1,14 @@ -# This source code is licensed under the license found in the -# LICENSE file in the root directory of this source tree. +#!/usr/bin/env python3 +"""Run the shared SiT trainer with the rotation-layer model implementation.""" -""" -A minimal training script for SiT using PyTorch DDP. -""" -import torch -# the first flag below was False when we tested this script but True makes A100 training a lot faster: -torch.backends.cuda.matmul.allow_tf32 = True -torch.backends.cudnn.allow_tf32 = True -import torch.distributed as dist -from torch.nn.parallel import DistributedDataParallel as DDP -from torch.utils.data import DataLoader -from torch.utils.data.distributed import DistributedSampler -from torchvision.datasets import ImageFolder -from torchvision import transforms -import numpy as np -from collections import OrderedDict -from PIL import Image -from copy import deepcopy -from glob import glob -from time import time -import argparse -import logging import os +import runpy -from models_rot_layer import SiT_models -from download import find_model -from transport import create_transport, Sampler -from diffusers.models import AutoencoderKL -from train_utils import parse_transport_args -import wandb_utils +# train.py imports its model registry at module load time. Set both variables +# first, and require the shared trainer to fail closed if another registry is +# ever selected accidentally. +os.environ["SIT_MODEL_MODULE"] = "models_rot_layer" +os.environ["SIT_EXPECTED_MODEL_MODULE"] = "models_rot_layer" -################################################################################# -# Training Helper Functions # -################################################################################# - -@torch.no_grad() -def update_ema(ema_model, model, decay=0.9999): - """ - Step the EMA model towards the current model. - """ - ema_params = OrderedDict(ema_model.named_parameters()) - model_params = OrderedDict(model.named_parameters()) - - for name, param in model_params.items(): - # TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed - ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay) - - -def requires_grad(model, flag=True): - """ - Set requires_grad flag for all parameters in a model. - """ - for p in model.parameters(): - p.requires_grad = flag - - -def cleanup(): - """ - End DDP training. - """ - dist.destroy_process_group() - - -def create_logger(logging_dir): - """ - Create a logger that writes to a log file and stdout. - """ - if dist.get_rank() == 0: # real logger - logging.basicConfig( - level=logging.INFO, - format='[\033[34m%(asctime)s\033[0m] %(message)s', - datefmt='%Y-%m-%d %H:%M:%S', - handlers=[logging.StreamHandler(), logging.FileHandler(f"{logging_dir}/log.txt")] - ) - logger = logging.getLogger(__name__) - else: # dummy logger (does nothing) - logger = logging.getLogger(__name__) - logger.addHandler(logging.NullHandler()) - return logger - - -def center_crop_arr(pil_image, image_size): - """ - Center cropping implementation from ADM. - https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126 - """ - while min(*pil_image.size) >= 2 * image_size: - pil_image = pil_image.resize( - tuple(x // 2 for x in pil_image.size), resample=Image.BOX - ) - - scale = image_size / min(*pil_image.size) - pil_image = pil_image.resize( - tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC - ) - - arr = np.array(pil_image) - crop_y = (arr.shape[0] - image_size) // 2 - crop_x = (arr.shape[1] - image_size) // 2 - return Image.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size]) - - -################################################################################# -# Training Loop # -################################################################################# - -def main(args): - """ - Trains a new SiT model. - """ - assert torch.cuda.is_available(), "Training currently requires at least one GPU." - - # Setup DDP: - dist.init_process_group("nccl") - assert args.global_batch_size % dist.get_world_size() == 0, f"Batch size must be divisible by world size." - rank = dist.get_rank() - device = rank % torch.cuda.device_count() - seed = args.global_seed * dist.get_world_size() + rank - torch.manual_seed(seed) - torch.cuda.set_device(device) - print(f"Starting rank={rank}, seed={seed}, world_size={dist.get_world_size()}.") - local_batch_size = int(args.global_batch_size // dist.get_world_size()) - - # Setup an experiment folder: - if rank == 0: - os.makedirs(args.results_dir, exist_ok=True) # Make results folder (holds all experiment subfolders) - experiment_index = len(glob(f"{args.results_dir}/*")) - model_string_name = args.model.replace("/", "-") # e.g., SiT-XL/2 --> SiT-XL-2 (for naming folders) - experiment_name = f"{experiment_index:03d}-{model_string_name}-rot-layer-" \ - f"{args.path_type}-{args.prediction}-{args.loss_weight}" - experiment_dir = f"{args.results_dir}/{experiment_name}" # Create an experiment folder - checkpoint_dir = f"{experiment_dir}/checkpoints" # Stores saved model checkpoints - os.makedirs(checkpoint_dir, exist_ok=True) - logger = create_logger(experiment_dir) - logger.info(f"Experiment directory created at {experiment_dir}") - - entity = os.environ["ENTITY"] - project = os.environ["PROJECT"] - if args.wandb: - wandb_utils.initialize(args, entity, experiment_name, project) - else: - logger = create_logger(None) - - # Create model: - assert args.image_size % 8 == 0, "Image size must be divisible by 8 (for the VAE encoder)." - latent_size = args.image_size // 8 - model = SiT_models[args.model]( - input_size=latent_size, - num_classes=args.num_classes - ) - - # Note that parameter initialization is done within the SiT constructor - ema = deepcopy(model).to(device) # Create an EMA of the model for use after training - - if args.ckpt is not None: - ckpt_path = args.ckpt - state_dict = find_model(ckpt_path) - model.load_state_dict(state_dict["model"]) - ema.load_state_dict(state_dict["ema"]) - opt.load_state_dict(state_dict["opt"]) - args = state_dict["args"] - - requires_grad(ema, False) - - model = DDP(model.to(device), device_ids=[device]) - transport = create_transport( - args.path_type, - args.prediction, - args.loss_weight, - args.train_eps, - args.sample_eps - ) # default: velocity; - transport_sampler = Sampler(transport) - vae = AutoencoderKL.from_pretrained(f"stabilityai/sd-vae-ft-{args.vae}").to(device) - logger.info(f"SiT Parameters: {sum(p.numel() for p in model.parameters()):,}") - - # Setup optimizer (we used default Adam betas=(0.9, 0.999) and a constant learning rate of 1e-4 in our paper): - opt = torch.optim.AdamW(model.parameters(), lr=2e-4, weight_decay=0) - - # Setup data: - transform = transforms.Compose([ - transforms.Lambda(lambda pil_image: center_crop_arr(pil_image, args.image_size)), - transforms.RandomHorizontalFlip(), - transforms.ToTensor(), - transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True) - ]) - dataset = ImageFolder(args.data_path, transform=transform) - sampler = DistributedSampler( - dataset, - num_replicas=dist.get_world_size(), - rank=rank, - shuffle=True, - seed=args.global_seed - ) - loader = DataLoader( - dataset, - batch_size=local_batch_size, - shuffle=False, - sampler=sampler, - num_workers=args.num_workers, - pin_memory=True, - drop_last=True - ) - logger.info(f"Dataset contains {len(dataset):,} images ({args.data_path})") - - # Prepare models for training: - update_ema(ema, model.module, decay=0) # Ensure EMA is initialized with synced weights - model.train() # important! This enables embedding dropout for classifier-free guidance - ema.eval() # EMA model should always be in eval mode - - # Variables for monitoring/logging purposes: - train_steps = 0 - log_steps = 0 - running_loss = 0 - start_time = time() - - # Labels to condition the model with (feel free to change): - ys = torch.randint(1000, size=(local_batch_size,), device=device) - use_cfg = args.cfg_scale > 1.0 - # Create sampling noise: - n = ys.size(0) - zs = torch.randn(n, 4, latent_size, latent_size, device=device) - - # Setup classifier-free guidance: - if use_cfg: - zs = torch.cat([zs, zs], 0) - y_null = torch.tensor([1000] * n, device=device) - ys = torch.cat([ys, y_null], 0) - sample_model_kwargs = dict(y=ys, cfg_scale=args.cfg_scale) - model_fn = ema.forward_with_cfg - else: - sample_model_kwargs = dict(y=ys) - model_fn = ema.forward - - logger.info(f"Training for {args.epochs} epochs...") - for epoch in range(args.epochs): - sampler.set_epoch(epoch) - logger.info(f"Beginning epoch {epoch}...") - for x, y in loader: - x = x.to(device) - y = y.to(device) - with torch.no_grad(): - # Map input images to latent space + normalize latents: - x = vae.encode(x).latent_dist.sample().mul_(0.18215) - model_kwargs = dict(y=y) - loss_dict = transport.training_losses(model, x, model_kwargs) - loss = loss_dict["loss"].mean() - opt.zero_grad() - loss.backward() - opt.step() - update_ema(ema, model.module) - - # Log loss values: - running_loss += loss.item() - log_steps += 1 - train_steps += 1 - if train_steps % args.log_every == 0: - # Measure training speed: - torch.cuda.synchronize() - end_time = time() - steps_per_sec = log_steps / (end_time - start_time) - # Reduce loss history over all processes: - avg_loss = torch.tensor(running_loss / log_steps, device=device) - dist.all_reduce(avg_loss, op=dist.ReduceOp.SUM) - avg_loss = avg_loss.item() / dist.get_world_size() - logger.info(f"(step={train_steps:07d}) Train Loss: {avg_loss:.4f}, Train Steps/Sec: {steps_per_sec:.2f}") - if args.wandb: - wandb_utils.log( - { "train loss": avg_loss, "train steps/sec": steps_per_sec }, - step=train_steps - ) - # Reset monitoring variables: - running_loss = 0 - log_steps = 0 - start_time = time() - - # Save SiT checkpoint: - if train_steps % args.ckpt_every == 0 and train_steps > 0: - if rank == 0: - checkpoint = { - "model": model.module.state_dict(), - "ema": ema.state_dict(), - "opt": opt.state_dict(), - "args": args - } - checkpoint_path = f"{checkpoint_dir}/{train_steps:07d}.pt" - torch.save(checkpoint, checkpoint_path) - logger.info(f"Saved checkpoint to {checkpoint_path}") - dist.barrier() - - if train_steps % args.sample_every == 0 and train_steps > 0: - logger.info("Generating EMA samples...") - with torch.no_grad(): - sample_fn = transport_sampler.sample_ode() # default to ode sampling - samples = sample_fn(zs, model_fn, **sample_model_kwargs)[-1] - dist.barrier() - - if use_cfg: #remove null samples - samples, _ = samples.chunk(2, dim=0) - samples = vae.decode(samples / 0.18215).sample - out_samples = torch.zeros((args.global_batch_size, 3, args.image_size, args.image_size), device=device) - dist.all_gather_into_tensor(out_samples, samples) - - if args.wandb: - wandb_utils.log_image(out_samples, train_steps) - logging.info("Generating EMA samples done.") - - model.eval() # important! This disables randomized embedding dropout - # do any sampling/FID calculation/etc. with ema (or model) in eval mode ... - - logger.info("Done!") - cleanup() - - -if __name__ == "__main__": - # Default args here will train SiT-XL/2 with the hyperparameters we used in our paper (except training iters). - parser = argparse.ArgumentParser() - parser.add_argument("--data-path", type=str, required=True) - parser.add_argument("--results-dir", type=str, default="results") - parser.add_argument("--model", type=str, choices=list(SiT_models.keys()), default="SiT-XL/2") - parser.add_argument("--image-size", type=int, choices=[256, 512], default=256) - parser.add_argument("--num-classes", type=int, default=1000) - parser.add_argument("--epochs", type=int, default=1400) - parser.add_argument("--global-batch-size", type=int, default=256) - parser.add_argument("--global-seed", type=int, default=0) - parser.add_argument("--vae", type=str, choices=["ema", "mse"], default="ema") # Choice doesn't affect training - parser.add_argument("--num-workers", type=int, default=4) - parser.add_argument("--log-every", type=int, default=100) - parser.add_argument("--ckpt-every", type=int, default=50_000) - parser.add_argument("--sample-every", type=int, default=10_000) - parser.add_argument("--cfg-scale", type=float, default=4.0) - parser.add_argument("--wandb", action="store_true") - parser.add_argument("--ckpt", type=str, default=None, - help="Optional path to a custom SiT checkpoint") - - parse_transport_args(parser) - args = parser.parse_args() - main(args) +runpy.run_module("train", run_name="__main__")