SiT-Complementary / code_snapshot /server_changes.patch
BlueSourceJY's picture
Backup final base/rotation-layer and latest conv-layer experiments
8b0b874 verified
Raw History Blame Contribute Delete
59.9 kB
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__")