| 20 | |
| 21 | |
| 22 | def get_args_parser(): |
| 23 | parser = argparse.ArgumentParser('Cache VAE latents', add_help=False) |
| 24 | parser.add_argument('--batch_size', default=128, type=int, |
| 25 | help='Batch size per GPU (effective batch size is batch_size * # gpus') |
| 26 | |
| 27 | # VAE parameters |
| 28 | parser.add_argument('--img_size', default=256, type=int, |
| 29 | help='images input size') |
| 30 | parser.add_argument('--vae_path', default="pretrained_models/vae/kl16.ckpt", type=str, |
| 31 | help='images input size') |
| 32 | parser.add_argument('--vae_embed_dim', default=16, type=int, |
| 33 | help='vae output embedding dimension') |
| 34 | # Dataset parameters |
| 35 | parser.add_argument('--data_path', default='./data/imagenet', type=str, |
| 36 | help='dataset path') |
| 37 | parser.add_argument('--device', default='cuda', |
| 38 | help='device to use for training / testing') |
| 39 | parser.add_argument('--seed', default=0, type=int) |
| 40 | |
| 41 | parser.add_argument('--num_workers', default=10, type=int) |
| 42 | parser.add_argument('--pin_mem', action='store_true', |
| 43 | help='Pin CPU memory in DataLoader for more efficient (sometimes) transfer to GPU.') |
| 44 | parser.add_argument('--no_pin_mem', action='store_false', dest='pin_mem') |
| 45 | parser.set_defaults(pin_mem=True) |
| 46 | |
| 47 | # distributed training parameters |
| 48 | parser.add_argument('--world_size', default=1, type=int, |
| 49 | help='number of distributed processes') |
| 50 | parser.add_argument('--local_rank', default=-1, type=int) |
| 51 | parser.add_argument('--dist_on_itp', action='store_true') |
| 52 | parser.add_argument('--dist_url', default='env://', |
| 53 | help='url used to set up distributed training') |
| 54 | |
| 55 | # caching latents |
| 56 | parser.add_argument('--cached_path', default='', help='path to cached latents') |
| 57 | |
| 58 | return parser |
| 59 | |
| 60 | |
| 61 | def main(args): |