| from curses import update_lines_cols |
| from math import comb, ceil |
| import time |
|
|
| import numpy as np |
| import torch |
| import torch.nn.functional as F |
| from torch.utils.data import DataLoader, TensorDataset |
|
|
| from LPNet import LPNet |
| from dciknn_cuda import DCI, MDCI |
| from torch.optim import AdamW |
| from helpers.utils import ZippedDataset |
| from models import parse_layer_string |
|
|
| class Sampler: |
| def __init__(self, H, sz, preprocess_fn): |
| 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).cuda() |
| self.H = H |
| self.latent_lr = H.latent_lr |
| self.entire_ds = torch.arange(sz) |
| self.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.neutral_snoise = [torch.zeros([self.H.imle_db_size, 1, s, s], dtype=torch.float32) for s in self.res] |
|
|
| if(H.use_snoise == True): |
| self.snoise_tmp = [torch.randn([self.H.imle_db_size, 1, s, s], dtype=torch.float32) for s in self.res] |
| self.selected_snoise = [torch.randn([sz, 1, s, s,], dtype=torch.float32) for s in self.res] |
| self.snoise_pool = [torch.randn([self.pool_size, 1, s, s], dtype=torch.float32) for s in self.res] |
| else: |
| self.snoise_tmp = [torch.zeros([self.H.imle_db_size, 1, s, s], dtype=torch.float32) for s in self.res] |
| self.selected_snoise = [torch.zeros([sz, 1, s, s,], dtype=torch.float32) for s in self.res] |
| self.snoise_pool = [torch.zeros([self.pool_size, 1, s, s], dtype=torch.float32) for s in self.res] |
| |
| self.selected_dists = torch.empty([sz], dtype=torch.float32).cuda() |
| self.selected_dists[:] = np.inf |
| self.selected_dists_tmp = torch.empty([sz], dtype=torch.float32).cuda() |
|
|
| self.selected_dists_lpips = torch.empty([sz], dtype=torch.float32).cuda() |
| self.selected_dists_lpips[:] = np.inf |
|
|
| self.selected_dists_l2 = torch.empty([sz], dtype=torch.float32).cuda() |
| 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 = torch.randn([self.pool_size, H.latent_dim], dtype=torch.float32) |
| self.sample_pool_usage = torch.ones([sz], dtype=torch.bool) |
|
|
| self.projections = [] |
| self.lpips_net = LPNet(pnet_type=H.lpips_net, path=H.lpips_path).cuda() |
| self.l2_projection = None |
|
|
| fake = torch.zeros(1, 3, H.image_size, H.image_size).cuda() |
| out, shapes = self.lpips_net(fake) |
| sum_dims = 0 |
|
|
| if(H.search_type == 'lpips'): |
| 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(len(out) - 1)] |
| dims.append(H.proj_dim - sum(dims)) |
| for ind, feat in enumerate(out): |
| self.projections.append(F.normalize(torch.randn(feat.shape[1], dims[ind]), p=2, dim=1).cuda()) |
| sum_dims = sum(dims) |
|
|
| elif(H.search_type == 'l2'): |
| interpolated = F.interpolate(fake,scale_factor = H.l2_search_downsample) |
| interpolated = interpolated.reshape(interpolated.shape[0],-1) |
| self.l2_projection = F.normalize(torch.randn(interpolated.shape[1], H.proj_dim), p=2, dim=1).cuda() |
| sum_dims = H.proj_dim |
|
|
| else: |
|
|
| projection_dim = H.proj_dim // 2 |
| dims = [int(projection_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] * (projection_dim / sm)) for feat_ind in range(len(out) - 1)] |
| dims.append(projection_dim - sum(dims)) |
| for ind, feat in enumerate(out): |
| self.projections.append(F.normalize(torch.randn(feat.shape[1], dims[ind]), p=2, dim=1).cuda()) |
|
|
| interpolated = F.interpolate(fake,scale_factor = H.l2_search_downsample) |
| interpolated = interpolated.reshape(interpolated.shape[0],-1) |
| self.l2_projection = F.normalize(torch.randn(interpolated.shape[1], H.proj_dim // 2), p=2, dim=1).cuda() |
| sum_dims = H.proj_dim |
|
|
| self.dci_dim = sum_dims |
| print('dci_dim', self.dci_dim) |
|
|
| self.temp_samples_proj = torch.empty([self.H.imle_db_size, sum_dims], dtype=torch.float32).cuda() |
| self.dataset_proj = torch.empty([sz, sum_dims], dtype=torch.float32) |
| self.pool_samples_proj = torch.empty([self.pool_size, sum_dims], dtype=torch.float32) |
| self.snoise_pool_samples_proj = torch.empty([sz * H.snoise_factor, sum_dims], dtype=torch.float32) |
|
|
| self.knn_ignore = H.knn_ignore |
| self.ignore_radius = H.ignore_radius |
|
|
| self.total_excluded = 0 |
| self.total_excluded_percentage = 0 |
| self.dataset_size = sz |
| self.db_iter = 0 |
|
|
| def get_projected(self, inp, permute=True): |
| if permute: |
| out, _ = self.lpips_net(inp.permute(0, 3, 1, 2).cuda()) |
| else: |
| out, _ = self.lpips_net(inp.cuda()) |
| 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) |
| lpips_feat = F.normalize(lpips_feat, p=2, dim=1) |
| return lpips_feat.cuda() |
| |
| 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) |
| interpolated = interpolated.reshape(interpolated.shape[0],-1) |
| interpolated = torch.mm(interpolated, self.l2_projection) |
| interpolated = F.normalize(interpolated, p=2, dim=1) |
| return interpolated.cuda() |
| |
| def get_combined_feature(self, inp, permute=True): |
| lpips_feat = self.get_projected(inp, permute) |
| l2_feat = self.get_l2_feature(inp, permute) |
| return torch.cat([lpips_feat, l2_feat], dim=1) |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| def init_projection(self, dataset): |
| for proj_mat in self.projections: |
| proj_mat[:] = F.normalize(torch.randn(proj_mat.shape), p=2, dim=1) |
|
|
| for ind, x in enumerate(DataLoader(TensorDataset(dataset), batch_size=self.H.n_batch)): |
| batch_slice = slice(ind * self.H.n_batch, ind * self.H.n_batch + x[0].shape[0]) |
| if(self.H.search_type == 'lpips'): |
| self.dataset_proj[batch_slice] = self.get_projected(self.preprocess_fn(x)[1]) |
| elif(self.H.search_type == 'l2'): |
| self.dataset_proj[batch_slice] = self.get_l2_feature(self.preprocess_fn(x)[1]) |
| else: |
| self.dataset_proj[batch_slice] = self.get_combined_feature(self.preprocess_fn(x)[1]) |
|
|
| def sample(self, latents, gen, snoise=None): |
| with torch.no_grad(): |
| nm = latents.shape[0] |
| if snoise is None: |
| for i in range(len(self.res)): |
| if(self.H.use_snoise == True): |
| self.snoise_tmp[i].normal_() |
| snoise = [s[:nm] for s in self.snoise_tmp] |
| px_z = gen(latents, snoise).permute(0, 2, 3, 1) |
| xhat = (px_z + 1.0) * 127.5 |
| xhat = xhat.detach().cpu().numpy() |
| xhat = np.minimum(np.maximum(0.0, xhat), 255.0).astype(np.uint8) |
| return xhat |
|
|
| def sample_from_out(self, px_z): |
| with torch.no_grad(): |
| px_z = px_z.permute(0, 2, 3, 1) |
| xhat = (px_z + 1.0) * 127.5 |
| xhat = xhat.detach().cpu().numpy() |
| xhat = np.minimum(np.maximum(0.0, xhat), 255.0).astype(np.uint8) |
| return xhat |
| |
| def calc_loss_projected(self, inp, tar): |
| inp_feat = self.get_projected(inp,False) |
| tar_feat = self.get_projected(tar,False) |
| res = torch.linalg.norm(inp_feat - tar_feat, dim=1) |
| return res |
| |
| def calc_loss_l2(self, inp, tar): |
| inp_feat = self.get_l2_feature(inp,False) |
| tar_feat = self.get_l2_feature(tar,False) |
| res = torch.linalg.norm(inp_feat - tar_feat, dim=1) |
| return res |
|
|
| def calc_loss(self, inp, tar, use_mean=True, logging=False): |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
|
|
| inp_feat, inp_shape = self.lpips_net(inp) |
| tar_feat, _ = self.lpips_net(tar) |
|
|
| if use_mean: |
| l2_loss = torch.mean(self.l2_loss(inp, tar), dim=[1, 2, 3]) |
| res = 0 |
| |
| for i, g_feat in enumerate(inp_feat): |
| lpips_feature_loss = (g_feat - tar_feat[i]) ** 2 |
| res += torch.sum(lpips_feature_loss, dim=1) / (inp_shape[i] ** 2) |
|
|
| loss = self.H.lpips_coef * res.mean() + self.H.l2_coef * l2_loss.mean() |
| if logging: |
| return loss, res.mean(), l2_loss.mean() |
| else: |
| return loss |
|
|
| else: |
| res = 0 |
| for i, g_feat in enumerate(inp_feat): |
| res += torch.sum((g_feat - tar_feat[i]) ** 2, dim=1) / (inp_shape[i] ** 2) |
| l2_loss = torch.mean(self.l2_loss(inp, tar), dim=[1, 2, 3]) |
| loss = self.H.lpips_coef * res + self.H.l2_coef * l2_loss |
| if logging: |
| return loss, res.mean(), l2_loss |
| else: |
| return loss |
|
|
|
|
| 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 snoise is None: |
| snoise = self.selected_snoise |
|
|
| if to_update is not None: |
| latents = latents[to_update] |
| dists = dists[to_update] |
| dataset_tensor = dataset_tensor[to_update] |
| snoise = [s[to_update] for s in snoise] |
|
|
| 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] |
| cur_snoise = [s[batch_slice] for s in snoise] |
| with torch.no_grad(): |
| out = gen(cur_latents, cur_snoise) |
| 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 calc_dists_existing_nn(self, dataset_tensor, gen, dists=None, latents=None, to_update=None, snoise=None): |
| if dists is None: |
| dists = self.selected_dists |
| if latents is None: |
| latents = self.selected_latents |
| if snoise is None: |
| snoise = self.selected_snoise |
|
|
| if to_update is not None: |
| latents = latents[to_update] |
| dists = dists[to_update] |
| dataset_tensor = dataset_tensor[to_update] |
| snoise = [s[to_update] for s in snoise] |
|
|
| 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] |
| cur_snoise = [s[batch_slice] for s in snoise] |
| with torch.no_grad(): |
| out = gen(cur_latents, cur_snoise) |
| if(self.H.search_type == 'lpips'): |
| dist = self.calc_loss_projected(target.permute(0, 3, 1, 2), out) |
| else: |
| dist = self.calc_loss_l2(target.permute(0, 3, 1, 2), out) |
| dists[batch_slice] = torch.squeeze(dist) |
| return dists |
|
|
| def imle_sample(self, dataset, gen, factor=None): |
| if factor is None: |
| factor = self.H.imle_factor |
| imle_pool_size = int(len(dataset) * factor) |
| t1 = time.time() |
| self.selected_dists_tmp[:] = self.selected_dists[:] |
| for i in range((imle_pool_size // self.H.imle_db_size)+1): |
| self.temp_latent_rnds.normal_() |
| for j in range(len(self.res)): |
| if(self.H.use_snoise == True): |
| self.snoise_tmp[j].normal_() |
| for j in range(self.H.imle_db_size // self.H.imle_batch): |
| batch_slice = slice(j * self.H.imle_batch, (j + 1) * self.H.imle_batch) |
| cur_latents = self.temp_latent_rnds[batch_slice] |
| cur_snoise = [x[batch_slice] for x in self.snoise_tmp] |
| with torch.no_grad(): |
| self.temp_samples[batch_slice] = gen(cur_latents, cur_snoise) |
| if(self.H.search_type == 'lpips'): |
| self.temp_samples_proj[batch_slice] = self.get_projected(self.temp_samples[batch_slice], False) |
| elif(self.H.search_type == 'l2'): |
| self.temp_samples_proj[batch_slice] = self.get_l2_feature(self.temp_samples[batch_slice], False) |
| else: |
| self.temp_samples_proj[batch_slice] = self.get_combined_feature(self.temp_samples[batch_slice], False) |
|
|
| if not gen.module.dci_db: |
| device_count = torch.cuda.device_count() |
| gen.module.dci_db = MDCI(self.temp_samples_proj.shape[1], num_comp_indices=self.H.num_comp_indices, |
| num_simp_indices=self.H.num_simp_indices, devices=[i for i in range(device_count)], ts=device_count) |
|
|
| |
| |
| gen.module.dci_db.add(self.temp_samples_proj) |
|
|
| t0 = time.time() |
| for ind, y in enumerate(DataLoader(dataset, batch_size=self.H.imle_batch)): |
| |
| _, target = self.preprocess_fn(y) |
| x = self.dataset_proj[ind * self.H.imle_batch:ind * self.H.imle_batch + target.shape[0]] |
| cur_batch_data_flat = x.float() |
| nearest_indices, _ = gen.module.dci_db.query(cur_batch_data_flat, num_neighbours=1) |
| nearest_indices = nearest_indices.long()[:, 0] |
|
|
| batch_slice = slice(ind * self.H.imle_batch, ind * self.H.imle_batch + x.size()[0]) |
| actual_selected_dists = self.calc_loss(target.permute(0, 3, 1, 2), |
| self.temp_samples[nearest_indices].cuda(), use_mean=False) |
| |
|
|
| to_update = torch.nonzero(actual_selected_dists < self.selected_dists[batch_slice], as_tuple=False) |
| to_update = torch.squeeze(to_update) |
| self.selected_dists[ind * self.H.imle_batch + to_update] = actual_selected_dists[to_update].clone() |
| self.selected_latents[ind * self.H.imle_batch + to_update] = self.temp_latent_rnds[nearest_indices[to_update]].clone() |
| for k in range(len(self.res)): |
| self.selected_snoise[k][ind * self.H.imle_batch + to_update] = self.snoise_tmp[k][nearest_indices[to_update]].clone() |
|
|
| del cur_batch_data_flat |
|
|
| gen.module.dci_db.clear() |
|
|
| |
| changed = torch.sum(self.selected_dists_tmp != self.selected_dists).item() |
| print("Samples and NN are calculated, time: {}, mean: {} # changed: {}, {}%".format(time.time() - t1, |
| self.selected_dists.mean(), |
| changed, (changed / len( |
| dataset)) * 100)) |
|
|
| def resample_pool(self, gen, ds): |
| |
| self.pool_latents.normal_() |
| for i in range(len(self.res)): |
| if(self.H.use_snoise == True): |
| self.snoise_pool[i].normal_() |
|
|
| for j in range(self.pool_size // self.H.imle_batch): |
| batch_slice = slice(j * self.H.imle_batch, (j + 1) * self.H.imle_batch) |
| cur_latents = self.pool_latents[batch_slice] |
|
|
| cur_snosie = [s[batch_slice] for s in self.snoise_pool] |
| with torch.no_grad(): |
| if(self.H.search_type == 'lpips'): |
| self.pool_samples_proj[batch_slice] = self.get_projected(gen(cur_latents, cur_snosie), False) |
| elif(self.H.search_type == 'l2'): |
| self.pool_samples_proj[batch_slice] = self.get_l2_feature(gen(cur_latents, cur_snosie), False) |
| else: |
| self.pool_samples_proj[batch_slice] = self.get_combined_feature(gen(cur_latents, cur_snosie), False) |
|
|
| def imle_sample_force(self, dataset, gen, to_update=None): |
| if to_update is None: |
| to_update = self.entire_ds |
| if to_update.shape[0] == 0: |
| return |
| |
| to_update = to_update.cpu() |
|
|
| t1 = time.time() |
| if torch.any(self.sample_pool_usage[to_update]): |
| self.resample_pool(gen, dataset) |
| self.sample_pool_usage[:] = False |
| print(f'resampling took {time.time() - t1}') |
|
|
| self.selected_dists_tmp[:] = np.inf |
| self.sample_pool_usage[to_update] = True |
|
|
| |
|
|
| total_rejected = 0 |
|
|
| if(self.H.use_rsimle): |
| with torch.no_grad(): |
| for i in range(self.pool_size // self.H.imle_db_size): |
| pool_slice = slice(i * self.H.imle_db_size, (i + 1) * self.H.imle_db_size) |
| if not gen.module.dci_db: |
| device_count = torch.cuda.device_count() |
| gen.module.dci_db = MDCI(self.dci_dim, num_comp_indices=self.H.num_comp_indices, |
| num_simp_indices=self.H.num_simp_indices, |
| devices=[i for i in range(device_count)]) |
| gen.module.dci_db.add(self.pool_samples_proj[pool_slice]) |
| pool_latents = self.pool_latents[pool_slice] |
| snoise_pool = [b[pool_slice] for b in self.snoise_pool] |
|
|
| rejected_flag = torch.zeros(self.H.imle_db_size, dtype=torch.bool) |
|
|
| for ind, y in enumerate(DataLoader(TensorDataset(dataset[to_update]), batch_size=self.H.imle_batch)): |
| _, target = self.preprocess_fn(y) |
| batch_slice = slice(ind * self.H.imle_batch, ind * self.H.imle_batch + target.shape[0]) |
| indices = to_update[batch_slice] |
| x = self.dataset_proj[indices] |
| nearest_indices, dci_dists = gen.module.dci_db.query(x.float(), num_neighbours=self.H.knn_ignore) |
| nearest_indices = nearest_indices.long() |
| check = dci_dists < self.H.eps_radius |
| easy_samples_list = torch.unique(nearest_indices[check]) |
| self.pool_samples_proj[pool_slice][easy_samples_list] = torch.tensor(float('inf')) |
| rejected_flag[easy_samples_list] = 1 |
|
|
| gen.module.dci_db.clear() |
| |
| total_rejected += rejected_flag.sum().item() |
| |
| self.total_excluded = total_rejected |
| self.total_excluded_percentage = (total_rejected * 1.0 / self.pool_size) * 100 |
|
|
| with torch.no_grad(): |
| for i in range(self.pool_size // self.H.imle_db_size): |
| pool_slice = slice(i * self.H.imle_db_size, (i + 1) * self.H.imle_db_size) |
| if not gen.module.dci_db: |
| device_count = torch.cuda.device_count() |
| gen.module.dci_db = MDCI(self.dci_dim, num_comp_indices=self.H.num_comp_indices, |
| num_simp_indices=self.H.num_simp_indices, devices=[i for i in range(device_count)]) |
| gen.module.dci_db.add(self.pool_samples_proj[pool_slice]) |
| pool_latents = self.pool_latents[pool_slice] |
| snoise_pool = [b[pool_slice] for b in self.snoise_pool] |
|
|
| t0 = time.time() |
| for ind, y in enumerate(DataLoader(TensorDataset(dataset[to_update]), batch_size=self.H.imle_batch)): |
| _, target = self.preprocess_fn(y) |
| batch_slice = slice(ind * self.H.imle_batch, ind * self.H.imle_batch + target.shape[0]) |
| indices = to_update[batch_slice] |
| x = self.dataset_proj[indices] |
| nearest_indices, dci_dists = gen.module.dci_db.query(x.float(), num_neighbours=1) |
| nearest_indices = nearest_indices.long()[:, 0] |
| nearest_indices = nearest_indices.cpu() |
| dci_dists = dci_dists[:, 0] |
|
|
| need_update = dci_dists < self.selected_dists_tmp[indices] |
| need_update = need_update.cpu() |
| global_need_update = indices[need_update] |
|
|
| self.selected_dists_tmp[global_need_update] = dci_dists[need_update].clone() |
| self.selected_latents_tmp[global_need_update] = pool_latents[nearest_indices[need_update]].clone() + self.H.imle_perturb_coef * torch.randn((need_update.sum(), self.H.latent_dim)) |
| for j in range(len(self.res)): |
| self.selected_snoise[j][global_need_update] = snoise_pool[j][nearest_indices[need_update]].clone() |
|
|
| gen.module.dci_db.clear() |
|
|
| if i % 100 == 0: |
| print("NN calculated for {} out of {} - {}".format((i + 1) * self.H.imle_db_size, self.pool_size, time.time() - t0)) |
| |
| self.selected_latents[to_update] = self.selected_latents_tmp[to_update] |
|
|
| print(f'Force resampling took {time.time() - t1}') |
|
|