JerMa88's picture
Upload folder using huggingface_hub
3ce19a2 verified
Raw
History Blame Contribute Delete
20.2 kB
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