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 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.warmup_iters = 10 fewshot.dataset = 'fewshot' fewshot.n_batch = 4 fewshot.ema_rate = 0.9999 HPARAMS_REGISTRY['fewshot'] = fewshot cifar10_hps = Hyperparams() cifar10_hps.width = 768 cifar10_hps.lr = 0.0008 cifar10_hps.wd = 0.01 cifar10_hps.dec_blocks = '1x1,4m1,4x2,8m4,8x5,16m8,16x5,32m16,32x5' cifar10_hps.warmup_iters = 100 cifar10_hps.dataset = 'cifar10' cifar10_hps.n_batch = 196 cifar10_hps.ema_rate = 0.9999 cifar10_hps.force_factor = 5 cifar10_hps.imle_force_resample = 5 cifar10_hps.search_type = 'lpips' cifar10_hps.imle_batch = 1024 cifar10_hps.image_size = 32 cifar10_hps.convnext_expansion = 6 cifar10_hps.use_se = True cifar10_hps.se_reduction = 16 cifar10_hps.dropout_p = 0.0 cifar10_hps.use_multi_res = True cifar10_hps.multi_res_scales = '8,12,16,24,28' cifar10_hps.mapping_lr_multiplier = 0.01 cifar10_hps.accumulation_steps = 1 cifar10_hps.epoch_per_save = 50 cifar10_hps.dino_coef = 1.0 cifar10_hps.l2_search_downsample = 1.0 cifar10_hps.latent_dim = 128 cifar10_hps.pixel_coef = 0.1 cifar10_hps.residual_ratio = -3.0 cifar10_hps.residual_type = 'convex' cifar10_hps.convnext_norm = 'rmsnorm' cifar10_hps.convnext_norm_eps = 1e-3 cifar10_hps.use_stopgrad_for_intermediate = False cifar10_hps.align_corners = False cifar10_hps.loss_type = 'l2' cifar10_hps.huber_delta = 0.05 cifar10_hps.loss_scale = 1.0 cifar10_hps.imle_db_topk = 10 cifar10_hps.nn_search_batch = 4096 cifar10_hps.ignore_radius = 0.0 cifar10_hps.resample_angle = 0.0 HPARAMS_REGISTRY['cifar10'] = cifar10_hps 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__) if isinstance(H.get('multi_res_scales'), str) and H['multi_res_scales']: H['multi_res_scales'] = [int(x) for x in H['multi_res_scales'].split(',')] 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='./') 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' # specify encoder blocks, e.g. '1x2,4m1,4x4,8m4,8x5,16m8,16x8,32m16,32x5,64m32,64x4,128m64,128x4,256m128' parser.add_argument('--enc_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('--dec_blocks', type=str, default=None) # width of encoder and decoder convs parser.add_argument('--width', type=int, default=512) parser.add_argument('--custom_width_str', type=str, default='') # custom width for each block # coefficient width of bottleneck layers, e.g. 0.25 means 1/4 of width parser.add_argument('--bottleneck_multiple', type=float, default=0.25) 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 # restore optimizer from checkpoint parser.add_argument('--restore_optimizer_path', type=str, default=None) # restore optimizer from scheduler parser.add_argument('--restore_scheduler_path', type=str, default=None) # restore nearest neighbour latent codes from checkpoint parser.add_argument('--restore_latent_path', type=str, default=None) # restore nearest neighbour thresholds, i.e., \tau_i, from checkpoint parser.add_argument('--restore_threshold_path', type=str, default=None) parser.add_argument('--restore_last_updated_path', type=str, default=None) parser.add_argument('--restore_times_updated_path', type=str, default=None) # exponential moving average rate parser.add_argument('--ema_rate', type=float, default=0.999) # number of iterations for warmup for scheduler parser.add_argument('--warmup_iters', type=float, default=0) # 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) parser.add_argument('--mapping_normalization', type=str, default='layernorm', choices=['none', 'rmsnorm', 'layernorm', 'pixelnorm']) parser.add_argument('--lr', type=float, default=0.0002) # learning rate parser.add_argument('--lr2', type=float, default=0.00005) parser.add_argument('--wd', type=float, default=0.00) # weight decay parser.add_argument('--num_epochs', type=int, default=15000) # number of epochs parser.add_argument('--n_batch', type=int, default=8) # batch size parser.add_argument('--adam_beta1', type=float, default=0.9) parser.add_argument('--adam_beta2', type=float, default=0.9) # number of iterations per checkpoint parser.add_argument('--iters_per_ckpt', type=int, default=100000) # number of iterations per saving the latest models parser.add_argument('--iters_per_save', type=int, default=1000) # number of iterations per sample save parser.add_argument('--iters_per_images', type=int, default=5000) parser.add_argument('--num_images_visualize', type=int, default=10) # number of images to visualize # number of rows to visualize, e.g. 3 means 3x8=24 images parser.add_argument('--num_rows_visualize', type=int, default=5) 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('--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 # imle soft-sampling factor parser.add_argument('--imle_factor', type=float, default=0.) # imle batch size used for sampling parser.add_argument('--imle_batch', type=int, default=16) # subset length for training -- random subset of the dataset. -1 means full dataset parser.add_argument('--subset_len', type=int, default=-1) parser.add_argument('--latent_dim', type=int, default=4096) # latent code dimension # imle perturbation coefficient to avoid same latent codes parser.add_argument('--imle_perturb_coef', type=float, default=0.001) parser.add_argument('--lpips_net', type=str, default='vgg') # lpips network type # projection dimension for nearest neighbour search parser.add_argument('--proj_dim', type=int, default=800) # whether to use projection proportional to the lpips feature dimensions for nearest neighbour search parser.add_argument('--proj_proportion', type=int, default=1) parser.add_argument('--lpips_coef', type=float, default=1.0) # lpips loss coefficient parser.add_argument('--pixel_coef', type=float, default=0.1) parser.add_argument('--l2_coef', type=float, default=0.1) # l2 loss coefficient # sampling factor for imle, i.e., force_factor * len(dataset) parser.add_argument('--force_factor', type=float, default=20) # mapping network layers parser.add_argument('--n_mpl', type=int, default=8) # number of iterations for reconstructing images using backtracking parser.add_argument('--reconstruct_iter_num', type=int, default=100000) # number of iterations to wait before ignoringthe threshold and resample anyway parser.add_argument('--imle_force_resample', type=int, default=30) parser.add_argument('--snoise_factor', type=int, default=8) # spatial noise factor # maximum hierarchy level for spatial noise, i.e., 64 means up to 64x64 spatial noise but not higher resolution parser.add_argument('--max_hierarchy', type=int, default=256) # whether to load checkpoints strict parser.add_argument('--load_strict', type=int, default=1) parser.add_argument('--lpips_path', type=str, default='./lpips') # path to lpips weights # image size of dataset -- possible to downsample the dataset parser.add_argument('--image_size', type=int, default=256) parser.add_argument('--num_images_to_generate', type=int, default=100) # mode of running, train, eval, reconstruct, generate parser.add_argument('--mode', type=str, default='train') # whether to use spatial noise parser.add_argument('--use_snoise', default=False, type=lambda x: bool(strtobool(x))) # search type for nearest neighbour search parser.add_argument('--search_type', type=str, default='l2', choices=['lpips', 'l2', 'combined']) # downsample factor for l2 search parser.add_argument('--l2_search_downsample', type=float, default=1.0) # RSIMLE specific arguments # whether to use spatial noise parser.add_argument('--use_rsimle', default=True, type=lambda x: bool(strtobool(x))) # rejection-sampling threshold for RS-IMLE parser.add_argument('--eps_radius', type=float, default=0.12) parser.add_argument('--knn_ignore', type=int, default=5) # knn ignore for RSIMLE # adaptive IMLE # whether to use adaptive imle parser.add_argument('--use_adaptive', default=False, type=lambda x: bool(strtobool(x))) # rate of change of the thresholds, tau_i parser.add_argument('--change_coef', type=float, default=0.04) parser.add_argument('--change_threshold', type=float, default=1) # starting threshold # imle staleness, i.e., number of iterations to wait before considering the thresholds, tau_i parser.add_argument('--imle_staleness', type=int, default=7) # wandb 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') # comet.ml 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 # comet.ml api key -- leave blank to disable comet.ml parser.add_argument('--comet_api_key', type=str, default='') parser.add_argument('--comet_experiment_key', type=str, default='') # learning rate for optimizing latent codes -- not used parser.add_argument('--latent_lr', type=float, default=0.0001) # learning rate decay for optimizing latent codes -- not used parser.add_argument('--latent_decay', type=float, default=0.0) # number of epochs for optimizing latent codes -- not used parser.add_argument('--latent_epoch', type=int, default=0) # 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=100, help="frequency of calculating fid") # Standalone FID-sample-dumping controls for --mode eval_fid parser.add_argument("--num_fid_samples", type=int, default=5000, help="number of samples to dump in --mode eval_fid") parser.add_argument("--eval_fid_subdir", type=str, default="fid", help="subdir under save_dir to dump eval_fid samples") parser.add_argument("--skip_cleanfid", default=False, type=lambda x: bool(strtobool(x)), help="skip cleanfid.compute_fid after dumping samples") # ConvNeXt/SE architecture arguments (for CIFAR-10 pipeline) parser.add_argument('--convnext_expansion', type=int, default=4) parser.add_argument('--convnext_norm', default='rmsnorm', choices=['layernorm', 'rmsnorm']) parser.add_argument('--convnext_norm_eps', type=float, default=1e-3) parser.add_argument('--use_convnext_bias', default=True, type=lambda x: bool(strtobool(x))) parser.add_argument('--use_convnext_weight', default=False, type=lambda x: bool(strtobool(x))) parser.add_argument('--use_se', default=True, type=lambda x: bool(strtobool(x))) parser.add_argument('--se_reduction', type=int, default=16) parser.add_argument('--dropout_p', type=float, default=0.0) parser.add_argument('--mapping_lr_multiplier', type=float, default=0.01) parser.add_argument('--compile', default=False, type=lambda x: bool(strtobool(x))) # Multi-resolution loss parser.add_argument('--use_multi_res', default=False, type=lambda x: bool(strtobool(x))) parser.add_argument('--align_corners', default=False, type=lambda x: bool(strtobool(x))) parser.add_argument('--use_resize_right', default=False, type=lambda x: bool(strtobool(x))) parser.add_argument('--frac_loss', default=False, type=lambda x: bool(strtobool(x))) parser.add_argument('--use_stopgrad_for_intermediate', default=False, type=lambda x: bool(strtobool(x))) parser.add_argument('--multi_res_scales', default='', type=str) parser.add_argument('--accumulation_steps', type=int, default=1) parser.add_argument('--epoch_per_save', type=int, default=50) # DINOv2 / combined search parser.add_argument('--dino_coef', type=float, default=0.0) parser.add_argument('--dino_cache_dir', type=str, default='./dinov2_cache') parser.add_argument('--imle_db_topk', type=int, default=10) parser.add_argument('--nn_search_batch', type=int, default=4096) parser.add_argument('--ignore_radius', type=float, default=0.0) parser.add_argument('--resample_angle', type=float, default=0.0) parser.add_argument('--loss_type', default='l2', choices=['l2', 'huber', 'welsch', 'mclure']) parser.add_argument('--huber_delta', type=float, default=0.05) parser.add_argument('--loss_scale', type=float, default=1.0) parser.add_argument('--adam_eps', type=float, default=1e-8) parser.add_argument('--total_iters', type=int, default=200000) # Scaler restore parser.add_argument('--restore_scaler_path', type=str, default=None) # 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