import os import sys import csv import json import time import math import signal try: from comet_ml import Experiment, ExistingExperiment except ImportError: Experiment = None ExistingExperiment = None import wandb import imageio import torch from torch.utils.data.distributed import DistributedSampler import torch.nn as nn from cleanfid import fid from torch.utils.data import DataLoader, TensorDataset import torch.nn.functional as F from models import IMLE import numpy as np from data import set_up_data from helpers.train_helpers import (load_imle, load_opt, save_model, set_up_hyperparams, update_ema, set_seed, restore_params, restore_log) from helpers.utils import ZippedDataset, init_distributed_mode, is_main_process, get_world_size, get_rank, safe_barrier from sampler import Sampler from visual.interpolate import random_interp from visual.utils import (generate_and_save, generate_for_NN, generate_visualization, get_sample_for_visualization) from helpers.improved_precision_recall import compute_prec_recall from torch import autocast import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP import torch.multiprocessing as mp import datetime import os import torch.distributed as dist def isValid(num): return math.isfinite(float(num)) def unwrap_model(model): seen = set() while id(model) not in seen: seen.add(id(model)) if hasattr(model, '_orig_mod'): # torch.compile wrapper model = model._orig_mod continue if hasattr(model, 'module') and isinstance(getattr(model, 'module', None), torch.nn.Module): model = model.module # DDP / DataParallel wrapper continue break return model def append_metrics_csv(save_dir, row: dict): csv_path = os.path.join(save_dir, "metrics.csv") file_exists = os.path.isfile(csv_path) with open(csv_path, "a", newline="") as f: writer = csv.DictWriter(f, fieldnames=list(row.keys())) if not file_exists: writer.writeheader() writer.writerow(row) def update_best_metrics(save_dir, row: dict): best_path = os.path.join(save_dir, "best_metrics.json") if os.path.isfile(best_path): with open(best_path, "r") as f: best = json.load(f) else: best = {} fid_val = row.get("fid") if fid_val is not None and (best.get("best_fid") is None or fid_val < best["best_fid"]): best["best_fid"] = fid_val best["best_fid_epoch"] = row.get("epoch") prec_val = row.get("precision") if prec_val is not None and (best.get("best_precision") is None or prec_val > best["best_precision"]): best["best_precision"] = prec_val best["best_precision_epoch"] = row.get("epoch") rec_val = row.get("recall") if rec_val is not None and (best.get("best_recall") is None or rec_val > best["best_recall"]): best["best_recall"] = rec_val best["best_recall_epoch"] = row.get("epoch") with open(best_path, "w") as f: json.dump(best, f, indent=2) def cleanup(): dist.destroy_process_group() def print_seed(device): cpu_seed = torch.initial_seed() cuda_seed = torch.cuda.initial_seed() print(f"Device {device} CPU seed = {cpu_seed}, GPU seed = {cuda_seed} \n") def training_step_imle(H, n, targets, latents, imle, ema_imle, optimizer, loss_fn, scaler, disable_amp=False): targets_nchw = targets.permute(0, 3, 1, 2) amp_ctx = autocast(device_type='cuda', enabled=not disable_amp) with amp_ctx: px_z = imle(latents, train=True) loss = loss_fn(px_z[-1], targets_nchw) loss_measure = loss.clone() if(H.use_multi_res): for i in range(2,len(px_z)-1): px_z_scale = px_z[i] targets_scale = F.interpolate(targets_nchw, size=(px_z_scale.shape[2], px_z_scale.shape[3]), antialias=True, mode='bicubic', align_corners=H.align_corners) loss_scale = loss_fn(px_z_scale, targets_scale) loss.add_(loss_scale) loss = loss / (H.accumulation_steps) if disable_amp: loss.backward() else: scaler.scale(loss).backward() return loss_measure.detach() def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, logprint, experiment=None): subset_len = len(data_train) if H.subset_len != -1: subset_len = H.subset_len optimizer, scheduler, scaler, best_fid, iterate, starting_epoch = load_opt(H, imle, logprint) H.ema_rate = torch.as_tensor(H.ema_rate) subset_len = H.subset_len if H.subset_len != -1 else len(data_train) sampler = Sampler(H, subset_len, preprocess_fn) safe_barrier() device = torch.device("cuda", torch.cuda.current_device()) epoch = starting_epoch sampler.init_projection(data_train) safe_barrier() viz_batch_original, _ = get_sample_for_visualization(data_train, preprocess_fn, H.num_images_visualize, H.dataset) latent_for_visualization = [] if(is_main_process()): latent_for_visualization = torch.randn(H.num_rows_visualize, H.num_images_visualize, H.latent_dim).to(device) mean_loss = float('inf') best_train_loss = float('inf') metrics = { 'mean_loss': mean_loss } _sigterm_received = [False] def _sigterm_handler(signum, frame): print(f'SIGTERM received at epoch {epoch}, saving checkpoint...', flush=True) _sigterm_received[0] = True if is_main_process(): try: fp = os.path.join(H.save_dir, 'latest') save_model(fp, imle, ema_imle, optimizer, scheduler, scaler, H) except Exception as e: print(f'WARNING: SIGTERM checkpoint save failed: {e}', flush=True) prev_handler = signal.signal(signal.SIGTERM, _sigterm_handler) while (epoch < H.num_epochs): if _sigterm_received[0]: signal.signal(signal.SIGTERM, prev_handler) os.kill(os.getpid(), signal.SIGTERM) return just_resampled = False if epoch % H.imle_force_resample == 0: torch.cuda.empty_cache() sampler.imle_sample_force(imle) torch.cuda.empty_cache() just_resampled = True if (not getattr(H, 'no_viz', False)) and (epoch % 20 == 0 and is_main_process()): latents = sampler.selected_latents[:H.num_images_visualize] raw = unwrap_model(imle) with torch.no_grad(): raw.eval() generate_for_NN(sampler, viz_batch_original, latents, viz_batch_original.shape, raw, f'{H.save_dir}/NN-samples_{epoch}-imle.png', logprint) raw.train() comb_dataset = ZippedDataset(data_train, TensorDataset(sampler.selected_latents)) train_sampler = DistributedSampler(comb_dataset, shuffle=True, num_replicas=H.world_size, rank=H.local_rank, seed=H.seed) data_loader = DataLoader(comb_dataset, batch_size=H.n_batch, sampler=train_sampler, pin_memory=True, num_workers=0, shuffle=False) train_sampler.set_epoch(epoch) if(is_main_process()): start_time = time.time() NORMAL_CLIP_NORM = 1.0 cur_clip_norm = NORMAL_CLIP_NORM epoch_loss_sum = torch.zeros(1, device=device) epoch_iter_count = 0 accum_counter = 0 imle.zero_grad(set_to_none=True) for cur, indices in data_loader: x = cur[0] latents = cur[1][0] _, target = preprocess_fn(x) target = target.to(device, non_blocking=True) latents = latents.to(device, non_blocking=True) loss = training_step_imle(H, target.shape[0], target, latents, imle, ema_imle, optimizer, sampler.calc_loss, scaler, disable_amp=False) epoch_loss_sum += loss epoch_iter_count += 1 accum_counter += 1 if accum_counter % H.accumulation_steps == 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(imle.parameters(), max_norm=cur_clip_norm) scaler.step(optimizer) scaler.update() scheduler.step() imle.zero_grad(set_to_none=True) update_ema(imle.module, ema_imle, H.ema_rate) if (not getattr(H, 'no_viz', False)) and iterate % H.iters_per_images == 0: if(is_main_process()): raw = unwrap_model(imle) raw.eval() with torch.no_grad(): generate_visualization(H, sampler, viz_batch_original, sampler.selected_latents[0: H.num_images_visualize], sampler.last_selected_latents[0: H.num_images_visualize], latent_for_visualization, viz_batch_original.shape, raw, f'{H.save_dir}/samples-{iterate}.png', logprint, experiment) raw.train() iterate += 1 if iterate % H.iters_per_ckpt == 0 and is_main_process(): fp = os.path.join(H.save_dir, f'iter-{iterate}') logprint(f'Saving model@ {iterate} to {fp}') save_model(fp, imle, ema_imle, optimizer, scheduler, scaler, H) if _sigterm_received[0]: signal.signal(signal.SIGTERM, prev_handler) os.kill(os.getpid(), signal.SIGTERM) return if accum_counter % H.accumulation_steps != 0 and epoch_iter_count > 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(imle.parameters(), max_norm=cur_clip_norm) scaler.step(optimizer) scaler.update() scheduler.step() imle.zero_grad(set_to_none=True) update_ema(imle.module, ema_imle, H.ema_rate) dist.all_reduce(epoch_loss_sum, op=dist.ReduceOp.SUM) total_batches_tensor = torch.tensor(epoch_iter_count, device=device) dist.all_reduce(total_batches_tensor, op=dist.ReduceOp.SUM) if total_batches_tensor.item() > 0: mean_loss = epoch_loss_sum.item() / total_batches_tensor.item() else: mean_loss = float('inf') metrics = { 'mean_loss': mean_loss, 'curr_lr': optimizer.param_groups[0]['lr'], } if (epoch > 0 and epoch % H.fid_freq == 0): torch.cuda.empty_cache() generate_and_save(H, unwrap_model(imle), sampler, min(5000, subset_len * H.fid_factor)) safe_barrier() torch.cuda.empty_cache() if(is_main_process()): cur_fid = fid.compute_fid(f'{H.data_root}/img', f'{H.save_dir}/fid/', verbose=False, use_dataparallel=False, num_workers=0, device=device) precision, recall = compute_prec_recall(f'{H.data_root}/img', f'{H.save_dir}/fid/') if cur_fid < best_fid: best_fid = cur_fid metrics.update({'fid': cur_fid, 'best_fid': best_fid, 'precision': precision, 'recall': recall}) csv_row = dict(epoch=epoch, fid=cur_fid, precision=precision, recall=recall) append_metrics_csv(H.save_dir, csv_row) update_best_metrics(H.save_dir, csv_row) if cur_fid == best_fid: fp = os.path.join(H.save_dir, 'best') logprint(f'Saving best model (fid={best_fid:.4f}) @ {iterate} to {fp}') logprint(model=H.desc, type='train_loss', epoch=epoch, step=iterate, **metrics) save_model(fp, imle, ema_imle, optimizer, scheduler, scaler, H) safe_barrier() if(is_main_process()): print(f'Epoch {epoch} took {time.time() - start_time} seconds') if epoch % 5 == 0: logprint(model=H.desc, type='train_loss', epoch=epoch, step=iterate, **metrics) if (not getattr(H, 'no_viz', False)) and (epoch % 5 == 0 and is_main_process()): raw = unwrap_model(imle) raw.eval() with torch.no_grad(): generate_visualization(H, sampler, viz_batch_original, sampler.selected_latents[0: H.num_images_visualize], sampler.last_selected_latents[0: H.num_images_visualize], latent_for_visualization, viz_batch_original.shape, raw, f'{H.save_dir}/latest.png', logprint, experiment) raw.train() if (epoch % 5 == 0 and experiment is not None and is_main_process()): experiment.log_metrics(metrics, epoch=epoch, step=iterate) if (epoch % 5 == 0 and is_main_process()): wandb.log(metrics, step=iterate) if is_main_process() and isValid(mean_loss) and epoch % H.epoch_per_save == 0: fp = os.path.join(H.save_dir, 'latest') logprint(f'Saving latest model@ {iterate} to {fp}') save_model(fp, imle, ema_imle, optimizer, scheduler, scaler, H) if mean_loss < best_train_loss: best_train_loss = mean_loss fp = os.path.join(H.save_dir, 'best_loss') logprint(f'New best train loss {best_train_loss:.6f} @ {iterate}, saving to {fp}') save_model(fp, imle, ema_imle, optimizer, scheduler, scaler, H) import shutil log_src = os.path.join(H.save_dir, 'latest-log.jsonl') log_dst = os.path.join(H.save_dir, 'best_loss-log.jsonl') if os.path.exists(log_src): shutil.copy2(log_src, log_dst) safe_barrier() epoch += 1 training_completed = (epoch >= H.num_epochs) if is_main_process(): if training_completed: print("Training complete. Saving final model.") fp = os.path.join(H.save_dir, 'final') logprint(f'Saving final model@ {iterate} to {fp}') save_model(fp, imle, ema_imle, optimizer, scheduler, scaler, H) fp = os.path.join(H.save_dir, 'latest') save_model(fp, imle, ema_imle, optimizer, scheduler, scaler, H) else: print(f"Training stopped early at epoch {epoch}/{H.num_epochs}. " f"Saving final checkpoint but preserving latest.") fp = os.path.join(H.save_dir, 'final') logprint(f'Saving final model@ {iterate} to {fp}') save_model(fp, imle, ema_imle, optimizer, scheduler, scaler, H) safe_barrier() return training_completed def main(): init_distributed_mode() H, logprint = set_up_hyperparams() H, data_train, data_valid_or_test, preprocess_fn = set_up_data(H) H.world_size = get_world_size() H.local_rank = get_rank() # imle, ema_imle = load_imle(H, logprint) experiment = None if(is_main_process()): print(H) if H.use_comet and H.comet_api_key: if(H.comet_experiment_key): print("Resuming experiment") experiment = ExistingExperiment( api_key=H.comet_api_key, previous_experiment=H.comet_experiment_key ) experiment.log_parameters(H) else: experiment = Experiment( api_key=H.comet_api_key, project_name=getattr(H, 'comet_project', 'adaptiveimle'), workspace=getattr(H, 'comet_workspace', None), ) experiment.set_name(H.comet_name) experiment.log_parameters(H) else: experiment = None wandb.init(project="rtm-latent-refinement", config=vars(H) if hasattr(H, '__dict__') else H) os.makedirs(f'{H.save_dir}/fid', exist_ok=True) safe_barrier() if(is_main_process()): logprint('training model', H.desc, 'on', H.dataset) imle, ema_imle = load_imle(H, logprint) if(is_main_process()): num_params = sum(p.numel() for p in imle.parameters()) print("Number of parameters in IMLE: ", num_params) logprint("Number of parameters in IMLE: ", num_params) H.num_params = num_params if(experiment is not None): experiment.log_parameter("num_params", num_params) if(H.mode == 'train'): training_completed = train_loop_imle(H, data_train, data_valid_or_test, preprocess_fn, imle, ema_imle, logprint, experiment) elif H.mode == 'eval_fid': subset_len = H.subset_len if subset_len == -1: subset_len = len(data_train) sampler = Sampler(H, len(data_train), preprocess_fn) safe_barrier() n_samp = getattr(H, 'num_fid_samples', 50000) subdir = getattr(H, 'eval_fid_subdir', 'fid') eval_model = ema_imle if ema_imle is not None else imle which = 'ema' if ema_imle is not None else 'main' raw = unwrap_model(eval_model) test_halt = getattr(H, 'test_refinement_steps', None) if test_halt is not None: mapper = raw.decoder.mapping_network old_halt = getattr(mapper, 'refinement_steps', None) mapper.refinement_steps = int(test_halt) if is_main_process(): print(f"[eval_fid] refinement_steps override: " f"{old_halt} -> {mapper.refinement_steps}") raw.eval() if is_main_process(): latent_std = float(getattr(H, 'eval_latent_std', 1.0) or 1.0) print(f"[eval_fid] dumping {n_samp} samples to " f"{H.save_dir}/{subdir}/ (using {which} model, " f"latent_std={latent_std})") generate_and_save(H, raw, sampler, n_samp, subdir=subdir) safe_barrier() # if(is_main_process()): # cur_fid = fid.compute_fid(f'{H.data_root}/img', f'{H.save_dir}/fid/', verbose=False) # print("FID: ", cur_fid) elif H.mode == 'interpolate': if(is_main_process()): print("Generating interpolations") os.makedirs(f'{H.save_dir}/interp', exist_ok=True) subset_len = H.subset_len if subset_len == -1: subset_len = len(data_train) raw = unwrap_model(imle) raw.eval() with torch.no_grad(): sampler = Sampler(H, subset_len, preprocess_fn) safe_barrier() rank = get_rank() world_size = get_world_size() for i in range(rank,H.num_images_to_generate, world_size): random_interp(H, sampler, (0, 256, 256, 3), raw, f'{H.save_dir}/interp/{i}.png', logprint) cleanup() if H.mode == 'train' and not training_completed: sys.exit(1) if __name__ == "__main__": mp.set_start_method("spawn", force=True) main()