1use 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#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
10pub enum TokenBatchLayout {
11 Packed,
13 Padded,
15}
16
17#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
19pub struct TokenBatchOptions {
20 pub maximum_tokens: usize,
22 pub maximum_examples: Option<usize>,
24 pub layout: TokenBatchLayout,
26}
27
28impl TokenBatchOptions {
29 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#[derive(Clone, Debug, PartialEq, Eq)]
43pub struct TokenBatchPlan {
44 pub indices: Vec<usize>,
46 pub real_tokens: usize,
48 pub maximum_sequence_length: usize,
50 pub token_slots: usize,
52}
53
54pub struct TokenBatch<I> {
56 pub plan: TokenBatchPlan,
58 pub items: Vec<I>,
60}
61
62#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
64pub struct TokenBatchState {
65 pub version: u32,
67 pub sampler: SamplerState,
69 pub lengths: Vec<usize>,
71 pub options: TokenBatchOptions,
73}
74
75impl TokenBatchState {
76 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#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
94pub struct ConsolidatedTokenBatchState {
95 pub version: u32,
97 pub sampler: ConsolidatedSamplerState,
99 pub lengths: Vec<usize>,
101 pub options: TokenBatchOptions,
103}
104
105impl ConsolidatedTokenBatchState {
106 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 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
140pub struct StatefulTokenBatchSampler {
148 sampler: StatefulShardSampler,
149 lengths: Vec<usize>,
150 options: TokenBatchOptions,
151}
152
153impl StatefulTokenBatchSampler {
154 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 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 pub fn state(&self) -> TokenBatchState {
171 TokenBatchState { version: 1, sampler: self.sampler.state(), lengths: self.lengths.clone(), options: self.options.clone() }
172 }
173
174 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 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 pub fn sampler(&self) -> &StatefulShardSampler { &self.sampler }
191
192 pub fn commit(&mut self, examples: usize) -> Result<(), SamplerError> { self.sampler.commit(examples) }
194
195 pub fn rewind_uncommitted(&mut self) { self.sampler.rewind_uncommitted(); }
197
198 pub fn start_epoch(&mut self, epoch: u64) -> Result<(), SamplerError> { self.sampler.start_epoch(epoch) }
200
201 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 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 #[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 #[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}