(args)
| 40 | |
| 41 | |
| 42 | def main(args): |
| 43 | local_rank = int(os.getenv("RANK", 0)) |
| 44 | world_size = int(os.getenv("WORLD_SIZE", 1)) |
| 45 | print("world_size", world_size, "local rank", local_rank) |
| 46 | |
| 47 | device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| 48 | torch.cuda.set_device(local_rank) |
| 49 | if not dist.is_initialized(): |
| 50 | dist.init_process_group( |
| 51 | backend="nccl", init_method="env://", world_size=world_size, rank=local_rank |
| 52 | ) |
| 53 | |
| 54 | os.makedirs(args.output_dir, exist_ok=True) |
| 55 | dataset = UniGenBenchDataset(args.prompt_dir) |
| 56 | sampler = DistributedSampler( |
| 57 | dataset, rank=local_rank, num_replicas=world_size, shuffle=False |
| 58 | ) |
| 59 | dataloader = DataLoader( |
| 60 | dataset, |
| 61 | sampler=sampler, |
| 62 | batch_size=args.batch_size, |
| 63 | num_workers=args.dataloader_num_workers, |
| 64 | ) |
| 65 | |
| 66 | transformer = FluxTransformer2DModel.from_pretrained( |
| 67 | args.model_path, |
| 68 | subfolder="transformer", |
| 69 | torch_dtype=torch.float16 |
| 70 | ).to(device) |
| 71 | |
| 72 | vae = AutoencoderKL.from_pretrained(args.model_path, subfolder="vae", torch_dtype=torch.float16).to(device) |
| 73 | text_encoder = CLIPTextModel.from_pretrained(args.model_path, subfolder="text_encoder", torch_dtype=torch.float16).to(device) |
| 74 | tokenizer = CLIPTokenizer.from_pretrained(args.model_path, subfolder="tokenizer") |
| 75 | text_encoder_2 = T5EncoderModel.from_pretrained(args.model_path, subfolder="text_encoder_2", torch_dtype=torch.float16).to(device) |
| 76 | tokenizer_2 = T5TokenizerFast.from_pretrained(args.model_path, subfolder="tokenizer_2") |
| 77 | scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(args.model_path, subfolder="scheduler") |
| 78 | |
| 79 | pipe = FluxPipeline( |
| 80 | scheduler=scheduler, |
| 81 | vae=vae, |
| 82 | text_encoder=text_encoder, |
| 83 | tokenizer=tokenizer, |
| 84 | text_encoder_2=text_encoder_2, |
| 85 | tokenizer_2=tokenizer_2, |
| 86 | transformer=transformer, |
| 87 | ) |
| 88 | pipe.to(device) |
| 89 | pipe.set_progress_bar_config(disable=False) |
| 90 | |
| 91 | for _, data in tqdm(enumerate(dataloader), disable=local_rank != 0): |
| 92 | try: |
| 93 | for j in range(4): |
| 94 | with torch.inference_mode(): |
| 95 | seed = 3407+j |
| 96 | prompt = data['caption'][0] |
| 97 | idx = data['idx'][0] |
| 98 | image = pipe( |
| 99 | prompt, |
no test coverage detected