| 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'): |
| model = model._orig_mod |
| continue |
| if hasattr(model, 'module') and isinstance(getattr(model, 'module', None), torch.nn.Module): |
| model = model.module |
| 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() |
| |
|
|
| 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() |
| |
| |
| |
|
|
| 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() |
|
|