Skip to main content

webdataset/
mix.rs

1//! Combining several datasets into one stream.
2//!
3//! Training on a mixture of sources — a large generic corpus plus a small
4//! domain-specific one, say — means interleaving several pipelines.
5//! [`RoundRobin`] takes from each in turn; [`RandomMix`] draws from them at
6//! given probabilities, which is how a small source is up-weighted without
7//! physically duplicating its shards.
8//!
9//! ```
10//! use webdataset::mix::RandomMix;
11//! use webdataset::pipeline::{DataPipeline, Samples};
12//! use webdataset_core::Sample;
13//!
14//! let a = DataPipeline::new().with(Samples::new((0..100).map(|i| Sample::with_key(format!("a{i}")))));
15//! let b = DataPipeline::new().with(Samples::new((0..100).map(|i| Sample::with_key(format!("b{i}")))));
16//!
17//! let mixed = DataPipeline::new().with(RandomMix::new(vec![a, b]).with_weights(&[0.9, 0.1])?);
18//! assert!(mixed.iter().count() > 0);
19//! # Ok::<(), webdataset_core::Error>(())
20//! ```
21
22use std::sync::Mutex;
23
24use rand::prelude::*;
25use rand::rngs::StdRng;
26use webdataset_core::error::{Error, Result};
27
28use crate::pipeline::{DataPipeline, SampleStream, Stage};
29
30/// Takes one sample from each dataset in turn.
31#[derive(Debug)]
32pub struct RoundRobin {
33    datasets: Vec<DataPipeline>,
34    longest: bool,
35}
36
37impl RoundRobin {
38    /// Interleave these datasets, stopping when the first one runs out.
39    pub fn new(datasets: Vec<DataPipeline>) -> RoundRobin {
40        RoundRobin { datasets, longest: false }
41    }
42
43    /// Keep going until every dataset has run out.
44    pub fn longest(mut self) -> RoundRobin {
45        self.longest = true;
46        self
47    }
48}
49
50impl Stage for RoundRobin {
51    fn apply(&self, _input: SampleStream) -> SampleStream {
52        let mut sources: Vec<SampleStream> = self.datasets.iter().map(DataPipeline::iter).collect();
53        let longest = self.longest;
54        let mut next = 0usize;
55
56        Box::new(std::iter::from_fn(move || {
57            while !sources.is_empty() {
58                next %= sources.len();
59                match sources[next].next() {
60                    Some(item) => {
61                        next += 1;
62                        return Some(item);
63                    }
64                    None if longest => {
65                        // Drop the exhausted source; the others continue.
66                        drop(sources.remove(next));
67                    }
68                    None => return None,
69                }
70            }
71            None
72        }))
73    }
74}
75
76/// Draws from several datasets at given probabilities.
77#[derive(Debug)]
78pub struct RandomMix {
79    datasets: Vec<DataPipeline>,
80    weights: Vec<f64>,
81    longest: bool,
82    seed: Option<u64>,
83}
84
85impl RandomMix {
86    /// Mix these datasets with equal weight.
87    pub fn new(datasets: Vec<DataPipeline>) -> RandomMix {
88        let weights = vec![1.0; datasets.len()];
89        RandomMix { datasets, weights, longest: false, seed: None }
90    }
91
92    /// Weight each dataset; weights need not sum to one.
93    pub fn with_weights(mut self, weights: &[f64]) -> Result<RandomMix> {
94        if weights.len() != self.datasets.len() {
95            return Err(Error::value(format!("got {} weights for {} datasets", weights.len(), self.datasets.len())));
96        }
97        if weights.iter().any(|w| *w < 0.0) || weights.iter().sum::<f64>() <= 0.0 {
98            return Err(Error::value("weights must be non-negative and not all zero"));
99        }
100        self.weights = weights.to_vec();
101        Ok(self)
102    }
103
104    /// Keep going until every dataset has run out, rather than stopping at the first.
105    pub fn longest(mut self) -> RandomMix {
106        self.longest = true;
107        self
108    }
109
110    /// Make the draw reproducible.
111    pub fn with_seed(mut self, seed: u64) -> RandomMix {
112        self.seed = Some(seed);
113        self
114    }
115}
116
117impl Stage for RandomMix {
118    fn apply(&self, _input: SampleStream) -> SampleStream {
119        let mut sources: Vec<SampleStream> = self.datasets.iter().map(DataPipeline::iter).collect();
120        let mut weights = self.weights.clone();
121        let longest = self.longest;
122        let rng = Mutex::new(match self.seed {
123            Some(seed) => StdRng::seed_from_u64(seed),
124            None => StdRng::seed_from_u64(rand::rng().random()),
125        });
126
127        Box::new(std::iter::from_fn(move || {
128            while !sources.is_empty() {
129                let index = {
130                    let mut rng = rng.lock().expect("rng lock");
131                    weighted_choice(&weights, rng.random::<f64>())
132                };
133                match sources[index].next() {
134                    Some(item) => return Some(item),
135                    None if longest => {
136                        drop(sources.remove(index));
137                        weights.remove(index);
138                        if weights.iter().sum::<f64>() <= 0.0 {
139                            return None;
140                        }
141                    }
142                    None => return None,
143                }
144            }
145            None
146        }))
147    }
148}
149
150/// Pick an index in proportion to `weights`, given a uniform draw in `0..1`.
151fn weighted_choice(weights: &[f64], uniform: f64) -> usize {
152    let total: f64 = weights.iter().sum();
153    let mut cumulative = 0.0;
154    for (i, weight) in weights.iter().enumerate() {
155        cumulative += weight / total;
156        if uniform < cumulative {
157            return i;
158        }
159    }
160    weights.len() - 1
161}
162
163#[cfg(test)]
164mod tests {
165    use super::*;
166    use crate::pipeline::Samples;
167    use webdataset_core::Sample;
168
169    fn dataset(prefix: &str, n: usize) -> DataPipeline {
170        let prefix = prefix.to_string();
171        DataPipeline::new().with(Samples::new((0..n).map(move |i| Sample::with_key(format!("{prefix}{i}")))))
172    }
173
174    fn keys(pipeline: &DataPipeline) -> Vec<String> {
175        pipeline.iter().map(|s| s.unwrap().key().unwrap().to_string()).collect()
176    }
177
178    #[test]
179    fn round_robin_interleaves() {
180        let mixed = DataPipeline::new().with(RoundRobin::new(vec![dataset("a", 3), dataset("b", 3)]));
181        assert_eq!(keys(&mixed), ["a0", "b0", "a1", "b1", "a2", "b2"]);
182    }
183
184    #[test]
185    fn round_robin_stops_at_the_shortest_by_default() {
186        let mixed = DataPipeline::new().with(RoundRobin::new(vec![dataset("a", 1), dataset("b", 5)]));
187        assert_eq!(keys(&mixed), ["a0", "b0"]);
188    }
189
190    #[test]
191    fn round_robin_can_drain_the_longest() {
192        let mixed = DataPipeline::new().with(RoundRobin::new(vec![dataset("a", 1), dataset("b", 3)]).longest());
193        assert_eq!(keys(&mixed), ["a0", "b0", "b1", "b2"]);
194    }
195
196    #[test]
197    fn random_mix_respects_weights() {
198        let mixed = DataPipeline::new().with(
199            RandomMix::new(vec![dataset("a", 10000), dataset("b", 10000)])
200                .with_weights(&[0.9, 0.1])
201                .unwrap()
202                .with_seed(1),
203        );
204
205        let drawn = keys(&mixed);
206        let from_a = drawn.iter().filter(|k| k.starts_with('a')).count();
207        let ratio = from_a as f64 / drawn.len() as f64;
208        assert!((0.85..0.95).contains(&ratio), "drew {ratio:.2} from a, expected about 0.9");
209    }
210
211    #[test]
212    fn random_mix_is_reproducible() {
213        let build =
214            || DataPipeline::new().with(RandomMix::new(vec![dataset("a", 100), dataset("b", 100)]).with_seed(7));
215        assert_eq!(keys(&build()), keys(&build()));
216    }
217
218    #[test]
219    fn random_mix_rejects_bad_weights() {
220        assert!(RandomMix::new(vec![dataset("a", 1)]).with_weights(&[1.0, 1.0]).is_err());
221        assert!(RandomMix::new(vec![dataset("a", 1)]).with_weights(&[0.0]).is_err());
222        assert!(RandomMix::new(vec![dataset("a", 1)]).with_weights(&[-1.0]).is_err());
223    }
224
225    #[test]
226    fn random_mix_can_drain_the_longest() {
227        let mixed =
228            DataPipeline::new().with(RandomMix::new(vec![dataset("a", 2), dataset("b", 50)]).longest().with_seed(3));
229        assert_eq!(keys(&mixed).len(), 52);
230    }
231
232    #[test]
233    fn weighted_choice_follows_the_cumulative_distribution() {
234        assert_eq!(weighted_choice(&[0.5, 0.5], 0.1), 0);
235        assert_eq!(weighted_choice(&[0.5, 0.5], 0.9), 1);
236        assert_eq!(weighted_choice(&[1.0, 0.0], 0.99), 0);
237        assert_eq!(weighted_choice(&[0.0, 1.0], 0.0), 1);
238    }
239}