1use rand::{SeedableRng, rngs::StdRng, seq::SliceRandom};
3use serde::{Deserialize, Serialize};
4use std::{collections::HashSet, error::Error, fmt};
5use crate::record::{PrecisionSettings, Record};
6use crate::tensor::backend::Backend;
7
8#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
10pub enum SampleTail {
11 Uneven,
13 Drop,
15}
16
17#[derive(Clone, Debug, PartialEq, Eq)]
19pub struct SamplerError(pub String);
20
21impl fmt::Display for SamplerError {
22 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { f.write_str(&self.0) }
23}
24impl Error for SamplerError {}
25
26fn invalid(message: impl Into<String>) -> SamplerError { SamplerError(message.into()) }
27fn quota(length: usize, rank: usize, world: usize) -> usize {
28 length / world + usize::from(rank < length % world)
29}
30fn check_indices(indices: &[usize], length: usize) -> Result<(), SamplerError> {
31 let mut seen = HashSet::with_capacity(indices.len());
32 if indices.iter().any(|&index| index >= length || !seen.insert(index)) {
33 return Err(invalid("epoch indices repeat or lie outside the immutable dataset"));
34 }
35 Ok(())
36}
37
38#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
44pub struct SamplerState {
45 pub version: u32,
47 pub source_id: String,
49 pub source_length: usize,
51 pub rank: usize,
53 pub world_size: usize,
55 pub epoch: u64,
57 pub shuffle_seed: Option<u64>,
59 pub tail: SampleTail,
61 pub epoch_size: usize,
63 pub carried_consumed: usize,
65 pub scope_size: usize,
67 pub indices: Vec<usize>,
69 pub committed: usize,
71}
72
73impl SamplerState {
74 pub fn validate(&self) -> Result<(), SamplerError> {
76 if self.version != 1 || self.source_id.is_empty() || self.world_size == 0 || self.rank >= self.world_size {
77 return Err(invalid("invalid sampler record identity or data topology"));
78 }
79 if self.epoch_size > self.source_length || self.carried_consumed.checked_add(self.scope_size) != Some(self.epoch_size)
80 || self.indices.len() != quota(self.scope_size, self.rank, self.world_size) || self.committed > self.indices.len() {
81 return Err(invalid("sampler record scope/order/consumption counts differ"));
82 }
83 check_indices(&self.indices, self.source_length)
84 }
85}
86
87impl<B: Backend> Record<B> for SamplerState {
88 type Item<P: PrecisionSettings> = Self;
89 fn into_item<P: PrecisionSettings>(self) -> Self { self }
90 fn from_item<P: PrecisionSettings>(item: Self, _: &B::Device) -> Self { item }
91}
92
93#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
95pub struct ConsolidatedSamplerState {
96 pub version: u32,
98 pub source_id: String,
100 pub source_length: usize,
102 pub epoch: u64,
104 pub shuffle_seed: Option<u64>,
106 pub tail: SampleTail,
108 pub epoch_size: usize,
110 pub consumed: usize,
112 pub remaining: Vec<usize>,
114}
115
116impl ConsolidatedSamplerState {
117 pub fn validate(&self) -> Result<(), SamplerError> {
119 if self.version != 1 || self.source_id.is_empty() || self.epoch_size > self.source_length
120 || self.consumed.checked_add(self.remaining.len()) != Some(self.epoch_size) {
121 return Err(invalid("invalid consolidated sampler identity or accounting"));
122 }
123 check_indices(&self.remaining, self.source_length)
124 }
125
126 pub fn consolidate(records: &[SamplerState]) -> Result<Self, SamplerError> {
129 let first = records.first().ok_or_else(|| invalid("consolidation needs every data rank"))?;
130 first.validate()?;
131 let mut ranks = vec![None; first.world_size];
132 for state in records {
133 state.validate()?;
134 if (state.source_id.as_str(), state.source_length, state.world_size, state.epoch, state.shuffle_seed,
135 state.tail, state.epoch_size, state.carried_consumed, state.scope_size)
136 != (first.source_id.as_str(), first.source_length, first.world_size, first.epoch, first.shuffle_seed,
137 first.tail, first.epoch_size, first.carried_consumed, first.scope_size) {
138 return Err(invalid("sampler source, epoch, policy or scope differs between data ranks"));
139 }
140 match ranks[state.rank] {
141 Some(old) if old != state => return Err(invalid("duplicate sampler coordinate has different committed state")),
142 _ => ranks[state.rank] = Some(state),
143 }
144 }
145 if ranks.iter().any(Option::is_none) { return Err(invalid("sampler consolidation is missing a data coordinate")); }
146 let ranks: Vec<&SamplerState> = ranks.into_iter().map(Option::unwrap).collect();
147 let mut seen = HashSet::with_capacity(first.scope_size);
148 let mut remaining = Vec::new();
149 for position in 0..first.scope_size {
150 let state = ranks[position % first.world_size];
151 let local_position = position / first.world_size;
152 let index = state.indices[local_position];
153 if !seen.insert(index) { return Err(invalid("data shards contain the same source sample")); }
154 if local_position >= state.committed { remaining.push(index); }
155 }
156 let consumed = first.carried_consumed + ranks.iter().map(|state| state.committed).sum::<usize>();
157 let result = Self { version: 1, source_id: first.source_id.clone(), source_length: first.source_length,
158 epoch: first.epoch, shuffle_seed: first.shuffle_seed, tail: first.tail, epoch_size: first.epoch_size,
159 consumed, remaining };
160 result.validate()?;
161 Ok(result)
162 }
163}
164
165impl<B: Backend> Record<B> for ConsolidatedSamplerState {
166 type Item<P: PrecisionSettings> = Self;
167 fn into_item<P: PrecisionSettings>(self) -> Self { self }
168 fn from_item<P: PrecisionSettings>(item: Self, _: &B::Device) -> Self { item }
169}
170
171pub struct StatefulShardSampler {
177 state: SamplerState,
178 issued: usize,
179}
180
181impl StatefulShardSampler {
182 pub fn new(source_id: String, source_length: usize, rank: usize, world_size: usize,
184 epoch: u64, shuffle_seed: Option<u64>, tail: SampleTail) -> Result<Self, SamplerError> {
185 if source_id.is_empty() || world_size == 0 || rank >= world_size {
186 return Err(invalid("supply dataset identity and a valid data rank/world size"));
187 }
188 let mut order: Vec<usize> = (0..source_length).collect();
189 if let Some(seed) = shuffle_seed { order.shuffle(&mut StdRng::seed_from_u64(seed.wrapping_add(epoch))); }
190 let epoch_size = match tail { SampleTail::Uneven => source_length, SampleTail::Drop => source_length / world_size * world_size };
191 order.truncate(epoch_size);
192 let indices = order.into_iter().skip(rank).step_by(world_size).collect();
193 let state = SamplerState { version: 1, source_id, source_length, rank, world_size, epoch, shuffle_seed,
194 tail, epoch_size, carried_consumed: 0, scope_size: epoch_size, indices, committed: 0 };
195 state.validate()?;
196 Ok(Self { state, issued: 0 })
197 }
198
199 pub fn from_state(state: SamplerState) -> Result<Self, SamplerError> {
201 state.validate()?;
202 let issued = state.committed;
203 Ok(Self { state, issued })
204 }
205
206 pub fn load_state(&mut self, state: SamplerState) -> Result<(), SamplerError> {
208 state.validate()?;
209 if (state.source_id.as_str(), state.source_length, state.rank, state.world_size, state.shuffle_seed, state.tail)
210 != (self.state.source_id.as_str(), self.state.source_length, self.state.rank, self.state.world_size,
211 self.state.shuffle_seed, self.state.tail) {
212 return Err(invalid("sampler source, configuration or topology changed; explicitly consolidate and reshard"));
213 }
214 self.issued = state.committed;
215 self.state = state;
216 Ok(())
217 }
218
219 pub fn reshard(state: &ConsolidatedSamplerState, rank: usize, world_size: usize) -> Result<Self, SamplerError> {
222 state.validate()?;
223 if world_size == 0 || rank >= world_size { return Err(invalid("invalid target data coordinate")); }
224 Self::from_state(SamplerState { version: 1, source_id: state.source_id.clone(), source_length: state.source_length,
225 rank, world_size, epoch: state.epoch, shuffle_seed: state.shuffle_seed, tail: state.tail,
226 epoch_size: state.epoch_size, carried_consumed: state.consumed, scope_size: state.remaining.len(),
227 indices: state.remaining.iter().copied().skip(rank).step_by(world_size).collect(), committed: 0 })
228 }
229
230 pub fn state(&self) -> SamplerState { self.state.clone() }
232
233 pub fn source_length(&self) -> usize { self.state.source_length }
235
236 pub fn committed(&self) -> usize { self.state.committed }
238
239 pub fn issued(&self) -> usize { self.issued }
241
242 pub fn remaining(&self) -> usize { self.state.indices.len() - self.state.committed }
244
245 pub fn unissued_indices(&self) -> &[usize] { &self.state.indices[self.issued..] }
247
248 pub fn commit(&mut self, count: usize) -> Result<(), SamplerError> {
250 let committed = self.state.committed.checked_add(count).ok_or_else(|| invalid("committed sample count overflow"))?;
251 if committed > self.issued { return Err(invalid("cannot commit samples that have not been issued")); }
252 self.state.committed = committed;
253 Ok(())
254 }
255
256 pub fn rewind_uncommitted(&mut self) { self.issued = self.state.committed; }
258
259 pub fn start_epoch(&mut self, epoch: u64) -> Result<(), SamplerError> {
261 *self = Self::new(self.state.source_id.clone(), self.state.source_length, self.state.rank,
262 self.state.world_size, epoch, self.state.shuffle_seed, self.state.tail)?;
263 Ok(())
264 }
265
266 pub fn next_indices(&mut self, maximum: usize) -> Result<Option<Vec<usize>>, SamplerError> {
268 if maximum == 0 { return Err(invalid("sample batch size must be positive")); }
269 let end = self.issued.saturating_add(maximum).min(self.state.indices.len());
270 if end == self.issued { return Ok(None); }
271 let indices = self.state.indices[self.issued..end].to_vec();
272 self.issued = end;
273 Ok(Some(indices))
274 }
275
276 #[cfg(feature = "dataset")]
279 pub fn load_next<I, D: super::dataset::Dataset<I>>(&mut self, dataset: &D, maximum: usize)
280 -> Result<Option<Vec<I>>, SamplerError> {
281 if maximum == 0 { return Err(invalid("sample batch size must be positive")); }
282 if dataset.len() != self.state.source_length { return Err(invalid("dataset length changed since sampler construction")); }
283 let end = self.issued.saturating_add(maximum).min(self.state.indices.len());
284 if end == self.issued { return Ok(None); }
285 let items = dataset.get_many(&self.state.indices[self.issued..end])
286 .ok_or_else(|| invalid("immutable dataset is missing a requested source index"))?;
287 self.issued = end;
288 Ok(Some(items))
289 }
290
291 #[cfg(feature = "dataset")]
293 pub fn load_batch<B: Backend, I, O, D, F>(&mut self, dataset: &D, batcher: &F,
294 device: &B::Device, maximum: usize) -> Result<Option<O>, SamplerError>
295 where D: super::dataset::Dataset<I>, F: super::dataloader::batcher::Batcher<B, I, O> {
296 Ok(self.load_next(dataset, maximum)?.map(|items| batcher.batch(items, device)))
297 }
298}
299
300impl Iterator for StatefulShardSampler {
301 type Item = usize;
302 fn next(&mut self) -> Option<usize> {
303 let item = self.state.indices.get(self.issued).copied()?;
304 self.issued += 1;
305 Some(item)
306 }
307 fn size_hint(&self) -> (usize, Option<usize>) {
308 let count = self.state.indices.len() - self.issued;
309 (count, Some(count))
310 }
311}
312impl ExactSizeIterator for StatefulShardSampler {}