| 14 | |
| 15 | |
| 16 | def parse_args(): |
| 17 | parser = argparse.ArgumentParser() |
| 18 | parser.add_argument("--finetuned_model_path", type=str, default=None) |
| 19 | parser.add_argument("--pretrained_model_name_or_path", type=str, default="inference/saved_pipeline/jodiffusion") |
| 20 | parser.add_argument("--pretrained_label_vae_path", type=str, default="output/dense-label-vae-light-ade20k-semantic-lr1e4") |
| 21 | parser.add_argument("--save_dir", type=str, default=None) |
| 22 | parser.add_argument("--dataset_name", type=str, default="ade20k_semantic") |
| 23 | parser.add_argument("--lightweight_label_vae", action="store_true") |
| 24 | parser.add_argument("--generate_mode", type=str, default="joint", choices=["text2img", "joint"]) |
| 25 | parser.add_argument("--num_images", type=int, default=100) |
| 26 | parser.add_argument("--seed", type=int, default=42) |
| 27 | return parser.parse_args() |
| 28 | |
| 29 | def generate(args, weight_dtype): |
| 30 | random.seed(args.seed) |