1use 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#[derive(Debug)]
32pub struct RoundRobin {
33 datasets: Vec<DataPipeline>,
34 longest: bool,
35}
36
37impl RoundRobin {
38 pub fn new(datasets: Vec<DataPipeline>) -> RoundRobin {
40 RoundRobin { datasets, longest: false }
41 }
42
43 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(sources.remove(next));
67 }
68 None => return None,
69 }
70 }
71 None
72 }))
73 }
74}
75
76#[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 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 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 pub fn longest(mut self) -> RandomMix {
106 self.longest = true;
107 self
108 }
109
110 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
150fn 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}