Skip to main content

webdataset/
dataset.rs

1//! The high-level entry point: [`WebDataset`].
2//!
3//! `WebDataset` assembles the standard pipeline — shard list, node and worker
4//! split, shard shuffle, archive reading, sample grouping — and then lets you
5//! append the rest with a fluid interface. It is the Rust counterpart of the
6//! Python class of the same name.
7//!
8//! ```no_run
9//! use webdataset::{Decoder, WebDataset};
10//! use webdataset::filters::{SampleIteratorExt, TupleIteratorExt};
11//!
12//! let dataset = WebDataset::builder("https://host/imagenet-{000000..000146}.tar")
13//!     .shard_shuffle(100)
14//!     .cache_dir("./_cache")
15//!     .build()?
16//!     .shuffle(1000)
17//!     .decode(Decoder::default());
18//!
19//! for batch in dataset.iter().to_tuple(["jpg;png", "cls"]).batched(32, true) {
20//!     let batch = batch?;
21//!     // batch[0] holds the images, batch[1] the labels
22//!     # let _ = batch;
23//! }
24//! # Ok::<(), webdataset_core::Error>(())
25//! ```
26//!
27//! ## Choosing how shards are distributed
28//!
29//! With `resampled(true)` each worker draws shards with replacement, so no
30//! worker ever runs dry and epoch length is whatever you ask for with
31//! [`with_epoch`](WebDataset::with_epoch). Otherwise shards are dealt out
32//! round-robin, which is exact but needs at least as many shards as workers.
33//! Without either, a pipeline run on several nodes refuses to start rather than
34//! silently training every node on the same data.
35
36use std::path::PathBuf;
37use std::sync::Arc;
38
39use webdataset_core::error::{Error, Result};
40use webdataset_core::handlers::{HandlerRef, reraise_exception};
41use webdataset_core::sample::Sample;
42use webdataset_core::value::Value;
43use webdataset_io::cache::FileCache;
44use webdataset_shard::reader::Selection;
45
46use crate::decode::Decoder;
47use crate::loader::DataLoader;
48use crate::pipeline::{DataPipeline, SampleStream, Stage};
49use crate::shardlists::{ResampledShards, SimpleShardList, SingleNodeOnly, SplitByNode, SplitByWorker};
50use crate::sources::{CachingOpener, Opener, ShardsToSamples, StreamingOpener};
51use crate::stages::{CheckEmpty, Decode, MapStage, Rename, SelectStage, Shuffle, Slice};
52
53/// How shards are divided between distributed processes.
54#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
55pub enum NodeSplit {
56    /// Refuse to run on more than one node without an explicit choice.
57    #[default]
58    Refuse,
59    /// Deal shards out round-robin across ranks.
60    ByNode,
61    /// Do nothing; every node sees every shard.
62    None,
63}
64
65/// Where a dataset's shards come from.
66#[derive(Debug, Clone)]
67enum SourceSpec {
68    /// Brace patterns and `::` lists to expand.
69    Patterns(Vec<String>),
70    /// URLs to use exactly as given.
71    Verbatim(Vec<String>),
72    /// A YAML multi-source specification.
73    #[cfg(feature = "yaml")]
74    Yaml(String),
75}
76
77/// Builds a [`WebDataset`].
78#[derive(Debug)]
79pub struct WebDatasetBuilder {
80    source: SourceSpec,
81    shard_shuffle: Option<usize>,
82    resampled: bool,
83    cache_dir: Option<PathBuf>,
84    cache_size: Option<u64>,
85    node_split: NodeSplit,
86    worker_split: bool,
87    selection: Selection,
88    empty_check: bool,
89    seed: Option<u64>,
90    deterministic: bool,
91    handler: HandlerRef,
92}
93
94impl WebDatasetBuilder {
95    fn new(source: SourceSpec) -> WebDatasetBuilder {
96        WebDatasetBuilder {
97            source,
98            shard_shuffle: None,
99            resampled: false,
100            cache_dir: None,
101            cache_size: None,
102            node_split: NodeSplit::Refuse,
103            worker_split: true,
104            selection: Selection::new(),
105            empty_check: true,
106            seed: None,
107            deterministic: false,
108            handler: reraise_exception(),
109        }
110    }
111
112    /// Shuffle the shard list through a buffer of this size each epoch.
113    ///
114    /// Shard-level shuffling is what decorrelates the data; the sample-level
115    /// [`shuffle`](WebDataset::shuffle) then mixes within that window.
116    pub fn shard_shuffle(mut self, bufsize: usize) -> WebDatasetBuilder {
117        self.shard_shuffle = Some(bufsize);
118        self
119    }
120
121    /// Draw shards with replacement instead of partitioning them.
122    pub fn resampled(mut self, resampled: bool) -> WebDatasetBuilder {
123        self.resampled = resampled;
124        self
125    }
126
127    /// Cache shards in this directory.
128    ///
129    /// The directory must already exist, so that a typo does not silently
130    /// scatter gigabytes into an unexpected place.
131    pub fn cache_dir(mut self, directory: impl Into<PathBuf>) -> WebDatasetBuilder {
132        self.cache_dir = Some(directory.into());
133        self
134    }
135
136    /// Evict cached shards once the cache exceeds this many bytes.
137    pub fn cache_size(mut self, bytes: u64) -> WebDatasetBuilder {
138        self.cache_size = Some(bytes);
139        self
140    }
141
142    /// Choose how shards are divided between distributed processes.
143    pub fn node_split(mut self, split: NodeSplit) -> WebDatasetBuilder {
144        self.node_split = split;
145        self
146    }
147
148    /// Whether to divide shards between loader workers. On by default.
149    pub fn worker_split(mut self, split: bool) -> WebDatasetBuilder {
150        self.worker_split = split;
151        self
152    }
153
154    /// Choose or rename the files read out of each shard.
155    pub fn selection(mut self, selection: Selection) -> WebDatasetBuilder {
156        self.selection = selection;
157        self
158    }
159
160    /// Whether to fail when the pipeline yields nothing. On by default.
161    pub fn empty_check(mut self, check: bool) -> WebDatasetBuilder {
162        self.empty_check = check;
163        self
164    }
165
166    /// Seed the shard shuffle and resampling.
167    pub fn seed(mut self, seed: u64) -> WebDatasetBuilder {
168        self.seed = Some(seed);
169        self
170    }
171
172    /// Make shuffling reproducible across runs while still varying per epoch.
173    pub fn deterministic(mut self, deterministic: bool) -> WebDatasetBuilder {
174        self.deterministic = deterministic;
175        self
176    }
177
178    /// Decide what happens when a shard or sample cannot be read.
179    pub fn handler(mut self, handler: HandlerRef) -> WebDatasetBuilder {
180        self.handler = handler;
181        self
182    }
183
184    /// Assemble the pipeline.
185    pub fn build(self) -> Result<WebDataset> {
186        let seed = self
187            .seed
188            .unwrap_or_else(|| std::env::var("WDS_SEED").ok().and_then(|v| v.parse().ok()).unwrap_or_else(rand_seed));
189
190        let mut pipeline = DataPipeline::new();
191
192        // 1. the shard list
193        match &self.source {
194            #[cfg(feature = "yaml")]
195            SourceSpec::Yaml(text) => {
196                let sample = crate::shardlists::MultiShardSample::from_yaml(text)?;
197                sample.set_seed(seed);
198                pipeline.push(sample);
199            }
200            SourceSpec::Patterns(patterns) if self.resampled => {
201                pipeline.push(ResampledShards::new(patterns)?.with_seed(seed).deterministic(self.deterministic));
202            }
203            SourceSpec::Verbatim(urls) if self.resampled => {
204                pipeline.push(ResampledShards::new(urls)?.with_seed(seed).deterministic(self.deterministic));
205            }
206            SourceSpec::Patterns(patterns) => pipeline.push(SimpleShardList::new(patterns)?),
207            SourceSpec::Verbatim(urls) => pipeline.push(SimpleShardList::verbatim(urls.clone())),
208        }
209
210        // 2. distribute shards across nodes and workers
211        match self.node_split {
212            NodeSplit::Refuse if !self.resampled => pipeline.push(SingleNodeOnly),
213            NodeSplit::ByNode => pipeline.push(SplitByNode),
214            _ => {}
215        }
216        if self.worker_split && !self.resampled {
217            pipeline.push(SplitByWorker);
218        }
219
220        // 3. shuffle the shard order
221        if let Some(bufsize) = self.shard_shuffle {
222            let shuffle = Shuffle::new(bufsize);
223            pipeline.push(match self.deterministic {
224                true => shuffle.deterministic(seed),
225                false => shuffle.with_seed(seed),
226            });
227        }
228
229        // 4. read the shards and group their files into samples
230        let opener: Arc<dyn Opener> = match &self.cache_dir {
231            Some(directory) => {
232                if !directory.exists() {
233                    return Err(Error::value(format!("cache directory {} does not exist", directory.display())));
234                }
235                let mut cache = FileCache::new(directory.clone());
236                if let Some(size) = self.cache_size {
237                    cache = cache.with_budget(size, Some(std::time::Duration::from_secs(30)));
238                }
239                Arc::new(CachingOpener::with_cache(Arc::new(cache)))
240            }
241            None => Arc::new(StreamingOpener),
242        };
243        pipeline.push(ShardsToSamples::new(opener).with_selection(self.selection).with_handler(self.handler.clone()));
244
245        if self.empty_check {
246            pipeline.push(CheckEmpty::default());
247        }
248
249        Ok(WebDataset { pipeline, handler: self.handler, seed })
250    }
251}
252
253fn rand_seed() -> u64 {
254    use rand::RngExt;
255    rand::rng().random()
256}
257
258/// A dataset read from WebDataset-format shards.
259///
260/// Build one with [`WebDataset::builder`], then append stages with the methods
261/// below. Anything that changes the item type — batching, projecting to tuples
262/// — happens on [`iter`](WebDataset::iter) using the adapters in
263/// [`filters`](crate::filters).
264#[derive(Debug, Clone)]
265pub struct WebDataset {
266    pipeline: DataPipeline,
267    handler: HandlerRef,
268    seed: u64,
269}
270
271impl WebDataset {
272    /// Start building a dataset from brace patterns or a `::`-separated list.
273    pub fn builder(urls: impl AsRef<str>) -> WebDatasetBuilder {
274        WebDatasetBuilder::new(SourceSpec::Patterns(vec![urls.as_ref().to_string()]))
275    }
276
277    /// Start building a dataset from several patterns.
278    pub fn builder_from<S: AsRef<str>>(patterns: impl IntoIterator<Item = S>) -> WebDatasetBuilder {
279        WebDatasetBuilder::new(SourceSpec::Patterns(patterns.into_iter().map(|p| p.as_ref().to_string()).collect()))
280    }
281
282    /// Start building a dataset from URLs that need no expansion.
283    pub fn builder_verbatim<S: Into<String>>(urls: impl IntoIterator<Item = S>) -> WebDatasetBuilder {
284        WebDatasetBuilder::new(SourceSpec::Verbatim(urls.into_iter().map(Into::into).collect()))
285    }
286
287    /// Start building a dataset from a YAML multi-source specification.
288    #[cfg(feature = "yaml")]
289    pub fn builder_from_yaml(spec: impl Into<String>) -> WebDatasetBuilder {
290        WebDatasetBuilder::new(SourceSpec::Yaml(spec.into()))
291    }
292
293    /// Build a dataset from `urls` with the default settings.
294    ///
295    /// Shards are not shuffled and the dataset refuses to run distributed; use
296    /// [`builder`](WebDataset::builder) to change either.
297    pub fn open(urls: impl AsRef<str>) -> Result<WebDataset> {
298        WebDataset::builder(urls).build()
299    }
300
301    /// Wrap an already assembled pipeline.
302    pub fn from_pipeline(pipeline: DataPipeline) -> WebDataset {
303        WebDataset { pipeline, handler: reraise_exception(), seed: 0 }
304    }
305
306    /// The seed the shard shuffle and resampling were built with.
307    pub fn seed(&self) -> u64 {
308        self.seed
309    }
310
311    /// The underlying pipeline.
312    pub fn pipeline(&self) -> &DataPipeline {
313        &self.pipeline
314    }
315
316    /// Take the underlying pipeline.
317    pub fn into_pipeline(self) -> DataPipeline {
318        self.pipeline
319    }
320
321    /// Append an arbitrary stage.
322    pub fn with(mut self, stage: impl Stage + 'static) -> WebDataset {
323        self.pipeline.push(stage);
324        self
325    }
326
327    /// Shuffle samples through a buffer of `bufsize`.
328    pub fn shuffle(self, bufsize: usize) -> WebDataset {
329        let seed = self.seed;
330        self.with(Shuffle::new(bufsize).deterministic(seed))
331    }
332
333    /// Decode every field of every sample.
334    pub fn decode(self, decoder: Decoder) -> WebDataset {
335        let handler = self.handler.clone();
336        self.with(Decode::new(decoder).with_handler(handler))
337    }
338
339    /// Decode with the default handler chain.
340    pub fn decode_basic(self) -> WebDataset {
341        self.decode(Decoder::default())
342    }
343
344    /// Decode images with the given `imagespec`, plus the default handlers.
345    #[cfg(feature = "image")]
346    pub fn decode_images(self, imagespec: &str) -> Result<WebDataset> {
347        let handler = crate::images::ImageHandler::parse(imagespec)?;
348        Ok(self.decode(Decoder::new(vec![Arc::new(handler)])))
349    }
350
351    /// Apply a function to each sample; returning `None` drops it.
352    pub fn map(self, f: impl Fn(Sample) -> Result<Option<Sample>> + Send + Sync + 'static) -> WebDataset {
353        let handler = self.handler.clone();
354        self.with(MapStage::new(f).with_handler(handler))
355    }
356
357    /// Apply a function to one field of each sample.
358    pub fn map_field(
359        self,
360        field: impl Into<String>,
361        f: impl Fn(Value) -> Result<Value> + Send + Sync + 'static,
362    ) -> WebDataset {
363        let field = field.into();
364        self.map(move |mut sample| {
365            let Some(value) = sample.remove(&field) else {
366                return Err(Error::MissingKey {
367                    wanted: vec![field.clone()],
368                    available: sample.keys().map(str::to_string).collect(),
369                });
370            };
371            sample.insert(field.clone(), f(value)?);
372            Ok(Some(sample))
373        })
374    }
375
376    /// Keep the samples a predicate accepts.
377    pub fn select(self, predicate: impl Fn(&Sample) -> bool + Send + Sync + 'static) -> WebDataset {
378        self.with(SelectStage::new(predicate))
379    }
380
381    /// Rename fields, resolving each source from a `;`-separated alternation.
382    pub fn rename<A: Into<String>, B: Into<String>>(self, renames: impl IntoIterator<Item = (A, B)>) -> WebDataset {
383        let handler = self.handler.clone();
384        self.with(Rename::new(renames).with_handler(handler))
385    }
386
387    /// Take `count` samples starting at `start`.
388    pub fn slice(self, start: usize, count: Option<usize>) -> WebDataset {
389        self.with(Slice::new(start, count))
390    }
391
392    /// Make an epoch exactly `nsamples` long, replaying the source as needed.
393    ///
394    /// Each [`DataLoader`] worker runs its own copy of the pipeline, so the
395    /// limit applies per worker: with four workers, `with_epoch(1000)` yields
396    /// 4000 samples per epoch. Divide by the worker count if you want the
397    /// total to match.
398    pub fn with_epoch(mut self, nsamples: usize) -> WebDataset {
399        self.pipeline = self.pipeline.with_epoch(nsamples);
400        self
401    }
402
403    /// Replay the dataset `epochs` times.
404    pub fn repeat(mut self, epochs: usize) -> WebDataset {
405        self.pipeline = self.pipeline.repeat(epochs);
406        self
407    }
408
409    /// Stop after `nsamples` in total.
410    pub fn take(mut self, nsamples: usize) -> WebDataset {
411        self.pipeline = self.pipeline.take(nsamples);
412        self
413    }
414
415    /// Start an epoch.
416    pub fn iter(&self) -> SampleStream {
417        self.pipeline.iter()
418    }
419
420    /// Start an epoch, discarding failures.
421    pub fn iter_ok(&self) -> impl Iterator<Item = Sample> + Send {
422        self.pipeline.iter_ok()
423    }
424
425    /// Read this dataset with a pool of worker threads.
426    pub fn loader(&self) -> DataLoader {
427        DataLoader::new(self.pipeline.clone())
428    }
429}
430
431impl IntoIterator for &WebDataset {
432    type Item = Result<Sample>;
433    type IntoIter = SampleStream;
434
435    fn into_iter(self) -> SampleStream {
436        self.iter()
437    }
438}
439
440#[cfg(test)]
441mod tests {
442    use super::*;
443    use crate::filters::SampleIteratorExt;
444
445    fn testdata(name: &str) -> String {
446        std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
447            .join("../../testdata")
448            .join(name)
449            .to_string_lossy()
450            .into_owned()
451    }
452
453    fn dataset() -> WebDataset {
454        WebDataset::builder_verbatim([testdata("imagenet-000000.tgz")]).build().unwrap()
455    }
456
457    #[test]
458    fn reads_a_shard_end_to_end() {
459        let samples: Vec<Sample> = dataset().iter().map(|s| s.unwrap()).collect();
460        assert_eq!(samples.len(), 47);
461        assert!(samples[0].contains_key("png"));
462        assert!(samples[0].contains_key("cls"));
463    }
464
465    #[test]
466    fn decodes_and_projects() {
467        let rows: Vec<Vec<Value>> =
468            dataset().decode_basic().iter().to_tuple(["png", "cls"]).map(|r| r.unwrap()).take(4).collect();
469
470        assert_eq!(rows.len(), 4);
471        assert!(rows[0][0].as_bytes().is_some(), "png stays raw without an image handler");
472        assert!(rows[0][1].as_i64().is_some(), "cls decodes to an integer");
473    }
474
475    #[test]
476    fn shuffles_reproducibly_for_a_fixed_seed() {
477        let keys = |seed: u64| -> Vec<String> {
478            WebDataset::builder_verbatim([testdata("imagenet-000000.tgz")])
479                .seed(seed)
480                .deterministic(true)
481                .build()
482                .unwrap()
483                .shuffle(20)
484                .iter()
485                .map(|s| s.unwrap().key().unwrap().to_string())
486                .collect()
487        };
488        assert_eq!(keys(1), keys(1));
489        assert_ne!(keys(1), keys(2));
490    }
491
492    #[test]
493    fn maps_and_selects() {
494        let samples: Vec<Sample> = dataset()
495            .decode_basic()
496            .select(|s| s.get("cls").and_then(Value::as_i64).unwrap_or(-1) >= 0)
497            .map_field("cls", |v| Ok(Value::Int(v.as_i64().unwrap_or(0) + 1000)))
498            .iter()
499            .map(|s| s.unwrap())
500            .collect();
501
502        assert_eq!(samples.len(), 47);
503        assert!(samples.iter().all(|s| s.get("cls").unwrap().as_i64().unwrap() >= 1000));
504    }
505
506    #[test]
507    fn renames_fields() {
508        let sample = dataset().rename([("image", "png;jpg"), ("label", "cls")]).iter().next().unwrap().unwrap();
509        assert!(sample.contains_key("image"));
510        assert!(sample.contains_key("label"));
511        assert!(!sample.contains_key("png"));
512    }
513
514    #[test]
515    fn honours_epoch_length() {
516        let dataset = dataset().with_epoch(100);
517        assert_eq!(dataset.iter().count(), 100, "the source replays to fill the epoch");
518    }
519
520    #[test]
521    fn batches_through_the_iterator_adapters() {
522        let batches: Vec<Sample> = dataset().decode_basic().iter().batched(16, true).map(|b| b.unwrap()).collect();
523        assert_eq!(batches.len(), 3, "47 samples in batches of 16");
524        assert_eq!(batches[0].get("cls").unwrap().as_tensor().unwrap().shape(), &[16]);
525    }
526
527    #[test]
528    fn refuses_a_missing_cache_directory() {
529        let err = WebDataset::builder_verbatim([testdata("sample.tgz")])
530            .cache_dir("/definitely/not/here")
531            .build()
532            .unwrap_err();
533        assert!(err.to_string().contains("does not exist"), "{err}");
534    }
535
536    #[test]
537    fn reports_a_dataset_that_yields_no_samples() {
538        // The check sits directly after shard reading, so it catches a starved
539        // source — the usual cause is having fewer shards than workers.
540        let starved =
541            || WebDataset::builder_verbatim([testdata("sample.tgz")]).selection(Selection::new().select(|_| false));
542
543        let outcome: Vec<_> = starved().build().unwrap().iter().collect();
544        assert!(matches!(outcome.as_slice(), [Err(Error::Empty(_))]), "{outcome:?}");
545
546        assert_eq!(starved().empty_check(false).build().unwrap().iter().count(), 0);
547    }
548
549    #[test]
550    fn filtering_everything_out_downstream_is_not_an_error() {
551        // The check runs before user stages, so a selective pipeline that keeps
552        // nothing is simply empty rather than a failure.
553        let dataset = WebDataset::builder_verbatim([testdata("sample.tgz")]).build().unwrap().select(|_| false);
554        assert_eq!(dataset.iter().count(), 0);
555    }
556
557    #[test]
558    fn resamples_shards_endlessly() {
559        let dataset =
560            WebDataset::builder_verbatim([testdata("sample.tgz")]).resampled(true).build().unwrap().with_epoch(500);
561        assert_eq!(dataset.iter().count(), 500);
562    }
563
564    #[test]
565    fn runs_over_a_worker_pool() {
566        let dataset =
567            WebDataset::builder_verbatim([testdata("sample.tgz"), testdata("imagenet-000000.tgz")]).build().unwrap();
568        let single = dataset.iter().count();
569        let pooled = dataset.loader().with_workers(2).iter().count();
570        assert_eq!(pooled, single, "splitting by worker must not change the sample count");
571    }
572
573    #[cfg(feature = "image")]
574    #[test]
575    fn decodes_images() {
576        let sample = dataset().decode_images("rgb8").unwrap().iter().next().unwrap().unwrap();
577        let image = sample.get("png").unwrap().as_tensor().expect("png should decode to a tensor");
578        assert_eq!(image.shape().len(), 3);
579        assert_eq!(image.shape()[2], 3, "rgb8 produces three channels");
580    }
581}