| from curses import update_lines_cols |
| from math import comb, ceil |
| import os |
| import time |
|
|
| import numpy as np |
| import torch |
| import torch.nn.functional as F |
| from torch.utils.data import DataLoader, TensorDataset |
| from transformers import AutoImageProcessor, AutoModel |
|
|
| from LPNet import LPNet |
| from helpers.utils import ( |
| configure_inductor_for_low_memory_compile, |
| is_dist_avail_and_initialized, |
| is_main_process, |
| get_world_size, |
| get_rank, |
| safe_barrier, |
| ) |
| from models import parse_layer_string |
| from helpers.angle_sampler import Angle_Generator |
| from torch import autocast |
| import faiss |
| from tqdm import tqdm |
|
|
| class Sampler: |
| def __init__(self, H, sz, preprocess_fn): |
| |
| self.device = torch.device("cuda", torch.cuda.current_device()) |
| self.world_size = get_world_size() |
| self.rank = get_rank() |
|
|
| self.pool_size = ceil(int(H.force_factor * sz) / H.imle_db_size) * H.imle_db_size |
| self.preprocess_fn = preprocess_fn |
| self.l2_loss = torch.nn.MSELoss(reduce=False).to(self.device) |
| self.H = H |
| self.latent_lr = H.latent_lr |
| self.sz = sz |
| self.entire_ds = torch.arange(sz) |
| self.selected_latents = torch.empty([sz, H.latent_dim], dtype=torch.float32) |
| self.last_selected_latents = torch.empty([sz, H.latent_dim], dtype=torch.float32) |
| self.selected_latents_tmp = torch.empty([sz, H.latent_dim], dtype=torch.float32) |
|
|
| blocks = parse_layer_string(H.dec_blocks) |
| self.block_res = [s[0] for s in blocks] |
| self.res = sorted(set([s[0] for s in blocks if s[0] <= H.max_hierarchy])) |
|
|
| self.selected_dists = torch.empty([sz], dtype=torch.float32) |
| self.selected_dists[:] = np.inf |
| self.selected_dists_tmp = torch.empty([sz], dtype=torch.float32) |
|
|
| self.selected_dists_lpips = torch.empty([sz], dtype=torch.float32) |
| self.selected_dists_lpips[:] = np.inf |
|
|
| self.selected_dists_l2 = torch.empty([sz], dtype=torch.float32) |
| self.selected_dists_l2[:] = np.inf |
|
|
| self.temp_latent_rnds = torch.empty([self.H.imle_db_size, self.H.latent_dim], dtype=torch.float32) |
| self.temp_samples = torch.empty([self.H.imle_db_size, H.image_channels, self.H.image_size, self.H.image_size], |
| dtype=torch.float32) |
|
|
| self.pool_latents = None |
|
|
| self.projections = [] |
| self.lpips_net = LPNet(pnet_type=H.lpips_net, path=H.lpips_path).to(self.device) |
| self.lpips_net.eval() |
| self.lpips_net.requires_grad_(False) |
|
|
| if self.H.compile: |
| configure_inductor_for_low_memory_compile() |
| self.lpips_net = torch.compile(self.lpips_net) |
|
|
| self._needs_dino = (H.image_size > 32) or (H.search_type == 'combined') |
|
|
| self.dino_mean = torch.tensor([0.48145466, 0.4578275, 0.40821073], device=self.device).view(1, 3, 1, 1) |
| self.dino_std = torch.tensor([0.26862954, 0.26130258, 0.27577711], device=self.device).view(1, 3, 1, 1) |
|
|
| if self._needs_dino: |
| dino_cache_dir = getattr(H, "dino_cache_dir", None) \ |
| or os.environ.get("DINO_CACHE_DIR") \ |
| or "./dinov2_cache" |
| snapshots_dir = os.path.join( |
| dino_cache_dir, |
| "models--facebook--dinov2-base", |
| "snapshots", |
| ) |
| local_snapshot = None |
| if os.path.isdir(snapshots_dir): |
| main_path = os.path.join(snapshots_dir, "main") |
| if os.path.isdir(main_path): |
| local_snapshot = main_path |
| else: |
| candidates = sorted( |
| d for d in os.listdir(snapshots_dir) |
| if os.path.isdir(os.path.join(snapshots_dir, d)) |
| ) |
| if candidates: |
| local_snapshot = os.path.join(snapshots_dir, candidates[0]) |
| if local_snapshot is not None: |
| dino_model = AutoModel.from_pretrained(local_snapshot, local_files_only=True) |
| else: |
| dino_model = AutoModel.from_pretrained( |
| "facebook/dinov2-base", |
| cache_dir=dino_cache_dir, |
| local_files_only=True, |
| ) |
| self.dino_encoder = dino_model.eval().to(self.device) |
| |
| if self.H.compile: |
| configure_inductor_for_low_memory_compile() |
| self.dino_encoder = torch.compile(self.dino_encoder) |
| else: |
| self.dino_encoder = None |
| |
| self.nn_search_batch = H.nn_search_batch |
|
|
|
|
|
|
| self.l2_projection = None |
|
|
| fake = torch.zeros(1, 3, H.image_size, H.image_size, device=self.device) |
|
|
| safe_barrier() |
| if(H.search_type == 'lpips'): |
| interpolated = F.interpolate(fake,scale_factor = H.l2_search_downsample, antialias=True, mode='bicubic') |
| out, shapes = self.lpips_net(interpolated) |
| sum_dims = 0 |
| dims = [int(H.proj_dim * 1. / len(out)) for _ in range(len(out))] |
| if H.proj_proportion: |
| sm = sum([dim.shape[1] for dim in out]) |
| dims = [int(out[feat_ind].shape[1] * (H.proj_dim / sm)) for feat_ind in range(1,len(out))] |
| dims.insert(0,H.proj_dim - sum(dims)) |
| for ind, feat in enumerate(out): |
| self.projections.append(F.normalize(torch.randn(feat.shape[1], dims[ind], device=self.device), p=2, dim=1)) |
| sum_dims = sum(dims) |
|
|
| elif(H.search_type == 'l2'): |
| interpolated = F.interpolate(fake,scale_factor = H.l2_search_downsample, antialias=True, mode='bicubic') |
| interpolated = interpolated.reshape(interpolated.shape[0],-1) |
| self.l2_projection = F.normalize(torch.randn(interpolated.shape[1], H.proj_dim, device=self.device), p=2, dim=1) |
| sum_dims = H.proj_dim |
| |
| elif(H.search_type == 'combined'): |
| interpolated = F.interpolate(fake,scale_factor = H.l2_search_downsample, antialias=True, mode='bicubic') |
| out, shapes = self.lpips_net(interpolated) |
| sum_dims = 0 |
| dims = [int(H.proj_dim * 1. / len(out)) for _ in range(len(out))] |
| if H.proj_proportion: |
| sm = sum([dim.shape[1] for dim in out]) |
| dims = [int(out[feat_ind].shape[1] * (H.proj_dim / sm)) for feat_ind in range(1,len(out))] |
| dims.insert(0,H.proj_dim - sum(dims)) |
| for ind, feat in enumerate(out): |
| self.projections.append(F.normalize(torch.randn(feat.shape[1], dims[ind], device=self.device), p=2, dim=1)) |
| sum_dims = sum(dims) |
|
|
| interpolated = self.preprocess_dino_tensor(fake) |
| with torch.no_grad(): |
| out = self.dino_encoder(pixel_values=interpolated) |
| out = out.last_hidden_state.mean(dim=1) |
| sum_dims += out.shape[-1] |
|
|
| else: |
| exit() |
|
|
| self.dci_dim = sum_dims |
|
|
| self.dataset_proj = torch.empty([sz, sum_dims], dtype=torch.float32, device='cpu') |
| self.pool_samples_proj = None |
|
|
| self.knn_ignore = H.knn_ignore |
| self.ignore_radius = H.ignore_radius |
| self.resample_angle = H.resample_angle |
|
|
| self.total_excluded = 0 |
| self.total_excluded_percentage = 0 |
|
|
| self.dataset_size = sz |
| self.db_iter = 0 |
| self.generator_seed = torch.Generator(device=self.device) |
| self.generator_seed.manual_seed(H.seed + self.rank) |
|
|
| self.faiss_res = faiss.StandardGpuResources() |
| index_flat = faiss.IndexFlatL2(self.dci_dim) |
| dev_id = torch.cuda.current_device() |
| self.gpu_index_flat = faiss.index_cpu_to_gpu(self.faiss_res, dev_id, index_flat) |
|
|
|
|
| def preprocess_dino_tensor(self, inp): |
| |
|
|
| x = (inp + 1.0) / 2.0 |
| x = torch.clamp(x, 0.0, 1.0) |
|
|
| x = F.interpolate(x, size=(224, 224), mode='bicubic', align_corners=False) |
| return (x - self.dino_mean) / self.dino_std |
|
|
|
|
| def get_projected(self, inp, permute=True): |
| if(permute): |
| inp = inp.permute(0, 3, 1, 2) |
| |
| interpolated = F.interpolate(inp,scale_factor = self.H.l2_search_downsample, antialias=True, mode='bicubic') |
| out, _ = self.lpips_net(interpolated.to(self.device)) |
| gen_feat = [] |
| for i in range(len(out)): |
| gen_feat.append(torch.mm(out[i], self.projections[i])) |
| |
| lpips_feat = torch.cat(gen_feat, dim=1) |
| |
| return lpips_feat |
| |
| def get_l2_feature(self, inp, permute=True): |
| if(permute): |
| inp = inp.permute(0, 3, 1, 2) |
| interpolated = F.interpolate(inp,scale_factor = self.H.l2_search_downsample, antialias=True, mode='bicubic') |
| interpolated = interpolated.reshape(interpolated.shape[0],-1) |
| interpolated = torch.mm(interpolated, self.l2_projection) |
| |
| return interpolated |
| |
| def get_dino_features(self, inp, permute=True, scale_factor=10): |
| if(permute): |
| inp = inp.permute(0, 3, 1, 2) |
| interpolated = self.preprocess_dino_tensor(inp) |
| with torch.no_grad(): |
| out = self.dino_encoder(pixel_values=interpolated) |
| out = out.last_hidden_state.mean(dim=1) |
| out = F.normalize(out, p=2, dim=1) |
| out = out * scale_factor |
| return out |
| |
| def get_combined_feature(self, inp, permute=True): |
| lpisps_feat = self.get_projected(inp, permute) |
| dino_feat = self.get_dino_features(inp, permute) |
| |
| |
| combined_feat = torch.cat((lpisps_feat, dino_feat), dim=1) |
| return combined_feat |
|
|
| def init_projection(self, dataset): |
|
|
| dataloader = DataLoader( |
| dataset, |
| batch_size=self.H.imle_batch, |
| ) |
|
|
| if(is_main_process()): |
| print("Starting Initialization") |
|
|
| for ind, x in tqdm(enumerate(dataloader), total=len(dataloader), desc="Initializing"): |
| batch_slice = slice(ind * self.H.imle_batch, ind * self.H.imle_batch + x[0].shape[0]) |
| if(self.H.search_type == 'lpips'): |
| self.dataset_proj[batch_slice] = self.get_projected(self.preprocess_fn(x)[1]).cpu() |
| elif(self.H.search_type == 'l2'): |
| self.dataset_proj[batch_slice] = self.get_l2_feature(self.preprocess_fn(x)[1]).cpu() |
| elif(self.H.search_type == 'combined'): |
| self.dataset_proj[batch_slice] = self.get_combined_feature(self.preprocess_fn(x)[1]).cpu() |
| else: |
| exit() |
|
|
| self.dataset_proj = self.dataset_proj.cpu().numpy().astype(np.float32) |
|
|
| def sample(self, latents, gen, snoise=None): |
| with torch.no_grad(): |
| with autocast(device_type='cuda'): |
| latents = latents.to(self.device) |
| px_z = gen(latents, None).permute(0, 2, 3, 1) |
| xhat = (px_z + 1.0) * 127.5 |
| xhat = xhat.detach().cpu().numpy() |
| |
| xhat = np.nan_to_num(xhat, nan=0.0, posinf=255.0, neginf=0.0) |
| xhat = np.minimum(np.maximum(0.0, xhat), 255.0).astype(np.uint8) |
| return xhat |
|
|
| def get_lpips_loss(self, inp, tar, use_mean=True): |
| res = 0 |
| if(inp.shape[2] < 32): |
| inp_interpolated = F.interpolate(inp, size=(32,32), mode='bicubic') |
| tar_interpolated = F.interpolate(tar, size=(32,32), mode='bicubic') |
| else: |
| inp_interpolated = inp |
| tar_interpolated = tar |
| inp_feat, inp_shape = self.lpips_net(inp_interpolated) |
| tar_feat, _ = self.lpips_net(tar_interpolated) |
| for i, g_feat in enumerate(inp_feat): |
| lpips_feature_residual = (g_feat - tar_feat[i]) |
|
|
| if self.H.loss_type == 'huber': |
| lpips_feature_loss = self.pseudo_huber(lpips_feature_residual) |
| elif self.H.loss_type == 'mclure': |
| lpips_feature_loss = self.mclure_loss(lpips_feature_residual) |
| elif self.H.loss_type == 'welsch': |
| lpips_feature_loss = self.welsch_loss(lpips_feature_residual) |
| else: |
| lpips_feature_loss = lpips_feature_residual.pow(2) |
|
|
| res += torch.sum(lpips_feature_loss, dim=1) / (inp_shape[i] ** 2) |
| |
| return res.mean() |
|
|
| |
| def get_dino_loss(self, inp, tar, use_mean=True): |
| dino_feat = self.get_dino_features(inp, scale_factor=1, permute=False) |
| tar_feat = self.get_dino_features(tar, scale_factor=1, permute=False) |
| dino_residual = (dino_feat - tar_feat) |
|
|
| if self.H.loss_type == 'huber': |
| dino_loss = self.pseudo_huber(dino_residual) |
| elif self.H.loss_type == 'mclure': |
| dino_loss = self.mclure_loss(dino_residual) |
| elif self.H.loss_type == 'welsch': |
| dino_loss = self.welsch_loss(dino_residual) |
| else: |
| dino_loss = dino_residual.pow(2) |
|
|
| return dino_loss.mean() |
| |
| def pseudo_huber(self, residual): |
| return self.H.huber_delta**2 * (torch.sqrt(1.0 + (residual / self.H.huber_delta) ** 2) - 1.0) |
|
|
| def mclure_loss(self, residual): |
| return (residual ** 2) / (residual ** 2 + self.H.loss_scale ** 2) |
|
|
| def welsch_loss(self, residual): |
| return 1 - torch.exp(-(residual / self.H.loss_scale)**2) |
|
|
| def calc_loss(self, inp, tar, use_mean=True, logging=False): |
|
|
| pixel_residual = (inp - tar) |
| |
| if self.H.loss_type == 'huber': |
| pixel_loss = self.pseudo_huber(pixel_residual).mean() |
| elif self.H.loss_type == 'mclure': |
| pixel_loss = self.mclure_loss(pixel_residual).mean() |
| elif self.H.loss_type == 'welsch': |
| pixel_loss = self.welsch_loss(pixel_residual).mean() |
| else: |
| pixel_loss = pixel_residual.pow(2).mean() |
|
|
| lpips_loss = self.get_lpips_loss(inp, tar, use_mean=True) |
|
|
| loss = self.H.lpips_coef * lpips_loss + self.H.pixel_coef * pixel_loss |
|
|
| if inp.shape[2] > 32: |
| dino_loss = self.get_dino_loss(inp, tar, use_mean=True) |
| loss = loss + self.H.dino_coef * dino_loss |
|
|
| return loss.mean() |
|
|
| def calc_dists_existing(self, dataset_tensor, gen, dists=None, dists_lpips = None, dists_l2 = None, latents=None, to_update=None, snoise=None, logging=False): |
| if dists is None: |
| dists = self.selected_dists |
| if dists_lpips is None: |
| dists_lpips = self.selected_dists_lpips |
| if dists_l2 is None: |
| dists_l2 = self.selected_dists_l2 |
| if latents is None: |
| latents = self.selected_latents |
|
|
| if to_update is not None: |
| latents = latents[to_update] |
| dists = dists[to_update] |
| dataset_tensor = dataset_tensor[to_update] |
|
|
| for ind, x in enumerate(DataLoader(TensorDataset(dataset_tensor), batch_size=self.H.n_batch)): |
| _, target = self.preprocess_fn(x) |
| batch_slice = slice(ind * self.H.n_batch, ind * self.H.n_batch + target.shape[0]) |
| cur_latents = latents[batch_slice] |
| with torch.no_grad(): |
| with autocast(device_type='cuda'): |
| out = gen(cur_latents, None) |
| if(logging): |
| dist, dist_lpips, dist_l2 = self.calc_loss(target.permute(0, 3, 1, 2), out, use_mean=False, logging=True) |
| dists[batch_slice] = torch.squeeze(dist) |
| dists_lpips[batch_slice] = torch.squeeze(dist_lpips) |
| dists_l2[batch_slice] = torch.squeeze(dist_l2) |
| else: |
| dist = self.calc_loss(target.permute(0, 3, 1, 2), out, use_mean=False) |
| dists[batch_slice] = torch.squeeze(dist) |
| |
| if(logging): |
| return dists, dists_lpips, dists_l2 |
| else: |
| return dists |
|
|
| def resample_pool(self, gen): |
|
|
| gen.eval() |
|
|
| |
| local_pool_size = ceil(self.pool_size / self.world_size) |
|
|
|
|
| |
| local_pool_latents = torch.randn((local_pool_size, self.H.latent_dim), |
| device=self.device, |
| generator=self.generator_seed) |
| |
|
|
| local_pool_proj = torch.empty((local_pool_size, self.dci_dim), device=self.device) |
|
|
| |
| for j in range(local_pool_size // self.H.imle_batch): |
| batch_slice = slice(j * self.H.imle_batch, (j + 1) * self.H.imle_batch) |
| cur_latents = local_pool_latents[batch_slice] |
| with torch.no_grad(): |
| with autocast(device_type='cuda'): |
| outputs = gen(cur_latents, None) |
| if self.H.search_type == 'lpips': |
| proj = self.get_projected(outputs, False) |
| elif self.H.search_type == 'l2': |
| proj = self.get_l2_feature(outputs, False) |
| elif self.H.search_type == 'combined': |
| proj = self.get_combined_feature(outputs, False) |
| else: |
| proj = self.get_combined_feature(outputs, False) |
| local_pool_proj[batch_slice] = proj |
|
|
| safe_barrier() |
| gathered_latents = [torch.empty_like(local_pool_latents) for _ in range(self.world_size)] |
| gathered_proj = [torch.empty_like(local_pool_proj) for _ in range(self.world_size)] |
|
|
| torch.distributed.all_gather(gathered_latents, local_pool_latents) |
| torch.distributed.all_gather(gathered_proj, local_pool_proj) |
|
|
| gen.train() |
|
|
| safe_barrier() |
| |
| self.pool_latents = torch.cat(gathered_latents, dim=0).to('cpu') |
| self.pool_samples_proj = torch.cat(gathered_proj, dim=0).to('cpu') |
| |
|
|
| def nn_search_batched(self, queries, dataset): |
|
|
| topk = self.H.imle_db_topk |
| tie_shuffle = True |
|
|
| Nq = queries.shape[0] |
| Nd = dataset.shape[0] |
| if Nq == 0: |
| return torch.empty(0, dtype=torch.float32), torch.empty(0, dtype=torch.long) |
|
|
| topk = int(min(max(1, topk), Nd)) |
|
|
| |
| self.gpu_index_flat.reset() |
| self.gpu_index_flat.add(dataset) |
|
|
| |
| |
| if Nd >= 2: |
| D2, _ = self.gpu_index_flat.search(queries, 2) |
| margin = D2[:, 0] |
| else: |
| margin = np.zeros(Nq, dtype=np.float32) |
|
|
| if tie_shuffle: |
| perm = np.random.permutation(Nq) |
| order = perm[np.argsort(margin[perm], kind="stable")] |
| else: |
| order = np.argsort(margin, kind="stable") |
|
|
| |
| D, I = self.gpu_index_flat.search(queries, topk) |
|
|
| |
| used = np.zeros(Nd, dtype=bool) |
|
|
| out_idx = np.empty(Nq, dtype=np.int64) |
| out_dst = np.empty(Nq, dtype=np.float32) |
|
|
| I_ordered = I[order] |
| D_ordered = D[order] |
|
|
| for i in range(len(order)): |
| qi = order[i] |
| cand = I_ordered[i] |
| cd = D_ordered[i] |
|
|
| mask = ~used[cand] |
| valid = np.flatnonzero(mask) |
| if len(valid) > 0: |
| k = valid[0] |
| chosen = cand[k] |
| used[chosen] = True |
| out_idx[qi] = chosen |
| out_dst[qi] = cd[k] |
| else: |
| out_idx[qi] = cand[0] |
| out_dst[qi] = cd[0] |
|
|
| |
| self.gpu_index_flat.reset() |
|
|
| return torch.from_numpy(out_dst), torch.from_numpy(out_idx) |
|
|
|
|
| def imle_sample_force(self, gen, to_update=None): |
| """ |
| Optimized force resampling routine using FAISS for batched nearest-neighbor search. |
| In a DDP setting, each process handles a different subset of the dataset features, |
| performs NN search locally, and then the results are merged and broadcast. |
| """ |
| if is_main_process(): |
| t1 = time.time() |
| print("Starting pool resampling...") |
|
|
| |
| |
| self.resample_pool(gen) |
| safe_barrier() |
|
|
| if(is_main_process()): |
| print(f"Resampling pool took {time.time() - t1:.2f} seconds") |
| |
| torch.cuda.empty_cache() |
|
|
| self.selected_dists_tmp[:] = np.inf |
|
|
| with torch.no_grad(): |
|
|
| if(is_main_process()): |
|
|
| local_ds_feats = np.ascontiguousarray(self.dataset_proj, dtype=np.float32) |
|
|
| |
| pool_feats = np.ascontiguousarray(self.pool_samples_proj.cpu().numpy().astype(np.float32), dtype=np.float32) |
|
|
| |
| local_distances, local_indices = self.nn_search_batched(local_ds_feats, pool_feats) |
|
|
| new_latents = self.pool_latents[local_indices].clone() |
| |
| safe_barrier() |
|
|
| if is_main_process(): |
| full_updated_latents = new_latents.to(self.device) |
| perturbation = self.H.imle_perturb_coef * torch.randn( |
| (self.sz, self.H.latent_dim), |
| device=self.device, |
| generator=self.generator_seed) |
| full_updated_latents += perturbation |
| else: |
| full_updated_latents = torch.empty(self.sz, self.H.latent_dim, dtype=torch.float32, device=self.device) |
|
|
| safe_barrier() |
|
|
| torch.distributed.broadcast(full_updated_latents, src=0) |
|
|
| safe_barrier() |
|
|
| |
| self.selected_latents_tmp = full_updated_latents.cpu() |
|
|
| |
| self.last_selected_latents = self.selected_latents.clone() |
| self.selected_latents = self.selected_latents_tmp.clone() |
|
|
| if is_main_process(): |
| print(f"Force resampling took {time.time() - t1:.2f} seconds") |
|
|
| safe_barrier() |
| self.gpu_index_flat.reset() |
|
|
|
|