Skip to main content

ruda_model/data/
sampler.rs

1//! Dataset-position recovery independent of GPU completion and loader prefetch.
2use 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/// Explicit handling of an epoch that does not divide the data rank count.
9#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
10pub enum SampleTail {
11    /// Unequal final shard lengths; never repeat or pad examples.
12    Uneven,
13    /// Exclude the final incomplete group of data-rank samples for this epoch.
14    Drop,
15}
16
17/// Invalid sampler configuration, incompatible continuation, or missing data.
18#[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/// Actual rank-local sample order and its committed cursor.
39///
40/// The order itself is recorded, so restoration does not depend on reconstructing
41/// a previous rand version's permutation. Outstanding prefetch is NOT committed:
42/// restore begins at `committed`, replaying issued-but-unconsumed samples.
43#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
44pub struct SamplerState {
45    /// Record format, currently one.
46    pub version: u32,
47    /// Caller-supplied identity of the immutable dataset/version.
48    pub source_id: String,
49    /// Original source length, not a rank-local length.
50    pub source_length: usize,
51    /// Data rank; tensor/pipeline ranks may use the same data coordinate.
52    pub rank: usize,
53    /// Data replicas assigning distinct examples.
54    pub world_size: usize,
55    /// Current caller-selected epoch number.
56    pub epoch: u64,
57    /// Private shuffle seed; None selects original source order.
58    pub shuffle_seed: Option<u64>,
59    /// Tail policy used when starting a fresh epoch.
60    pub tail: SampleTail,
61    /// Size of the originally selected epoch, after its initial tail policy.
62    pub epoch_size: usize,
63    /// Global consumption preceding the most recent topology change.
64    pub carried_consumed: usize,
65    /// Remaining global scope assigned at that topology change.
66    pub scope_size: usize,
67    /// Actual local order; positions are original dataset indices.
68    pub indices: Vec<usize>,
69    /// Samples actually consumed by this rank from this scope.
70    pub committed: usize,
71}
72
73impl SamplerState {
74    /// Validate rank geometry, ordering, bounds and committed accounting.
75    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/// Topology-independent continuation containing ONLY unconsumed epoch examples.
94#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
95pub struct ConsolidatedSamplerState {
96    /// Record format, currently one.
97    pub version: u32,
98    /// Immutable dataset identity.
99    pub source_id: String,
100    /// Original dataset length.
101    pub source_length: usize,
102    /// Epoch whose remaining order is retained.
103    pub epoch: u64,
104    /// Shuffle configuration for subsequent epochs.
105    pub shuffle_seed: Option<u64>,
106    /// Tail policy for subsequent epochs, not reapplied to this continuation.
107    pub tail: SampleTail,
108    /// Original effective epoch size.
109    pub epoch_size: usize,
110    /// Global samples already committed across all previous data layouts.
111    pub consumed: usize,
112    /// Remaining source indices in their original relative global order.
113    pub remaining: Vec<usize>,
114}
115
116impl ConsolidatedSamplerState {
117    /// Validate an explicit continuation before loading it on any topology.
118    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    /// Join every data rank, accepting identical duplicate TP/PP copies.
127    /// Conflicting copies, missing ranks or overlapping source indices fail.
128    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
171/// A data-shard iterator with separate issued and actually-consumed positions.
172///
173/// Use `commit` after the caller's batch consumption boundary, not when workers
174/// prefetch. State captures only that boundary. This sampler never pads examples,
175/// automatically advances epochs, guesses source identity, or alters GPU state.
176pub struct StatefulShardSampler {
177    state: SamplerState,
178    issued: usize,
179}
180
181impl StatefulShardSampler {
182    /// Start an explicit epoch on a caller-selected data coordinate.
183    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    /// Restore the recorded actual order; outstanding prefetch is replayed.
200    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    /// Replace a matching source/configuration/data-coordinate continuation.
207    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    /// Assign unconsumed samples on a new data topology, preserving their order.
220    /// The original epoch tail is retained; no additional samples are dropped.
221    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    /// Capture a lossless sampler record for TrainingRecord's application state.
231    pub fn state(&self) -> SamplerState { self.state.clone() }
232
233    /// Immutable source length without cloning the recorded epoch order.
234    pub fn source_length(&self) -> usize { self.state.source_length }
235
236    /// Number of samples actually consumed in this local scope.
237    pub fn committed(&self) -> usize { self.state.committed }
238
239    /// Number of samples issued, including uncommitted prefetch.
240    pub fn issued(&self) -> usize { self.issued }
241
242    /// Actual rank-local work remaining to be consumed, including prefetch.
243    pub fn remaining(&self) -> usize { self.state.indices.len() - self.state.committed }
244
245    /// Actual unissued sample indices in this rank's recorded epoch order.
246    pub fn unissued_indices(&self) -> &[usize] { &self.state.indices[self.issued..] }
247
248    /// Commit only a prefix of samples already issued by this sampler.
249    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    /// Explicitly discard outstanding prefetch and replay it from the commit.
257    pub fn rewind_uncommitted(&mut self) { self.issued = self.state.committed; }
258
259    /// Explicitly start a fresh epoch using the current data coordinate/policy.
260    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    /// Issue at most the requested number of actual indices; no padded batch.
267    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    /// Fetch actual dataset items before advancing the issued position.
277    /// Missing records or a changed source length leave the cursor unchanged.
278    #[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    /// Fetch and batch on the explicit backend device; caller still owns commit.
292    #[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 {}