File size: 20,167 Bytes
3ce19a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
from distutils.util import strtobool

HPARAMS_REGISTRY = {}


class Hyperparams(dict):
    def __getattr__(self, attr):
        try:
            return self[attr]
        except KeyError:
            return None

    def __setattr__(self, attr, value):
        self[attr] = value

cifar10 = Hyperparams()
cifar10.width = 768
cifar10.lr = 0.0002
cifar10.wd = 0.01
cifar10.dec_blocks = "1x1,4m1,4x2,8m4,8x5,16m8,16x5,32m16,32x5"
cifar10.dataset = 'cifar10'
cifar10.n_batch = 196
cifar10.imle_batch = 32 
cifar10.ema_rate = 0.9999
cifar10.l2_search_downsample = 1.0
cifar10.multi_res_scales = '16,20,24,28'
cifar10.convnext_expansion = 4
HPARAMS_REGISTRY['cifar10'] = cifar10

imagenet32 = Hyperparams()
imagenet32.width = 512
imagenet32.lr = 0.0002
imagenet32.wd = 0.01
imagenet32.dec_blocks = "1x1,4m1,4x8,8m4,8x16,16m8,16x16,32m16,32x21"
imagenet32.dataset = 'imagenet32'
imagenet32.n_batch = 32
imagenet32.imle_batch = 32
imagenet32.ema_rate = 0.9999
imagenet32.l2_search_downsample = 1.0
imagenet32.multi_res_scales = '8,12,16,24,28'
imagenet32.convnext_expansion = 6
HPARAMS_REGISTRY['imagenet32'] = imagenet32


stl10 = Hyperparams()
stl10.width = 384
stl10.lr = 0.0002
stl10.wd = 0.01
stl10.dec_blocks = "1x2,4m1,4x3,8m4,8x7,16m8,16x15,32m16,32x31,64m32,64x12"
# stl10.dec_blocks = "1x1,4m1,4x8,8m4,8x10,16m8,16x10,32m16,32x10,64m32,64x10"
stl10.dataset = 'stl10'
stl10.n_batch = 8
stl10.imle_batch = 32 
stl10.ema_rate = 0.9999
stl10.l2_search_downsample = 0.5
stl10.multi_res_scales = '16,32,48'
stl10.convnext_expansion = 4
HPARAMS_REGISTRY['stl10'] = stl10

lsun = Hyperparams()
lsun.width = 384
lsun.lr = 0.0002
lsun.wd = 0.01
lsun.dec_blocks = '1x4,4m1,4x4,8m4,8x4,16m8,16x3,32m16,32x2,64m32,64x2,128m64,128x2,256m128'
# lsun.dec_blocks = '1x2,4m1,4x3,8m4,8x4,16m8,16x9,32m16,32x21,64m32,64x13,128m64,128x7,256m128'
lsun.dataset = 'lsun'
lsun.n_batch = 4
lsun.ema_rate = 0.9999
lsun.l2_search_downsample = 0.125
lsun.multi_res_scales = '8,12,16,24,32,48,64,96,128,150,200,230'
HPARAMS_REGISTRY['lsun'] = lsun

fewshot = Hyperparams()
fewshot.width = 384
fewshot.lr = 0.0002
fewshot.wd = 0.01
fewshot.dec_blocks = '1x4,4m1,4x4,8m4,8x4,16m8,16x3,32m16,32x2,64m32,64x2,128m64,128x2,256m128'
# fewshot.dec_blocks = '1x2,4m1,4x3,8m4,8x4,16m8,16x9,32m16,32x21,64m32,64x13,128m64,128x7,256m128'
fewshot.dataset = 'fewshot'
fewshot.n_batch = 4
fewshot.ema_rate = 0.9999
fewshot.l2_search_downsample = 0.125
fewshot.multi_res_scales = '8,12,16,24,32,48,64,96,128,150,200,230'
HPARAMS_REGISTRY['fewshot'] = fewshot


fewshot64 = Hyperparams()
fewshot64.width = 384
fewshot64.lr = 0.0002
fewshot64.wd = 0.01
fewshot64.image_size = 64
fewshot64.dec_blocks = '1x2,4m1,4x3,8m4,8x7,16m8,16x8,32m16,32x8,64m32,64x8'
# fewshot.dec_blocks = '1x2,4m1,4x3,8m4,8x4,16m8,16x9,32m16,32x21,64m32,64x13,128m64,128x7,256m128'
fewshot64.dataset = 'fewshot'
fewshot64.n_batch = 8
fewshot64.ema_rate = 0.9999
fewshot64.l2_search_downsample = 1.0
fewshot64.multi_res_scales = '8,12,16,24,32,48'
HPARAMS_REGISTRY['fewshot64'] = fewshot64

# CelebA-HQ-256 entry; the dataset / dec_blocks / latent_dim / RTM knobs are
# overridden on the command line by `scripts/eval_celebahq256.sh`, so this
# block only needs to exist as a registry key.
celebahq256 = Hyperparams()
celebahq256.width = 384
celebahq256.lr = 0.0002
celebahq256.wd = 0.01
celebahq256.image_size = 256
celebahq256.dec_blocks = '1x1,4m1,4x2,8m4,8x4,16m8,16x5,32m16,32x5,64m32,64x5,128m64,128x4,256m128,256x1'
celebahq256.dataset = 'celebahq256'
celebahq256.n_batch = 48
celebahq256.imle_batch = 256
celebahq256.ema_rate = 0.9999
celebahq256.l2_search_downsample = 0.125
celebahq256.multi_res_scales = '8,12,16,24,32,48,64,96,128,150,200,230'
HPARAMS_REGISTRY['celebahq256'] = celebahq256

def parse_args_and_update_hparams(H, parser, s=None):
    args = parser.parse_args(s)
    valid_args = set(args.__dict__.keys())
    hparam_sets = [x for x in args.hparam_sets.split(',') if x]
    for hp_set in hparam_sets:
        hps = HPARAMS_REGISTRY[hp_set]
        for k in hps:
            if k not in valid_args:
                raise ValueError(f"{k} not in default args")
        parser.set_defaults(**hps)
    H.update(parser.parse_args(s).__dict__)

    try:
        value = H['multi_res_scales']
        list_value = value.split(',')
        list_value_int = [int(x) for x in list_value]
        H['multi_res_scales'] = list_value_int
    except:
        pass

def add_imle_arguments(parser):
    parser.add_argument('--seed', type=int, default=0)
    parser.add_argument('--save_dir', type=str, default='./saved_models')
    parser.add_argument('--data_root', type=str, default='./datasets/ffhq/')
    parser.add_argument('--desc', type=str, default='train')
    parser.add_argument('--dataset', type=str, default='cifar10')  # path to dataset
    parser.add_argument('--hparam_sets', '--hps', type=str)  # e.g. 'fewshot'
    parser.add_argument('--enc_blocks', type=str, default=None)  # specify encoder blocks, e.g. '1x2,4m1,4x4,8m4,8x5,16m8,16x8,32m16,32x5,64m32,64x4,128m64,128x4,256m128'
    parser.add_argument('--dec_blocks', type=str, default=None)  # specify decoder blocks, e.g. '256x4,128m64,128x4,64m32,64x4,32m16,32x5,16m8,16x8,8m4,8x5,4m1,4x4,1x2'
    parser.add_argument('--width', type=int, default=512)  # width of encoder and decoder convs
    parser.add_argument('--custom_width_str', type=str, default='')  # custom width for each block
    parser.add_argument('--bottleneck_multiple', type=float, default=0.25)  # coefficient width of bottleneck layers, e.g. 0.25 means 1/4 of width

    parser.add_argument('--restore_path', type=str, default=None)  # restore from checkpoint
    parser.add_argument('--restore_ema_path', type=str, default=None)  # restore ema from checkpoint
    parser.add_argument('--restore_log_path', type=str, default=None)  # restore log from checkpoint
    parser.add_argument('--restore_optimizer_path', type=str, default=None)  # restore optimizer from checkpoint
    parser.add_argument('--restore_scheduler_path', type=str, default=None)  # restore optimizer from scheduler
    parser.add_argument('--restore_scaler_path', type=str, default=None)  # restore optimizer from scheduler

    parser.add_argument('--restore_latent_path', type=str, default=None)  # restore nearest neighbour latent codes from checkpoint
    parser.add_argument('--restore_threshold_path', type=str, default=None)  # restore nearest neighbour thresholds, i.e., \tau_i, from checkpoint
    parser.add_argument('--ema_rate', type=float, default=0.999)  # exponential moving average rate
    parser.add_argument('--warmup_iters', type=float, default=2000)  # number of iterations for warmup for scheduler
    parser.add_argument('--lr_decay_iters', type=float, default=4000)  # number of iterations for warmup for scheduler
    parser.add_argument('--lr_decay_rate', type=float, default=0.25)  # number of iterations for warmup for scheduler

    parser.add_argument('--mapping_lr_multiplier', type=float, default=1.0)  # weight decay
    parser.add_argument('--mapping_normalization', type=str, default='layernorm', choices=['none', 'rmsnorm', 'layernorm', 'pixelnorm'])  # mapping network normalization type


    parser.add_argument(
        '--compile',
        default=False,
        type=lambda x: bool(strtobool(x)),
    )  # torch.compile (Inductor); default off — autotune can OOM large CIFAR jobs

    parser.add_argument('--lr', type=float, default=0.00015)  # learning rate
    parser.add_argument('--lr2', type=float, default=0.00005)  # learning rate

    parser.add_argument('--wd', type=float, default=0.00)  # weight decay
    parser.add_argument('--num_epochs', type=int, default=10000)  # number of epochs
    parser.add_argument('--n_batch', type=int, default=4)  # batch size
    parser.add_argument('--adam_beta1', type=float, default=0.9)
    parser.add_argument('--adam_beta2', type=float, default=0.9)
    parser.add_argument('--adam_eps', type=float, default=1e-8)

    parser.add_argument('--iters_per_ckpt', type=int, default=5000)  # number of iterations per checkpoint
    parser.add_argument('--iters_per_save', type=int, default=1000)  # number of iterations per saving the latest models
    parser.add_argument('--epoch_per_save', type=int, default=50)  # number of epochs per saving the latest models
    parser.add_argument('--iters_per_images', type=int, default=1000)  # number of iterations per sample save
    parser.add_argument('--num_images_visualize', type=int, default=10)  # number of images to visualize
    parser.add_argument('--num_rows_visualize', type=int, default=9)  # number of rows to visualize, e.g. 3 means 3x8=24 images
    # When True, all per-epoch / per-iter image dumps (NN-samples, samples-N, latest.png)
    # are skipped on rank 0. This avoids costly Lustre PNG writes that can stall an
    # entire epoch (and even cause DDP hangs while the other ranks wait).
    parser.add_argument('--no_viz', default=False,
                        type=lambda x: bool(strtobool(x)))

    parser.add_argument('--residual_ratio', type=float, default=-3.0)
    parser.add_argument('--residual_type', type=str, default='convex', choices=['normal', 'convex'])

    parser.add_argument('--accumulation_steps', type=int, default=1)  # accumulation steps
    parser.add_argument('--num_comp_indices', type=int, default=2)  # dci number of components
    parser.add_argument('--num_simp_indices', type=int, default=7)  # dci number of simplices
    parser.add_argument('--imle_db_size', type=int, default=1024)  # imle database size
    parser.add_argument('--imle_factor', type=float, default=0.)  # imle soft-sampling factor
    parser.add_argument('--imle_staleness', type=int, default=7)  # imle staleness, i.e., number of iterations to wait before considering the thresholds, tau_i
    parser.add_argument('--imle_batch', type=int, default=32)  # imle batch size used for sampling
    parser.add_argument('--subset_len', type=int, default=-1)  # subset length for training -- random subset of the dataset. -1 means full dataset
    parser.add_argument('--latent_dim', type=int, default=128)  # latent code dimension
    parser.add_argument('--imle_perturb_coef', type=float, default=0.001)  # imle perturbation coefficient to avoid same latent codes
    parser.add_argument('--lpips_net', type=str, default='vgg')  # lpips network type
    parser.add_argument('--proj_dim', type=int, default=800)  # projection dimension for nearest neighbour search
    parser.add_argument('--proj_proportion', type=int, default=1)  # whether to use projection proportional to the lpips feature dimensions for nearest neighbour search
    parser.add_argument('--lpips_coef', type=float, default=1.0)  # lpips loss coefficient
    parser.add_argument('--pixel_coef', type=float, default=0.1)  # pixel loss coefficient
    parser.add_argument('--dino_coef', type=float, default=1.0)  # dino loss coefficient
    parser.add_argument('--dino_cache_dir', type=str, default='./dinov2_cache')
    parser.add_argument('--force_factor', type=float, default=5)  # sampling factor for imle, i.e., force_factor * len(dataset)
    parser.add_argument('--change_coef', type=float, default=0.04)  # rate of change of thresholds tau_i
    parser.add_argument('--change_threshold', type=float, default=1)  # starting threshold
    parser.add_argument('--n_mpl', type=int, default=8)  # mapping network layers
    parser.add_argument('--latent_lr', type=float, default=0.0001)  # learning rate for optimizing latent codes -- not used
    parser.add_argument('--latent_decay', type=float, default=0.0)  # learning rate decay for optimizing latent codes -- not used
    parser.add_argument('--latent_epoch', type=int, default=0)  # number of epochs for optimizing latent codes -- not used
    parser.add_argument('--reconstruct_iter_num', type=int, default=100000)  # number of iterations for reconstructing images using backtracking
    parser.add_argument('--imle_force_resample', type=int, default=5)  # number of iterations to wait before ignoringthe threshold and resample anyway
    parser.add_argument('--snoise_factor', type=int, default=8)  # spatial noise factor
    parser.add_argument('--max_hierarchy', type=int, default=256)  # maximum hierarchy level for spatial noise, i.e., 64 means up to 64x64 spatial noise but not higher resolution
    parser.add_argument('--load_strict', type=int, default=1)  # whether to load checkpoints strict
    parser.add_argument('--lpips_path', type=str, default='./lpips')  # path to lpips weights
    parser.add_argument('--image_size', type=int, default=256)  # image size of dataset -- possible to downsample the dataset
    parser.add_argument('--num_images_to_generate', type=int, default=100)
    parser.add_argument('--mode', type=str, default='train')  # mode of running, train, eval, reconstruct, generate
    
    parser.add_argument('--use_adaptive', default=False, type=lambda x: bool(strtobool(x)))  # whether to use adaptive imle
    parser.add_argument('--zero_init', default=True, type=lambda x: bool(strtobool(x)))  # whether to use adaptive imle

    parser.add_argument('--angle', type=float, default=0.0)  # angle to splatter
    parser.add_argument('--use_splatter', default=False, type=lambda x: bool(strtobool(x)))  # whether to use splatter
    
    parser.add_argument('--use_gaussian', default=False, type=lambda x: bool(strtobool(x)))  # whether to use splatter
    parser.add_argument('--gaussian_std', type=float, default=0.1)  # gaussian std
    # parser.add_argument('--mode', type=str, default='lpips', choices=['lpips', 'l2', 'combined']) # search type for nearest neighbour search

    parser.add_argument('--use_multi_res', default=True, type=lambda x: bool(strtobool(x)))  # whether to use nearest neighbour search
    parser.add_argument('--align_corners', default=False, type=lambda x: bool(strtobool(x)))  # whether to use nearest neighbour search
    parser.add_argument('--use_resize_right', default=False, type=lambda x: bool(strtobool(x)))  # whether to use resize_right for resizing
    parser.add_argument('--frac_loss', default=False, type=lambda x: bool(strtobool(x)))  # whether to use fractional loss scaling
    parser.add_argument('--use_stopgrad_for_intermediate', default=False, type=lambda x: bool(strtobool(x)))  # whether to use stopgrad for intermediate targets

    parser.add_argument('--multi_res_scales', default='', type=str)  # extra multi-res dimension

    # parser.add_argument('--use_splatter_snoise', default=False, type=lambda x: bool(strtobool(x)))  # whether to use splatter snoise

    parser.add_argument('--use_snoise', default=False, type=lambda x: bool(strtobool(x)))  # whether to use spatial noise

    parser.add_argument('--search_type', type=str, default='lpips', choices=['lpips', 'l2', 'combined', 'vae']) # search type for nearest neighbour search
    parser.add_argument('--l2_search_downsample', type=float, default=0.125) # downsample factor for l2 search

    parser.add_argument('--wandb_name', type=str, default='AdaptiveIMLE')  # used for wandb
    parser.add_argument('--wandb_project', type=str, default='AdaptiveIMLE')  # used for wandb
    parser.add_argument('--use_wandb', type=int, default=0)
    parser.add_argument('--wandb_mode', type=str, default='online')

    parser.add_argument('--use_comet', default=False, type=lambda x: bool(strtobool(x)))
    parser.add_argument('--comet_name', type=str, default='AdaptiveIMLE')  # used in comet.ml
    parser.add_argument('--comet_api_key', type=str, default='')  # comet.ml api key -- leave blank to disable comet.ml
    parser.add_argument('--comet_experiment_key', type=str, default='')

    parser.add_argument("--convnext_expansion", type=int, default=4, help="expansion factor for convnext")
    parser.add_argument("--convnext_norm", default='rmsnorm',choices=["layernorm", "rmsnorm"], help="norm type for convnext block")
    parser.add_argument("--convnext_norm_eps", type=float, default=1e-3, help="epsilon for convnext norm")
    parser.add_argument("--use_convnext_bias", default=True, type=lambda x: bool(strtobool(x)))  # whether to use se block
    parser.add_argument("--use_convnext_weight", default=False, type=lambda x: bool(strtobool(x)))  # whether to use se block

    parser.add_argument("--use_se", default=True, type=lambda x: bool(strtobool(x)))  # whether to use se block
    parser.add_argument("--se_reduction", type=int, default=16, help="reduction factor for se block")
    parser.add_argument("--dropout_p", type=float, default=0.0, help="dropout rate for convnext block")

    parser.add_argument('--imle_db_topk', type=int, default=10)  # top-k for imle database search

    parser.add_argument("--loss_type", default='l2',choices=["l2", "huber", "welsch", "mclure"], help="type of loss")
    parser.add_argument("--huber_delta", type=float, default=0.05, help="delta for huber loss")
    parser.add_argument("--loss_scale", type=float, default=1.0, help="scale for general robust losses, e.g. pseudo-huber, pseudo-l1, cauchy")
    # some metric args
    parser.add_argument("--space", choices=["z", "w"], help="space that PPL calculated with")
    parser.add_argument("--batch", type=int, default=16, help="batch size for the models")
    parser.add_argument("--n_sample", type=int, default=5000, help="number of the samples for calculating PPL",)
    parser.add_argument("--size", type=int, default=256, help="output image sizes of the generator")
    parser.add_argument("--eps", type=float, default=1e-4, help="epsilon for numerical stability")
    parser.add_argument("--ppl_snoise", type=int, default=0, help="whether to interpolate spatial noise in PPL")
    parser.add_argument("--sampling", default="end", choices=["end", "full"], help="set endpoint sampling method",)
    parser.add_argument("--step", type=float, default=0.1, help="step size for interpolation")
    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=500, help="frequency of calculating fid")
    # Standalone FID-sample-dumping controls for --mode eval_fid
    parser.add_argument("--num_fid_samples", type=int, default=50000,
                        help="number of samples to dump in --mode eval_fid")
    parser.add_argument("--eval_fid_subdir", type=str, default="fid_eval_200k",
                        help="subdir under save_dir/train to dump eval_fid samples")
    parser.add_argument("--skip_cleanfid", default=True,
                        type=lambda x: bool(strtobool(x)),
                        help="skip cleanfid.compute_fid after dumping samples")

    # Inference-time knobs (read only in --mode eval_fid)
    parser.add_argument("--test_refinement_steps", type=int, default=None,
                        help="override TRM refinement_steps at eval time")
    parser.add_argument("--eval_latent_std", type=float, default=1.0,
                        help="scale of latent noise at eval time (1.0 = standard N(0,I))")

    # RTM mapper arguments
    parser.add_argument('--use_rtm', default=False,
                        type=lambda x: bool(strtobool(x)),
                        help='Use the Recursive Token Mapper instead of the single-pass MLP mapper.')
    parser.add_argument('--rtm_with_grad', default=False,
                        type=lambda x: bool(strtobool(x)))
    parser.add_argument('--H_cycles', type=int, default=1)
    parser.add_argument('--L_cycles', type=int, default=1)
    parser.add_argument('--L_layers', type=int, default=2)
    parser.add_argument('--H_layers', type=int, default=2)
    parser.add_argument('--refinement_steps', type=int, default=1)
    parser.add_argument('--num_tokens', type=int, default=1)
    parser.add_argument('--rtm_hidden_size', type=int, default=256)
    parser.add_argument('--rtm_expansion', type=float, default=4.0)
    parser.add_argument(
        '--rtm_cycle_noise_std', type=float, default=0.0,
        help='Optional Gaussian noise std added per H-cycle during training for mode coverage',
    )

    return parser