Skip to main content

ruda_model/data/
token_batch.rs

1//! Variable-sized batches with exact token cost and recoverable sample cursors.
2use super::sampler::{ConsolidatedSamplerState, SamplerError, SamplerState, StatefulShardSampler};
3use crate::{record::{PrecisionSettings, Record}, tensor::backend::Backend};
4use serde::{Deserialize, Serialize};
5
6fn invalid(message: impl Into<String>) -> SamplerError { SamplerError(message.into()) }
7
8/// The actual token-axis layout used by the caller's batcher.
9#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
10pub enum TokenBatchLayout {
11    /// Flat document-isolated tokens; cost is the sum of sequence lengths.
12    Packed,
13    /// Rectangular rows; cost is example count times maximum sequence length.
14    Padded,
15}
16
17/// Explicit batch limits; these do not truncate, repeat, reorder or pad samples.
18#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
19pub struct TokenBatchOptions {
20    /// Maximum token slots in one batch under the selected layout.
21    pub maximum_tokens: usize,
22    /// Optional independent maximum number of actual examples in one batch.
23    pub maximum_examples: Option<usize>,
24    /// Packed or rectangular cost accounting; the batcher must honor this layout.
25    pub layout: TokenBatchLayout,
26}
27
28impl TokenBatchOptions {
29    /// Reject zero limits and empty/oversized sequence metadata before issuance.
30    pub fn validate_lengths(&self, lengths: &[usize]) -> Result<(), SamplerError> {
31        if self.maximum_tokens == 0 || self.maximum_examples == Some(0) {
32            return Err(invalid("token batch limits must be positive"));
33        }
34        if let Some(index) = lengths.iter().position(|&length| length == 0 || length > self.maximum_tokens) {
35            return Err(invalid(format!("sequence {index} is empty or exceeds the token budget; preprocess explicitly")));
36        }
37        Ok(())
38    }
39}
40
41/// One actual batch's original source indices and exact work accounting.
42#[derive(Clone, Debug, PartialEq, Eq)]
43pub struct TokenBatchPlan {
44    /// Source indices, in the distributed sampler's recorded relative order.
45    pub indices: Vec<usize>,
46    /// Real tokens in these examples, before any caller-side rectangular padding.
47    pub real_tokens: usize,
48    /// Maximum actual example length.
49    pub maximum_sequence_length: usize,
50    /// Layout-dependent token slots: sum for packed, rows times maximum for padded.
51    pub token_slots: usize,
52}
53
54/// Actual examples paired with the plan used to load them.
55pub struct TokenBatch<I> {
56    /// The exact source indices and token cost, not inferred from padded tensors.
57    pub plan: TokenBatchPlan,
58    /// Loaded source examples; no filler rows.
59    pub items: Vec<I>,
60}
61
62/// Complete batch continuation, including the actual immutable length table.
63#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
64pub struct TokenBatchState {
65    /// Record format, currently one.
66    pub version: u32,
67    /// Committed sample order and explicit immutable source identity.
68    pub sampler: SamplerState,
69    /// One caller-supplied encoded length per original source example.
70    pub lengths: Vec<usize>,
71    /// Exact token accounting/limits used to partition that order.
72    pub options: TokenBatchOptions,
73}
74
75impl TokenBatchState {
76    /// Check source geometry and all length/limit metadata before restoration.
77    pub fn validate(&self) -> Result<(), SamplerError> {
78        self.sampler.validate()?;
79        if self.version != 1 || self.lengths.len() != self.sampler.source_length {
80            return Err(invalid("token batch checkpoint length table/source geometry differs"));
81        }
82        self.options.validate_lengths(&self.lengths)
83    }
84}
85
86impl<B: Backend> Record<B> for TokenBatchState {
87    type Item<P: PrecisionSettings> = Self;
88    fn into_item<P: PrecisionSettings>(self) -> Self { self }
89    fn from_item<P: PrecisionSettings>(item: Self, _: &B::Device) -> Self { item }
90}
91
92/// Topology-independent remaining epoch, retaining the exact length table/policy.
93#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
94pub struct ConsolidatedTokenBatchState {
95    /// Record format, currently one.
96    pub version: u32,
97    /// Only actual unconsumed original source examples.
98    pub sampler: ConsolidatedSamplerState,
99    /// Immutable encoded lengths indexed by original source position.
100    pub lengths: Vec<usize>,
101    /// Batch partitioning policy; no additional epoch tail is dropped.
102    pub options: TokenBatchOptions,
103}
104
105impl ConsolidatedTokenBatchState {
106    /// Validate a continuation without changing any data-rank cursor.
107    pub fn validate(&self) -> Result<(), SamplerError> {
108        self.sampler.validate()?;
109        if self.version != 1 || self.lengths.len() != self.sampler.source_length {
110            return Err(invalid("consolidated token batch metadata differs from its source"));
111        }
112        self.options.validate_lengths(&self.lengths)
113    }
114
115    /// Consolidate every data coordinate, accepting identical TP/PP copies.
116    pub fn consolidate(states: &[TokenBatchState]) -> Result<Self, SamplerError> {
117        let first = states.first().ok_or_else(|| invalid("supply every data-rank token batch state"))?;
118        for state in states {
119            state.validate()?;
120            if state.lengths != first.lengths || state.options != first.options {
121                return Err(invalid("data ranks use different encoded lengths or token batch policies"));
122            }
123        }
124        let sampler_states: Vec<_> = states.iter().map(|state| state.sampler.clone()).collect();
125        let result = Self {
126            version: 1, sampler: ConsolidatedSamplerState::consolidate(&sampler_states)?,
127            lengths: first.lengths.clone(), options: first.options.clone(),
128        };
129        result.validate()?;
130        Ok(result)
131    }
132}
133
134impl<B: Backend> Record<B> for ConsolidatedTokenBatchState {
135    type Item<P: PrecisionSettings> = Self;
136    fn into_item<P: PrecisionSettings>(self) -> Self { self }
137    fn from_item<P: PrecisionSettings>(item: Self, _: &B::Device) -> Self { item }
138}
139
140/// Greedy contiguous token-budget batches over the actual distributed order.
141///
142/// This changes batch boundaries only. Encoded lengths must describe the actual
143/// examples returned by the immutable source. No tokenization, truncation,
144/// length sorting, duplicate tail rows or cross-document supervision is inferred.
145/// Different DP coordinates can have different batch counts; the caller must
146/// use its training coordinator's actual participation/accumulation contract.
147pub struct StatefulTokenBatchSampler {
148    sampler: StatefulShardSampler,
149    lengths: Vec<usize>,
150    options: TokenBatchOptions,
151}
152
153impl StatefulTokenBatchSampler {
154    /// Wrap a committed distributed order with exact original-source lengths.
155    pub fn new(sampler: StatefulShardSampler, lengths: Vec<usize>, options: TokenBatchOptions) -> Result<Self, SamplerError> {
156        if lengths.len() != sampler.source_length() {
157            return Err(invalid("supply one encoded length per original source example"));
158        }
159        options.validate_lengths(&lengths)?;
160        Ok(Self { sampler, lengths, options })
161    }
162
163    /// Restore committed samples; all prefetched-but-unconsumed rows are replayed.
164    pub fn from_state(state: TokenBatchState) -> Result<Self, SamplerError> {
165        state.validate()?;
166        Self::new(StatefulShardSampler::from_state(state.sampler)?, state.lengths, state.options)
167    }
168
169    /// Capture the committed continuation without saving an issued prefetch cursor.
170    pub fn state(&self) -> TokenBatchState {
171        TokenBatchState { version: 1, sampler: self.sampler.state(), lengths: self.lengths.clone(), options: self.options.clone() }
172    }
173
174    /// Replace a matching continuation; reject silently changed lengths/limits.
175    pub fn load_state(&mut self, state: TokenBatchState) -> Result<(), SamplerError> {
176        state.validate()?;
177        if state.lengths != self.lengths || state.options != self.options {
178            return Err(invalid("token batch restoration requires unchanged encoded lengths and limits"));
179        }
180        self.sampler.load_state(state.sampler)
181    }
182
183    /// Assign unconsumed examples to a new data topology and repartition batches.
184    pub fn reshard(state: &ConsolidatedTokenBatchState, rank: usize, world_size: usize) -> Result<Self, SamplerError> {
185        state.validate()?;
186        Self::new(StatefulShardSampler::reshard(&state.sampler, rank, world_size)?, state.lengths.clone(), state.options.clone())
187    }
188
189    /// Underlying actual committed/issued order, without mutable aliasing.
190    pub fn sampler(&self) -> &StatefulShardSampler { &self.sampler }
191
192    /// Commit actual consumed examples, never token slots or prefetched batches.
193    pub fn commit(&mut self, examples: usize) -> Result<(), SamplerError> { self.sampler.commit(examples) }
194
195    /// Replay all outstanding prefetch when constructing a fresh loader iterator.
196    pub fn rewind_uncommitted(&mut self) { self.sampler.rewind_uncommitted(); }
197
198    /// Start the next explicitly selected epoch on the current data topology.
199    pub fn start_epoch(&mut self, epoch: u64) -> Result<(), SamplerError> { self.sampler.start_epoch(epoch) }
200
201    /// Inspect the next exact batch without advancing either sample cursor.
202    pub fn peek(&self) -> Result<Option<TokenBatchPlan>, SamplerError> {
203        let indices = self.sampler.unissued_indices();
204        if indices.is_empty() { return Ok(None); }
205        let mut real_tokens = 0usize;
206        let mut maximum_sequence_length = 0usize;
207        let mut count = 0usize;
208        let mut token_slots = 0usize;
209        for &index in indices {
210            if self.options.maximum_examples.is_some_and(|limit| count == limit) { break; }
211            let length = self.lengths[index];
212            let maximum = maximum_sequence_length.max(length);
213            let total = real_tokens.checked_add(length);
214            let slots = match self.options.layout {
215                TokenBatchLayout::Packed => total,
216                TokenBatchLayout::Padded => (count + 1).checked_mul(maximum),
217            };
218            let Some(slots) = slots.filter(|&slots| slots <= self.options.maximum_tokens) else { break; };
219            real_tokens = total.ok_or_else(|| invalid("real token count overflow"))?;
220            maximum_sequence_length = maximum;
221            token_slots = slots;
222            count += 1;
223        }
224        if count == 0 { return Err(invalid("next actual sample cannot fit the explicit token budget")); }
225        Ok(Some(TokenBatchPlan { indices: indices[..count].to_vec(), real_tokens, maximum_sequence_length, token_slots }))
226    }
227
228    /// Issue the next batch of real indices; no lookahead sample is prematurely issued.
229    pub fn next_plan(&mut self) -> Result<Option<TokenBatchPlan>, SamplerError> {
230        let Some(plan) = self.peek()? else { return Ok(None); };
231        self.sampler.next_indices(plan.indices.len())?;
232        Ok(Some(plan))
233    }
234
235    /// Load every actual example before advancing the issued cursor.
236    #[cfg(feature = "dataset")]
237    pub fn load_next<I, D: super::dataset::Dataset<I>>(&mut self, dataset: &D) -> Result<Option<TokenBatch<I>>, SamplerError> {
238        let Some(plan) = self.peek()? else { return Ok(None); };
239        let items = self.sampler.load_next(dataset, plan.indices.len())?
240            .ok_or_else(|| invalid("source order changed while loading a token batch"))?;
241        Ok(Some(TokenBatch { plan, items }))
242    }
243
244    /// Batch the loaded actual records on the explicit backend device.
245    #[cfg(feature = "dataset")]
246    pub fn load_batch<B: Backend, I, O, D, F>(&mut self, dataset: &D, batcher: &F, device: &B::Device)
247        -> Result<Option<(TokenBatchPlan, O)>, SamplerError>
248    where D: super::dataset::Dataset<I>, F: super::dataloader::batcher::Batcher<B, I, O> {
249        Ok(self.load_next(dataset)?.map(|batch| (batch.plan, batcher.batch(batch.items, device))))
250    }
251}