| 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.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.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.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' |
| |
| 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 |
|
|
| |
| |
| |
| 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') |
| parser.add_argument('--hparam_sets', '--hps', type=str) |
| parser.add_argument('--enc_blocks', type=str, default=None) |
| parser.add_argument('--dec_blocks', type=str, default=None) |
| parser.add_argument('--width', type=int, default=512) |
| parser.add_argument('--custom_width_str', type=str, default='') |
| parser.add_argument('--bottleneck_multiple', type=float, default=0.25) |
|
|
| parser.add_argument('--restore_path', type=str, default=None) |
| parser.add_argument('--restore_ema_path', type=str, default=None) |
| parser.add_argument('--restore_log_path', type=str, default=None) |
| parser.add_argument('--restore_optimizer_path', type=str, default=None) |
| parser.add_argument('--restore_scheduler_path', type=str, default=None) |
| parser.add_argument('--restore_scaler_path', type=str, default=None) |
|
|
| parser.add_argument('--restore_latent_path', type=str, default=None) |
| parser.add_argument('--restore_threshold_path', type=str, default=None) |
| parser.add_argument('--ema_rate', type=float, default=0.999) |
| parser.add_argument('--warmup_iters', type=float, default=2000) |
| parser.add_argument('--lr_decay_iters', type=float, default=4000) |
| parser.add_argument('--lr_decay_rate', type=float, default=0.25) |
|
|
| parser.add_argument('--mapping_lr_multiplier', type=float, default=1.0) |
| parser.add_argument('--mapping_normalization', type=str, default='layernorm', choices=['none', 'rmsnorm', 'layernorm', 'pixelnorm']) |
|
|
|
|
| parser.add_argument( |
| '--compile', |
| default=False, |
| type=lambda x: bool(strtobool(x)), |
| ) |
|
|
| parser.add_argument('--lr', type=float, default=0.00015) |
| parser.add_argument('--lr2', type=float, default=0.00005) |
|
|
| parser.add_argument('--wd', type=float, default=0.00) |
| parser.add_argument('--num_epochs', type=int, default=10000) |
| parser.add_argument('--n_batch', type=int, default=4) |
| 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) |
| parser.add_argument('--iters_per_save', type=int, default=1000) |
| parser.add_argument('--epoch_per_save', type=int, default=50) |
| parser.add_argument('--iters_per_images', type=int, default=1000) |
| parser.add_argument('--num_images_visualize', type=int, default=10) |
| parser.add_argument('--num_rows_visualize', type=int, default=9) |
| |
| |
| |
| 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) |
| parser.add_argument('--num_comp_indices', type=int, default=2) |
| parser.add_argument('--num_simp_indices', type=int, default=7) |
| parser.add_argument('--imle_db_size', type=int, default=1024) |
| parser.add_argument('--imle_factor', type=float, default=0.) |
| parser.add_argument('--imle_staleness', type=int, default=7) |
| parser.add_argument('--imle_batch', type=int, default=32) |
| parser.add_argument('--subset_len', type=int, default=-1) |
| parser.add_argument('--latent_dim', type=int, default=128) |
| parser.add_argument('--imle_perturb_coef', type=float, default=0.001) |
| parser.add_argument('--lpips_net', type=str, default='vgg') |
| parser.add_argument('--proj_dim', type=int, default=800) |
| parser.add_argument('--proj_proportion', type=int, default=1) |
| parser.add_argument('--lpips_coef', type=float, default=1.0) |
| parser.add_argument('--pixel_coef', type=float, default=0.1) |
| parser.add_argument('--dino_coef', type=float, default=1.0) |
| parser.add_argument('--dino_cache_dir', type=str, default='./dinov2_cache') |
| parser.add_argument('--force_factor', type=float, default=5) |
| parser.add_argument('--change_coef', type=float, default=0.04) |
| parser.add_argument('--change_threshold', type=float, default=1) |
| parser.add_argument('--n_mpl', type=int, default=8) |
| parser.add_argument('--latent_lr', type=float, default=0.0001) |
| parser.add_argument('--latent_decay', type=float, default=0.0) |
| parser.add_argument('--latent_epoch', type=int, default=0) |
| parser.add_argument('--reconstruct_iter_num', type=int, default=100000) |
| parser.add_argument('--imle_force_resample', type=int, default=5) |
| parser.add_argument('--snoise_factor', type=int, default=8) |
| parser.add_argument('--max_hierarchy', type=int, default=256) |
| parser.add_argument('--load_strict', type=int, default=1) |
| parser.add_argument('--lpips_path', type=str, default='./lpips') |
| parser.add_argument('--image_size', type=int, default=256) |
| parser.add_argument('--num_images_to_generate', type=int, default=100) |
| parser.add_argument('--mode', type=str, default='train') |
| |
| parser.add_argument('--use_adaptive', default=False, type=lambda x: bool(strtobool(x))) |
| parser.add_argument('--zero_init', default=True, type=lambda x: bool(strtobool(x))) |
|
|
| parser.add_argument('--angle', type=float, default=0.0) |
| parser.add_argument('--use_splatter', default=False, type=lambda x: bool(strtobool(x))) |
| |
| parser.add_argument('--use_gaussian', default=False, type=lambda x: bool(strtobool(x))) |
| parser.add_argument('--gaussian_std', type=float, default=0.1) |
| |
|
|
| parser.add_argument('--use_multi_res', default=True, 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('--use_snoise', default=False, type=lambda x: bool(strtobool(x))) |
|
|
| parser.add_argument('--search_type', type=str, default='lpips', choices=['lpips', 'l2', 'combined', 'vae']) |
| parser.add_argument('--l2_search_downsample', type=float, default=0.125) |
|
|
| parser.add_argument('--wandb_name', type=str, default='AdaptiveIMLE') |
| parser.add_argument('--wandb_project', type=str, default='AdaptiveIMLE') |
| 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') |
| parser.add_argument('--comet_api_key', type=str, default='') |
| 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))) |
| 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, 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) |
|
|
| 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") |
| |
| 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") |
| |
| 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") |
|
|
| |
| 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))") |
|
|
| |
| 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 |
|
|