Configuring Multi-Scale WebDataset Training with TAR Files in NVIDIA Sana

Use SanaWebDatasetMS with properly structured TAR archives and a wids-meta.json index to train on heterogeneous image resolutions while dynamically selecting aspect ratios and sampling captions by ClipScore.

Configuring multi-scale WebDataset training with TAR files enables efficient, high-throughput training on large-scale image-text datasets without uniform preprocessing. In the NVlabs/Sana repository, the SanaWebDatasetMS class extends the standard WebDataset pipeline to support dynamic resolution switching, external caption weighting, and VAE latent caching directly from shard archives.

Understanding the Multi-Scale WebDataset Architecture

The multi-scale implementation centers on SanaWebDatasetMS, defined in diffusion/data/datasets/sana_data_multi_scale.py. This subclass inherits from SanaWebDataset (base implementation in diffusion/data/datasets/sana_data.py) and injects aspect-ratio-aware resizing logic into the standard data loading pipeline.

Dynamic Aspect Ratio Selection

Instead of fixed-size cropping, SanaWebDatasetMS.__getitem__ calculates the closest target aspect ratio for each sample by comparing original dimensions (ori_h, ori_w) against a user-provided ratio dictionary. The implementation at lines 49-55 of sana_data_multi_scale.py uses the get_closest_ratio helper to select the matching bucket, then resizes images using either BICUBIC or LANCZOS interpolation based on the configured interpolate_model.

Shard-Level Metadata Caching

Under the hood, ShardListDatasetMulti (instantiated in the _initialize_dataset method at lines 100-114 of sana_data.py) manages TAR access. It caches shard indexes under ~/.cache/_wids_cache/<user>-<uuid>, enabling O(1) sample retrieval without re-scanning archives on subsequent training runs.

Preparing TAR Archives and Metadata

Before training, you must structure your WebDataset shards and generate the required metadata index.

Required Shard Structure

Each TAR archive must contain:

  • Image files: .png or .jpg (or .npy for pre-computed VAE latents)
  • Side-car metadata: <sample_key>.json containing height, width, and caption fields

Generating wids-meta.json

The dataset expects a wids-meta.json file adjacent to your TAR shards that maps sample keys to byte offsets within archives. Generate this using the provided utility:

python tools/create_wids_metadata.py \
    --input_dir /path/to/tar/shards \
    --output_dir /path/to/tar/shards

This metadata enables streaming reads without extracting archives, which is critical for training on petabyte-scale datasets.

Implementing Multi-Scale Training Configuration

Instantiate the dataset by passing the aspect ratio dictionary as a string (parsed via eval at line 94 of sana_data_multi_scale.py) along with your caption sampling preferences.

Minimal Configuration Example

import torch
from torch.utils.data import DataLoader
from diffusion.data.datasets.sana_data_multi_scale import SanaWebDatasetMS
from diffusion.data.transforms import get_transform

# Define target aspect ratios as width/height floats

ASPECT_RATIOS = {
    "0.57": {"name": "portrait"},   # 9:16 aspect

    "1.00": {"name": "square"},
    "1.78": {"name": "landscape"},  # 16:9 aspect

}

# Initialize transform pipeline

transform = get_transform("default_train", image_size=256)

# Configure multi-scale WebDataset

train_dataset = SanaWebDatasetMS(
    data_dir="/path/to/tar/shards",
    meta_path=None,                           # Auto-detect wids-meta.json

    cache_dir="/tmp/sana-wids-cache",
    resolution=256,                           # Longest edge after resize

    aspect_ratio_type=str(ASPECT_RATIOS),     # Passed as string, parsed internally

    transform=transform,
    max_length=300,
    caption_selection_type="clipscore",       # Alternative: "proportion"

    external_caption_suffixes=["_alt"],
    external_clipscore_suffixes=["_clip"],
    clip_thr=0.0,
    clip_thr_temperature=1.0,
)

# Multi-process DataLoader is fully supported

loader = DataLoader(
    train_dataset,
    batch_size=8,
    shuffle=True,
    num_workers=4,
    pin_memory=True,
)

# Training loop yields structured tuples

for batch in loader:
    img, txt_fea, attn_mask, data_info, idx, cap_type, shard_info, clipscore = batch
    # Training logic here

Key configuration parameters:

  • aspect_ratio_type: Must be a string representation of the ratio dictionary. The class parses this to determine available resolution buckets.
  • caption_selection_type: Set to "clipscore" for weighted sampling based on external ClipScore JSONs, or "proportion" for fixed probability distribution via caption_proportion.
  • external_caption_suffixes: List of suffixes (e.g., ["_alt"]) to search for additional caption files named <shard>_<suffix>.json.

Advanced Features: External Captions and VAE Caching

The SanaWebDatasetMS pipeline supports sophisticated caption strategies and feature caching to maximize GPU utilization.

Weighted Caption Sampling by ClipScore

When caption_selection_type="clipscore", the dataset loads external JSONs specified by external_clipscore_suffixes (lines 84-106 of sana_data_multi_scale.py). The weighted_sample_clipscore method (lines 158-176 of sana_data.py) converts scores to sampling weights using the transformation weights ** (1/clip_thr_temperature), then probabilistically selects among available caption variants for each sample.

For fixed probability sampling without ClipScore data, the weighted_sample_fix_prob method (lines 57-63 of sana_data.py) applies the proportions defined in caption_proportion.

Pre-computed VAE Features

To bypass on-the-fly VAE encoding during training, set load_vae_feat=True. When enabled, __getitem__ loads .npy files from the TAR archives instead of decoding images (lines 73-81 of sana_data_multi_scale.py). This reduces CPU load and eliminates the VAE forward pass from the training loop, increasing throughput by 2-3x depending on model size.

Performance Optimizations and Caching

The WebDataset implementation includes aggressive caching strategies to minimize I/O bottlenecks.

Shard index caching: ShardListDatasetMulti maintains persistent caches in ~/.cache/_wids_cache, storing the mapping between sample indices and TAR byte offsets. This eliminates metadata regeneration overhead across training restarts.

Automatic meta-detection: When meta_path=None, the dataset automatically locates wids-meta.json files next to TAR archives, falling back to on-the-fly metadata construction if necessary.

Summary

  • SanaWebDatasetMS extends the base WebDataset class with dynamic aspect-ratio bucketing and multi-resolution training support.
  • Configure aspect ratios by passing a dictionary (as a string) to aspect_ratio_type, which the dataset evaluates at line 94 of sana_data_multi_scale.py.
  • Generate wids-meta.json using tools/create_wids_metadata.py to enable fast random access into TAR archives.
  • Use caption_selection_type="clipscore" with external_clipscore_suffixes to implement weighted caption sampling based on aesthetic or alignment scores.
  • Enable load_vae_feat=True and store .npy latents in TARs to eliminate VAE encoding overhead during training.
  • Shard indexes are automatically cached under ~/.cache/_wids_cache for accelerated dataset initialization.

Frequently Asked Questions

What is the difference between SanaWebDataset and SanaWebDatasetMS?

SanaWebDataset (defined in diffusion/data/datasets/sana_data.py) provides the base WebDataset functionality for streaming image-caption pairs from TAR archives. SanaWebDatasetMS (in diffusion/data/datasets/sana_data_multi_scale.py) inherits from this base class and adds multi-scale logic including the get_closest_ratio method for dynamic aspect ratio bucketing and resolution-specific cropping.

How do I configure custom aspect ratios for multi-scale training?

Define your target ratios as a dictionary where keys are width/height floats (e.g., "1.78" for 16:9) and values are configuration placeholders. Pass this dictionary as a string to the aspect_ratio_type parameter when instantiating SanaWebDatasetMS. The dataset evaluates this string and matches each sample to the closest aspect ratio bucket before applying the transform.

Can I use pre-computed VAE latents instead of raw images with WebDataset TAR files?

Yes. Set load_vae_feat=True in your dataset configuration. Your TAR archives must contain .npy files (pre-computed VAE latents) alongside the JSON metadata files. When enabled, lines 73-81 of sana_data_multi_scale.py bypass image decoding and load the latent arrays directly, significantly reducing CPU memory bandwidth and preprocessing time.

How does the weighted caption sampling work with external ClipScore files?

When caption_selection_type="clipscore", the dataset searches for JSON files matching your external_clipscore_suffixes (e.g., shard_clip.json). It loads ClipScore values for each caption variant, applies temperature scaling using clip_thr_temperature, and samples according to the distribution defined in weighted_sample_clipscore at lines 158-176 of sana_data.py. Captions below clip_thr are automatically filtered from the sampling pool.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →