(args)
| 20 | |
| 21 | |
| 22 | def main(args): |
| 23 | assert torch.cuda.is_available(), "Requires at least one GPU" |
| 24 | |
| 25 | try: |
| 26 | dist.init_process_group("nccl") |
| 27 | rank, world_size = dist.get_rank(), dist.get_world_size() |
| 28 | device = rank % torch.cuda.device_count() |
| 29 | seed = args.seed + rank |
| 30 | if rank == 0: |
| 31 | print(f"rank={rank}, seed={seed}, world_size={world_size}") |
| 32 | except: |
| 33 | rank, device, world_size, seed = 0, 0, 1, args.seed |
| 34 | |
| 35 | torch.manual_seed(seed) |
| 36 | torch.cuda.set_device(device) |
| 37 | |
| 38 | # Determine output directory based on model name |
| 39 | model_name = os.path.basename(args.hf_model_path.rstrip('/')) |
| 40 | output_dir = os.path.join(args.output_path, 'latents', model_name, f'imgnet{args.image_size}_norm{args.normalize_type}') |
| 41 | if rank == 0: |
| 42 | os.makedirs(output_dir, exist_ok=True) |
| 43 | print(f"Output directory: {output_dir}") |
| 44 | |
| 45 | # Create tokenizer |
| 46 | tokenizer = VTP_Tokenizer( |
| 47 | hf_model_path=args.hf_model_path, |
| 48 | img_size=args.image_size, |
| 49 | horizon_flip=0.0, |
| 50 | fp16=args.fp16, |
| 51 | normalize_type=args.normalize_type |
| 52 | ) |
| 53 | |
| 54 | datasets = [ |
| 55 | ImageFolder(args.data_path, transform=tokenizer.img_transform(p_hflip=p)) |
| 56 | for p in [0.0, 1.0] |
| 57 | ] |
| 58 | samplers = [ |
| 59 | DistributedSampler(ds, num_replicas=world_size, rank=rank, shuffle=False, seed=args.seed) |
| 60 | for ds in datasets |
| 61 | ] |
| 62 | loaders = [ |
| 63 | DataLoader(ds, batch_size=args.batch_size, shuffle=False, sampler=s, |
| 64 | num_workers=args.num_workers, pin_memory=True, drop_last=False) |
| 65 | for ds, s in zip(datasets, samplers) |
| 66 | ] |
| 67 | |
| 68 | if rank == 0: |
| 69 | print(f"Total data: {len(loaders[0].dataset)}") |
| 70 | |
| 71 | run_images = saved_files = 0 |
| 72 | latents, latents_flip, labels = [], [], [] |
| 73 | |
| 74 | for batch_idx, batch_data in enumerate(zip(*loaders)): |
| 75 | run_images += batch_data[0][0].shape[0] |
| 76 | if run_images % 100 == 0 and rank == 0: |
| 77 | print(f'{datetime.now()} processing {run_images}/{len(loaders[0].dataset)}') |
| 78 | |
| 79 | for loader_idx, (x, y) in enumerate(batch_data): |
no test coverage detected