diff --git a/cosmodiff/optim.py b/cosmodiff/optim.py index 5c8746b..9b44196 100644 --- a/cosmodiff/optim.py +++ b/cosmodiff/optim.py @@ -1,5 +1,6 @@ import os import pickle +import yaml import numpy as np import torch import torch.nn.functional as F @@ -14,14 +15,24 @@ from . import utils +def _to_yaml_safe(obj): + """Recursively convert tuples to lists so yaml.safe_load can round-trip.""" + if isinstance(obj, dict): + return {k: _to_yaml_safe(v) for k, v in obj.items()} + if isinstance(obj, (list, tuple)): + return [_to_yaml_safe(v) for v in obj] + return obj + + def train( dataset, - model, + model=None, *, noise_scheduler=None, # None → DDPMScheduler(num_train_timesteps=1000) optimizer=None, # None → AdamW lr_scheduler=None, # None → ConstantLR() output_dir: str = "checkpoints", + resume_from_checkpoint: Optional[str] = None, num_epochs: int = 50, batch_size: int = 16, shuffle: bool = True, @@ -36,11 +47,10 @@ def train( ): """Train a diffusers diffusion model. - To resume from a checkpoint, load it first with ``load_checkpoint()`` and - pass the returned objects directly into this function. Augmentations are - expected to be built into the dataset object directly (e.g. via - ``ArrayDataset.augmentations``), and are checkpointed automatically if the - dataset exposes an ``augmentations`` attribute. + ``model`` is optional when ``resume_from_checkpoint`` is set — the model + (and any unspecified scheduler/optimizer/augmentations) are loaded from the + checkpoint automatically. Augmentations are checkpointed automatically if + the dataset exposes an ``augmentations`` attribute. The model's forward call is dispatched automatically: if the batch contains ``"labels"``, they are passed as a keyword argument (for @@ -54,8 +64,9 @@ def train( ``"labels"`` key (LongTensor of shape ``(batch_size,)``) for class-conditional DiT training. Augmentations should be applied inside the dataset's ``__getitem__``. - model (nn.Module): Pre-instantiated diffusers model (e.g. - ``UNet2DModel``, ``DiTTransformer2DModel``). + model (nn.Module, optional): Pre-instantiated diffusers model (e.g. + ``UNet2DModel``, ``DiTTransformer2DModel``). May be omitted when + ``resume_from_checkpoint`` is provided. noise_scheduler (optional): Pre-instantiated diffusers noise scheduler. Defaults to ``DDPMScheduler(num_train_timesteps=1000)``. optimizer (torch.optim.Optimizer, optional): Optimizer for ``model``. @@ -65,6 +76,12 @@ def train( fixed learning rate for the entire run. output_dir (str): Root directory for checkpoints and TensorBoard logs. Defaults to ``"checkpoints"``. + resume_from_checkpoint (str, optional): Path to a checkpoint directory + produced by a previous call to ``train()``. Objects not explicitly + passed (model, noise_scheduler, optimizer, lr_scheduler) are loaded + from the checkpoint. After ``accelerator.prepare()`` the full + training state (optimizer moments, grad scaler, RNG) is restored + via ``accelerator.load_state()``. num_epochs (int): Total number of training epochs. Defaults to ``50``. batch_size (int): Per-device batch size. Defaults to ``16``. shuffle (bool): Shuffle the dataset each epoch. Defaults to ``True``. @@ -104,23 +121,29 @@ def train( # dataset must return dicts with "images" and "labels" keys train(my_dataset, model) - Resume from a checkpoint:: + Resume from a checkpoint (model loaded automatically):: - model, noise_scheduler, optimizer, lr_scheduler, augmentations = ( - load_checkpoint("checkpoints/checkpoint-epoch-10") - ) - dataset.augmentations = augmentations - train( - my_dataset, - model, - noise_scheduler=noise_scheduler, - optimizer=optimizer, - lr_scheduler=lr_scheduler, - ) + train(my_dataset, resume_from_checkpoint="checkpoints/checkpoint-epoch-0010") """ # ------------------------------------------------------------------ # - # 1. Defaults # + # 1. Defaults / checkpoint loading # # ------------------------------------------------------------------ # + start_epoch = 0 + + if model is None and resume_from_checkpoint is None: + raise ValueError( + "Either `model` or `resume_from_checkpoint` must be provided." + ) + + if resume_from_checkpoint is not None: + model, noise_scheduler, optimizer, lr_scheduler, _aug = ( + utils.load_checkpoint(resume_from_checkpoint) + ) + if isinstance(dataset, utils.ArrayDataset): + dataset.augmentations = _aug + + start_epoch = int(resume_from_checkpoint.split('-')[-1]) + 1 + if noise_scheduler is None: noise_scheduler = DDPMScheduler(num_train_timesteps=1000) @@ -142,6 +165,26 @@ def train( ) accelerator.init_trackers(project_name="cosmodiff") + # Register hooks so save_state() delegates model serialisation to + # save_pretrained() and load_state() restores via from_pretrained(), + # avoiding a redundant second copy of the weights on disk. + def _save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + for m in models: + m.save_pretrained(output_dir) + weights.clear() + + def _load_model_hook(models, input_dir): + for _ in range(len(models)): + m = models.pop() + loaded = m.__class__.from_pretrained(input_dir) + m.register_to_config(**loaded.config) + m.load_state_dict(loaded.state_dict()) + del loaded + + accelerator.register_save_state_pre_hook(_save_model_hook) + accelerator.register_load_state_pre_hook(_load_model_hook) + # ------------------------------------------------------------------ # # 3. DataLoader # # ------------------------------------------------------------------ # @@ -160,6 +203,10 @@ def train( model, optimizer, dataloader, lr_scheduler = accelerator.prepare( model, optimizer, dataloader, lr_scheduler ) + + if resume_from_checkpoint is not None: + accelerator.load_state(resume_from_checkpoint) + # ------------------------------------------------------------------ # # 5. Training loop # # ------------------------------------------------------------------ # @@ -174,10 +221,10 @@ def train( "epoch_lr": [], } - for epoch in range(num_epochs): + for epoch in range(start_epoch, start_epoch + num_epochs): progress = tqdm( dataloader, - desc=f"Epoch {epoch}/{num_epochs - 1}", + desc=f"Epoch {epoch}/{start_epoch + num_epochs - 1}", disable=not verbose or not accelerator.is_local_main_process, ) @@ -258,17 +305,35 @@ def train( # ---------------------------------------------------------------- # # 6. Checkpointing # # ---------------------------------------------------------------- # - if (epoch + 1) % checkpoint_every_n_epochs == 0 or epoch == num_epochs - 1: + if (epoch + 1) % checkpoint_every_n_epochs == 0 or epoch == (start_epoch + num_epochs - 1): if accelerator.is_main_process: ckpt_save_path = os.path.join(output_dir, f"checkpoint-epoch-{epoch:04d}") + + # Noise scheduler config (needed by SchedulerClass.from_pretrained) + noise_scheduler.save_pretrained(ckpt_save_path) + + # Class names and constructor kwargs for fresh reconstruction: + # this is only needed when resuming training from a checkpoint. + raw_opt = optimizer.optimizer + raw_sched = lr_scheduler.scheduler + ckpt_cfg = { + "noise_scheduler": { + "class": f"{noise_scheduler.__class__.__module__}.{noise_scheduler.__class__.__name__}", + }, + "optimizer": { + "class": f"{raw_opt.__class__.__module__}.{raw_opt.__class__.__name__}", + }, + "lr_scheduler": { + "class": f"{raw_sched.__class__.__module__}.{raw_sched.__class__.__name__}", + "kwargs": utils._get_lr_scheduler_kwargs(raw_sched), + }, + } + with open(os.path.join(ckpt_save_path, "checkpoint_config.yaml"), "w") as f: + yaml.dump(_to_yaml_safe(ckpt_cfg), f) + + # Model weights (via hook) + optimizer moments + grad scaler + RNG accelerator.save_state(ckpt_save_path) - accelerator.unwrap_model(model).save_pretrained(ckpt_save_path) - with open(os.path.join(ckpt_save_path, "optimizer.pkl"), "wb") as f: - pickle.dump(optimizer.optimizer, f) - with open(os.path.join(ckpt_save_path, "noise_scheduler.pkl"), "wb") as f: - pickle.dump(noise_scheduler, f) - with open(os.path.join(ckpt_save_path, "lr_scheduler.pkl"), "wb") as f: - pickle.dump(lr_scheduler.scheduler, f) + if hasattr(dataset, "augmentations") and dataset.augmentations is not None: with open(os.path.join(ckpt_save_path, "augmentations.pkl"), "wb") as f: pickle.dump(dataset.augmentations, f) @@ -356,10 +421,7 @@ def generate( else: noise_pred = model(images, timesteps, return_dict=False)[0] - step_kwargs = {} - if "generator" in noise_scheduler.step.__code__.co_varnames: - step_kwargs["generator"] = generator - images = noise_scheduler.step(noise_pred, t, images, **step_kwargs).prev_sample + images = noise_scheduler.step(noise_pred, t, images, generator=generator).prev_sample if renorm is not None: images = renorm(images) diff --git a/cosmodiff/tests/test_optim.py b/cosmodiff/tests/test_optim.py index 0374bf2..8cf7fad 100644 --- a/cosmodiff/tests/test_optim.py +++ b/cosmodiff/tests/test_optim.py @@ -3,7 +3,7 @@ import numpy as np import torch from diffusers import UNet2DModel, DDPMScheduler, DDIMScheduler, DiTTransformer2DModel -from cosmodiff.utils import load_checkpoint, ArrayDataset +from cosmodiff.utils import load_checkpoint, ArrayDataset, find_latest_checkpoint from cosmodiff.optim import train, generate, compute_fid, compute_kid, build_pca_encoder from cosmodiff.augment import RandomRoll, RandomFlip @@ -42,9 +42,9 @@ def test_train_basic(): ckpt_path = os.path.join(tmp_dir, "checkpoint-epoch-0001") assert os.path.isdir(ckpt_path) - assert os.path.exists(os.path.join(ckpt_path, "noise_scheduler.pkl")) - assert os.path.exists(os.path.join(ckpt_path, "optimizer.pkl")) - assert os.path.exists(os.path.join(ckpt_path, "lr_scheduler.pkl")) + assert os.path.exists(os.path.join(ckpt_path, "config.json")) + assert os.path.exists(os.path.join(ckpt_path, "scheduler_config.json")) + assert os.path.exists(os.path.join(ckpt_path, "checkpoint_config.yaml")) assert os.path.exists(os.path.join(ckpt_path, "augmentations.pkl")) assert os.path.exists(os.path.join(ckpt_path, "metrics.json")) @@ -64,12 +64,11 @@ def test_train_basic(): assert _lr_scheduler is not None assert _augmentations is not None - # continue training + # continue training from checkpoint initial_weights = model.conv_in.weight.data.clone() metrics = train( dataset, - model, - noise_scheduler=DDPMScheduler(num_train_timesteps=10), + resume_from_checkpoint=ckpt_path, num_epochs=2, batch_size=4, checkpoint_every_n_epochs=2, @@ -79,10 +78,17 @@ def test_train_basic(): verbose=False, ) + # get new checkpoint: ensure it is epoch-0003 + ckpt_path2 = find_latest_checkpoint(tmp_dir) + assert int(ckpt_path2.split('-')[-1]) == 3 + _model2, _noise_scheduler2, _optimizer2, _lr_scheduler2, _augmentations2 = ( + load_checkpoint(ckpt_path2) + ) + # training checks: finite output, and weights changed assert all(torch.isfinite(torch.tensor(v)) for v in metrics["loss"]) assert all(torch.isfinite(torch.tensor(v)) for v in metrics["epoch_loss"]) - assert not torch.allclose(model.conv_in.weight.data, initial_weights) + assert not torch.allclose(_model2.conv_in.weight.data, initial_weights) def test_train_conditional_dit(): @@ -121,9 +127,9 @@ def test_train_conditional_dit(): ckpt_path = os.path.join(tmp_dir, "checkpoint-epoch-0001") assert os.path.isdir(ckpt_path) - assert os.path.exists(os.path.join(ckpt_path, "noise_scheduler.pkl")) - assert os.path.exists(os.path.join(ckpt_path, "optimizer.pkl")) - assert os.path.exists(os.path.join(ckpt_path, "lr_scheduler.pkl")) + assert os.path.exists(os.path.join(ckpt_path, "config.json")) + assert os.path.exists(os.path.join(ckpt_path, "scheduler_config.json")) + assert os.path.exists(os.path.join(ckpt_path, "checkpoint_config.yaml")) assert os.path.exists(os.path.join(ckpt_path, "augmentations.pkl")) assert os.path.exists(os.path.join(ckpt_path, "metrics.json")) diff --git a/cosmodiff/utils.py b/cosmodiff/utils.py index d4cb0c1..050fb88 100644 --- a/cosmodiff/utils.py +++ b/cosmodiff/utils.py @@ -1,6 +1,8 @@ import os import glob import pickle +import importlib +import inspect import numpy as np import torch import torch.nn.functional as F @@ -195,17 +197,32 @@ def read_images(path): return images, labels +def _import_class(qualified_name: str): + module_name, class_name = qualified_name.rsplit(".", 1) + return getattr(importlib.import_module(module_name), class_name) + + +def _get_lr_scheduler_kwargs(scheduler) -> dict: + sig = inspect.signature(scheduler.__class__.__init__) + kwargs = {} + for name in sig.parameters: + if name in ("self", "optimizer", "last_epoch", "verbose"): + continue + if hasattr(scheduler, name): + kwargs[name] = getattr(scheduler, name) + return kwargs + + def load_checkpoint(ckpt_path: str): """Reconstruct model, noise_scheduler, optimizer, lr_scheduler, and optionally an augmentation pipeline from a saved checkpoint directory produced by ``train()``. - All objects are returned fully reconstructed and ready to be passed - directly into ``train()``. The lr_scheduler and noise_scheduler are - stored as pickles so no knowledge of the original class is required. - The augmentation pipeline is restored from ``augmentations.pt`` if - present in the checkpoint directory and should be re-attached to the - dataset before passing it to ``train()``. + The returned objects are freshly constructed from the checkpoint config + and are ready to be passed directly into ``train()``. Pass + ``resume_from_checkpoint=ckpt_path`` to ``train()`` so it can call + ``accelerator.load_state()`` after ``accelerator.prepare()`` to restore + optimizer moments, the grad scaler, and RNG state. Args: ckpt_path (str): Path to a checkpoint directory produced by ``train()``. @@ -216,33 +233,34 @@ def load_checkpoint(ckpt_path: str): saved in the checkpoint. Example: - Resume without augmentations:: - - model, noise_scheduler, optimizer, lr_scheduler, _ = load_checkpoint( - "checkpoints/checkpoint-epoch-10" - ) - train(my_dataset, model, noise_scheduler=noise_scheduler, - optimizer=optimizer, lr_scheduler=lr_scheduler) - - Resume with augmentations:: + Resume training:: model, noise_scheduler, optimizer, lr_scheduler, augmentations = ( load_checkpoint("checkpoints/checkpoint-epoch-10") ) - dataset.augmentations = augmentations - train(my_dataset, model, noise_scheduler=noise_scheduler, - optimizer=optimizer, lr_scheduler=lr_scheduler) + if augmentations is not None: + dataset.augmentations = augmentations + train( + my_dataset, model, + noise_scheduler=noise_scheduler, + optimizer=optimizer, + lr_scheduler=lr_scheduler, + resume_from_checkpoint="checkpoints/checkpoint-epoch-10", + ) """ model = AutoModel.from_pretrained(ckpt_path) - with open(os.path.join(ckpt_path, "noise_scheduler.pkl"), "rb") as f: - noise_scheduler = pickle.load(f) + with open(os.path.join(ckpt_path, "checkpoint_config.yaml")) as f: + cfg = yaml.safe_load(f) - with open(os.path.join(ckpt_path, "optimizer.pkl"), "rb") as f: - optimizer = pickle.load(f) + noise_scheduler_cls = _import_class(cfg["noise_scheduler"]["class"]) + noise_scheduler = noise_scheduler_cls.from_pretrained(ckpt_path) - with open(os.path.join(ckpt_path, "lr_scheduler.pkl"), "rb") as f: - lr_scheduler = pickle.load(f) + optimizer_cls = _import_class(cfg["optimizer"]["class"]) + optimizer = optimizer_cls(model.parameters()) + + lr_scheduler_cls = _import_class(cfg["lr_scheduler"]["class"]) + lr_scheduler = lr_scheduler_cls(optimizer, **cfg["lr_scheduler"]["kwargs"]) augmentations_path = os.path.join(ckpt_path, "augmentations.pkl") if os.path.exists(augmentations_path): @@ -503,15 +521,10 @@ def plot_metrics(metrics: dict | str, save_dir: str = None, show: bool = False) """Plot training metrics from a dictionary or JSON file produced by ``write_metrics()``. - Produces four plots: + Produces three plots: - Batch loss over training steps - Epoch loss over epochs - Learning rate over epochs - - Epoch wall time - - matplotlib is imported lazily inside this function so it is not a strict - dependency of cosmodiff. If plotting on a remote server with no display, - use ``save_dir`` and set ``show=False``. Args: metrics (dict or str): Metrics dictionary returned by ``train()``, or @@ -542,58 +555,58 @@ def plot_metrics(metrics: dict | str, save_dir: str = None, show: bool = False) os.makedirs(save_dir, exist_ok=True) # --- batch loss ----------------------------------------------------- - fig, ax = plt.subplots(figsize=(8, 4)) + fig1, ax = plt.subplots(figsize=(8, 4)) ax.plot(metrics["loss"], alpha=0.7) ax.set_xlabel("Step") ax.set_ylabel("Loss") ax.set_title("Batch Loss") ax.set_yscale('log') - fig.tight_layout() + fig1.tight_layout() if save_dir is not None: - fig.savefig(os.path.join(save_dir, "batch_loss.png"), dpi=150, bbox_inches="tight") + fig1.savefig(os.path.join(save_dir, "batch_loss.png"), dpi=150, bbox_inches="tight") if show: plt.show() - plt.close(fig) # --- epoch loss ----------------------------------------------------- - fig, ax = plt.subplots(figsize=(8, 4)) + fig2, ax = plt.subplots(figsize=(8, 4)) ax.plot(metrics["epoch_loss"], marker="o") ax.set_xlabel("Epoch") ax.set_ylabel("Loss") ax.set_title("Epoch Loss") ax.set_yscale('log') - fig.tight_layout() + fig2.tight_layout() if save_dir is not None: - fig.savefig(os.path.join(save_dir, "epoch_loss.png"), dpi=150, bbox_inches="tight") + fig2.savefig(os.path.join(save_dir, "epoch_loss.png"), dpi=150, bbox_inches="tight") if show: plt.show() - plt.close(fig) # --- learning rate -------------------------------------------------- - fig, ax = plt.subplots(figsize=(8, 4)) + fig3, ax = plt.subplots(figsize=(8, 4)) ax.plot(metrics["epoch_lr"], marker="o") ax.set_xlabel("Epoch") ax.set_ylabel("Learning Rate") ax.set_title("Learning Rate Schedule") - fig.tight_layout() + ax.set_yscale('log') + fig3.tight_layout() if save_dir is not None: - fig.savefig(os.path.join(save_dir, "learning_rate.png"), dpi=150, bbox_inches="tight") + fig3.savefig(os.path.join(save_dir, "learning_rate.png"), dpi=150, bbox_inches="tight") if show: plt.show() - plt.close(fig) # --- epoch times ---------------------------------------------------- - fig, ax = plt.subplots(figsize=(8, 4)) - ax.plot(metrics["epoch_times"], marker="o") - ax.set_xlabel("Epoch") - ax.set_ylabel("Time (s)") - ax.set_title("Epoch Wall Time") - fig.tight_layout() - if save_dir is not None: - fig.savefig(os.path.join(save_dir, "epoch_times.png"), dpi=150, bbox_inches="tight") - if show: - plt.show() - plt.close(fig) + #fig, ax = plt.subplots(figsize=(8, 4)) + #ax.plot(metrics["epoch_times"], marker="o") + #ax.set_xlabel("Epoch") + #ax.set_ylabel("Time (s)") + #ax.set_title("Epoch Wall Time") + #fig.tight_layout() + #if save_dir is not None: + # fig.savefig(os.path.join(save_dir, "epoch_times.png"), dpi=150, bbox_inches="tight") + #if show: + # plt.show() + #plt.close(fig) + + return fig1, fig2, fig3 def find_latest_checkpoint(output_dir: str) -> str | None: diff --git a/scripts/cosmodiff_train.py b/scripts/cosmodiff_train.py index cf96dbc..8467239 100644 --- a/scripts/cosmodiff_train.py +++ b/scripts/cosmodiff_train.py @@ -71,34 +71,32 @@ def main(): shutil.copy2(args.config, config_dest) from cosmodiff.optim import train - from cosmodiff.utils import load_checkpoint, parse_config_model, parse_config_data, write_metrics, find_latest_checkpoint + from cosmodiff.utils import parse_config_model, parse_config_data, write_metrics, find_latest_checkpoint # --- check for existing checkpoint ---------------------------------- latest_ckpt = find_latest_checkpoint(output_dir) + dataset = parse_config_data(config) if latest_ckpt is not None: print(f"Resuming from checkpoint: {latest_ckpt}") - model, noise_scheduler, optimizer, lr_scheduler, augmentations = ( - load_checkpoint(latest_ckpt) + metrics = train( + dataset, + resume_from_checkpoint=latest_ckpt, + output_dir=output_dir, + **config["train"], ) - dataset = parse_config_data(config) - if augmentations is not None: - dataset.augmentations = augmentations - else: print("No checkpoint found, training from scratch.") model, optimizer, noise_scheduler, lr_scheduler = parse_config_model(config) - dataset = parse_config_data(config) - - metrics = train( - dataset, - model, - optimizer=optimizer, - noise_scheduler=noise_scheduler, - lr_scheduler=lr_scheduler, - output_dir=output_dir, - **config["train"], - ) + metrics = train( + dataset, + model, + optimizer=optimizer, + noise_scheduler=noise_scheduler, + lr_scheduler=lr_scheduler, + output_dir=output_dir, + **config["train"], + ) print(f"Training complete.") print(f"Final epoch loss: {metrics['epoch_loss'][-1]:.4f}")