MCPcopy Create free account
hub / github.com/Alpha-VLLM/LLaMA2-Accessory / main

Function main

Large-DiT-ImageNet/sample.py:26–113  ·  view source on GitHub ↗
(args, rank, master_port)

Source from the content-addressed store, hash-verified

24
25
26def 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,

Callers 1

sample.pyFile · 0.70

Calls 6

create_diffusionFunction · 0.90
printFunction · 0.85
from_pretrainedMethod · 0.80
decodeMethod · 0.80
load_state_dictMethod · 0.45
p_sample_loopMethod · 0.45

Tested by

no test coverage detected