Encode pixel data to VAE latents. Parallelization and Memory Management: Use --gpus to control the number of GPUs used for parallel encoding (default: 8). Use --batch-size to control how many images each GPU processes in each batch (default: 100). Images are distributed acr
(
model_url: str,
source: str,
dest: str,
max_images: Optional[int],
gpus: int,
batch_size: int,
)
| 449 | @click.option('--batch-size', help='Number of images per GPU in each batch', metavar='INT', type=int, default=100, show_default=True) |
| 450 | |
| 451 | def encode( |
| 452 | model_url: str, |
| 453 | source: str, |
| 454 | dest: str, |
| 455 | max_images: Optional[int], |
| 456 | gpus: int, |
| 457 | batch_size: int, |
| 458 | ): |
| 459 | """Encode pixel data to VAE latents. |
| 460 | |
| 461 | Parallelization and Memory Management: |
| 462 | |
| 463 | Use --gpus to control the number of GPUs used for parallel encoding |
| 464 | (default: 8). Use --batch-size to control how many images each GPU |
| 465 | processes in each batch (default: 100). Images are distributed across |
| 466 | GPUs in round-robin fashion and processed in batches to avoid loading |
| 467 | all images into memory at once. |
| 468 | |
| 469 | Example: |
| 470 | \b |
| 471 | python dataset_tool.py encode --source=datasets/img64.zip \\ |
| 472 | --dest=datasets/img64_encoded.zip --gpus=8 --batch-size=50 |
| 473 | """ |
| 474 | PIL.Image.init() |
| 475 | if dest == '': |
| 476 | raise click.ClickException('--dest output filename or directory must not be an empty string') |
| 477 | |
| 478 | num_files, input_iter = open_dataset(source, max_images=max_images) |
| 479 | archive_root_dir, save_bytes, close_dest = open_dest(dest) |
| 480 | |
| 481 | # Process images in batches across GPUs to avoid loading everything into memory |
| 482 | labels = [] |
| 483 | |
| 484 | print(f"Processing {num_files} images in batches of {batch_size * gpus} across {gpus} GPUs...") |
| 485 | |
| 486 | with mp.Pool(gpus) as pool: |
| 487 | batch = [] |
| 488 | gpu_batches = [[] for _ in range(gpus)] |
| 489 | |
| 490 | for idx, image in tqdm(enumerate(input_iter), total=num_files, desc="Encoding images"): |
| 491 | # Distribute images across GPUs in round-robin fashion |
| 492 | gpu_id = idx % gpus |
| 493 | gpu_batches[gpu_id].append((idx, image)) |
| 494 | |
| 495 | # Process when any GPU batch is full or we've reached the end |
| 496 | max_batch_size = max(len(gpu_batch) for gpu_batch in gpu_batches) |
| 497 | if max_batch_size >= batch_size or idx == num_files - 1: |
| 498 | # Prepare arguments for each GPU |
| 499 | gpu_args = [] |
| 500 | for gpu_id, gpu_batch in enumerate(gpu_batches): |
| 501 | if gpu_batch: # Only process non-empty batches |
| 502 | gpu_args.append((gpu_id, gpu_batch, model_url)) |
| 503 | |
| 504 | # Process current batches in parallel across GPUs |
| 505 | if gpu_args: |
| 506 | batch_results = pool.map(encode_image_worker, gpu_args) |
| 507 | |
| 508 | # Flatten results and sort by index |
nothing calls this directly
no test coverage detected