use super::sampler::{ConsolidatedSamplerState, SamplerError, SamplerState, StatefulShardSampler};
use crate::{record::{PrecisionSettings, Record}, tensor::backend::Backend};
use serde::{Deserialize, Serialize};
fn invalid(message: impl Into<String>) -> SamplerError { SamplerError(message.into()) }
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum TokenBatchLayout {
Packed,
Padded,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct TokenBatchOptions {
pub maximum_tokens: usize,
pub maximum_examples: Option<usize>,
pub layout: TokenBatchLayout,
}
impl TokenBatchOptions {
pub fn validate_lengths(&self, lengths: &[usize]) -> Result<(), SamplerError> {
if self.maximum_tokens == 0 || self.maximum_examples == Some(0) {
return Err(invalid("token batch limits must be positive"));
}
if let Some(index) = lengths.iter().position(|&length| length == 0 || length > self.maximum_tokens) {
return Err(invalid(format!("sequence {index} is empty or exceeds the token budget; preprocess explicitly")));
}
Ok(())
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TokenBatchPlan {
pub indices: Vec<usize>,
pub real_tokens: usize,
pub maximum_sequence_length: usize,
pub token_slots: usize,
}
pub struct TokenBatch<I> {
pub plan: TokenBatchPlan,
pub items: Vec<I>,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct TokenBatchState {
pub version: u32,
pub sampler: SamplerState,
pub lengths: Vec<usize>,
pub options: TokenBatchOptions,
}
impl TokenBatchState {
pub fn validate(&self) -> Result<(), SamplerError> {
self.sampler.validate()?;
if self.version != 1 || self.lengths.len() != self.sampler.source_length {
return Err(invalid("token batch checkpoint length table/source geometry differs"));
}
self.options.validate_lengths(&self.lengths)
}
}
impl<B: Backend> Record<B> for TokenBatchState {
type Item<P: PrecisionSettings> = Self;
fn into_item<P: PrecisionSettings>(self) -> Self { self }
fn from_item<P: PrecisionSettings>(item: Self, _: &B::Device) -> Self { item }
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ConsolidatedTokenBatchState {
pub version: u32,
pub sampler: ConsolidatedSamplerState,
pub lengths: Vec<usize>,
pub options: TokenBatchOptions,
}
impl ConsolidatedTokenBatchState {
pub fn validate(&self) -> Result<(), SamplerError> {
self.sampler.validate()?;
if self.version != 1 || self.lengths.len() != self.sampler.source_length {
return Err(invalid("consolidated token batch metadata differs from its source"));
}
self.options.validate_lengths(&self.lengths)
}
pub fn consolidate(states: &[TokenBatchState]) -> Result<Self, SamplerError> {
let first = states.first().ok_or_else(|| invalid("supply every data-rank token batch state"))?;
for state in states {
state.validate()?;
if state.lengths != first.lengths || state.options != first.options {
return Err(invalid("data ranks use different encoded lengths or token batch policies"));
}
}
let sampler_states: Vec<_> = states.iter().map(|state| state.sampler.clone()).collect();
let result = Self {
version: 1, sampler: ConsolidatedSamplerState::consolidate(&sampler_states)?,
lengths: first.lengths.clone(), options: first.options.clone(),
};
result.validate()?;
Ok(result)
}
}
impl<B: Backend> Record<B> for ConsolidatedTokenBatchState {
type Item<P: PrecisionSettings> = Self;
fn into_item<P: PrecisionSettings>(self) -> Self { self }
fn from_item<P: PrecisionSettings>(item: Self, _: &B::Device) -> Self { item }
}
pub struct StatefulTokenBatchSampler {
sampler: StatefulShardSampler,
lengths: Vec<usize>,
options: TokenBatchOptions,
}
impl StatefulTokenBatchSampler {
pub fn new(sampler: StatefulShardSampler, lengths: Vec<usize>, options: TokenBatchOptions) -> Result<Self, SamplerError> {
if lengths.len() != sampler.source_length() {
return Err(invalid("supply one encoded length per original source example"));
}
options.validate_lengths(&lengths)?;
Ok(Self { sampler, lengths, options })
}
pub fn from_state(state: TokenBatchState) -> Result<Self, SamplerError> {
state.validate()?;
Self::new(StatefulShardSampler::from_state(state.sampler)?, state.lengths, state.options)
}
pub fn state(&self) -> TokenBatchState {
TokenBatchState { version: 1, sampler: self.sampler.state(), lengths: self.lengths.clone(), options: self.options.clone() }
}
pub fn load_state(&mut self, state: TokenBatchState) -> Result<(), SamplerError> {
state.validate()?;
if state.lengths != self.lengths || state.options != self.options {
return Err(invalid("token batch restoration requires unchanged encoded lengths and limits"));
}
self.sampler.load_state(state.sampler)
}
pub fn reshard(state: &ConsolidatedTokenBatchState, rank: usize, world_size: usize) -> Result<Self, SamplerError> {
state.validate()?;
Self::new(StatefulShardSampler::reshard(&state.sampler, rank, world_size)?, state.lengths.clone(), state.options.clone())
}
pub fn sampler(&self) -> &StatefulShardSampler { &self.sampler }
pub fn commit(&mut self, examples: usize) -> Result<(), SamplerError> { self.sampler.commit(examples) }
pub fn rewind_uncommitted(&mut self) { self.sampler.rewind_uncommitted(); }
pub fn start_epoch(&mut self, epoch: u64) -> Result<(), SamplerError> { self.sampler.start_epoch(epoch) }
pub fn peek(&self) -> Result<Option<TokenBatchPlan>, SamplerError> {
let indices = self.sampler.unissued_indices();
if indices.is_empty() { return Ok(None); }
let mut real_tokens = 0usize;
let mut maximum_sequence_length = 0usize;
let mut count = 0usize;
let mut token_slots = 0usize;
for &index in indices {
if self.options.maximum_examples.is_some_and(|limit| count == limit) { break; }
let length = self.lengths[index];
let maximum = maximum_sequence_length.max(length);
let total = real_tokens.checked_add(length);
let slots = match self.options.layout {
TokenBatchLayout::Packed => total,
TokenBatchLayout::Padded => (count + 1).checked_mul(maximum),
};
let Some(slots) = slots.filter(|&slots| slots <= self.options.maximum_tokens) else { break; };
real_tokens = total.ok_or_else(|| invalid("real token count overflow"))?;
maximum_sequence_length = maximum;
token_slots = slots;
count += 1;
}
if count == 0 { return Err(invalid("next actual sample cannot fit the explicit token budget")); }
Ok(Some(TokenBatchPlan { indices: indices[..count].to_vec(), real_tokens, maximum_sequence_length, token_slots }))
}
pub fn next_plan(&mut self) -> Result<Option<TokenBatchPlan>, SamplerError> {
let Some(plan) = self.peek()? else { return Ok(None); };
self.sampler.next_indices(plan.indices.len())?;
Ok(Some(plan))
}
#[cfg(feature = "dataset")]
pub fn load_next<I, D: super::dataset::Dataset<I>>(&mut self, dataset: &D) -> Result<Option<TokenBatch<I>>, SamplerError> {
let Some(plan) = self.peek()? else { return Ok(None); };
let items = self.sampler.load_next(dataset, plan.indices.len())?
.ok_or_else(|| invalid("source order changed while loading a token batch"))?;
Ok(Some(TokenBatch { plan, items }))
}
#[cfg(feature = "dataset")]
pub fn load_batch<B: Backend, I, O, D, F>(&mut self, dataset: &D, batcher: &F, device: &B::Device)
-> Result<Option<(TokenBatchPlan, O)>, SamplerError>
where D: super::dataset::Dataset<I>, F: super::dataloader::batcher::Batcher<B, I, O> {
Ok(self.load_next(dataset)?.map(|batch| (batch.plan, batcher.batch(batch.items, device))))
}
}