JerMa88's picture
Upload folder using huggingface_hub
3ce19a2 verified
Raw
History Blame Contribute Delete
19.8 kB
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()