1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
55pub enum NodeSplit {
56 #[default]
58 Refuse,
59 ByNode,
61 None,
63}
64
65#[derive(Debug, Clone)]
67enum SourceSpec {
68 Patterns(Vec<String>),
70 Verbatim(Vec<String>),
72 #[cfg(feature = "yaml")]
74 Yaml(String),
75}
76
77#[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 pub fn shard_shuffle(mut self, bufsize: usize) -> WebDatasetBuilder {
117 self.shard_shuffle = Some(bufsize);
118 self
119 }
120
121 pub fn resampled(mut self, resampled: bool) -> WebDatasetBuilder {
123 self.resampled = resampled;
124 self
125 }
126
127 pub fn cache_dir(mut self, directory: impl Into<PathBuf>) -> WebDatasetBuilder {
132 self.cache_dir = Some(directory.into());
133 self
134 }
135
136 pub fn cache_size(mut self, bytes: u64) -> WebDatasetBuilder {
138 self.cache_size = Some(bytes);
139 self
140 }
141
142 pub fn node_split(mut self, split: NodeSplit) -> WebDatasetBuilder {
144 self.node_split = split;
145 self
146 }
147
148 pub fn worker_split(mut self, split: bool) -> WebDatasetBuilder {
150 self.worker_split = split;
151 self
152 }
153
154 pub fn selection(mut self, selection: Selection) -> WebDatasetBuilder {
156 self.selection = selection;
157 self
158 }
159
160 pub fn empty_check(mut self, check: bool) -> WebDatasetBuilder {
162 self.empty_check = check;
163 self
164 }
165
166 pub fn seed(mut self, seed: u64) -> WebDatasetBuilder {
168 self.seed = Some(seed);
169 self
170 }
171
172 pub fn deterministic(mut self, deterministic: bool) -> WebDatasetBuilder {
174 self.deterministic = deterministic;
175 self
176 }
177
178 pub fn handler(mut self, handler: HandlerRef) -> WebDatasetBuilder {
180 self.handler = handler;
181 self
182 }
183
184 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 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 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 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 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#[derive(Debug, Clone)]
265pub struct WebDataset {
266 pipeline: DataPipeline,
267 handler: HandlerRef,
268 seed: u64,
269}
270
271impl WebDataset {
272 pub fn builder(urls: impl AsRef<str>) -> WebDatasetBuilder {
274 WebDatasetBuilder::new(SourceSpec::Patterns(vec![urls.as_ref().to_string()]))
275 }
276
277 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 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 #[cfg(feature = "yaml")]
289 pub fn builder_from_yaml(spec: impl Into<String>) -> WebDatasetBuilder {
290 WebDatasetBuilder::new(SourceSpec::Yaml(spec.into()))
291 }
292
293 pub fn open(urls: impl AsRef<str>) -> Result<WebDataset> {
298 WebDataset::builder(urls).build()
299 }
300
301 pub fn from_pipeline(pipeline: DataPipeline) -> WebDataset {
303 WebDataset { pipeline, handler: reraise_exception(), seed: 0 }
304 }
305
306 pub fn seed(&self) -> u64 {
308 self.seed
309 }
310
311 pub fn pipeline(&self) -> &DataPipeline {
313 &self.pipeline
314 }
315
316 pub fn into_pipeline(self) -> DataPipeline {
318 self.pipeline
319 }
320
321 pub fn with(mut self, stage: impl Stage + 'static) -> WebDataset {
323 self.pipeline.push(stage);
324 self
325 }
326
327 pub fn shuffle(self, bufsize: usize) -> WebDataset {
329 let seed = self.seed;
330 self.with(Shuffle::new(bufsize).deterministic(seed))
331 }
332
333 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 pub fn decode_basic(self) -> WebDataset {
341 self.decode(Decoder::default())
342 }
343
344 #[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 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 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 pub fn select(self, predicate: impl Fn(&Sample) -> bool + Send + Sync + 'static) -> WebDataset {
378 self.with(SelectStage::new(predicate))
379 }
380
381 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 pub fn slice(self, start: usize, count: Option<usize>) -> WebDataset {
389 self.with(Slice::new(start, count))
390 }
391
392 pub fn with_epoch(mut self, nsamples: usize) -> WebDataset {
399 self.pipeline = self.pipeline.with_epoch(nsamples);
400 self
401 }
402
403 pub fn repeat(mut self, epochs: usize) -> WebDataset {
405 self.pipeline = self.pipeline.repeat(epochs);
406 self
407 }
408
409 pub fn take(mut self, nsamples: usize) -> WebDataset {
411 self.pipeline = self.pipeline.take(nsamples);
412 self
413 }
414
415 pub fn iter(&self) -> SampleStream {
417 self.pipeline.iter()
418 }
419
420 pub fn iter_ok(&self) -> impl Iterator<Item = Sample> + Send {
422 self.pipeline.iter_ok()
423 }
424
425 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 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 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}