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:
.pngor.jpg(or.npyfor pre-computed VAE latents) - Side-car metadata:
<sample_key>.jsoncontainingheight,width, andcaptionfields
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 viacaption_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
SanaWebDatasetMSextends 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 ofsana_data_multi_scale.py. - Generate
wids-meta.jsonusingtools/create_wids_metadata.pyto enable fast random access into TAR archives. - Use
caption_selection_type="clipscore"withexternal_clipscore_suffixesto implement weighted caption sampling based on aesthetic or alignment scores. - Enable
load_vae_feat=Trueand store.npylatents in TARs to eliminate VAE encoding overhead during training. - Shard indexes are automatically cached under
~/.cache/_wids_cachefor 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →