MCPcopy Create free account
hub / github.com/MiniMax-AI/VTP / main

Function main

generation/tools/extract_features_vtp.py:22–126  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

20
21
22def 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):

Callers 1

Calls 3

img_transformMethod · 0.95
encode_imagesMethod · 0.95
VTP_TokenizerClass · 0.90

Tested by

no test coverage detected