From dde68810073b89e6824c76ef1430b9c78e61d41d Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Tue, 14 Nov 2023 19:23:32 -0800 Subject: [PATCH 01/12] keep --- sampler.py | 36 ++++++++++-------------------------- 1 file changed, 10 insertions(+), 26 deletions(-) diff --git a/sampler.py b/sampler.py index f4bb40b..ff4065e 100644 --- a/sampler.py +++ b/sampler.py @@ -23,6 +23,7 @@ def __init__(self, H, sz, preprocess_fn): 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_indices = torch.empty([sz], dtype=torch.int64) self.selected_latents_tmp = torch.empty([sz, H.latent_dim], dtype=torch.float32) blocks = parse_layer_string(H.dec_blocks) @@ -200,9 +201,9 @@ def imle_sample(self, dataset, gen, factor=None): def resample_pool(self, gen, ds): # self.init_projection(ds) - self.pool_latents.normal_() - for i in range(len(self.res)): - self.snoise_pool[i].normal_() + # self.pool_latents.normal_() + # for i in range(len(self.res)): + # 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) @@ -226,6 +227,7 @@ def imle_sample_force(self, dataset, gen, to_update=None): self.selected_dists_tmp[:] = np.inf self.sample_pool_usage[to_update] = True + with torch.no_grad(): for i in range(self.pool_size // self.H.imle_db_size): @@ -251,6 +253,7 @@ def imle_sample_force(self, dataset, gen, to_update=None): need_update = dci_dists < self.selected_dists_tmp[indices] 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)): @@ -262,29 +265,10 @@ def imle_sample_force(self, dataset, gen, to_update=None): print("NN calculated for {} out of {} - {}".format((i + 1) * self.H.imle_db_size, self.pool_size, time.time() - t0)) - if self.H.latent_epoch > 0: - for param in gen.parameters(): - param.requires_grad = False - updatable_latents = self.selected_latents_tmp[to_update].clone().requires_grad_(True) - latent_optimizer = AdamW([updatable_latents], lr=self.latent_lr) - comb_dataset = ZippedDataset(TensorDataset(dataset[to_update]), TensorDataset(updatable_latents)) - - for gd_epoch in range(self.H.latent_epoch): - losses = [] - for cur, _ in DataLoader(comb_dataset, batch_size=self.H.n_batch): - x = cur[0] - latents = cur[1][0] - _, target = self.preprocess_fn(x) - gen.zero_grad() - px_z = gen(latents) # TODO fix this - loss = self.calc_loss(px_z, target.permute(0, 3, 1, 2)) - loss.backward() - latent_optimizer.step() - updatable_latents.grad.zero_() - - losses.append(loss.detach()) - print('avg loss', gd_epoch, sum(losses) / len(losses)) - self.selected_latents[to_update] = updatable_latents.detach().clone() + self.selected_latents[to_update] = self.updatable_latents_tmp[to_update].detach() + self.pool_latents[self.selected_indices[to_update]].normal_() + for i in range(len(self.res)): + self.snoise_pool[i][self.selected_indices[to_update]].normal_() if self.H.latent_epoch > 0: for param in gen.parameters(): From ba1277d3119db0b9a0a78a0eed962e22dec7f65f Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Tue, 14 Nov 2023 19:48:56 -0800 Subject: [PATCH 02/12] keep --- sampler.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sampler.py b/sampler.py index ff4065e..2465f3c 100644 --- a/sampler.py +++ b/sampler.py @@ -253,7 +253,7 @@ def imle_sample_force(self, dataset, gen, to_update=None): need_update = dci_dists < self.selected_dists_tmp[indices] global_need_update = indices[need_update] - + self.selected_indices[global_need_update] = nearest_indices[need_update].clone() 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)): From 74156b6535bdc6401c9346d911b2d976f026ee43 Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Wed, 15 Nov 2023 11:13:51 -0800 Subject: [PATCH 03/12] keep --- sampler.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/sampler.py b/sampler.py index 2465f3c..a1918a6 100644 --- a/sampler.py +++ b/sampler.py @@ -1,5 +1,4 @@ from curses import update_lines_cols -from math import comb import time import numpy as np @@ -265,7 +264,7 @@ def imle_sample_force(self, dataset, gen, to_update=None): 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.updatable_latents_tmp[to_update].detach() + self.selected_latents[to_update] = self.selected_latents_tmp[to_update].clone() self.pool_latents[self.selected_indices[to_update]].normal_() for i in range(len(self.res)): self.snoise_pool[i][self.selected_indices[to_update]].normal_() From 7e5bf22cb4a4e01e309e695702a9745bc4f030ae Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Wed, 15 Nov 2023 22:42:58 -0800 Subject: [PATCH 04/12] new scheduler --- hps.py | 1 + train.py | 25 ++++++++++++++++++++++++- 2 files changed, 25 insertions(+), 1 deletion(-) diff --git a/hps.py b/hps.py index da7a595..600a991 100644 --- a/hps.py +++ b/hps.py @@ -118,4 +118,5 @@ def add_imle_arguments(parser): parser.add_argument('--ppl_save_name', type=str, default='ppl') parser.add_argument("--fid_factor", type=int, default=5, help="number of the samples for calculating FID") parser.add_argument("--fid_freq", type=int, default=5, help="frequency of calculating fid") + parser.add_argument("--num_steps", type=int, default=200000, help="frequency of calculating fid") return parser diff --git a/train.py b/train.py index f8c7a22..6b9a93c 100644 --- a/train.py +++ b/train.py @@ -25,6 +25,7 @@ from visual.utils import (generate_and_save, generate_for_NN, generate_images_initial, get_sample_for_visualization) +from torch.optim.lr_scheduler import LambdaLR def training_step_imle(H, n, targets, latents, snoise, imle, ema_imle, optimizer, loss_fn): @@ -41,6 +42,26 @@ def training_step_imle(H, n, targets, latents, snoise, imle, ema_imle, optimizer stats.update(skipped_updates=0, iter_time=time.time() - t0, grad_norm=0) return stats +class DecayLR: + def __init__(self, tmax=100000): + self.tmax = int(tmax) + assert self.tmax > 0 + self.lr_step = (0 - 1) / self.tmax + + def step(self, step): + lr = 1 + self.lr_step * step + lr = max(1e-6, min(1.0, lr)) + return lr + +def get_lrschedule(args, optimizer): + if args.lr_schedule: + scheduler = DecayLR(tmax=args.num_steps) + lr_scheduler = LambdaLR(optimizer, lambda x: scheduler.step(x)) + else: + lr_scheduler = LambdaLR(optimizer, lambda x: 1.0) + return lr_scheduler + + def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, logprint): subset_len = len(data_train) @@ -51,6 +72,7 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo break optimizer, scheduler, _, iterate, _ = load_opt(H, imle, logprint) + lr_scheduler = get_lrschedule(H, optimizer) stats = [] H.ema_rate = torch.as_tensor(H.ema_rate) @@ -141,7 +163,8 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo 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) - scheduler.step() + # scheduler.step() + lr_scheduler.step() if iterate % H.iters_per_images == 0: with torch.no_grad(): From 7d677bd9a18d8285dfbec4b145ac2e9f99858345 Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Wed, 22 Nov 2023 16:55:32 -0800 Subject: [PATCH 05/12] keep --- sampler.py | 2 +- train.py | 28 ++++++++++++++++++---------- 2 files changed, 19 insertions(+), 11 deletions(-) diff --git a/sampler.py b/sampler.py index a1918a6..0dc6c1b 100644 --- a/sampler.py +++ b/sampler.py @@ -246,7 +246,7 @@ def imle_sample_force(self, dataset, gen, to_update=None): 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.long()[:, 0].cpu() dci_dists = dci_dists[:, 0] need_update = dci_dists < self.selected_dists_tmp[indices] diff --git a/train.py b/train.py index 6b9a93c..0554abb 100644 --- a/train.py +++ b/train.py @@ -43,23 +43,30 @@ def training_step_imle(H, n, targets, latents, snoise, imle, ema_imle, optimizer return stats class DecayLR: - def __init__(self, tmax=100000): + def __init__(self, tmax=100000, staleness=10): self.tmax = int(tmax) + self.staleness = staleness assert self.tmax > 0 self.lr_step = (0 - 1) / self.tmax def step(self, step): + per = step % self.staleness lr = 1 + self.lr_step * step + lr = lr + ((0 - 1) / self.staleness) * per lr = max(1e-6, min(1.0, lr)) return lr def get_lrschedule(args, optimizer): - if args.lr_schedule: - scheduler = DecayLR(tmax=args.num_steps) - lr_scheduler = LambdaLR(optimizer, lambda x: scheduler.step(x)) - else: - lr_scheduler = LambdaLR(optimizer, lambda x: 1.0) - return lr_scheduler + # if args.lr_schedule: + # scheduler = DecayLR(tmax=args.num_steps) + # lr_scheduler = LambdaLR(optimizer, lambda x: scheduler.step(x)) + # else: + # lr_scheduler = LambdaLR(optimizer, lambda x: 1.0) + # return lr_scheduler + scheduler = DecayLR(tmax=args.num_steps, staleness=args.imle_staleness) + return LambdaLR(optimizer, lambda x: scheduler.step(x)) + # return LambdaLR(optimizer, lambda x: 1.0) + @@ -163,8 +170,6 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo 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) - # scheduler.step() - lr_scheduler.step() if iterate % H.iters_per_images == 0: with torch.no_grad(): @@ -188,6 +193,8 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo save_latents(H, iterate, split_ind, change_thresholds, name='threshold') save_snoise(H, iterate, sampler.selected_snoise) + lr_scheduler.step() + cur_dists = torch.empty([subset_len], dtype=torch.float32).cuda() cur_dists[:] = sampler.calc_dists_existing(split_x_tensor, imle, dists=cur_dists) torch.save(cur_dists, f'{H.save_dir}/latent/dists-{epoch}.npy') @@ -197,6 +204,7 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo 'std_loss': torch.std(cur_dists).item(), 'max_loss': torch.max(cur_dists).item(), 'min_loss': torch.min(cur_dists).item(), + 'epoch': epoch, } if epoch % H.fid_freq == 0: @@ -214,7 +222,7 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo metrics['best_fid'] = best_fid - logprint(model=H.desc, type='train_loss', epoch=epoch, step=iterate, **metrics) + logprint(model=H.desc, type='train_loss', step=iterate, **metrics) if H.use_wandb: wandb.log(metrics, step=iterate) From 61a60f3d218976144168a340868743d1db1e0d40 Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Wed, 22 Nov 2023 17:28:22 -0800 Subject: [PATCH 06/12] keep --- hps.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/hps.py b/hps.py index 600a991..759f945 100644 --- a/hps.py +++ b/hps.py @@ -118,5 +118,5 @@ def add_imle_arguments(parser): parser.add_argument('--ppl_save_name', type=str, default='ppl') parser.add_argument("--fid_factor", type=int, default=5, help="number of the samples for calculating FID") parser.add_argument("--fid_freq", type=int, default=5, help="frequency of calculating fid") - parser.add_argument("--num_steps", type=int, default=200000, help="frequency of calculating fid") + parser.add_argument("--num_steps", type=int, default=4000, help="frequency of calculating fid") return parser From dc0a4d67b9ac578cbbc6e7a65c90b5bd8c8a0bd0 Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Thu, 23 Nov 2023 15:39:25 -0800 Subject: [PATCH 07/12] partial resample --- hps.py | 1 + sampler.py | 110 ++++++++++++++++++----------------------------------- train.py | 12 +++--- 3 files changed, 44 insertions(+), 79 deletions(-) diff --git a/hps.py b/hps.py index 759f945..bc7999f 100644 --- a/hps.py +++ b/hps.py @@ -119,4 +119,5 @@ def add_imle_arguments(parser): parser.add_argument("--fid_factor", type=int, default=5, help="number of the samples for calculating FID") parser.add_argument("--fid_freq", type=int, default=5, help="frequency of calculating fid") parser.add_argument("--num_steps", type=int, default=4000, help="frequency of calculating fid") + parser.add_argument("--pool_staleness", type=int, default=3, help="frequency of calculating fid") return parser diff --git a/sampler.py b/sampler.py index 0dc6c1b..e167c2a 100644 --- a/sampler.py +++ b/sampler.py @@ -62,6 +62,7 @@ def __init__(self, H, sz, preprocess_fn): 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.pool_last_updated = torch.zeros([self.pool_size], dtype=torch.int64) def get_projected(self, inp, permute=True): if permute: @@ -139,90 +140,48 @@ def calc_dists_existing(self, dataset_tensor, gen, dists=None, latents=None, to_ 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): - self.temp_latent_rnds.normal_() - for j in range(len(self.res)): - 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) - self.temp_samples_proj[batch_slice] = self.get_projected(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 = DCI(self.temp_samples_proj.shape[1], num_comp_indices=self.H.num_comp_indices, - # num_simp_indices=self.H.num_simp_indices) - 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)): - # t2 = time.time() - _, 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) - # actual_selected_dists = torch.squeeze(actual_selected_dists) - - 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() - - # adding perturbation - 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.init_projection(ds) - # self.pool_latents.normal_() - # for i in range(len(self.res)): - # self.snoise_pool[i].normal_() + def resample_pool(self, gen, to_update, rnd=True): + self.pool_last_updated[to_update] = 0 + if rnd: + self.pool_latents[to_update].normal_() + for i in range(len(self.res)): + self.snoise_pool[i][to_update].normal_() + + for j in range(to_update.shape[0] // self.H.imle_batch + 1): + sl = slice(j * self.H.imle_batch, (j + 1) * self.H.imle_batch) + batch_slice = to_update[sl] + if batch_slice.shape[0] == 0: + continue - 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(): - self.pool_samples_proj[batch_slice] = self.get_projected(gen(cur_latents, cur_snosie), False) + self.pool_samples_proj[batch_slice] = self.get_projected(gen(cur_latents, cur_snosie), False).cpu() + # self.get_projected(gen(cur_latents, cur_snosie), False) + def imle_sample_force(self, dataset, gen, to_update=None): + self.pool_last_updated += 1 if to_update is None: to_update = self.entire_ds + # resample all pool + self.resample_pool(gen, torch.arange(self.pool_size)) if to_update.shape[0] == 0: return + + pool_acceptable_stal = self.H.pool_staleness + # resample those that are too old + pool_old_indices = torch.where(self.pool_last_updated > pool_acceptable_stal)[0] + self.resample_pool(gen, pool_old_indices, rnd=False) t1 = time.time() print(torch.any(self.sample_pool_usage[to_update]), torch.any(self.sample_pool_usage)) - 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}') + # 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}') + # to_update_indices = self.selected_indices[to_update] + # self.resample_pool(gen, to_update_indices) self.selected_dists_tmp[:] = np.inf self.sample_pool_usage[to_update] = True @@ -265,9 +224,12 @@ def imle_sample_force(self, dataset, gen, to_update=None): self.selected_latents[to_update] = self.selected_latents_tmp[to_update].clone() - self.pool_latents[self.selected_indices[to_update]].normal_() - for i in range(len(self.res)): - self.snoise_pool[i][self.selected_indices[to_update]].normal_() + # self.pool_latents[self.selected_indices[to_update]].normal_() + # for i in range(len(self.res)): + # self.snoise_pool[i][self.selected_indices[to_update]].normal_() + + to_update_indices = self.selected_indices[to_update].cuda() + self.resample_pool(gen, to_update_indices) if self.H.latent_epoch > 0: for param in gen.parameters(): diff --git a/train.py b/train.py index 0554abb..494be74 100644 --- a/train.py +++ b/train.py @@ -118,6 +118,7 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo in_threshold = torch.logical_and(dists_in_threshold, updated_enough) all_conditions = torch.logical_or(in_threshold, updated_too_much) to_update = torch.nonzero(all_conditions, as_tuple=False).squeeze(1) + change_thresholds[to_update] = sampler.selected_dists[to_update].clone() * (1 - H.change_coef) if epoch == 0: if os.path.isfile(str(H.restore_latent_path)): @@ -138,20 +139,21 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo change_thresholds[:] = threshold[:] print('loaded thresholds', torch.mean(change_thresholds)) else: - to_update = sampler.entire_ds + to_update = None - change_thresholds[to_update] = sampler.selected_dists[to_update].clone() * (1 - H.change_coef) + print(to_update) sampler.imle_sample_force(split_x_tensor, imle, to_update) - last_updated[to_update] = 0 - times_updated[to_update] = times_updated[to_update] + 1 + if to_update is not None: + 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') - if to_update.shape[0] >= H.num_images_visualize: + if to_update is not None and to_update.shape[0] >= H.num_images_visualize: 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, From bb9ff8ecb5d52e429973edc7af0c4d096d0ce9be Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Fri, 24 Nov 2023 16:41:39 -0800 Subject: [PATCH 08/12] partial resample --- sampler.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sampler.py b/sampler.py index e167c2a..8d2791b 100644 --- a/sampler.py +++ b/sampler.py @@ -205,7 +205,7 @@ def imle_sample_force(self, dataset, gen, to_update=None): 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].cpu() + nearest_indices = nearest_indices.long()[:, 0].cpu() + pool_slice.start dci_dists = dci_dists[:, 0] need_update = dci_dists < self.selected_dists_tmp[indices] From a281b1eb380f178db8a7ab3c9125cafc7292a0d7 Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Fri, 24 Nov 2023 16:44:50 -0800 Subject: [PATCH 09/12] partial resample --- sampler.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/sampler.py b/sampler.py index 8d2791b..f24694e 100644 --- a/sampler.py +++ b/sampler.py @@ -205,13 +205,15 @@ def imle_sample_force(self, dataset, gen, to_update=None): 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].cpu() + pool_slice.start + nearest_indices = nearest_indices.long()[:, 0].cpu() dci_dists = dci_dists[:, 0] need_update = dci_dists < self.selected_dists_tmp[indices] global_need_update = indices[need_update] - self.selected_indices[global_need_update] = nearest_indices[need_update].clone() + real_nearest_indices = nearest_indices[need_update] + pool_slice.start + self.selected_indices[global_need_update] = nearest_indices[real_nearest_indices].clone() + 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)): From 227ffe106cbb621d6ddfb8a955ffd39de47b2183 Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Sat, 25 Nov 2023 15:35:10 -0800 Subject: [PATCH 10/12] eps decay --- sampler.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sampler.py b/sampler.py index f24694e..9185b64 100644 --- a/sampler.py +++ b/sampler.py @@ -212,7 +212,7 @@ def imle_sample_force(self, dataset, gen, to_update=None): global_need_update = indices[need_update] real_nearest_indices = nearest_indices[need_update] + pool_slice.start - self.selected_indices[global_need_update] = nearest_indices[real_nearest_indices].clone() + self.selected_indices[global_need_update] = real_nearest_indices.clone() 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)) From e2cdca7257bb636717c01ce9230fbc6024d4cda8 Mon Sep 17 00:00:00 2001 From: mehranagh20 Date: Sun, 26 Nov 2023 09:47:57 -0800 Subject: [PATCH 11/12] eps decay --- train.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/train.py b/train.py index 494be74..8abba77 100644 --- a/train.py +++ b/train.py @@ -79,7 +79,8 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo break optimizer, scheduler, _, iterate, _ = load_opt(H, imle, logprint) - lr_scheduler = get_lrschedule(H, optimizer) + # lr_scheduler = get_lrschedule(H, optimizer) + lr_scheduler = scheduler stats = [] H.ema_rate = torch.as_tensor(H.ema_rate) @@ -172,6 +173,7 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo 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) + scheduler.step() if iterate % H.iters_per_images == 0: with torch.no_grad(): @@ -195,7 +197,7 @@ def train_loop_imle(H, data_train, data_valid, preprocess_fn, imle, ema_imle, lo save_latents(H, iterate, split_ind, change_thresholds, name='threshold') save_snoise(H, iterate, sampler.selected_snoise) - lr_scheduler.step() + # lr_scheduler.step() cur_dists = torch.empty([subset_len], dtype=torch.float32).cuda() cur_dists[:] = sampler.calc_dists_existing(split_x_tensor, imle, dists=cur_dists) From b8f55f04e1b28308881d023669b6028d6cb30146 Mon Sep 17 00:00:00 2001 From: Mehran Aghabozorgi Date: Thu, 30 Nov 2023 21:17:16 -0500 Subject: [PATCH 12/12] remove pixel norm --- mapping_network.py | 2 +- test.sh | 30 ++++++++++++++++++++++++++++++ 2 files changed, 31 insertions(+), 1 deletion(-) create mode 100755 test.sh diff --git a/mapping_network.py b/mapping_network.py index 21d6480..851ebe4 100644 --- a/mapping_network.py +++ b/mapping_network.py @@ -62,7 +62,7 @@ class MappingNetowrk(nn.Module): def __init__(self, code_dim=512, n_mlp=8): super().__init__() - layers = [PixelNorm()] + layers = [] for i in range(n_mlp): layers.append(EqualLinear(code_dim, code_dim)) layers.append(nn.LeakyReLU(0.2)) diff --git a/test.sh b/test.sh new file mode 100755 index 0000000..1c84e40 --- /dev/null +++ b/test.sh @@ -0,0 +1,30 @@ +#!/bin/bash + + +name=100-shot-panda +change=0.0 +factor=20 +force=10 +lr=0.00001 +stal=10 +wand_name="$name-chg-${change}-fac-${factor}-frc-${force}-lr-${lr}-stl-${stal}" + +save_dir=/home/mehranag/scratch/saved_models/vdimle-reproduce/$wand_name +data_root=/home/mehranag/projects/rrg-keli/data/few-shot-images/${name} +restore_latent_path=/home/mehranag/scratch/saved_models/archived/4-ada3-2048-50p-full/test/latent/0-latest.npy +restore_path=/home/mehranag/scratch/saved_models/panda-naive/test/iter-450000- + +#cd dciknn_cuda +#python setup.py install +#cd .. +cp /home/mehranag/inception-2015-12-05.pt /tmp + +ssh -D 9050 -q -C -N narval1 & +python train.py --hps fewshot --save_dir $save_dir --data_root $data_root --lpips_coef 1 --l2_coef 0.1 \ + --change_threshold 1 --change_coef $change --force_factor $factor --imle_db_size 5000 --imle_staleness $stal \ + --imle_force_resample $force --latent_epoch 0 --latent_lr 0.0 --imle_factor 0 --lr $lr --n_batch 4 \ + --proj_dim 800 --imle_batch 20 --iters_per_save 1000 --iters_per_images 500 --image_size 256 \ + --proj_proportion 1 --latent_dim 1024 --iters_per_ckpt 5000 \ + --dec_blocks '1x4,4m1,4x4,8m4,8x4,16m8,16x3,32m16,32x2,64m32,64x2,128m64,128x2,256m128' \ + --max_hierarchy 256 --image_size 256 --use_wandb 1 --wandb_project $name-keep --wandb_name $wand_name --wandb_mode offline \ + --fid_freq 10 --fid_factor 5