(args, rank, master_port)
| 24 | |
| 25 | |
| 26 | def main(args, rank, master_port): |
| 27 | # Setup PyTorch: |
| 28 | torch.manual_seed(args.seed) |
| 29 | torch.set_grad_enabled(False) |
| 30 | |
| 31 | os.environ["RANK"] = str(rank) |
| 32 | os.environ["WORLD_SIZE"] = str(args.num_gpus) |
| 33 | os.environ["MASTER_PORT"] = str(master_port) |
| 34 | os.environ["MASTER_ADDR"] = "127.0.0.1" |
| 35 | |
| 36 | dist.init_process_group("nccl") |
| 37 | fs_init.initialize_model_parallel(args.num_gpus) |
| 38 | torch.cuda.set_device(rank) |
| 39 | |
| 40 | train_args = torch.load(os.path.join(args.ckpt, "model_args.pth")) |
| 41 | |
| 42 | if dist.get_rank() == 0: |
| 43 | print("Model arguments used for inference:", |
| 44 | json.dumps(train_args.__dict__, indent=2)) |
| 45 | |
| 46 | # Load model: |
| 47 | latent_size = train_args.image_size // 8 |
| 48 | model = DiT_models[train_args.model]( |
| 49 | input_size=latent_size, |
| 50 | num_classes=train_args.num_classes, |
| 51 | qk_norm=train_args.qk_norm, |
| 52 | ) |
| 53 | |
| 54 | torch_dtype = { |
| 55 | "fp32": torch.float, "tf32": torch.float, |
| 56 | "bf16": torch.bfloat16, "fp16": torch.float16, |
| 57 | }[args.precision] |
| 58 | model.to(torch_dtype).cuda() |
| 59 | if args.precision == "tf32": |
| 60 | torch.backends.cuda.matmul.allow_tf32 = True |
| 61 | torch.backends.cudnn.allow_tf32 = True |
| 62 | |
| 63 | assert train_args.model_parallel_size == args.num_gpus |
| 64 | ckpt = torch.load(os.path.join( |
| 65 | args.ckpt, |
| 66 | f"consolidated{'_ema' if args.ema else ''}." |
| 67 | f"{rank:02d}-of-{args.num_gpus:02d}.pth", |
| 68 | ), map_location="cpu") |
| 69 | model.load_state_dict(ckpt, strict=True) |
| 70 | |
| 71 | model.eval() # important! |
| 72 | diffusion = create_diffusion(str(args.num_sampling_steps)) |
| 73 | vae = AutoencoderKL.from_pretrained( |
| 74 | f"stabilityai/sd-vae-ft-{train_args.vae}" |
| 75 | if args.local_diffusers_model_root is None else |
| 76 | os.path.join(args.local_diffusers_model_root, |
| 77 | f"stabilityai/sd-vae-ft-{train_args.vae}") |
| 78 | ).cuda() |
| 79 | |
| 80 | # Create sampling noise: |
| 81 | n = len(args.class_labels) |
| 82 | z = torch.randn( |
| 83 | n, 4, latent_size, latent_size, |
no test coverage detected