| 175 | return |
| 176 | |
| 177 | def get_args_parser(): |
| 178 | parser = argparse.ArgumentParser('Double Conditioning LDM Finetuning', add_help=False) |
| 179 | # project parameters |
| 180 | parser.add_argument('--seed', type=int) |
| 181 | parser.add_argument('--root_path', type=str, default = '../dreamdiffusion/') |
| 182 | parser.add_argument('--pretrain_mbm_path', type=str) |
| 183 | parser.add_argument('--checkpoint_path', type=str) |
| 184 | parser.add_argument('--crop_ratio', type=float) |
| 185 | parser.add_argument('--dataset', type=str) |
| 186 | |
| 187 | # finetune parameters |
| 188 | parser.add_argument('--batch_size', type=int) |
| 189 | parser.add_argument('--lr', type=float) |
| 190 | parser.add_argument('--num_epoch', type=int) |
| 191 | parser.add_argument('--precision', type=int) |
| 192 | parser.add_argument('--accumulate_grad', type=int) |
| 193 | parser.add_argument('--global_pool', type=bool) |
| 194 | |
| 195 | # diffusion sampling parameters |
| 196 | parser.add_argument('--pretrain_gm_path', type=str) |
| 197 | parser.add_argument('--num_samples', type=int) |
| 198 | parser.add_argument('--ddim_steps', type=int) |
| 199 | parser.add_argument('--use_time_cond', type=bool) |
| 200 | parser.add_argument('--eval_avg', type=bool) |
| 201 | |
| 202 | # # distributed training parameters |
| 203 | # parser.add_argument('--local_rank', type=int) |
| 204 | |
| 205 | return parser |
| 206 | |
| 207 | def update_config(args, config): |
| 208 | for attr in config.__dict__: |