Skip to content

Repository files navigation

Speed up model training by fixing data loading

LitData

Transform

✅ Parallelize data processing
✅ Create vector embeddings
✅ Run distributed inference
✅ Scrape websites at scale

Optimize / Stream

✅ Stream raw files with no prep
✅ Stream large cloud datasets
✅ Accelerate training by 20x
✅ Pause and resume data streaming
✅ Use remote data without local loading


PyPIDownloadsLicenseDiscord

Lightning AIQuick startOptimize dataTransform dataModalityFeaturesStream raw filesPaths & cloud URLsBenchmarksTemplatesUsed bySkillsCommunity

Get started

Why LitData?

Speeding up model training involves more than kernel tuning. Data loading frequently slows down training, because datasets are too large to fit on disk, consist of millions of small files, or stream slowly from the cloud.

LitData provides tools to preprocess and optimize datasets into a format that streams efficiently from any cloud or local source. It also includes a map operator for distributed data processing before optimization. This makes data pipelines faster, cloud-agnostic, and can improve training throughput by up to 20×.

Looking for GPUs?

Over 340,000 developers use Lightning Cloud - purpose-built for PyTorch and PyTorch Lightning.

Quick start

First, install LitData:

pip install litdata

Choose your workflow:

🚀 Speed up model training
🚀 Transform datasets

Advanced install

Install all the extras

pip install 'litdata[extras]'

On Linux/macOS, [extras] includes optional uvloop for a faster asyncio event loop used by StreamingRawDataset (stdlib asyncio is the fallback when it is not installed).

AI agent skill (Cursor, Claude Code, …)

Install the LitData expert skill so coding agents know the full API, path resolver, optimize/stream recipes, and internals. Full file map → Skills.

npx skills add Lightning-AI/litData

Source: .claude/skills/litdata/ in this repository (skills CLI).


Speed up model training

Stream datasets directly from cloud storage without local downloads. Choose the approach that fits your workflow:

Option 1: Stream existing files as-is ⚡⚡ — StreamingRawDataset

No optimize step. Point LitData at a folder of images, audio, text, or any files (local or cloud) and train with a normal PyTorch DataLoader. Downloads are fully asynchronous and batched; cloud clients include built-in retries. You receive raw bytes — decode, parse, or transform however you want.

Details → Stream raw files.

fromlitdataimportStreamingRawDatasetfromtorch.utils.dataimportDataLoaderfromPILimportImageimportiodataset=StreamingRawDataset(
"s3://my-bucket/raw-images/", # or gs://, azure://, /teamspace/s3_connections/..., local pathtransform=lambdab: Image.open(io.BytesIO(b)).convert("RGB"), # optional — default is raw bytes
)
loader=DataLoader(dataset, batch_size=32, num_workers=8)
forbatchinloader:
train_step(batch)

Key benefits:

Zero preprocess: No chunking job — use the files you already have.
Raw bytes, your rules: Each sample is file bytes; decode with PIL, torchaudio, json, or any custom logic (transform= optional).
Fully async + batched: Concurrent downloads via asyncio / __getitems__ (not one-file-at-a-time).
Built-in retries: Cloud downloads retry transient failures (adaptive client retries).
Cloud-native: S3 / GCS / Azure / Studio connections; same path resolver as optimized streaming.
Grouped samples: Override setup() to yield image+mask, audio+transcript, etc.
Indexed once:index.json.zstd cached locally and on the bucket for fast restarts.
Upgrade path: When I/O becomes the bottleneck, optimizeStreamingDataset for max throughput.

Option 2: Optimize for maximum performance ⚡⚡⚡

Accelerate model training (20x faster) by optimizing datasets for streaming directly from cloud storage. Work with remote data without local downloads with features like loading data subsets, accessing individual samples, and resumable streaming.

Step 1: Optimize your data (one-time setup)

Transform raw data into optimized chunks for maximum streaming speed. This step formats the dataset for fast loading by writing data in an efficient chunked binary format.

importnumpyasnpimportlitdataaslddefrandom_images(index):
# Replace with your files: Image(path="photo.jpg") or Image(bytes=...).# Wrappers pick the serializer (a caption string is not an image).# quality/format encode JPEG — not uncompressed PIL RAW.array=np.random.randint(0, 256, (32, 32, 3), dtype=np.uint8)
return {
"index": index,
"image": ld.Image(array=array, quality=95, format="jpeg"),
"class": np.random.randint(10),
}
if__name__=="__main__":
# Exactly one of chunk_bytes or chunk_sizeld.optimize(
fn=random_images, # the function applied to each inputinputs=list(range(1000)), # the inputs to the function (here it's a list of numbers)output_dir="fast_data", # optimized data is stored herenum_workers=4, # the number of workers on the same machinechunk_bytes="64MB"# default; see FAQ for larger samples
)

Step 2: Put the data on the cloud

Upload the data to a Lightning Studio (backed by S3) or your own S3 bucket:

aws s3 cp --recursive fast_data s3://my-bucket/fast_data

Step 3: Stream the data during training

Load the data by replacing the PyTorch Dataset and DataLoader with the StreamingDataset and StreamingDataLoader.

importlitdataaslddataset=ld.StreamingDataset(
's3://my-bucket/fast_data',
shuffle=True,
drop_last=True, # important for multi-GPU so every rank sees the same lengthseed=42,
)
# Custom collate function to handle the batch (optional)defcollate_fn(batch):
return {
"image": [sample["image"] forsampleinbatch],
"class": [sample["class"] forsampleinbatch],
}
dataloader=ld.StreamingDataLoader(dataset, batch_size=64, collate_fn=collate_fn)
forsampleindataloader:
img, cls=sample["image"], sample["class"]

Keyed lookup and in-place patches (needs polars and optimize(..., key_fn=...) or build_keys_index):

ld.optimize(fn=fn, inputs=inputs, output_dir="fast_data", chunk_bytes="64MB", key_fn=lambdas: s["id"])
ds=ld.StreamingDataset("fast_data")
sample=ds["entity-id"] # str keyssample=ds.get_by_key(42) # int entity keys; ds[42] is still positionalwithld.dataset_update("fast_data") asupdate: # local directory onlyupdate["entity-id"] = {"id": "entity-id", "x": 1}
update.commit()

mode="append" continues chunk numbering. use_checkpoint=True tries to resume an interrupted optimize. They are not the same.

Key benefits:

Accelerate training: Optimized datasets load 20x faster.
Stream cloud datasets: Work with cloud data without downloading it.
PyTorch-first: Works with PyTorch libraries like PyTorch Lightning, Lightning Fabric, Hugging Face.
Easy collaboration: Share and access datasets in the cloud, streamlining team projects.
Scale across GPUs: Streamed data automatically scales to all GPUs.
Flexible storage: Use S3, GCS, Azure, or your own cloud account for data storage.
Compression: Reduce your data footprint by using advanced compression algorithms.
Run local or cloud: Run on your own machines or auto-scale to 1000s of cloud GPUs with Lightning Studios.
Enterprise security: Self host or process data on your cloud account with Lightning Studios.


Transform datasets

Accelerate data processing tasks (data scraping, image resizing, embedding creation, distributed inference) by parallelizing (map) the work across many machines at once.

Here's an example that resizes and crops a large image dataset:

fromPILimportImageimportlitdataasld# use a local or S3 folderinput_dir="my_large_images"# or "s3://my-bucket/my_large_images"output_dir="my_resized_images"# or "s3://my-bucket/my_resized_images"inputs= [os.path.join(input_dir, f) forfinos.listdir(input_dir)]
# resize the input imagedefresize_image(image_path, output_dir):
output_image_path=os.path.join(output_dir, os.path.basename(image_path))
Image.open(image_path).resize((224, 224)).save(output_image_path)
ld.map(
fn=resize_image,
inputs=inputs,
output_dir="output_dir",
)

Key benefits:

✅ Parallelize processing: Reduce processing time by transforming data across multiple machines simultaneously.
✅ Scale to large data: Increase the size of datasets you can efficiently handle.
✅ Flexible usecases: Resize images, create embeddings, scrape the internet, etc...
✅ Run local or cloud: Run on your own machines or auto-scale to 1000s of cloud GPUs with Lightning Studios.
✅ Enterprise security: Self host or process data on your cloud account with Lightning Studios.


Modality

Wrap each file so a caption is not treated as a path: Text(path=...), Image(path=...), Audio(path=...). Path and raw bytes are stored as-is; array / image / mesh encode.

TypeWriteStream
Text
TextText(path="a.txt")
Text(bytes=utf8)
Text(text="a caption")
text # str
TokensTensor(array=token_ids)
optimize(..., item_loader=TokensLoader())
tokens # Tensor, length block_size — LLM training
Image
ImageImage(path="a.jpg")
Image(bytes=jpeg)
Image(array=hwc, quality=95, format="jpeg")
image.shape # Tensor CHW
JpegJpeg(path="a.jpg")
Jpeg(array=hwc, quality=95)
image.shape # Tensor CHW
JpegArrayJpegArray(images=[Jpeg(path=p) for p in frames])images[0].shape # Tensor CHW
PilPil(path="a.png")
Pil(image=pil_img, mode="RGB")
pil_img.size # PIL.Image
TiffTiff(path="a.tif")
Tiff(array=hw)
array.shape # NumPy
Audio and Video
AudioAudio(path="a.wav")
Audio(bytes=wav)
Audio(array=wave, sampling_rate=16000)
audio["array"]
audio["sampling_rate"]
VideoVideo(path="c.mp4")
Video(bytes=mp4)
Video(array=frames, fps=25)
video.get_frames_at(0)
video.get_frames_in_range(0, 8)
File
FileFile(path="doc.bin")
File(bytes=blob)
sidecar # raw bytes
PdfPdf(path="p.pdf")
Pdf(pdf=pdfplumber_doc)
pdf.pages[0] # Pdfplumber
3D and volume
MeshMesh(path="m.glb")
Mesh(mesh=trimesh_obj, file_type="glb")
mesh.vertices # Trimesh
NiftiNifti(path="v.nii.gz")
Nifti(array=vol, affine=np.eye(4))
nifti.get_fdata() # Nibabel
Array and Graph
Numpynp.load("a.npy")
np.zeros((3, 4, 4))
array # NumPy
TensorTensor(array=torch.randn(3, 4, 4))feat # Tensor — 1-D token ids use TokensLoader under Text
GraphData(x=…, edge_index=…, y=…)
Graph(x=…, edge_index=…, y=…)
Graph(data=pyg_data)
graph.x, graph.edge_index # PyG Data or Graph — PyG graphs
Parquet
Parquetfolder of .parquet files
StreamingDataset(..., item_loader=ParquetLoader())
row["col"] # dict of columns — stream parquet

Examples (path on disk → optimize → batch): examples/modality.


Key Features

Features for optimizing and streaming datasets for model training

✅ Stream raw files as-is (no optimize) — StreamingRawDataset 🔗

StreamingRawDataset streams your existing files from local disk or cloud storage with no conversion step. It is a map-style torch.utils.data.Dataset: use a standard PyTorch DataLoader (not StreamingDataLoader).

You get raw bytes. LitData does not impose a sample schema — open images with PIL, parse JSONL, decode audio, run your own tokenizer, or pass a transform= if you prefer. Grouped items yield list[bytes] (e.g. image + mask).

Downloads are fully asynchronous and batched: when the DataLoader requests a batch, __getitems__ fetches those files concurrently with asyncio.gather. Cloud clients include built-in retries for transient network errors.

Use it when you want to train or prototype on JPEGs, masks, audio, JSONL, etc. immediately. Switch to optimizeStreamingDataset later if you need maximum cloud training throughput.

StreamingRawDatasetStreamingDataset (optimized)
PrepNone — point at a folderOne-time optimizechunk-*.bin + index.json
ItemRaw file bytes (you decide how to decode)Deserialized samples (dict/tensor/…)
I/OFully async, batched downloads + retriesChunk prefetch / cache pipeline
Loadertorch.utils.data.DataLoaderPrefer StreamingDataLoader (shuffle, resume)
Best forInstant start, full control over bytesHighest sustained training I/O

Install (cloud)

pip install "litdata[extra]" s3fs # Amazon S3
pip install "litdata[extra]" gcsfs # Google Cloud Storage# Azure / Studio connections: see Paths & cloud URLs

Quick start

fromtorch.utils.dataimportDataLoaderfromlitdataimportStreamingRawDatasetfromPILimportImageimportiodefto_image(data: bytes):
returnImage.open(io.BytesIO(data)).convert("RGB")
dataset=StreamingRawDataset(
"s3://my-bucket/images/", # also: gs://, azure://, /teamspace/s3_connections/..., local pathtransform=to_image, # optional; default yields raw bytesstorage_options={}, # optional cloud credentials / endpoint
)
loader=DataLoader(dataset, batch_size=32, num_workers=8)
forbatchinloader:
train_step(batch)

Constructor knobs

ArgDefaultPurpose
input_dirrequiredFolder URL/path (same resolver as optimized streaming)
cache_dirLitData default cacheWhere the file index (and optional file cache) live
cache_filesFalseIf True, keep downloaded files on disk under cache_dir (mirror remote layout)
recompute_indexFalseForce re-scan when remote files changed
transformNonefn(bytes) -> Any or fn(list[bytes]) -> Any for grouped items
storage_options{}Cloud client options
indexerFileIndexer()Custom discovery (subclass BaseIndexer)
max_concurrent_downloadsNone (adaptive)Per-worker in-flight downloads. None = size-aware budget (bandwidth; Little’s-law only for medians <~8 MiB) split across workers; single-process capped at 128. An explicit int is used exactly (no silent clamp)
max_prefetch16Per-worker sequential look-ahead after each batch (default on). When num_workers > 1, effective look-ahead is min(max_prefetch, 64 // num_workers) so aggregate stays ~64 items. Pass 0 to disable
prefetch_cache_sizeautoLRU cap for prefetched items (defaults from max_prefetch)
hedge_delay0Seconds before a hedged duplicate GET for a slow download (0 = off, default; opt-in)
range_parallel_threshold0Objects ≥ this many bytes use parallel ranged GETs (0 = whole-object only; opt-in)
item_type"bytes""bytes" buffers in RAM; "path" returns local cache paths (cache_files=True required)

Group related files (setup)

Default: one file = one sample. Override setup to filter or group (image + mask, audio + transcript, …). Return either a list of FileMetadata or a list of groups (list[list[FileMetadata]]).

fromcollectionsimportdefaultdictfromtorch.utils.dataimportDataLoaderfromlitdataimportStreamingRawDatasetfromlitdata.raw.indexerimportFileMetadataclassSegmentationRawDataset(StreamingRawDataset):
defsetup(self, files: list[FileMetadata]) ->list[list[FileMetadata]]:
# Pair img_001.jpg with img_001.png (mask) by stemby_stem: dict[str, dict[str, FileMetadata]] =defaultdict(dict)
forfinfiles:
name=f.path.rsplit("/", 1)[-1]
stem, _, ext=name.rpartition(".")
by_stem[stem][ext.lower()] =fitems= []
forstem, partsinsorted(by_stem.items()):
if"jpg"inpartsand"png"inparts:
items.append([parts["jpg"], parts["png"]])
returnitemsdataset=SegmentationRawDataset(
"s3://bucket/seg/",
transform=lambdapair: (pair[0], pair[1]), # list[bytes]: [image, mask]
)
loader=DataLoader(dataset, batch_size=16, num_workers=4)
forimages, masksinloader:
...

Index caching (index.json.zstd)

First open scans the tree and writes a compressed file list:

  • Local cache under your LitData cache dir (fast restart on the same machine)
  • Remote copy next to the data when possible (e.g. s3://bucket/files/index.json.zstd) so every machine skips the scan
# After adding/removing files on the bucket:dataset=StreamingRawDataset("s3://bucket/files/", recompute_index=True)

Do not confuse this with optimized LitData’s index.json (chunk metadata). Raw indexing only lists files.

How downloads work

  1. DataLoader asks for a batch of indices → __getitems__.
  2. LitData asynchronously downloads those files in parallel (asyncio.gather + adownload_fileobj).
  3. Cloud SDKs apply retries on transient failures (e.g. S3 adaptive retries).
  4. Each item is returned as bytes (or list[bytes] if setup grouped files), then optional transform.

Your training loop stays normal PyTorch — no async/await in user code.

# Default: you own the bytesdataset=StreamingRawDataset("s3://bucket/files/")
raw: bytes=dataset[0]
# e.g. Image.open(io.BytesIO(raw)), json.loads(raw), np.frombuffer(raw), ...

Tips

  • Prefer num_workers > 0 so worker processes overlap async batch downloads with training. Scale workers toward host vCPUs for network-bound JPEG-sized objects — avoid saturating every vCPU.
  • On Linux, after any parent-process dataset I/O, use DataLoader(..., multiprocessing_context="spawn", persistent_workers=True) — default fork can hang S3 clients in workers.
  • Default max_prefetch=16 enables sequential look-ahead per DataLoader worker; shuffled access disables it. Pass 0 to turn off. When num_workers > 1, look-ahead and download concurrency both scale down with worker count so aggregate in-flight work stays bounded.
  • Prefer an s3:// / gs:// URL or /teamspace/s3_connections/... so LitData hits the bucket directly (resolver) — avoid reading through FUSE.
  • Leave range_parallel_threshold=0 (default) for typical JPEGs; raise it only for large objects where parallel ranged GETs help.
  • Best for medium/large files. Tiny objects (≲100 KB) are request-overhead bound — pack with optimizeStreamingDataset when I/O plateaus.

Throughput

On ImageNet val raw over S3 (50 k JPEGs, batch size 64, spawn workers), throughput gains are clearest at low worker counts / notebooks (+20–80% at ≤8 workers). At high workers (≥16), results are roughly parity within run-to-run noise.

workersbeforeafterΔ
0543735+35%
28161475+81%
848415718+18%
16+~6k~6k~parity

Useful knobs: num_workers, max_prefetch (default 16; worker-aware), download_timeout (batch-level hang protection). Ranged parallel downloads stay opt-in (range_parallel_threshold=0).

✅ Stream large cloud datasets 🔗

Use data stored on the cloud without needing to download it all to your computer, saving time and space.

Imagine you're working on a project with a huge amount of data stored online. Instead of waiting hours to download it all, you can start working with the data almost immediately by streaming it.

Once you've optimized the dataset with LitData, stream it as follows:

fromlitdataimportStreamingDataset, StreamingDataLoaderdataset=StreamingDataset('s3://my-bucket/my-data', shuffle=True)
dataloader=StreamingDataLoader(dataset, batch_size=64)
forbatchindataloader:
process(batch) # Replace with your data processing logic

Additionally, you can inject client connection settings for S3 or GCP when initializing your dataset. This is useful for specifying custom endpoints and credentials per dataset.

fromlitdataimportStreamingDataset# boto3 compatible storage options for a custom S3-compatible endpointstorage_options= {
"endpoint_url": "your_endpoint_url",
"aws_access_key_id": "your_access_key_id",
"aws_secret_access_key": "your_secret_access_key",
}
dataset=StreamingDataset('s3://my-bucket/my-data', storage_options=storage_options)

Also, you can specify a custom cache directory when initializing your dataset. This is useful when you want to store the cache in a specific location.

fromlitdataimportStreamingDataset# Initialize the StreamingDataset with the custom cache directorydataset=StreamingDataset('s3://my-bucket/my-data', cache_dir="/path/to/cache")

Any local path, s3:// / gs:// / r2:// / azure:// / hf://, local: network drive, or Lightning /teamspace/... connection works — see Resolve any path or cloud URL.

✅ Optimize images as JPEG (not raw PIL) 🔗

How you return images from optimize controls storage size and streaming speed.

What you returnSerializerResult
litdata.Image(path=...) / Image(array=..., quality=95, format="jpeg")imageCompressed bytes — preferred
litdata.Jpeg(path=...) / Jpeg(array=..., quality=95)jpegJPEG bytes
PIL.JpegImageFile (e.g. PIL.Image.open("x.jpg"))jpegCompressed bytes
Plain PIL.Image / Image.fromarray(...)pilUncompressed pixels — often 10×+ larger

Best practice: wrap with Image / Jpeg at quality ≈ 95, or keep existing .jpg files via Image(path=...). Resize when helpful.

importlitdataaslddefload_image(path):
return {"image": ld.Image(path=path, quality=95, format="jpeg"), "id": path}
if__name__=="__main__":
ld.optimize(fn=load_image, inputs=list_of_paths, output_dir="fast_data", chunk_bytes="64MB", num_workers=8)

Ready-made ImageNet optimize/stream scripts: benchmarks/litdata/ (--write_mode jpeg --quality 90).

✅ Custom serializers 🔗

LitData serializes each pytree leaf with a pluggable registry. Built-ins (tried in order) include: str, bool, int, float, video, audio, image, nifti, mesh, pdf, tifffile, file, pil, jpeg, jpeg_array, bytes, numpy / tensor (and no-header variants), graph, and pickle (fallback).

Prefer typed media wrappers (Audio, Video, Image, Graph, …) so a filepath is not confused with a caption. For images, Image(..., quality=95, format="jpeg") or a JpegImageFile stores JPEG; a plain PIL.Image selects pil (raw pixels). See Optimize images as JPEG.

Pass custom serializers when streaming (and when using the lower-level Cache writer):

fromlitdataimportStreamingDatasetfromlitdata.streaming.serializersimportSerializerclassMyTypeSerializer(Serializer):
defserialize(self, item):
returnitem.to_bytes(), None# (bytes, optional metadata string)defdeserialize(self, data: bytes):
returnMyType.from_bytes(data)
defcan_serialize(self, item) ->bool:
returnisinstance(item, MyType)
dataset=StreamingDataset(
"s3://bucket/data",
serializers={"my_type": MyTypeSerializer()}, # merged on top of built-ins
)

Keys you pass are tried before the defaults (so they win over pickle). optimize() uses the built-in registry based on the Python types your fn returns — prefer typed wrappers / JPEG / numpy / tensor leaves for best results.

✅ Stream PyG graphs 🔗

Store PyTorch GeometricData / HeteroData as packed tensors (to_dict()), not torch.save / pickle. On read, LitData reconstructs with from_dict when torch-geometric is installed. optimize needs a top-level function (spawn).

StreamingDataLoader uses litdata_collate by default: graph samples become a DataBatch (Batch.from_data_list); everything else uses PyTorch default_collate. For follow_batch / exclude_keys, use torch_geometric.loader.DataLoader. Without PyG, graph batches stay a list of Graph.

Homogeneous Data + GCN

importtorchimporttorch.nn.functionalasFfromtorch_geometric.dataimportDatafromtorch_geometric.nnimportGCNConv, global_mean_poolfromlitdataimportStreamingDataLoader, StreamingDataset, optimizedefmake_graph(i: int) ->Data:
n=8+i%5src=torch.randint(0, n, (12,), dtype=torch.long)
dst=torch.randint(0, n, (12,), dtype=torch.long)
returnData(
x=torch.randn(n, 8),
edge_index=torch.stack([src, dst], 0),
y=torch.tensor(i%3),
train_mask=torch.ones(n, dtype=torch.bool),
num_nodes=n,
)
optimize(make_graph, inputs=list(range(1024)), output_dir="graphs", chunk_size=64)
dataset=StreamingDataset("graphs")
sample=dataset[0] # Data when PyG is installed, else Graphloader=StreamingDataLoader(dataset, batch_size=32, shuffle=True)
batch=next(iter(loader)) # DataBatchclassNet(torch.nn.Module):
def__init__(self):
super().__init__()
self.conv=GCNConv(8, 16)
self.lin=torch.nn.Linear(16, 3)
defforward(self, data):
x=F.relu(self.conv(data.x, data.edge_index))
returnself.lin(global_mean_pool(x, data.batch))

Graph wrapper (no PyG at write time)

fromlitdataimportGraph, optimizedefmake_graph(i: int) ->Graph:
n=6returnGraph(
x=torch.randn(n, 4),
edge_index=torch.tensor([[0, 1, 2], [1, 2, 0]], dtype=torch.long),
y=torch.tensor(i%2),
data={"num_nodes": n}, # extra tensors/scalars; field kwargs override data=
)
optimize(make_graph, inputs=list(range(256)), output_dir="graphs")
# later: sample.to_pyg() if the stream returned Graph

Graph(data=pyg_data) uses pyg_data.to_dict(). Do not mix tensor fields with an opaque NetworkX data=.

Heterogeneous HeteroData

fromtorch_geometric.dataimportHeteroDatafromlitdataimportStreamingDataLoader, StreamingDataset, optimizedefmake_hetero(i: int) ->HeteroData:
data=HeteroData()
data["paper"].x=torch.randn(8, 16)
data["author"].x=torch.randn(4, 8)
data["author", "writes", "paper"].edge_index=torch.tensor(
[[0, 1, 2, 3], [0, 2, 4, 6]], dtype=torch.long
)
data.y=torch.tensor(i%3)
returndataoptimize(make_hetero, inputs=list(range(512)), output_dir="hetero", chunk_size=32)
dataset=StreamingDataset("hetero")
sample=dataset[0] # HeteroDataprint(sample["paper"].x.shape, sample["author", "writes", "paper"].edge_index.shape)
loader=StreamingDataLoader(dataset, batch_size=16)
batch=next(iter(loader)) # HeteroDataBatch# batch["paper"].x, batch["paper"].batch, batch["author", "writes", "paper"].edge_index

Graph plus metadata in one sample

defmake_row(i: int) ->dict:
return {"id": i, "graph": make_graph(i)}
optimize(make_row, inputs=list(range(1024)), output_dir="rows")
loader=StreamingDataLoader(StreamingDataset("rows"), batch_size=8)
batch=next(iter(loader))
# batch["id"] is a tensor; batch["graph"] is a DataBatch

Sample subgraphs first, then stream

NeighborLoader needs one in-memory graph. To stream, run the sampler in optimize and store each subgraph as a Data:

# sampler = NeighborSampler(big_graph, num_neighbors=[10, 10])defsample_seed(seed: int) ->Data:
out=sampler.sample_from_nodes(torch.tensor([seed]))
returnData(x=out.x, edge_index=out.edge_index, y=out.y)
optimize(sample_seed, inputs=train_seeds.tolist(), output_dir="subgraphs")

NetworkX (or any non-tensor object) uses Graph(data=nx_graph)graph:pickle. Do not torch.save a graph into the sample.

✅ Stream MosaicML MDS datasets 🔗

If you already have datasets written in MosaicML Streaming MDS (Mosaic Data Shard) format, you can stream them directly with LitData—no re-optimization or conversion required!

LitData's default PyTreeLoader natively understands the MDS binary layout, so you can read existing MDS shards using the familiar StreamingDataset and StreamingDataLoader APIs.

Assumption:

Your dataset directory contains MDS shard files (e.g. shard.00000.mds, ...) along with an index.json describing the shards and their column_sizes/column_names.

Stream the MDS dataset:

importlitdataasld# point to your MDS dataset stored locally or in the cloudmds_dataset_uri="s3://my-bucket/my-mds-data"# or a local path# LitData automatically detects and deserializes the MDS formatdataset=ld.StreamingDataset(mds_dataset_uri)
print("Sample", dataset[0])
dataloader=ld.StreamingDataLoader(dataset, batch_size=4)
forsampleindataloader:
pass

How it works:

  • LitData reads the format field from the dataset config. When it's set to "mds", the item loader uses MDS-aware deserialization (mds_deserialize) that respects the per-column sizes stored in each shard.
  • Fixed-size columns are read directly, while variable-size columns are prefixed with a uint32 length header—exactly as in the MosaicML MDS spec.
  • Each sample is reconstructed into its original Python structure via LitData's data_spec.

Key benefits:

Zero conversion: Reuse existing MDS shards as-is.
Drop-in APIs: Use the same StreamingDataset / StreamingDataLoader you already know.
Cloud-native: Stream MDS shards directly from S3, GCS, or Azure.
Easy migration: Move from MosaicML Streaming to LitData without re-optimizing.

Note: Encrypted data loading is not currently supported for the MDS format.

✅ Stream Hugging Face 🤗 datasets 🔗

To use your favorite Hugging Face dataset with LitData, simply pass its URL to StreamingDataset.

How to get HF dataset URI?
how-to-use-hf-dataset.mov

Prerequisites:

pip install 'litdata[extras]' huggingface_hub
# Optional: faster downloads on high-bandwidth networks
pip install hf_transfer
export HF_HUB_ENABLE_HF_TRANSFER=1

Supported for HF: datasets stored as Parquet only. Gated datasets: set HF_TOKEN.

Stream Hugging Face dataset (auto-index + auto ParquetLoader):

importlitdataasldhf_dataset_uri="hf://datasets/leonardPKU/clevr_cogen_a_train/data"dataset=ld.StreamingDataset(hf_dataset_uri) # indexes on first use; caches index.json locallyprint("Sample", dataset[0]) # dict of columns# With workers on Linux, use spawn (same as other ParquetLoader usage)dataloader=ld.StreamingDataLoader(
dataset, batch_size=4, num_workers=4, multiprocessing_context="spawn"
)
forsampleindataloader:
pass

Unlike local/S3 parquet (stream parquet), hf://automatically indexes (if needed) and selects ParquetLoader.

Indexing the HF dataset (optional, faster cold start)

importlitdataasld# Returns the local cache directory that contains index.jsoncache_dir=ld.index_hf_dataset("hf://datasets/leonardPKU/clevr_cogen_a_train/data")

Or control the index path explicitly:

importlitdataasldfromlitdata.streaming.item_loaderimportParquetLoaderuri="hf://datasets/open-thoughts/OpenThoughts-114k/data"ld.index_parquet_dataset(uri, "hf-index-dir") # writes index under hf-index-dirdataset=ld.StreamingDataset(uri, item_loader=ParquetLoader(), index_path="hf-index-dir")
forbatchinld.StreamingDataLoader(dataset, batch_size=4, multiprocessing_context="spawn"):
pass

See also Stream parquet datasets for ParquetLoader knobs, wildcards, and stream-vs-optimize.

LitData Optimize v/s Parquet

Below is the benchmark for the Imagenet dataset (155 GB), demonstrating that optimizing the dataset using LitData is faster and results in smaller output size compared to raw Parquet files.

OperationSize (GB)Time (seconds)Throughput (images/sec)
LitData Optimize Dataset45283.174000-4700
Parquet Optimize Dataset51465.963600-3900
Index Parquet Dataset (overhead)N/A6N/A
✅ Streams on multi-GPU, multi-node 🔗

Data optimized and loaded with Lightning automatically streams efficiently in distributed training across GPUs or multi-node.

The StreamingDataset and StreamingDataLoader automatically make sure each rank receives the same quantity of varied batches of data, so it works out of the box with your favorite frameworks (PyTorch Lightning, Lightning Fabric, or PyTorch) to do distributed training.

Here you can see an illustration showing how the Streaming Dataset works with multi node / multi gpu under the hood.

fromlitdataimportStreamingDataset, StreamingDataLoader# For the training dataset, don't forget to enable shuffle and drop_last !!! train_dataset=StreamingDataset('s3://my-bucket/my-train-data', shuffle=True, drop_last=True)
train_dataloader=StreamingDataLoader(train_dataset, batch_size=64)
forbatchintrain_dataloader:
process(batch) # Replace with your data processing logicval_dataset=StreamingDataset('s3://my-bucket/my-val-data', shuffle=False, drop_last=False)
val_dataloader=StreamingDataLoader(val_dataset, batch_size=64)
forbatchinval_dataloader:
process(batch) # Replace with your data processing logic

An illustration showing how the Streaming Dataset works with multi node.

✅ Shuffle, seed, and drop_last 🔗

Shuffling is deterministic and designed for distributed training:

  1. Chunks are assigned (and possibly split) across ranks/workers.
  2. Items inside each chunk are permuted.

The permutation depends on seed, the epoch, and chunk metadata — the same settings always yield the same order (required for resumable state_dict).

Object storage (s3://, gs://, …) globally permutes chunks (FullShuffle). Random chunk order is cheap once files are already copied into the local cache.

POSIX-fast (automatic for any local path) mmaps chunks in place. Vast / NFS / Lustre / GPFS (and LITDATA_POSIX_FAST=1) use WindowShuffle: each worker gets whole chunks in a sequential stripe, then shuffles only inside a sliding window (default 16, LITDATA_POSIX_SHUFFLE_WINDOW) for both chunk order and in-chunk items. Local disks (ext4/xfs) keep global FullShuffle. Object URLs stay on FullShuffle. LITDATA_POSIX_FAST=0 disables in-place mmap.

fromlitdataimportStreamingDataset, StreamingDataLoadertrain=StreamingDataset(
"s3://my-bucket/train",
shuffle=True,
drop_last=True, # keep every rank/worker at the same length (default True under DDP)seed=42, # default is 42; keep stable when resuming
)
loader=StreamingDataLoader(train, batch_size=64, num_workers=8)
# shuffle=/drop_last= on the loader override the datasetloader=StreamingDataLoader(train, batch_size=64, shuffle=True, drop_last=True)

Notes

  • Val/test: usually shuffle=False, drop_last=False.
  • If drop_last=False under multi-GPU, LitData warns — collectives can hang when ranks see different lengths.
  • Resume with loader.state_dict() / load_state_dict(). To deliberately ignore checkpointed shuffle settings, set force_override_state_dict=True on the dataset.
✅ FAQ: chunk size & shuffle before optimize 🔗

What chunk_bytes should I use?

Default is 64MB — a good starting point for typical small/medium samples.

When each datapoint is large (e.g. a few MB), prefer a larger chunk (practical range often 256–512MB) so each chunk holds more samples and intra-chunk batch randomization has a bigger pool. Tradeoff: larger chunks take longer to download before they can be used.

This is expert guidance (recommended-range mindset), not a published chunk-size sweep.

Is StreamingDataset shuffle enough if my source data is ordered?

Not always. LitData handles distributed sampling and bucket sampling within chunks automatically (shuffle=True randomizes chunk order and item order inside each chunk). That is not a substitute for a fully shuffled file-level DataLoader when the source has strong structure (same subject/set contiguous, class blocks, etc.).

If ordered data would make chunked sampling problematic and you cannot embed the grouping as the sample unit:

  • Shuffle the list of samples beforeoptimize so chunks mix well, or
  • Use StreamingRawDataset (per-file random access via a standard PyTorch DataLoader with shuffle=True) instead of optimize → StreamingDataset.

FUSE vs LitData (Lightning Studios)

/teamspace/s3_connections (and related mounts) are FUSE — fine for browsing, not for training I/O. Under load they are very slow and can crash. Pass the same path into LitData (StreamingRawDataset / StreamingDataset / optimize): LitData resolves it and talks directly to the bucket (Resolve any path).

Rough ImageNet order-of-magnitude on a Studio (not hard guarantees; right tuning for raw): FUSE hand-read ~600 images/s · StreamingRawDataset ~6–7k · optimized StreamingDataset (64MB chunks) ~11k.

✅ StreamingDataset & StreamingDataLoader knobs 🔗

StreamingDataset

ArgumentDefaultDescription
input_dirrequiredLocal path, cloud URI, Dir, or parquet path (basename wildcards OK)
cache_dirLITDATA_CACHE_DIR or ~/.lightning/chunksWhere chunks are cached
item_loaderfrom indexTokensLoader, ParquetLoader, …
shuffleFalseDeterministic shuffle (see Shuffle)
drop_lastTrue if distributed else FalseEqual length across ranks
seed42Shuffle / subsample RNG
serializersbuilt-insCustom serialize/deserialize map
max_cache_sizeNoneEvict consumed chunks beyond this size. Default: 75% of free disk, leaving ≥50GB when possible. Pin with "100G" / "50GB", a fraction (0.90), or MAX_CACHE_SIZE.
max_pre_download2Chunks each worker may prefetch (raise for throughput; watch disk / RAM)
subsample1.0Fraction of data (0.01) or upsample (2.5)
encryptionNoneFernetEncryption / RSAEncryption / custom
storage_options{}Cloud client options
session_options{}boto3 session options (S3)
index_pathNoneParquet/HF index.json file or directory
force_override_state_dictFalseLocal ctor args override loaded checkpoint
transformNoneCallable or list of callables per sample

Peak disk ≈ num_workers × max_pre_download × mean_chunk_size.

On Vast / NFS / local disk, POSIX-fast is on by default (LITDATA_POSIX_FAST=0 to disable). WILLNEED prefetch and num_workers are capped when they would exceed about half of MemAvailable. Idle hugepages (common on GPU nodes) do not count as available RAM — drop unused nr_hugepages if MemAvailable looks tiny next to MemTotal.

StreamingDataLoader

ArgumentDescription
All usual torch.utils.data.DataLoader kwargsbatch_size, num_workers, collate_fn, pin_memory, …
shuffle / drop_lastForwarded to the streaming dataset
profile_batchesint / True / False — viztracer worker trace (see Profile data loading)
profile_cprofileTrue — stdlib cProfile of the main process + worker 0 (see Profile data loading)
profile_skip_batches / profile_dirWarm-up skip count; output dir for viztracer / cProfile files
multiprocessing_contextUse "spawn" (or "forkserver") with ParquetLoader + num_workers>0 on Linux

Prefer StreamingDataLoader over a plain PyTorch DataLoader for optimized / combined / parallel datasets (resume + correct batch metadata).

✅ Stream from multiple cloud providers 🔗

The StreamingDataset provides support for reading optimized datasets from common cloud storage providers like AWS S3, Google Cloud Storage (GCS), and Azure Blob Storage. Below are examples of how to use StreamingDataset with each cloud provider.

importosimportlitdataasld# Read data from AWS S3 using boto3aws_storage_options={
"aws_access_key_id": os.environ['AWS_ACCESS_KEY_ID'],
"aws_secret_access_key": os.environ['AWS_SECRET_ACCESS_KEY'],
}
# You can also pass the session options. (for boto3 only)aws_session_options= {
"profile_name": os.environ['AWS_PROFILE_NAME'], # Required only for custom profiles"region_name": os.environ['AWS_REGION_NAME'], # Required only for custom regions
}
dataset=ld.StreamingDataset("s3://my-bucket/my-data", storage_options=aws_storage_options, session_options=aws_session_options)
# Read Data from AWS S3 with Unsigned Request using boto3aws_storage_options={
"config": botocore.config.Config(
retries={"max_attempts": 1000, "mode": "adaptive"}, # Configure retries for S3 operationssignature_version=botocore.UNSIGNED, # Use unsigned requests
)
}
dataset=ld.StreamingDataset("s3://my-bucket/my-data", storage_options=aws_storage_options)
aws_storage_options={
"AWS_ACCESS_KEY_ID": os.environ['AWS_ACCESS_KEY_ID'],
"AWS_SECRET_ACCESS_KEY": os.environ['AWS_SECRET_ACCESS_KEY'],
"S3_ENDPOINT_URL": os.environ['AWS_ENDPOINT_URL'], # Required only for custom endpoints
}
dataset=ld.StreamingDataset("s3://my-bucket/my-data", storage_options=aws_storage_options)
dataset=ld.StreamingDataset("s3://my-bucket/my-data", storage_options=aws_storage_options)
# Read data from GCSgcp_storage_options={
"project": os.environ['PROJECT_ID'],
}
dataset=ld.StreamingDataset("gs://my-bucket/my-data", storage_options=gcp_storage_options)
# Read data from Azureazure_storage_options={
"account_url": f"https://{os.environ['AZURE_ACCOUNT_NAME']}.blob.core.windows.net",
"credential": os.environ['AZURE_ACCOUNT_ACCESS_KEY']
}
dataset=ld.StreamingDataset("azure://my-bucket/my-data", storage_options=azure_storage_options)
✅ Pause, resume data streaming 🔗

Stream data during long training, if interrupted, pick up right where you left off without any issues.

LitData provides a stateful Streaming DataLoader e.g. you can pause and resume your training whenever you want.

Info: The Streaming DataLoader was used by Lit-GPT to pretrain LLMs. Restarting from an older checkpoint was critical to get to pretrain the full model due to several failures (network, CUDA Errors, etc..).

importosimporttorchfromlitdataimportStreamingDataset, StreamingDataLoaderdataset=StreamingDataset("s3://my-bucket/my-data", shuffle=True)
dataloader=StreamingDataLoader(dataset, num_workers=os.cpu_count(), batch_size=64)
# Restore the dataLoader state if it existsifos.path.isfile("dataloader_state.pt"):
state_dict=torch.load("dataloader_state.pt")
dataloader.load_state_dict(state_dict)
# Iterate over the dataforbatch_idx, batchinenumerate(dataloader):
# Store the state every 1000 batchesifbatch_idx%1000==0:
torch.save(dataloader.state_dict(), "dataloader_state.pt")

Same seed and shuffle are required. For a StreamingDataset, num_workers and world_size may change: LitData drops a global sample_in_epoch prefix and restripes the rest (never duplicates remaining IDs). For a matching loss curve keep global batch size (world_size * batch_size) constant and DDP ranks in lockstep. num_canonical_nodes (default: first-run world_size) is frozen in the checkpoint. POSIX WindowShuffle resumes whole remaining chunks. CombinedStreamingDataset and ParallelStreamingDataset resume only with the same world_size, num_workers, and batch_size.

dataset=StreamingDataset("s3://my-bucket/my-data", shuffle=True, num_canonical_nodes=8)
✅ Use shared queue for Optimizing 🔗

optimize / map default to a shared per-node queue (keep_data_ordered=False). Work is packed per node, then every worker on that node pulls the next item, so a slow worker does not leave others idle. Set keep_data_ordered=True to keep a static per-worker slice (required for use_checkpoint and align_chunking).

Local output_dir writes chunks in place. Remote inputs and outputs use the streaming downloader (adownload_file / aupload_file, obstore when available).

importnumpyasnpimportlitdataaslddefrandom_images(index):
array=np.random.randint(0, 256, (32, 32, 3), dtype=np.uint8)
return {
"index": index,
"image": ld.Image(array=array, quality=95, format="jpeg"),
"class": np.random.randint(10),
}
if__name__=="__main__":
# The optimize function writes data in an optimized format.ld.optimize(
fn=random_images, # the function applied to each inputinputs=list(range(1000)), # the inputs to the function (here it's a list of numbers)output_dir="fast_data", # optimized data is stored herenum_workers=4, # The number of workers on the same machinechunk_bytes="64MB" , # size of each chunkkeep_data_ordered=False, # default: shared queue (set True to keep input order)
)

Shared queue vs ordered (skewed local files)

scripts/bench/bench_node_queue.py --files 4000 --workers 8 (first 500 files are 1 MiB). On main, unordered optimize sat on a 200s empty-queue timeout after work finished.

TreeModeTimeThroughput
main (old default)keep_data_ordered=True23.7s169 files/s
mainkeep_data_ordered=False223.6s18 files/s
this treekeep_data_ordered=True22.9s175 files/s
this tree (new default)keep_data_ordered=False18.8s213 files/s

Shared-queue before → after: ~12×. New default vs old ordered default: 1.22×.

Local / remote input × output

python scripts/bench/bench_node_queue.py --files 200 --workers 4 --io-matrix (first 50 files are 1 MiB).

TopologyOrderedSharedSpeedup
local → local6.55s2.96s2.21×
remote → local7.44s3.48s2.14×
local → remote10.98s6.51s1.69×
remote → remote11.48s8.55s1.34×

Shared queue balances uneven workers. It does not change later StreamingDataset throughput.

✅ Use a Queue as input for optimizing data 🔗

Sometimes you don’t have a static list of inputs to optimize — instead, you have a stream of data coming in over time. In such cases, you can use a multiprocessing.Queue to feed data into the optimize() function.

  • This is especially useful when you're collecting data from a remote source like a web scraper, socket, or API.

  • You can also use this setup to store replay buffer data during reinforcement learning and later stream it back for training.

frommultiprocessingimportProcess, Queuefromlitdata.processing.data_processorimportALL_DONEimportlitdataasldimporttimedefyield_numbers():
foriinrange(1000):
time.sleep(0.01)
yield (i, i**2)
defdata_producer(q: Queue):
foriteminyield_numbers():
q.put(item)
q.put(ALL_DONE) # Sentinel value to signal completiondeffn(index):
returnindex# Identity function for demoif__name__=="__main__":
q=Queue(maxsize=100)
producer=Process(target=data_producer, args=(q,))
producer.start()
ld.optimize(
fn=fn, # Function to process each itemqueue=q, # 👈 Stream data from this queueoutput_dir="fast_data", # Where to store optimized datanum_workers=2,
chunk_size=100,
mode="overwrite",
)
producer.join()

📌 Note: Using queues to optimize your dataset impacts optimization time, not streaming speed.

Irrespective of number of workers, you only need to put one sentinel value to signal completion.

It'll be handled internally by LitData.

✅ LLM Pre-training 🔗

LitData is highly optimized for LLM pre-training. First, we need to tokenize the entire dataset and then we can consume it.

importjsonfrompathlibimportPathimportzstandardaszstdfromlitdataimportoptimize, TokensLoaderfromtokenizerimportTokenizerfromfunctoolsimportpartial# 1. Define a function to convert the text within the jsonl files into tokensdeftokenize_fn(filepath, tokenizer=None):
withzstd.open(open(filepath, "rb"), "rt", encoding="utf-8") asf:
forrowinf:
text=json.loads(row)["text"]
ifjson.loads(row)["meta"]["redpajama_set_name"] =="RedPajamaGithub":
continue# exclude the GitHub data since it overlaps with starcodertext_ids=tokenizer.encode(text, bos=False, eos=True)
yieldtext_idsif__name__=="__main__":
# 2. Generate the inputs (we are going to optimize all the compressed json files from SlimPajama dataset )input_dir="./slimpajama-raw"inputs= [str(file) forfileinPath(f"{input_dir}/SlimPajama-627B/train").rglob("*.zst")]
# 3. Store the optimized data wherever you want under "/teamspace/datasets" or "/teamspace/s3_connections"outputs=optimize(
fn=partial(tokenize_fn, tokenizer=Tokenizer(f"{input_dir}/checkpoints/Llama-2-7b-hf")), # Note: You can use HF tokenizer or any othersinputs=inputs,
output_dir="./slimpajama-optimized",
chunk_size=(2049*8012),
# This is important to inform LitData that we are encoding contiguous 1D array (tokens). # LitData skips storing metadata for each sample e.g all the tokens are concatenated to form one large tensor.item_loader=TokensLoader(),
)
importosfromlitdataimportStreamingDataset, StreamingDataLoader, TokensLoaderfromtqdmimporttqdm# Increase by one because we need the next word as welldataset=StreamingDataset(
input_dir=f"./slimpajama-optimized/train",
item_loader=TokensLoader(block_size=2048+1),
shuffle=True,
drop_last=True,
)
train_dataloader=StreamingDataLoader(dataset, batch_size=8, pin_memory=True, num_workers=os.cpu_count())
# Iterate over the SlimPajama datasetforbatchintqdm(train_dataloader):
pass
✅ Filter illegal data 🔗

Sometimes, you have bad data that you don't want to include in the optimized dataset. With LitData, yield only the good data sample to include.

fromlitdataimportoptimize, StreamingDatasetdefshould_keep(index) ->bool:
# Replace with your own logicreturnindex%2==0deffn(data):
ifshould_keep(data):
yielddataif__name__=="__main__":
optimize(
fn=fn,
inputs=list(range(1000)),
output_dir="only_even_index_optimized",
chunk_bytes="64MB",
num_workers=1
)
dataset=StreamingDataset("only_even_index_optimized")
data=list(dataset)
print(data)
# [0, 2, 4, 6, 8, 10, ..., 992, 994, 996, 998]

You can even use try/expect.

fromlitdataimportoptimize, StreamingDatasetdeffn(data):
try:
yield1/dataexcept:
passif__name__=="__main__":
optimize(
fn=fn,
inputs=[0, 0, 0, 1, 2, 4, 0],
output_dir="only_defined_ratio_optimized",
chunk_bytes="64MB",
num_workers=1
)
dataset=StreamingDataset("only_defined_ratio_optimized")
data=list(dataset)
# The 0 are filtered out as they raise a division by zero print(data)
# [1.0, 0.5, 0.25] 
✅ Combine datasets 🔗

Mix and match different sets of data to experiment and create better models.

Combine datasets with CombinedStreamingDataset. As an example, this mixture of Slimpajama & StarCoder was used in the TinyLLAMA project to pretrain a 1.1B Llama model on 3 trillion tokens.

fromlitdataimportStreamingDataset, CombinedStreamingDataset, StreamingDataLoader, TokensLoaderfromtqdmimporttqdmimportostrain_datasets= [
StreamingDataset(
input_dir="s3://tinyllama-template/slimpajama/train/",
item_loader=TokensLoader(block_size=2048+1), # Optimized loader for tokens used by LLMsshuffle=True,
drop_last=True,
),
StreamingDataset(
input_dir="s3://tinyllama-template/starcoder/",
item_loader=TokensLoader(block_size=2048+1), # Optimized loader for tokens used by LLMsshuffle=True,
drop_last=True,
),
]
# Mix SlimPajama data and Starcoder data with these proportions:weights= (0.693584, 0.306416)
combined_dataset=CombinedStreamingDataset(
datasets=train_datasets,
seed=42,
weights=weights,
iterate_over_all=False, # required when passing weights (see below)
)
train_dataloader=StreamingDataLoader(combined_dataset, batch_size=8, pin_memory=True, num_workers=os.cpu_count())
# Iterate over the combined datasetsforbatchintqdm(train_dataloader):
pass

iterate_over_all vs weights (important)

ModeBehavior
iterate_over_all=True (default)Iterate until all datasets are exhausted. Do not pass weights — LitData derives them from dataset lengths (raises ValueError if you pass both).
iterate_over_all=FalseStop when any dataset is exhausted. Pass explicit weights for your mixture (e.g. TinyLlama). Length may be None (variable).

Batching Methods (batching_method)

Stratified (default): each batch mixes samples from multiple datasets according to the weights.

combined_dataset=CombinedStreamingDataset(
datasets=[dataset1, dataset2],
batching_method="stratified", # default
)

Per-stream: each batch comes from only one randomly selected dataset (useful when shapes/dtypes differ).

combined_dataset=CombinedStreamingDataset(
datasets=[dataset1, dataset2],
batching_method="per_stream",
)

Other knobs: seed (default 42), force_override_state_dict=True to let local ctor args override a loaded checkpoint.

✅ Parallel streaming 🔗

While CombinedDataset allows to fetch a sample from one of the datasets it wraps at each iteration, ParallelStreamingDataset can be used to fetch a sample from all the wrapped datasets at each iteration:

fromlitdataimportStreamingDataset, ParallelStreamingDataset, StreamingDataLoaderfromtqdmimporttqdmparallel_dataset=ParallelStreamingDataset(
[
StreamingDataset(input_dir="input_dir_1"),
StreamingDataset(input_dir="input_dir_2"),
],
)
dataloader=StreamingDataLoader(parallel_dataset)
forbatch_1, batch_2intqdm(dataloader):
pass

This is useful to generate new data on-the-fly using a sample from each dataset. To do so, provide a transform function to ParallelStreamingDataset:

deftransform(samples: Tuple[Any]):
sample_1, sample_2=samples# as many samples as wrapped datasetsreturnsample_1+sample_2# example transformationparallel_dataset=ParallelStreamingDataset([dset_1, dset_2], transform=transform)
dataloader=StreamingDataLoader(parallel_dataset)
fortransformed_batchintqdm(dataloader):
pass

If the transformation requires random number generation, internal random number generators provided by ParallelStreamingDataset can be used. These are seeded using the current dataset state at the beginning of each epoch, which allows for reproducible and resumable data transformation. To use them, define a transform which takes a dictionary of random number generators as its second argument:

deftransform(samples: Tuple[Any], rngs: Dict[str, Any]):
sample_1, sample_2=samples# as many samples as wrapped datasetsrng=rngs["random"] # "random", "numpy" and "torch" keys availablereturnrng.random() *sample_1+rng.random() *sample_2# example transformationparallel_dataset=ParallelStreamingDataset([dset_1, dset_2], transform=transform)
✅ Cycle datasets 🔗

ParallelStreamingDataset can also be used to cycle a StreamingDataset. This allows to dissociate the epoch length from the number of samples in the dataset.

To do so, set the length option to the desired number of samples to yield per epoch. If length is greater than the number of samples in the dataset, the dataset is cycled. At the beginning of a new epoch, the dataset resumes from where it left off at the end of the previous epoch.

fromlitdataimportStreamingDataset, ParallelStreamingDataset, StreamingDataLoaderfromtqdmimporttqdmdataset=StreamingDataset(input_dir="input_dir")
cycled_dataset=ParallelStreamingDataset([dataset], length=100)
print(len(cycled_dataset))) # 100dataloader=StreamingDataLoader(cycled_dataset)
forbatch, intqdm(dataloader):
pass

You can even set length to float("inf") for an infinite dataset!

✅ Merge datasets 🔗

Merge multiple optimized datasets into one.

importnumpyasnpfromlitdataimportImage, StreamingDataset, merge_datasets, optimizedefrandom_images(index):
array=np.random.randint(0, 256, (32, 32, 3), dtype=np.uint8)
return {
"index": index,
"image": Image(array=array, quality=95, format="jpeg"),
"class": np.random.randint(10),
}
if__name__=="__main__":
out_dirs= ["fast_data_1", "fast_data_2", "fast_data_3", "fast_data_4"] # or ["s3://my-bucket/fast_data_1", etc.]"forout_dirinout_dirs:
optimize(fn=random_images, inputs=list(range(250)), output_dir=out_dir, num_workers=4, chunk_bytes="64MB")
merged_out_dir="merged_fast_data"# or "s3://my-bucket/merged_fast_data"merge_datasets(input_dirs=out_dirs, output_dir=merged_out_dir)
dataset=StreamingDataset(merged_out_dir)
print(len(dataset))
# out: 1000

If you wrote chunks yourself (Cache / BinaryWriter) and only have {rank}.index.json shards, finish the dataset:

fromlitdataimportcomplete_dataset, StreamingDatasetcomplete_dataset("my_chunks") # no-op if index.json already existsStreamingDataset("my_chunks") # also tries this automatically
✅ Transform datasets while Streaming 🔗

Transform datasets on-the-fly while streaming them, allowing for efficient data processing without the need to store intermediate results.

  • You can use the transform argument in StreamingDataset to apply a transformation function or a list of transformation functions to each sample as it is streamed.
# Define a simple transform functiontorch_transform=transforms.Compose([
transforms.Resize((256, 256)), # Resize to 256x256transforms.ToTensor(), # Convert to PyTorch tensor (C x H x W)transforms.Normalize( # Normalize using ImageNet statsmean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]
)
])
deftransform_fn(x, *args, **kwargs):
"""Define your transform function."""returntorch_transform(x) # Apply the transform to the input image# Create dataset with appropriate configurationdataset=StreamingDataset(data_dir, cache_dir=str(cache_dir), shuffle=shuffle, transform=[transform_fn])

Or, you can create a subclass of StreamingDataset and override its transform method to apply custom transformations to each sample.

classStreamingDatasetWithTransform(StreamingDataset):
"""A custom dataset class that inherits from StreamingDataset and applies a transform."""def__init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.torch_transform=transforms.Compose([
transforms.Resize((256, 256)), # Resize to 256x256transforms.ToTensor(), # Convert to PyTorch tensor (C x H x W)transforms.Normalize( # Normalize using ImageNet statsmean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]
)
])
# Define your transform methoddeftransform(self, x, *args, **kwargs):
"""A simple transform function."""returnself.torch_transform(x)
dataset=StreamingDatasetWithTransform(data_dir, cache_dir=str(cache_dir), shuffle=shuffle)
✅ Split datasets for train, val, test 🔗

Split a dataset into train, val, test splits with train_test_split.

fromlitdataimportStreamingDataset, train_test_splitdataset=StreamingDataset("s3://my-bucket/my-data") # data are stored in the cloudprint(len(dataset)) # display the length of your data# out: 100,000train_dataset, val_dataset, test_dataset=train_test_split(dataset, splits=[0.3, 0.2, 0.5])
print(train_dataset)
# out: 30,000print(val_dataset)
# out: 20,000print(test_dataset)
# out: 50,000

Or pick exact indices with StreamingDataset.subset:

train=dataset.subset(range(0, 30_000))
# dataset.subset(slice(0, 1000))
✅ Load a subset of the remote dataset 🔗

Work on a smaller, manageable portion of your data to save time and resources.

fromlitdataimportStreamingDataset, train_test_splitdataset=StreamingDataset("s3://my-bucket/my-data", subsample=0.01) # data are stored in the cloudprint(len(dataset)) # display the length of your data# out: 1000# or a list / slice of global indicessmall=dataset.subset([0, 10, 20])
✅ Upsample from your source datasets 🔗

Use to control the size of one iteration of a StreamingDataset using repeats. Contains floor(N) possibly shuffled copies of the source data, then a subsampling of the remainder.

fromlitdataimportStreamingDatasetdataset=StreamingDataset("s3://my-bucket/my-data", subsample=2.5, shuffle=True)
print(len(dataset)) # display the length of your data# out: 250000
✅ Easily modify optimized cloud datasets 🔗

Add new data to an existing dataset or start fresh if needed, providing flexibility in data management.

LitData optimized datasets are assumed to be immutable. However, you can make the decision to modify them by changing the mode to either append or overwrite.

fromlitdataimportoptimize, StreamingDatasetdefcompress(index):
returnindex, index**2if__name__=="__main__":
# Add some dataoptimize(
fn=compress,
inputs=list(range(100)),
output_dir="./my_optimized_dataset",
chunk_bytes="64MB",
)
# Later on, you add more dataoptimize(
fn=compress,
inputs=list(range(100, 200)),
output_dir="./my_optimized_dataset",
chunk_bytes="64MB",
mode="append",
)
ds=StreamingDataset("./my_optimized_dataset")
assertlen(ds) ==200assertds[:] == [(i, i**2) foriinrange(200)]

The overwrite mode will delete the existing data and start from fresh.

✅ Stream parquet datasets 🔗

Stream existing Parquet files with LitData without converting them to LitData chunks — or convert them when you need LitData’s optimized binary format. Hugging Face parquet datasets are covered in Stream Hugging Face datasets.

Stream vs optimize vs map

GoalUse
Train on parquet as-is (no conversion)index_parquet_datasetStreamingDataset + ParquetLoader
Faster I/O / tokenize / custom sample shapeoptimize(fn) that yields rows from parquet (reduce memory)
Reshard huge parquet files while mappingmap(..., reader=ParquetReader(cache_folder, num_rows=...))

Each sample from ParquetLoader is a dict (column name → value).

Prerequisites

pip install 'litdata[extras]'# includes polars + pyarrow# Cloud listing/index extras as needed:
pip install s3fs # s3://
pip install gcsfs # gs://

Index a parquet directory

importlitdataasldld.index_parquet_dataset(
"s3://my-bucket/my-parquet-data", # local path, s3://, gs://, or hf://cache_dir=None, # see table belowstorage_options={}, # cloud credentials / endpointsnum_workers=4, # parallel metadata reads
)
SchemeWhere index.json is written
Local directoryNext to the files, or under cache_dir if set
s3:// / gs://Uploaded to the bucket at {url}/index.json (needs write access)
hf://Localcache_dir (required for HF indexing via this helper)

Indexing notes

  • Lists top-level.parquet files only (not recursive subfolders).
  • All files must share the same schema.
  • Supported for indexing today: local, s3://, gs://, hf:// (not r2:// / azure:// yet).
  • For HF, prefer index_hf_dataset(uri) (returns a local cache dir) or auto-index via StreamingDataset("hf://...") — see HF section.

Stream with ParquetLoader

If the folder looks like parquet and has no index.json, StreamingDataset now builds the index automatically. You still need ParquetLoader for local/S3/GCS (it must match index.json). hf:// already auto-indexes and selects the loader.

importlitdataasldfromlitdata.streaming.item_loaderimportParquetLoaderuri="s3://my-bucket/my-parquet-data"dataset=ld.StreamingDataset(
uri,
item_loader=ParquetLoader(low_memory=True), # default: row-group streaming# index_path="/path/to/index.json", # optional if index lives elsewhere
)
# Basename wildcards when the path ends with .parquet:# dataset = ld.StreamingDataset("s3://bucket/data/train-*.parquet", item_loader=ParquetLoader())print(dataset[0]) # dict of columns# Linux + num_workers>0: use spawn (Polars + fork deadlocks)dataloader=ld.StreamingDataLoader(
dataset,
batch_size=4,
num_workers=4,
multiprocessing_context="spawn",
)
forbatchindataloader:
pass

ParquetLoader knobs

ArgDefaultMeaning
low_memoryTrueStream by row group (lower RAM). False loads each whole file into memory (warns).
pre_load_chunkFalsePrefetch full DataFrame — only effective when low_memory=False.

Import: from litdata.streaming.item_loader import ParquetLoader (not re-exported at litdata top level).

Reshard parquet for map / optimize

fromlitdataimportmapfromlitdata.processing.readersimportParquetReaderdefprocess(pq_file, output_dir):
# pq_file is a pyarrow.parquet.ParquetFile
...
map(
fn=process,
inputs=list_of_parquet_paths,
output_dir="s3://bucket/out",
reader=ParquetReader(cache_folder="/tmp/pq-shards", num_rows=65536),
)

ParquetReader splits inputs that exceed num_rows into smaller cached files before your fn runs.

✅ Use compression 🔗

Reduce your data footprint by using advanced compression algorithms.

importlitdataaslddefcompress(index):
returnindex, index**2if__name__=="__main__":
# Add some datald.optimize(
fn=compress,
inputs=list(range(100)),
output_dir="./my_optimized_dataset",
chunk_bytes="64MB",
num_workers=1,
compression="zstd"
)

Using zstd, you can achieve high compression ratio like 4.34x for this simple example.

WithoutWith
2.8kb646b
✅ Access samples without full data download 🔗

Look at specific parts of a large dataset without downloading the whole thing or loading it on a local machine.

fromlitdataimportStreamingDatasetdataset=StreamingDataset("s3://my-bucket/my-data") # data are stored in the cloudprint(len(dataset)) # display the length of your dataprint(dataset[42]) # show the 42th element of the dataset
✅ Use any data transforms 🔗

Customize how your data is processed to better fit your needs.

Subclass the StreamingDataset and override its __getitem__ method to add any extra data transformations.

fromlitdataimportStreamingDataset, StreamingDataLoaderimporttorchvision.transforms.v2.functionalasFclassImagenetStreamingDataset(StreamingDataset):
def__getitem__(self, index):
image=super().__getitem__(index)
returnF.resize(image, (224, 224))
dataset=ImagenetStreamingDataset(...)
dataloader=StreamingDataLoader(dataset, batch_size=4)
forbatchindataloader:
print(batch.shape)
# Out: (4, 3, 224, 224)
✅ Profile data loading speed 🔗

StreamingDataLoader can record a viztracer Chrome trace of the DataLoader worker loop so you can see where time goes (fetch, deserialize, collate, IPC).

Prerequisites

pip install viztracer

Profiling requires num_workers >= 1 (raises otherwise). On multi-GPU, only global rank 0 installs the worker profiler.

Usage

fromlitdataimportStreamingDataset, StreamingDataLoaderdataset=StreamingDataset("s3://my-bucket/my-data", shuffle=True, drop_last=True)
loader=StreamingDataLoader(
dataset,
batch_size=64,
num_workers=4,
profile_batches=20, # record this many batches (int), or True for the whole runprofile_skip_batches=5, # warm up / skip cold-start batches before recordingprofile_dir="./profiles", # where to write result.json (default: cwd)
)
forbatchinloader:
train_step(batch)
# after profile_batches (+ skip) complete, worker 0 saves the trace and prints the path
ArgDefaultMeaning
profile_batchesFalseint → stop after that many recorded batches; True → profile until the iterator ends; False → off
profile_skip_batches0Batches to skip before the tracer starts (useful to skip cache cold-start)
profile_dircurrent working directoryDirectory for result.json (overwrites an existing file)

Only worker 0 is instrumented. When an int is used, the tracer wraps fetcher.fetch and stops after profile_skip_batches + profile_batches fetch calls. When True, tracing runs for the lifetime of that worker loop.

View the trace

# Option A — Chrome# open chrome://tracing and load profiles/result.json# Option B — Perfetto (often better for large traces)# open https://ui.perfetto.dev and load the same file

Tips

  • Delete or change profile_dir between runs — LitData removes an existing result.json before starting.
  • Pair with a wiped chunk cache if you care about cold epoch behavior (litdata cache clear).

cProfile (main + worker 0)

Stdlib statistical profile of the parent process (queue wait, unpickle, collate handoff) and worker 0 (fetch, decode, transforms). No extra package.

loader=StreamingDataLoader(
dataset,
batch_size=64,
num_workers=4,
profile_cprofile=True,
profile_dir="./profiles",
)
forbatchinloader:
train_step(batch)
# writes profiles/cprofile_main.prof + cprofile_worker0.prof (and .txt summaries)
python -m pstats profiles/cprofile_worker0.prof
# then: sort tottime stats 30

Do not set profile_cprofile and profile_batches together — both install a sys.setprofile hook. With num_workers=0 only the main file is written. The parent profiler starts after workers spawn so fork does not inherit an active cProfile.

  • For deeper LitData internals (download / read / delete timeline), use enable_tracer() + Litracer instead — see Debug & Profile LitData. That path is complementary: viztracer = DataLoader worker CPU timeline; cProfile = function self/cum time; Litracer = LitData pipeline events.
✅ Reduce memory use for large files 🔗

Handle large data files efficiently without using too much of your computer's memory.

Optimize from parquet (convert into LitData chunks) when you need tokenization or LitData’s binary format. To stream parquet without converting, see Stream parquet datasets.

When processing large parquet files, yield one item at a time to keep memory low:

frompathlibimportPathimportpyarrow.parquetaspqfromlitdataimportoptimizefromtokenizerimportTokenizerfromfunctoolsimportpartial# 1. Define a function to convert the text within the parquet files into tokensdeftokenize_fn(filepath, tokenizer=None):
parquet_file=pq.ParquetFile(filepath)
# Process per batch to reduce RAM usageforbatchinparquet_file.iter_batches(batch_size=8192, columns=["content"]):
fortextinbatch.to_pandas()["content"]:
yieldtokenizer.encode(text, bos=False, eos=True)
# 2. Generate the inputsinput_dir="/teamspace/s3_connections/tinyllama-template"inputs= [str(file) forfileinPath(f"{input_dir}/starcoderdata").rglob("*.parquet")]
# 3. Store the optimized data wherever you want under "/teamspace/datasets" or "/teamspace/s3_connections"outputs=optimize(
fn=partial(tokenize_fn, tokenizer=Tokenizer(f"{input_dir}/checkpoints/Llama-2-7b-hf")), # Note: Use HF tokenizer or any othersinputs=inputs,
output_dir="/teamspace/datasets/starcoderdata",
chunk_size=(2049*8012), # Number of tokens to store by chunks. This is roughly 64MB of tokens per chunk.
)
✅ Limit local cache space 🔗

Control how much disk the local chunk cache may use. Downloaded chunks are deleted after use once the cache exceeds the limit.

Default max_cache_size is None: LitData uses 75% of currently free disk and still leaves ≥50GB free when the volume is large enough for checkpoints. Pass "100G" / "50GB" for a fixed budget, or a float (0.90) for that fraction of currently free space. MAX_CACHE_SIZE overrides the constructor (size or fraction).

Peak disk in flight is roughly:

num_workers × max_pre_download × mean_chunk_size

Keep max_cache_size comfortably above that peak. For remote datasets, async chunk prefetch often raises max_pre_download to ≥4 automatically — see async prefetch & environment variables.

fromlitdataimportStreamingDatasetdataset=StreamingDataset(
"s3://my-bucket/my-data",
max_cache_size="10GB",
max_pre_download=4, # chunks each worker may prefetch (default 2; async may floor to 4)
)
✅ Async chunk prefetch & environment variables 🔗

Async chunk prefetch

LitData can overlap remote chunk downloads with training using asyncio inside each DataLoader worker’s prepare thread. This is not an async DataLoader — your loop stays:

forbatchinStreamingDataLoader(dataset, batch_size=64, num_workers=8):
train_step(batch)
SituationAsync prefetch
Remote dataset (s3://, gs://, …)On by default
Local-only datasetOff by default
LITDATA_ASYNC_CHUNK_PREFETCH=1Force on
LITDATA_ASYNC_CHUNK_PREFETCH=0Force off

When async is on, LitData raises max_pre_download to at least 4 so asyncio.gather has enough in-flight downloads (override with LITDATA_ASYNC_MIN_PRE_DOWNLOAD; set 0 to disable the floor). Concurrent GETs are capped by LITDATA_ASYNC_DOWNLOAD_CONCURRENCY (default 8). The prepare thread drains the prefetch queue up to that gather width whenever a cache slot is free. Peak disk ≈ num_workers × max_pre_download × chunk_size — size max_cache_size accordingly.

The reader does not poll the cache directory for each chunk. After a download finishes, the downloader atomically os.replaces the temp file and sets an in-process Event. Compressed chunks set that Event only after decompress publishes the readable .bin. Other DataLoader workers still fall back to a short filesystem check (Events are per process).

# Debugging download/delete races — force synchronous downloadsexport LITDATA_ASYNC_CHUNK_PREFETCH=0
# Keep max_pre_download=2 even with async enabledexport LITDATA_ASYNC_MIN_PRE_DOWNLOAD=0
# Cap overlapping remote GETs (default 8)export LITDATA_ASYNC_DOWNLOAD_CONCURRENCY=4

Common environment variables

VariableDefaultPurpose
LITDATA_CACHE_DIR~/.lightning/chunksDefault chunk cache directory
LITDATA_ASYNC_CHUNK_PREFETCHon for remote0/1 force async chunk download overlap
LITDATA_ASYNC_MIN_PRE_DOWNLOAD4Floor for max_pre_download when async is on (0 = no floor)
LITDATA_ASYNC_DOWNLOAD_CONCURRENCY8Max in-flight chunk GETs per gather
LITDATA_OBSTORE_STREAM_MIN_CHUNK_MIB8S3 obstore stream chunk size (MiB)
MAX_WAIT_TIME120Seconds to wait for a chunk before error
FORCE_DOWNLOAD_TIME30Seconds before force re-download of a missing chunk
LITDATA_CHECK_UPDATESunset1 enables the PyPI upgrade tip (off by default)
LITDATA_DISABLE_VERSION_CHECKon unless updates enabled1 skips the upgrade tip
HF_TOKENGated Hugging Face datasets
DEBUG_LITDATA / PRINT_DEBUG_LOGS0Internal debug / stdout logs
LITDATA_LOG_FILElitdata_debug.logenable_tracer() output path
LITDATA_TRACE_LEVELunsetbatch / chunk / sample / debug / off (see Debug & Profile)
LITDATA_TRACE_CATEGORIESfrom levelComma-separated cats, e.g. download,read,delete

Multi-node optimize/map on Studios also uses DATA_OPTIMIZER_* (set by the platform). Full catalog (debug logs, Studio injects, torchrun): see the LitData skill reference/env-vars.md when using agent skills, or the source modules constants.py / async_prefetch.py.

✅ Change cache directory path 🔗

Specify where cached chunk files are stored.

fromlitdataimportStreamingDatasetfromlitdata.streaming.cacheimportDir# Simple: dedicated cache directorydataset=StreamingDataset("s3://my-bucket/my_optimized_dataset", cache_dir="/path/to/your/cache")
# Or when cache path and remote URL should differ:dataset=StreamingDataset(input_dir=Dir(path="/path/to/your/cache", url="s3://my-bucket/my_optimized_dataset"))

Global default without passing cache_dir every time:

export LITDATA_CACHE_DIR=/path/to/your/cache

CLI:

litdata cache path # print the active cache directory
litdata cache clear # delete cached chunks
✅ Optimize loading on networked drives 🔗

Optimize data handling for computers on a local network to improve performance for on-site setups.

On-prem compute nodes can mount and use a network drive. A network drive is a shared storage device on a local area network. In order to reduce their network overload, the StreamingDataset supports caching the data chunks.

fromlitdataimportStreamingDatasetdataset=StreamingDataset(input_dir="local:/data/shared-drive/some-data")
✅ Optimize / map across multiple machines (Lightning Studios) 🔗

On Lightning Studios, num_nodes and machine scale optimize / map across many machines. This is not the same as num_workers (processes on one machine).

How it works

  1. You call optimize(..., num_nodes=N, machine=...) (or map) inside a Studio.
  2. LitData starts a data-prep job that re-runs your script on N machines.
  3. Each machine processes a shard of the inputs (num_nodes × num_workers total workers). The last node merges chunk indexes into a single index.json.
  4. Your local call blocks until the job finishes; open the printed Runs URL to monitor.

Outside Studio, passing num_nodes / machine raises an error (create a Studio account to use multi-node).

fromlitdataimportoptimize, Machinedefcompress(index):
return (index, index**2)
if__name__=="__main__":
optimize(
fn=compress,
inputs=list(range(100)),
num_workers=8, # processes per machineoutput_dir="/teamspace/s3_connections/my-data/optimized-v1", # durable bucket (recommended)chunk_bytes="64MB",
num_nodes=32, # machines in the jobmachine=Machine.DATA_PREP, # or omit to use the current Studio machine type
)

Where outputs land

output_dirResult
/teamspace/s3_connections/..., /teamspace/datasets/..., s3://..., gs://...Written directly to that store (recommended)
Local or /teamspace/studios/this_studio/...Remapped to the job’s artifacts storage; the Studio UI may also expose it under /teamspace/jobs/<job>/...
fromlitdataimportStreamingDataset# Prefer the same connection / cloud URL you wrote to:dataset=StreamingDataset("/teamspace/s3_connections/my-data/optimized-v1")

The same num_nodes / machine pattern works with map. See also Parallelize transforms and data optimization.

✅ Encrypt, decrypt data at chunk/sample level 🔗

Encrypt optimized data at sample or chunk level. Built-ins: FernetEncryption and RSAEncryption (litdata.utilities.encryption). Requires the cryptography package. Not supported for Mosaic MDS.

levelMeaning
"sample" (default)Encrypt each sample independently
"chunk"Encrypt whole chunks

Fernet (symmetric)

fromlitdataimportoptimize, StreamingDatasetfromlitdata.utilities.encryptionimportFernetEncryptionfernet=FernetEncryption(password="your_secure_password", level="sample") # or level="chunk"data_dir="s3://my-bucket/optimized_data"deffn(index):
return {"index": index, "value": index**2}
if__name__=="__main__":
optimize(
fn=fn,
inputs=list(range(5)),
num_workers=1,
output_dir=data_dir,
chunk_bytes="64MB",
encryption=fernet,
)
fernet.save("fernet.pem") # persist salt/level; keep the password safe# Later — load key material with the same passwordfernet=FernetEncryption.load("fernet.pem", password="your_secure_password")
ds=StreamingDataset(input_dir=data_dir, encryption=fernet)

RSA (asymmetric)

fromlitdata.utilities.encryptionimportRSAEncryptionrsa=RSAEncryption(password="your_secure_password", level="sample") # or "chunk"optimize(fn=fn, inputs=list(range(5)), output_dir=data_dir, chunk_bytes="64MB", encryption=rsa)
rsa.save("rsa.pem")
rsa=RSAEncryption.load("rsa.pem", password="your_secure_password")
ds=StreamingDataset(input_dir=data_dir, encryption=rsa)

Custom algorithm — subclass Encryption and implement encrypt / decrypt / save / load / state_dict / algorithm.

✅ Debug & Profile LitData with logs & Litracer 🔗

enable_tracer() records the streaming pipeline (download vs read vs delete vs batch) as one-line events. Litracer converts that log into a Chrome / Perfetto trace.

This is complementary to profile_batches (viztracer = DataLoader worker CPU; Litracer = LitData pipeline events).

431247797-0e955e71-2f9a-4aad-b7c1-a8218fed2e2e
importlitdataasldfromlitdata.debuggerimportenable_tracer# Call once per process, before the DataLoader. Delete an existing log before re-tracing (append).enable_tracer(level="chunk", log_file="litdata_debug.log")
# level="batch" | "chunk" (default) | "sample" | "debug" | "off"# enable_tracer(categories=["download", "read", "delete"])if__name__=="__main__":
dataset=ld.StreamingDataset("s3://my-bucket/my-data", shuffle=True)
forbatchinld.StreamingDataLoader(dataset, batch_size=64, num_workers=8):
...
LevelEvents
batchEpoch + per-batch spans, plus crashes
chunk (default)+ download, read, delete, decompress, prefetch
sample+ per-item __getitem__ (high volume)
debug+ .cnt lock refcount spans
offDisable

Event names are stable (download, read, delete, batch, sample, crash). Chunk / sample indexes live in args so Perfetto groups all downloads together. Each line is key: value; pairs with Chrome microsecond timestamps. Crashes are a one-line instant (ph: I, name: crash); the Python traceback is printed to stderr, not the log file (a multi-line logger.exception would break Litracer).

Env overrides: LITDATA_LOG_FILE, LITDATA_TRACE_LEVEL, LITDATA_TRACE_CATEGORIES (comma-separated). Tracer calls are no-ops when tracing is off.

  1. Generate the log:

    python train.py # writes litdata_debug.log
  2. Install Litracer (Go 1.23+):

    git clone https://github.com/Lightning-AI/litracer.git
    cd litracer && go build -o litracer .

    Or go install github.com/deependujha/litracer@latest (published Go module path). Until go.mod is renamed, go install github.com/Lightning-AI/litracer@latest does not work. Release binaries: GitHub Releases.

  3. Convert and open in Perfetto:

    litracer --quiet --validate -o litdata_trace.json.gz litdata_debug.log
    litracer --quiet --cat download,read,delete -o io.json.gz litdata_debug.log
    # open the .json.gz at https://ui.perfetto.dev (preferred) or chrome://tracing

    --quiet prints a one-line summary (per-category durations, unmatched B/E, crashes). --cat keeps only those categories. Matched B/E pairs become complete (ph: X) spans unless --no-complete. Default output is gzip Chrome JSON (.json.gz) — both Perfetto and chrome://tracing open it; pass -o file.json for uncompressed.

  • For trace files > 2GB, see Perfetto large traces.
  • If you connect Perfetto to the RPC server, prefer Chrome over Brave (Brave often does not autodetect the RPC server).

Multi-worker s3://FileNotFoundError after ~120s:num_workers=0 working while num_workers>0 fails usually means the DataLoader parent started obstore (tokio) before fork and worker GETs hung. Current LitData fetches index.json with boto3 so workers can lazy-init obstore; they fall back to boto3 if the parent already started the runtime. On Studio R2 / lightning_storage, the same symptom can be a prefetch-thread crash (data_connection_id / endpoint_url into boto3.Session) — look for [litdata] PrepareChunksThread CRASHED on stderr and a crash instant in the trace.

✅ Resolve any path or cloud URL (local, S3, GCS, R2, Azure, HF, Studio) 🔗

LitData resolves every dataset path you pass to StreamingDataset, StreamingRawDataset, optimize, map, and related APIs. You write one path string; LitData figures out whether to read locally, download from object storage, or (inside Lightning Studios) talk directly to the bucket behind a /teamspace/... mount instead of going through slow FUSE I/O.

Supported URI schemes

SchemeExampleUse when
Local path./data or /data/imagenetFiles on disk
s3://s3://my-bucket/optimizedAWS S3
gs://gs://my-bucket/optimizedGoogle Cloud Storage
r2://r2://my-bucket/optimizedCloudflare R2
azure://azure://container/optimizedAzure Blob Storage
hf://hf://datasets/org/name/dataHugging Face datasets (parquet)
local:local:/mnt/nfs/datasetNetwork / shared drive (LitData still caches chunks locally to reduce NAS load)
fromlitdataimportStreamingDataset, optimize# Same APIs — only the path changesStreamingDataset("s3://my-bucket/fast_data", shuffle=True, drop_last=True)
StreamingDataset("gs://my-bucket/fast_data")
StreamingDataset("r2://my-bucket/fast_data", storage_options={...})
StreamingDataset("azure://my-container/fast_data", storage_options={...})
StreamingDataset("hf://datasets/org/name/data")
StreamingDataset("local:/data/shared-drive/some-data")
StreamingDataset("/var/data/fast_data") # plain local directory

Pass cloud credentials with storage_options (and optional session_options for boto3 profiles/regions). See Stream from multiple cloud providers.

Cache directory vs remote URL

By default LitData caches downloaded chunks under ~/.lightning/chunks (override with cache_dir= or LITDATA_CACHE_DIR). When the cache location and the dataset URL must differ, use Dir:

fromlitdataimportStreamingDatasetfromlitdata.streaming.resolverimportDirdataset=StreamingDataset(
Dir(path="/fast-ssd/cache/run-1", url="s3://my-bucket/fast_data")
)
# Equivalent:dataset=StreamingDataset("s3://my-bucket/fast_data", cache_dir="/fast-ssd/cache/run-1")
export LITDATA_CACHE_DIR=/fast-ssd/cache
litdata cache path # show active cache directory
litdata cache clear # wipe cached chunks

Date/time path templates

Embed a strftime pattern in {...} and LitData expands it to the current time (useful for versioned output_dirs):

# e.g. on 2025-05-05 → ".../run_2025-05-05"optimize(
fn=fn,
inputs=inputs,
output_dir="s3://my-bucket/datasets/run_{%Y-%m-%d}",
chunk_bytes="64MB",
)

Lightning Studio /teamspace/... paths (direct bucket I/O)

In Lightning Studios, data connections appear under /teamspace/.... Prefer these paths in LitData — optimize/map uploads and StreamingDataset downloads use the backing object store URL (and temporary credentials when needed), which is much faster than reading every file through the FUSE mount.

Path prefixWhat LitData does
/teamspace/studios/this_studio/...Local Studio workspace disk (not a cloud URL)
/teamspace/studios/<other_studio>/...Resolves to that Studio’s content bucket (s3:// or gs://)
/teamspace/s3_connections/<name>/...Direct S3 to the connection’s bucket
/teamspace/gcs_connections/<name>/...Direct GCS
/teamspace/s3_folders/<name>/...S3 folder connection
/teamspace/gcs_folders/<name>/...GCS folder connection
/teamspace/lightning_storage/<name>/...Lightning-managed object storage (R2-style)
/teamspace/datasets/...Teamspace datasets mount → project datasets bucket
fromlitdataimportStreamingDataset, StreamingRawDataset, optimize# Stream optimized data from an attached S3 connection (direct bucket download)dataset=StreamingDataset("/teamspace/s3_connections/my-data-1/fast_data", shuffle=True, drop_last=True)
# Stream raw files from a connectionraw=StreamingRawDataset("/teamspace/s3_connections/my-bucket-1/raw")
# Optimize *into* a connection — chunks upload straight to the bucketdefshould_keep(data):
ifdata%2==0:
yielddataif__name__=="__main__":
optimize(
fn=should_keep,
inputs=list(range(1000)),
output_dir="/teamspace/s3_connections/my-data-1/output",
chunk_bytes="64MB",
num_workers=1,
)

Tips

  • Version remote outputs (.../v2, .../run_{%Y-%m-%d}). Optimized datasets are immutable unless you pass mode="append" or mode="overwrite".
  • Outside Studio, use s3:// / gs:// / … with your own credentials — /teamspace/... resolution needs Lightning Studio environment variables.
  • optimize / map with num_nodes launch a Studio job (not local multi-process). Prefer a connection / cloud output_dir; local / this_studio optimize outputs go to job artifacts (UI may show /teamspace/jobs/...). Details: distributed optimization.

Features for transforming datasets

✅ Parallelize data transformations (map) 🔗

Apply the same change to different parts of the dataset at once to save time and effort.

The map operator applies a function over a list of inputs. fn must write into output_dir and return None. Guard with if __name__ == "__main__" when using multiple workers.

importosfromlitdataimportmapfromPILimportImageinput_dir="my_large_images"# or s3://...inputs= [os.path.join(input_dir, f) forfinos.listdir(input_dir)]
defresize_image(image_path, output_dir):
output_image_path=os.path.join(output_dir, os.path.basename(image_path))
Image.open(image_path).resize((224, 224)).save(output_image_path)
if__name__=="__main__":
map(
fn=resize_image,
inputs=inputs,
output_dir="s3://my-bucket/my_resized_images",
num_workers=8,
)

map arguments

ArgumentDefaultDescription
fnrequiredfn(input, output_dir) -> None
inputsrequiredSequence or StreamingDataLoader
output_dirrequiredLocal or cloud path (resolver)
input_dirNoneRoot for remote inputs (background download while processing)
weightsNonePer-input weights to balance workers
num_workersCPU countLocal process workers
fast_dev_runFalseProcess only a few items (True → small default, or an int)
num_nodes / machineNoneScale out on Lightning Studios
num_downloaders / num_uploadersautoI/O concurrency per worker
reorder_filesTruePack by file size for balance; False preserves order
error_when_not_emptyFalseError if output_dir already has files
readerdefaultCustom reader for inputs
batch_sizeNoneGroup inputs into batches for fn
start_methodspawn†Multiprocessing start method (†spawn unless IPython)
optimize_dnsNoneOptimized DNS (Studio / cloud)
storage_options{}Cloud credentials / endpoints
keep_data_orderedFalseShared work queue (faster for uneven workers). True keeps a static per-worker slice. Forced True with use_checkpoint / align_chunking.
optimize arguments reference 🔗

Full knob list for litdata.optimize (see Quick start for the minimal recipe). Exactly one of chunk_bytes or chunk_size. Use if __name__ == "__main__".

ArgumentDefaultDescription
fnrequiredMaps each input → sample (or yield samples / skip bad ones)
inputsNoneSequence or StreamingDataLoader (ignored if queue is set)
queueNonemultiprocessing.Queue of live inputs; send oneALL_DONE when finished
output_dir"optimized_data"Local or cloud (resolver); version remote prefixes
input_dirNoneRemote input root for background download
weightsNonePer-input weights to balance workers
chunk_bytesNoneMax bytes per chunk (e.g. "64MB"; see FAQ for larger samples)
chunk_sizeNoneMax items (or tokens with TokensLoader) per chunk
align_chunkingFalseMatch single-worker chunk boundaries (needs chunk_size; uneven load)
compressionNone"zstd" today
encryptionNoneFernetEncryption / RSAEncryption / custom (encrypt)
num_workersCPU countLocal workers
fast_dev_runFalseSmoke a subset of inputs
num_nodes / machineNoneMulti-node on Lightning Studios
num_downloaders / num_uploadersautoI/O concurrency per worker
reorder_filesTrueSize-based packing; False preserves order
readerdefaultCustom input reader
batch_sizeNoneGroup inputs for fn
modeNone"append" or "overwrite" existing dataset; default treats data as immutable
use_checkpointFalseResume an interrupted optimize from .checkpoints
item_loaderNonee.g. TokensLoader() for contiguous tokens
start_methodspawn†Multiprocessing start method
optimize_dnsNoneOptimized DNS
storage_options{}Cloud credentials / endpoints
keep_data_orderedFalseShared queue among workers. True keeps input order. Forced True with use_checkpoint / align_chunking.
verboseTrueProgress logging

Related features: shared queue, queue input, append/overwrite, compression, TokensLoader / LLM, filter.

✅ Cloud-optimized walk (list files at scale) 🔗

litdata.walk is a threaded, cloud-friendly alternative to os.walk for building large inputs= lists (especially on Lightning Studios). Yields (dirpath, dirnames, filenames) like os.walk, but order is not depth-first.

fromlitdataimportwalk, optimizeinputs= []
forroot, dirs, filesinwalk("/teamspace/s3_connections/my-data/raw", max_workers=32):
fornameinfiles:
ifname.endswith(".jpg"):
inputs.append(f"{root}/{name}")
if__name__=="__main__":
optimize(fn=load_image, inputs=inputs, output_dir="...", chunk_bytes="64MB")

Prints a warning outside Lightning Studio — it is optimized for that environment; elsewhere prefer os.walk or your cloud SDK’s listing API.


Benchmarks

In this section we show benchmarks for speed to optimize a dataset and the resulting streaming speed (Reproduce the benchmark).

Streaming speed

LitData Chunks

Data optimized and streamed with LitData achieves a 20x speed up over non optimized data and 2x speed up over other streaming solutions.

Speed to stream Imagenet 1.2M from AWS S3:

FrameworkImages / sec 1st Epoch (float32)Images / sec 2nd Epoch (float32)Images / sec 1st Epoch (torch16)Images / sec 2nd Epoch (torch16)
LitData5839669262827221
Web Dataset3134392433434424
Mosaic ML2898509928095158
Benchmark details
  • Imagenet-1.2M dataset contains 1,281,167 images.
  • To align with other benchmarks, we measured the streaming speed (images per second) loaded from AWS S3 for several frameworks.

Speed to stream Imagenet 1.2M from other cloud storage providers:

Storage ProviderFrameworkImages / sec 1st Epoch (float32)Images / sec 2nd Epoch (float32)
Cloudflare R2LitData53355630

Speed to stream a synthetic ImageNet-scale set from Vast NFS with POSIX-fast (mmap in place, decode only, no transforms):

WorkersImages / sec
4818.2k
20835.8k

Raw Dataset

Speed to stream raw Imagenet 1.2M from different cloud storage providers:

StorageImages / s (without transform)Images / s (with transform)
AWS S3~6400 +/- 100~3200 +/- 100
Google Cloud Storage~5650 +/- 100~3100 +/- 100

Also see:StreamingRawDataset streams existing files with no optimize step (great default to start). Use StreamingDataset after optimize when you need the highest sustained training throughput.

Time to optimize data

LitData optimizes the Imagenet dataset for fast training 3-5x faster than other frameworks:

Time to optimize 1.2 million ImageNet images (Faster is better):

FrameworkTrain Conversion TimeVal Conversion TimeDataset Size# Files
LitData10:05 min00:30 min143.1 GB2.339
Web Dataset32:36 min01:22 min147.8 GB1.144
Mosaic ML49:49 min01:04 min143.1 GB2.298


Parallelize transforms and data optimization on cloud machines

Lightning

Parallelize data transforms

Transformations with LitData are linearly parallelizable across machines on Lightning Studios (see distributed optimization for how the job launch works).

For example, let's say that it takes 56 hours to embed a dataset on a single A10G machine. With LitData, this can be speed up by adding more machines in parallel

Number of machinesHours
156
228
414
......
640.875
fromlitdataimportmap, Machinemap(
...
num_nodes=32,
machine=Machine.DATA_PREP, # or omit to inherit the Studio machine# Prefer output_dir on /teamspace/s3_connections/... or s3://...
)

Parallelize data optimization

Same Studio job launch as mapnum_nodes machines × num_workers processes; last node merges the index.

fromlitdataimportoptimize, Machineoptimize(
...
num_nodes=32,
machine=Machine.DATA_PREP,
output_dir="/teamspace/s3_connections/my-data/optimized-v1",
)

Example: Process the LAION 400 million image dataset in 2 hours on 32 machines, each with 32 CPUs.


Start from a template

Below are templates for real-world applications of LitData at scale.

Templates: Transform datasets

StudioData typeTime (minutes)MachinesDataset
Download LAION-400MILLION datasetImage & Text12032LAION-400M
Tokenize 2M Swedish Wikipedia ArticlesText74Swedish Wikipedia
Embed English Wikipedia under 5 dollarsText153English Wikipedia

Templates: Optimize + stream data

StudioData typeTime (minutes)MachinesDataset
Benchmark cloud data-loading librariesImage & Label101Imagenet 1M
Optimize GeoSpatial data for model trainingImage & Mask12032Chesapeake Roads Spatial Context
Optimize TinyLlama 1T dataset for trainingText24032SlimPajama & StarCoder
Optimize parquet files for model trainingParquet Files1216Randomly Generated data


Used by

ProjectDescription
PyHealthDeep-learning toolkit for clinical prediction (MIMIC, eICU, OMOP, sleep, CXR). set_task() writes processed samples with LitData; SampleDataset subclasses StreamingDataset so training streams chunked EHR tensors instead of holding the cohort in RAM.
LBSTERProtein and biological-sequence language models from Prescient Design (Genentech). Pre-training and concept-bottleneck models (fitness, embeddings, guided generation) stream large sequence corpora through LitData.
OpenSynthOpen toolkit for synthetic smart-meter / energy time series. Generated or historical meter traces are optimized and streamed for model training.
biomed-multi-viewIBM BiomedSciAI multi-view biomedical models. LitData is used to cache and stream paired modalities during training.
DEMPhenotype and gene-mining pipeline (biodem). Large genomic / trait tables are packed into LitData chunks for repeated training passes.
deeptanGraph multi-task models for multi-omics trait-associated networks. Guide graphs and expression tables are converted to LitData chunks before GNN training.
dataraxData tooling with an optional cloud-streaming extra that uses LitData to read remote datasets without a full local copy.
fasrSpeech ASR framework. The LitData extra streams audio and transcripts for training instead of random-access file lists.

Skills

Coding agents (Cursor, Claude Code, and others) should load the LitData skill instead of guessing the API.

npx skills add Lightning-AI/litData

Useful options: -g (user-global), -a cursor (Cursor only), -y (non-interactive). In this repo the skill already lives at .claude/skills/litdata/. Installer: skills CLI.

Start at SKILL.md, then load reference/using-litdata.md before writing examples.

FileWhen to load
SKILL.mdTriggers, public API, traps
using-litdata.mdOptimize / stream / raw / modality cookbook
streaming.mdRead path, shuffle, resume, serializers
processing.mdoptimize / map orchestration
data-movement.mdDownload / upload / FUSE vs direct I/O
multi-node.mdStudio num_nodes jobs
resolver.mdPaths, URLs, Studio mounts
storage-format.mdChunks, index.json, writer / reader
cache-and-chunk-lifecycle.mdPrefetch and eviction
env-vars.mdLITDATA_* and DATA_OPTIMIZER_*
keyed-lookup.mdkey_fn, dataset_update
debugging.mdenable_tracer, Litracer
benchmarking.mdFair benches
lightning-studio.mdStudio env and credentials
testing.mdPytest / CI
contributing.mdPR / lint path

Offline streaming what-if (not the Python package): simulator/ (litsim).

Community

LitData is a community project accepting contributions - Let's make the world's most advanced AI data processing framework.

💬 Get help on Discord
📋 License: Apache 2.0


Citation

@misc{litdata2023,
author = {Thomas Chaton and Lightning AI},
title = {LitData: Transform datasets at scale. Optimize datasets for fast AI model training.},
year = {2023},
howpublished = {\url{https://github.com/Lightning-AI/litdata}}
}

Papers with LitData

Papers that train or stream with LitData (optimize / StreamingDataset). Scholar hits for “litdata streaming” are often weather lightning data, or mention LitData only as example source code.

PaperVenueHow LitData is used
Towards Interpretable Protein Structure Prediction with Sparse Autoencoders (code)ICLR 2025 GEMoptimize shards ESM-2 embeddings; StreamingDataset streams from S3 for multi-GPU SAE training
TinyLlama: An Open-Source Small Language Model (code)arXiv 20241.1B pretrain on SlimPajama + StarCoder via Lit-GPT’s lightning.data stack (now LitData): CombinedStreamingDataset + TokensLoader

Governance

Maintainers

Emeritus Maintainers

Alumni

About

Speed up model training by fixing data loading.

Resources

Contributing

Stars

611 stars

Watchers

11 watching

Forks

Releases

Packages

Used by

Contributors

Languages