JerMa88's picture
Upload folder using huggingface_hub
3ce19a2 verified
Raw
History Blame Contribute Delete
39.2 kB
import os as _os, sys as _sys
_sys.path.append(_os.path.dirname(_os.path.dirname(_os.path.abspath(__file__))))
from helpers.improved_precision_recall import compute_prec_recall
from visual.utils import (generate_and_save, generate_for_NN,
generate_images_initial,
get_sample_for_visualization)
from visual.spatial_visual import spatial_vissual
from visual.nn_interplate import nn_interp
from visual.interpolate import random_interp
from visual.generate_sample_nn import generate_sample_nn
from visual.generate_rnd_nn import generate_rnd_nn
from visual.generate_rnd import generate_rnd
from sampler import Sampler
from metrics.ppl_uniform import calc_ppl_uniform
from metrics.ppl import calc_ppl
from helpers.utils import ZippedDataset, get_cpu_stats_over_ranks
from helpers.train_helpers_fewshot import (load_imle, load_opt, save_latents,
save_latents_latest, save_model,
save_snoise, set_up_hyperparams, update_ema)
from helpers.imle_helpers import backtrack, reconstruct
from data import set_up_data
from torch.utils.data import DataLoader, TensorDataset
import csv
import json
import os
import signal
import shutil
import time
from pathlib import Path
try:
from comet_ml import Experiment, ExistingExperiment
except ImportError:
Experiment = None
ExistingExperiment = None
import imageio
import torch
import torch.nn.functional as F
import torchvision
import wandb
import torch.nn as nn
from cleanfid import fid
import cleanfid.features as _cleanfid_feat
import cleanfid.inception_torchscript as _cleanfid_incept
_inception_cache = os.path.join(os.path.expanduser("~"), ".cache", "cleanfid")
os.makedirs(_inception_cache, exist_ok=True)
_orig_feature_extractor = _cleanfid_feat.feature_extractor
def _patched_feature_extractor(name="torchscript_inception", device=torch.device("cuda"),
resize_inside=False, use_dataparallel=True):
if name == "torchscript_inception":
model = _cleanfid_incept.InceptionV3W(
_inception_cache, download=True, resize_inside=resize_inside).to(device)
model.eval()
if use_dataparallel:
model = torch.nn.DataParallel(model)
return lambda x: model(x)
return _orig_feature_extractor(name, device, resize_inside, use_dataparallel)
_cleanfid_feat.feature_extractor = _patched_feature_extractor
# ---------- metrics CSV / best_metrics helpers ----------
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 = {}
updated = False
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")
updated = True
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)
return updated
def save_images_to_dir(images_nchw: torch.Tensor, out_dir: Path):
out_dir.mkdir(parents=True, exist_ok=True)
images_nchw = images_nchw.clamp(0.0, 1.0)
for i in range(images_nchw.shape[0]):
img = (images_nchw[i].permute(
1, 2, 0).cpu().numpy() * 255).astype("uint8")
imageio.imwrite(str(out_dir / f"{i}.png"), img)
def save_grid_image(grid_chw: torch.Tensor, out_path: str):
# Save visualization grids directly under checkpoint save_dir.
grid_img = torchvision.transforms.functional.to_pil_image(
grid_chw.detach().cpu().clamp(0.0, 1.0))
grid_img.save(out_path)
def slerp(a: torch.Tensor, b: torch.Tensor, t: torch.Tensor) -> torch.Tensor:
a = F.normalize(a, dim=-1)
b = F.normalize(b, dim=-1)
dot = torch.sum(a * b, dim=-1, keepdim=True).clamp(-1.0, 1.0)
omega = torch.acos(dot)
sin_omega = torch.sin(omega)
t = t.view(-1, 1)
factor1 = torch.sin((1.0 - t) * omega) / sin_omega
factor2 = torch.sin(t * omega) / sin_omega
return factor1 * a + factor2 * b
def training_step_imle(H, n, targets, latents, snoise, imle, ema_imle, optimizer, loss_fn):
t0 = time.time()
imle.zero_grad()
cur_batch_latents = latents
px_z = imle(cur_batch_latents, snoise)
loss = loss_fn(px_z, targets.permute(0, 3, 1, 2))
loss.backward()
optimizer.step()
if ema_imle is not None:
update_ema(imle, ema_imle, H.ema_rate)
stats = get_cpu_stats_over_ranks(dict(loss_nans=0, loss=loss))
stats.update(skipped_updates=0, iter_time=time.time() - t0, grad_norm=0)
return stats
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
for data_train in DataLoader(data_train, batch_size=subset_len):
data_train = TensorDataset(data_train[0])
break
optimizer, scheduler, _, iterate, starting_epoch = load_opt(
H, imle, logprint)
print("Starting epoch: ", starting_epoch)
print("Starting iteration: ", iterate)
stats = []
H.ema_rate = torch.as_tensor(H.ema_rate)
subset_len = H.subset_len
if subset_len == -1:
subset_len = len(data_train)
sampler = Sampler(H, subset_len, preprocess_fn)
last_updated = torch.zeros(subset_len, dtype=torch.int16).cuda()
times_updated = torch.zeros(subset_len, dtype=torch.int8).cuda()
change_thresholds = torch.empty(subset_len).cuda()
change_thresholds[:] = H.change_threshold
best_fid = 100000
epoch = starting_epoch - 1
_sigterm_received = [False]
def _save_full_checkpoint(split_ind, sampler, imle, ema_imle, optimizer,
scheduler, change_thresholds, last_updated,
times_updated):
fp = os.path.join(H.save_dir, 'latest')
logprint(f'SIGTERM checkpoint save @ epoch {epoch} iter {iterate} to {fp}')
save_model(fp, imle, ema_imle, optimizer, scheduler, H)
save_latents_latest(H, split_ind, sampler.selected_latents)
save_latents_latest(H, split_ind, change_thresholds, name='threshold_latest')
save_latents_latest(H, split_ind, last_updated, name='last_updated')
save_latents_latest(H, split_ind, times_updated, name='times_updated')
def _sigterm_handler(signum, frame):
print(f'SIGTERM received at epoch {epoch}, saving checkpoint...')
_sigterm_received[0] = True
def _safe_restore_tensor(path, label):
if not path or not os.path.isfile(str(path)):
return None
try:
return torch.load(path, map_location='cpu')
except Exception as e:
logprint(f'WARNING: Failed to restore {label} from {path}: {type(e).__name__}: {e}. Ignoring this restore artifact.')
return None
prev_handler = signal.signal(signal.SIGTERM, _sigterm_handler)
for split_ind, split_x_tensor in enumerate(DataLoader(data_train, batch_size=subset_len, pin_memory=True)):
split_x_tensor = split_x_tensor[0].contiguous()
split_x = TensorDataset(split_x_tensor)
sampler.init_projection(split_x_tensor)
viz_batch_original, _ = get_sample_for_visualization(
split_x, preprocess_fn, H.num_images_visualize, H.dataset)
print('Outer batch - {}'.format(split_ind, len(split_x)))
while (epoch < H.num_epochs):
if _sigterm_received[0]:
_save_full_checkpoint(
split_ind, sampler, imle, ema_imle, optimizer,
scheduler, change_thresholds, last_updated, times_updated)
signal.signal(signal.SIGTERM, prev_handler)
os.kill(os.getpid(), signal.SIGTERM)
return
epoch += 1
last_updated[:] = last_updated + 1
restored_latents = False
if (epoch == starting_epoch):
latents = _safe_restore_tensor(H.restore_latent_path, 'latents')
if latents is not None:
sampler.selected_latents[:] = latents.to(sampler.selected_latents.device)[:]
restored_latents = True
print(f'loaded latest latents (shape={latents.shape})')
threshold = _safe_restore_tensor(H.restore_threshold_path, 'thresholds')
if threshold is not None:
change_thresholds[:] = threshold.to(change_thresholds.device)[:]
print('loaded thresholds', torch.mean(change_thresholds))
lu_path = getattr(H, 'restore_last_updated_path', None)
restored_last_updated = _safe_restore_tensor(lu_path, 'last_updated')
if restored_last_updated is not None:
last_updated[:] = restored_last_updated.to(last_updated.device)
last_updated[:] = last_updated + 1
print('loaded last_updated', torch.mean(last_updated.float()))
tu_path = getattr(H, 'restore_times_updated_path', None)
restored_times_updated = _safe_restore_tensor(tu_path, 'times_updated')
if restored_times_updated is not None:
times_updated[:] = restored_times_updated.to(times_updated.device)
print('loaded times_updated', torch.mean(times_updated.float()))
sampler.selected_dists[:] = sampler.calc_dists_existing(
split_x_tensor, imle, dists=sampler.selected_dists)
dists_in_threshold = sampler.selected_dists < change_thresholds
updated_enough = last_updated >= H.imle_staleness
updated_too_much = last_updated >= H.imle_force_resample
in_threshold = torch.logical_and(
dists_in_threshold, updated_enough)
if (H.use_adaptive):
all_conditions = torch.logical_or(
in_threshold, updated_too_much)
else:
all_conditions = updated_too_much
to_update = torch.nonzero(
all_conditions, as_tuple=False).squeeze(1)
if (epoch == starting_epoch):
if not restored_latents:
to_update = sampler.entire_ds
else:
for x in DataLoader(split_x, batch_size=H.num_images_visualize, pin_memory=True):
break
batch_slice = slice(0, x[0].size()[0])
latents_viz = sampler.selected_latents[batch_slice]
with torch.no_grad():
snoise = [s[batch_slice]
for s in sampler.selected_snoise]
generate_for_NN(sampler, x[0], latents_viz, snoise, viz_batch_original.shape, imle,
f'{H.save_dir}/NN-samples_{epoch}-{split_ind}-imle.png', logprint)
change_thresholds[to_update] = sampler.selected_dists[to_update].clone(
) * (1 - H.change_coef)
sampler.imle_sample_force(split_x_tensor, imle, to_update)
to_update = to_update.cpu()
last_updated[to_update] = 0
times_updated[to_update] = times_updated[to_update] + 1
save_latents_latest(H, split_ind, sampler.selected_latents)
save_latents_latest(
H, split_ind, change_thresholds, name='threshold_latest')
save_latents_latest(H, split_ind, last_updated, name='last_updated')
save_latents_latest(H, split_ind, times_updated, name='times_updated')
if to_update.shape[0] >= H.num_images_visualize + 8 and epoch % H.fid_freq == 0:
latents = sampler.selected_latents[to_update[:H.num_images_visualize]]
with torch.no_grad():
generate_for_NN(sampler, split_x_tensor[to_update[:H.num_images_visualize]], latents,
[s[to_update[:H.num_images_visualize]]
for s in sampler.selected_snoise],
viz_batch_original.shape, imle,
f'{H.save_dir}/NN-samples_{epoch}-imle.png', logprint)
comb_dataset = ZippedDataset(
split_x, TensorDataset(sampler.selected_latents))
data_loader = DataLoader(comb_dataset, batch_size=H.n_batch, pin_memory=True,
shuffle=False, num_workers=4, persistent_workers=False)
start_time = time.time()
for cur, indices in data_loader:
x = cur[0]
latents = cur[1][0]
_, target = preprocess_fn(x)
# if(H.use_snoise):
cur_snoise = [s[indices] for s in sampler.selected_snoise]
for i in range(len(H.res)):
cur_snoise[i].zero_()
# else:
# cur_snoise = [s[indices] for s in sampler.selected_snoise]
stat = training_step_imle(
H, target.shape[0], target, latents, cur_snoise, imle, ema_imle, optimizer, sampler.calc_loss)
stats.append(stat)
if (iterate <= H.warmup_iters):
# print("Warmup iteration: ", iterate)
scheduler.step()
iterate += 1
if iterate % H.iters_per_save == 0:
fp = os.path.join(H.save_dir, 'latest')
logprint(f'Saving model@ {iterate} to {fp}')
save_model(fp, imle, ema_imle, optimizer, scheduler, H)
save_latents_latest(H, split_ind, sampler.selected_latents)
save_latents_latest(
H, split_ind, change_thresholds, name='threshold_latest')
save_latents_latest(H, split_ind, last_updated, name='last_updated')
save_latents_latest(H, split_ind, times_updated, name='times_updated')
if iterate % H.iters_per_ckpt == 0:
save_model(os.path.join(
H.save_dir, f'iter-{iterate}'), imle, ema_imle, optimizer, scheduler, H)
save_latents(H, iterate, split_ind,
sampler.selected_latents)
save_latents(H, iterate, split_ind,
change_thresholds, name='threshold')
save_snoise(H, iterate, sampler.selected_snoise)
if _sigterm_received[0]:
_save_full_checkpoint(
split_ind, sampler, imle, ema_imle, optimizer,
scheduler, change_thresholds, last_updated,
times_updated)
signal.signal(signal.SIGTERM, prev_handler)
os.kill(os.getpid(), signal.SIGTERM)
return
print(f'Epoch {epoch} took {time.time() - start_time} seconds')
if (iterate > H.warmup_iters):
scheduler.step()
cur_dists = torch.empty([subset_len], dtype=torch.float32).cuda()
cur_dists_lpips = torch.empty(
[subset_len], dtype=torch.float32).cuda()
cur_dists_l2 = torch.empty(
[subset_len], dtype=torch.float32).cuda()
cur_dists[:], cur_dists_lpips[:], cur_dists_l2[:] = sampler.calc_dists_existing(split_x_tensor, imle,
dists=cur_dists,
dists_lpips=cur_dists_lpips,
dists_l2=cur_dists_l2,
logging=True)
# torch.save(cur_dists, f'{H.save_dir}/latent/dists-{epoch}.npy')
metrics = {
'mean_loss': torch.mean(cur_dists).item(),
'std_loss': torch.std(cur_dists).item(),
'max_loss': torch.max(cur_dists).item(),
'min_loss': torch.min(cur_dists).item(),
'mean_loss_lpips': torch.mean(cur_dists_lpips).item(),
'std_loss_lpips': torch.std(cur_dists_lpips).item(),
'max_loss_lpips': torch.max(cur_dists_lpips).item(),
'min_loss_lpips': torch.min(cur_dists_lpips).item(),
'mean_loss_l2': torch.mean(cur_dists_l2).item(),
'std_loss_l2': torch.std(cur_dists_l2).item(),
'max_loss_l2': torch.max(cur_dists_l2).item(),
'min_loss_l2': torch.min(cur_dists_l2).item(),
'total_excluded': sampler.total_excluded,
'total_excluded_percentage': sampler.total_excluded_percentage,
}
if (epoch > 0 and epoch % H.fid_freq == 0):
print("Learning rate: ", optimizer.param_groups[0]['lr'])
# --- FID: 5000 generated samples (temp dir, cleaned after) ---
generate_and_save(H, imle, sampler, 5000)
print(f'{H.data_root}/img', f'{H.save_dir}/fid/')
cur_fid = fid.compute_fid(f'{H.data_root}/img', f'{H.save_dir}/fid/',
verbose=False, num_workers=0)
# --- Save a 4x4 sample grid from the generated images ---
with torch.no_grad():
grid_z = torch.randn(16, H.latent_dim).cuda()
grid_sn = [s[:16].normal_() for s in sampler.snoise_tmp]
grid_np = sampler.sample(grid_z, imle, grid_sn)
grid_t = torch.from_numpy(grid_np).float() / 255.0
grid_t = grid_t.permute(0, 3, 1, 2)
grid_img = torchvision.utils.make_grid(
grid_t, nrow=4, padding=2)
grid_pil = torchvision.transforms.functional.to_pil_image(
grid_img)
grid_pil.save(f'{H.save_dir}/samples_epoch_{epoch}.png')
shutil.rmtree(f'{H.save_dir}/fid', ignore_errors=True)
os.makedirs(f'{H.save_dir}/fid', exist_ok=True)
is_new_best = False
if cur_fid < best_fid:
best_fid = cur_fid
is_new_best = True
fp = os.path.join(H.save_dir, 'best')
logprint(
f'Saving best model (fid={best_fid:.4f}) @ {iterate} to {fp}')
save_model(fp, imle, ema_imle, optimizer, scheduler, H)
# --- P/R: 1000 generated samples (temp dir, cleaned after) ---
os.makedirs(f'{H.save_dir}/prec_rec', exist_ok=True)
generate_and_save(H, imle, sampler, 1000, subdir='prec_rec')
precision, recall = compute_prec_recall(
f'{H.data_root}/img', f'{H.save_dir}/prec_rec/')
shutil.rmtree(f'{H.save_dir}/prec_rec', ignore_errors=True)
metrics['fid'] = cur_fid
metrics['best_fid'] = best_fid
metrics['precision'] = precision
metrics['recall'] = recall
# --- Save metrics CSV and best_metrics.json ---
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)
# --- Periodic checkpoint ---
ckpt_fp = os.path.join(H.save_dir, f'epoch_{epoch}')
logprint(f'Saving periodic ckpt @ epoch {epoch} to {ckpt_fp}')
save_model(ckpt_fp, imle, ema_imle, optimizer, scheduler, H)
# --- Slerp interpolation ---
interp_num_pairs = 5
interp_steps = 10
with torch.no_grad():
z1 = torch.randn(interp_num_pairs, H.latent_dim).cuda()
z2 = torch.randn(interp_num_pairs, H.latent_dim).cuda()
t_vals = torch.linspace(0.0, 1.0, interp_steps).cuda()
all_rows = []
for pi in range(interp_num_pairs):
z_interp = slerp(
z1[pi:pi + 1].repeat(interp_steps, 1),
z2[pi:pi + 1].repeat(interp_steps, 1),
t_vals,
)
snoise_tmp = [s[:interp_steps].normal_()
for s in sampler.snoise_tmp]
preds = sampler.sample(z_interp, imle, snoise_tmp)
preds_t = torch.from_numpy(preds).float() / 255.0
preds_t = preds_t.permute(0, 3, 1, 2)
all_rows.append(preds_t)
interp_grid = torchvision.utils.make_grid(
torch.cat(all_rows, dim=0), nrow=interp_steps)
save_grid_image(
interp_grid, f'{H.save_dir}/interp_slerp_{epoch}.png')
# --- Latent walk ---
n_perturb_seeds = 5
n_perturb_steps = 10
perturb_max_sigma = 1.0
with torch.no_grad():
perturb_rows = []
for _ in range(n_perturb_seeds):
z_seed = torch.randn(1, H.latent_dim).cuda()
direction = torch.randn(1, H.latent_dim).cuda()
direction = direction / \
direction.norm(dim=1, keepdim=True)
magnitudes = torch.linspace(
0.0, perturb_max_sigma, n_perturb_steps).cuda()
z_walk = z_seed + magnitudes.unsqueeze(1) * direction
snoise_tmp = [s[:n_perturb_steps].normal_()
for s in sampler.snoise_tmp]
preds = sampler.sample(z_walk, imle, snoise_tmp)
preds_t = torch.from_numpy(preds).float() / 255.0
preds_t = preds_t.permute(0, 3, 1, 2)
perturb_rows.append(preds_t)
perturb_all = torch.cat(perturb_rows, dim=0)
perturb_grid = torchvision.utils.make_grid(
perturb_all, nrow=n_perturb_steps, padding=2)
save_grid_image(
perturb_grid, f'{H.save_dir}/latent_walk_{epoch}.png')
# --- Visual recall (LPIPS nearest) ---
try:
import lpips as lpips_module
vr_num_real = 10
vr_num_fake = 200
vr_topk = 5
with torch.no_grad():
vr_real_imgs = split_x_tensor[:vr_num_real]
_, vr_real_processed = preprocess_fn([vr_real_imgs])
if vr_real_processed.dim() == 4 and vr_real_processed.shape[-1] in (1, 3):
vr_real_processed = vr_real_processed.permute(
0, 3, 1, 2)
z_vr = torch.randn(vr_num_fake, H.latent_dim).cuda()
snoise_tmp = [s[:min(H.imle_batch, vr_num_fake)].normal_()
for s in sampler.snoise_tmp]
vr_fake_np = sampler.sample(z_vr[:min(H.imle_batch, vr_num_fake)],
imle, snoise_tmp)
all_fake = [torch.from_numpy(
vr_fake_np).float() / 255.0]
for fi in range(H.imle_batch, vr_num_fake, H.imle_batch):
bsz = min(H.imle_batch, vr_num_fake - fi)
sn = [s[:bsz].normal_()
for s in sampler.snoise_tmp]
f_np = sampler.sample(z_vr[fi:fi + bsz], imle, sn)
all_fake.append(
torch.from_numpy(f_np).float() / 255.0)
vr_fake_all = torch.cat(all_fake, dim=0)
vr_fake_all = vr_fake_all.permute(0, 3, 1, 2)
lpips_fn = lpips_module.LPIPS(net='vgg').cuda()
vr_real_for_grid = (vr_real_processed + 1.0) / 2.0
for qi in range(vr_num_real):
dists = []
real_qi = vr_real_processed[qi:qi + 1].cuda()
for ci in range(0, vr_fake_all.shape[0], 8):
chunk = vr_fake_all[ci:ci + 8].cuda()
d = lpips_fn(real_qi.expand(chunk.shape[0], -1, -1, -1),
chunk * 2 - 1)
dists.append(d.view(-1).cpu())
dists = torch.cat(dists)
topk_idx = torch.topk(-dists,
k=min(vr_topk, len(dists))).indices
row = torch.cat([
vr_real_for_grid[qi:qi + 1].cpu(),
vr_fake_all[topk_idx].cpu()
], dim=0)
row_grid = torchvision.utils.make_grid(
row, nrow=vr_topk + 1)
save_grid_image(
row_grid, f'{H.save_dir}/visual_recall_{qi}_{epoch}.png')
del lpips_fn
torch.cuda.empty_cache()
except ImportError:
print("lpips not installed, skipping visual recall")
if (to_update.shape[0] != 0):
metrics['mean_loss_resample'] = torch.mean(cur_dists).item()
metrics['std_loss_resample'] = torch.std(cur_dists).item()
metrics['max_loss_resample'] = torch.max(cur_dists).item()
metrics['min_loss_resample'] = torch.min(cur_dists).item()
log_metrics = {k: v for k, v in metrics.items()
if not k.startswith("viz/")}
logprint(model=H.desc, type='train_loss',
epoch=epoch, step=iterate, **log_metrics)
if epoch % H.fid_freq == 0:
with torch.no_grad():
generate_images_initial(H, sampler, viz_batch_original,
sampler.selected_latents[0: H.num_images_visualize],
[s[0: H.num_images_visualize]
for s in sampler.selected_snoise],
viz_batch_original.shape, imle, ema_imle,
f'{H.save_dir}/latest.png', logprint, experiment)
if H.use_wandb:
wandb.log(metrics, step=iterate)
if experiment is not None:
experiment.log_metrics(metrics, epoch=epoch, step=iterate)
def main(H=None):
H_cur, logprint = set_up_hyperparams()
if not H:
H = H_cur
H, data_train, data_valid_or_test, preprocess_fn = set_up_data(H)
imle, ema_imle = load_imle(H, logprint)
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
if H.use_wandb:
wandb.init(
name=H.wandb_name,
project=H.wandb_project,
config=H,
mode=H.wandb_mode,
dir=H.save_dir,
)
os.makedirs(f'{H.save_dir}/fid', exist_ok=True)
if H.mode == 'eval':
os.makedirs(f'{H.save_dir}/eval', exist_ok=True)
print(H)
with torch.no_grad():
# Generating
sampler = Sampler(H, len(data_train), preprocess_fn)
n_samp = H.n_batch
temp_latent_rnds = torch.randn(
[n_samp, H.latent_dim], dtype=torch.float32).cuda()
for i in range(0, H.num_images_to_generate // n_samp):
if (i % 10 == 0):
print(i * n_samp)
temp_latent_rnds.normal_()
tmp_snoise = [s[:n_samp].normal_() for s in sampler.snoise_tmp]
torch.save(temp_latent_rnds,
f'{H.save_dir}/eval/temp_latent_rnds_{i}.pt')
torch.save(tmp_snoise, f'{H.save_dir}/eval/tmp_snoise_{i}.pt')
samp = sampler.sample(temp_latent_rnds, imle, tmp_snoise)
for j in range(n_samp):
imageio.imwrite(
f'{H.save_dir}/eval/{i * n_samp + j}.png', samp[j])
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)
n_samp = getattr(H, 'num_fid_samples', 5000)
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'
print(f'[eval_fid] dumping {n_samp} samples to {H.save_dir}/{subdir}/ (using {which} model)')
generate_and_save(H, eval_model, sampler, n_samp, subdir=subdir)
if getattr(H, 'skip_cleanfid', False):
print('[eval_fid] skip_cleanfid=True, not running cleanfid.compute_fid')
else:
try:
print(f'{H.data_root}/img', f'{H.save_dir}/{subdir}/')
cur_fid = fid.compute_fid(
f'{H.data_root}/img', f'{H.save_dir}/{subdir}/', verbose=False)
print("FID: ", cur_fid)
except Exception as e:
print(f'[eval_fid] cleanfid.compute_fid failed: {e!r} '
'(ignored; samples already dumped)')
elif H.mode == 'reconstruct':
subset_len = H.subset_len
if subset_len == -1:
subset_len = len(data_train)
ind = 0
for split_ind, split_x_tensor in enumerate(DataLoader(data_train, batch_size=H.subset_len, pin_memory=True)):
if (ind == 14):
break
split_x = TensorDataset(split_x_tensor[0])
ind += 1
for param in imle.parameters():
param.requires_grad = False
viz_batch_original, _ = get_sample_for_visualization(split_x, preprocess_fn,
H.num_images_visualize, H.dataset)
if os.path.isfile(str(H.restore_latent_path)):
latents = torch.tensor(torch.load(
H.restore_latent_path), requires_grad=True)
else:
latents = torch.randn(
[viz_batch_original.shape[0], H.latent_dim], requires_grad=True)
sampler = Sampler(H, subset_len, preprocess_fn)
reconstruct(H, sampler, imle, preprocess_fn, viz_batch_original,
latents, 'reconstruct', logprint, training_step_imle)
elif H.mode == 'backtrack':
for param in imle.parameters():
param.requires_grad = False
for split_x in DataLoader(data_train, batch_size=H.subset_len):
split_x = split_x[0]
pass
print(f'split shape is {split_x.shape}')
sampler = Sampler(H, H.subset_len, preprocess_fn)
backtrack(H, sampler, imle, preprocess_fn,
split_x, logprint, training_step_imle)
elif H.mode == 'train':
print(H)
train_loop_imle(H, data_train, data_valid_or_test,
preprocess_fn, imle, ema_imle, logprint, experiment)
elif H.mode == 'ppl':
sampler = Sampler(H, H.subset_len, preprocess_fn)
calc_ppl(H, imle, sampler)
elif H.mode == 'ppl_uniform':
sampler = Sampler(H, H.subset_len, preprocess_fn)
calc_ppl_uniform(H, imle, sampler)
elif H.mode == 'interpolate':
subset_len = H.subset_len
if subset_len == -1:
subset_len = len(data_train)
with torch.no_grad():
for split_x in DataLoader(data_train, batch_size=subset_len):
split_x = split_x[0]
viz_batch_original, _ = get_sample_for_visualization(split_x, preprocess_fn,
H.num_images_visualize, H.dataset)
sampler = Sampler(H, subset_len, preprocess_fn)
for i in range(H.num_images_to_generate):
random_interp(H, sampler, (0, 256, 256, 3), imle,
f'{H.save_dir}/interp-{i}.png', logprint)
elif H.mode == 'spatial_visual':
with torch.no_grad():
for split_x in DataLoader(data_train, batch_size=H.subset_len):
split_x = split_x[0]
viz_batch_original, _ = get_sample_for_visualization(split_x, preprocess_fn,
H.num_images_visualize, H.dataset)
sampler = Sampler(H, H.subset_len, preprocess_fn)
for i in range(H.num_images_to_generate):
print(H.num_images_to_generate, i)
spatial_vissual(H, sampler, (0, 256, 256, 3),
imle, f'{H.save_dir}/interp-{i}.png', logprint)
elif H.mode == 'generate_rnd':
with torch.no_grad():
for split_x in DataLoader(data_train, batch_size=H.subset_len):
split_x = split_x[0]
viz_batch_original, _ = get_sample_for_visualization(split_x, preprocess_fn,
H.num_images_visualize, H.dataset)
sampler = Sampler(H, H.subset_len, preprocess_fn)
generate_rnd(H, sampler, (0, 256, 256, 3), imle,
f'{H.save_dir}/rnd.png', logprint)
elif H.mode == 'generate_rnd_nn':
with torch.no_grad():
for split_x in DataLoader(data_train, batch_size=len(data_train)):
split_x = split_x[0]
viz_batch_original, _ = get_sample_for_visualization(split_x, preprocess_fn,
H.num_images_visualize, H.dataset)
sampler = Sampler(H, H.subset_len, preprocess_fn)
generate_rnd_nn(H, split_x, sampler, (0, 256, 256, 3),
imle, f'{H.save_dir}', logprint, preprocess_fn)
elif H.mode == 'nn_interp':
with torch.no_grad():
for split_x in DataLoader(data_train, batch_size=len(data_train)):
split_x = split_x[0]
viz_batch_original, _ = get_sample_for_visualization(split_x, preprocess_fn,
H.num_images_visualize, H.dataset)
sampler = Sampler(H, H.subset_len, preprocess_fn)
nn_interp(H, split_x, sampler, (0, 256, 256, 3), imle,
f'{H.save_dir}', logprint, preprocess_fn)
elif H.mode == 'generate_sample_nn':
with torch.no_grad():
for split_x in DataLoader(data_train, batch_size=len(data_train)):
split_x = split_x[0]
viz_batch_original, _ = get_sample_for_visualization(split_x, preprocess_fn,
H.num_images_visualize, H.dataset)
sampler = Sampler(H, H.subset_len, preprocess_fn)
generate_sample_nn(H, split_x, sampler, (0, 256, 256, 3),
imle, f'{H.save_dir}/rnd2.png', logprint, preprocess_fn)
elif H.mode == 'backtrack_interpolate':
subset_len = H.subset_len
if subset_len == -1:
subset_len = len(data_train)
with torch.no_grad():
for split_x in DataLoader(data_train, batch_size=subset_len):
split_x = split_x[0]
viz_batch_original, _ = get_sample_for_visualization(split_x, preprocess_fn,
H.num_images_visualize, H.dataset)
sampler = Sampler(H, subset_len, preprocess_fn)
latents = torch.tensor(torch.load(
f'{H.restore_latent_path}'), requires_grad=True, dtype=torch.float32, device='cuda')
for i in range(latents.shape[0] - 1):
lat0 = latents[i:i+1]
lat1 = latents[i+1:i+2]
sn1 = None
sn2 = None
random_interp(H, sampler, (0, 256, 256, 3), imle,
f'{H.save_dir}/back-interp-{i}.png', logprint, lat0, lat1, sn1, sn2)
elif H.mode == 'prec_rec':
os.makedirs(f'{H.save_dir}/prec_rec', exist_ok=True)
subset_len = H.subset_len
if subset_len == -1:
subset_len = len(data_train)
sampler = Sampler(H, len(data_train), preprocess_fn)
# generate_and_save(H, imle, sampler, 5000)
print("Generating images")
generate_and_save(H, imle, sampler, 1000, subdir='prec_rec')
print(f'{H.data_root}/img', f'{H.save_dir}/prec_rec/')
precision, recall = compute_prec_recall(
f'{H.data_root}/img', f'{H.save_dir}/prec_rec/')
print("Precision: ", precision)
print("Recall: ", recall)
if __name__ == "__main__":
main()