Skip to main content

ruda_model/data/
causal.rs

1//! Causal data collation independent of architecture and tokenizer selection.
2use crate::tensor::{Bool, DType, Int, Tensor, TensorData, backend::Backend};
3use serde::{Deserialize, Serialize};
4use std::{error::Error, fmt};
5
6/// Invalid actual input/label geometry or explicit batch configuration.
7#[derive(Clone, Debug, PartialEq, Eq, crate::record::Record)]
8pub struct CausalDataError(pub String);
9
10impl fmt::Display for CausalDataError {
11    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { formatter.write_str(&self.0) }
12}
13impl Error for CausalDataError {}
14fn invalid(message: impl Into<String>) -> CausalDataError { CausalDataError(message.into()) }
15
16/// One already-encoded document with caller-declared supervision.
17#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
18pub struct CausalExample {
19    /// Nonnegative actual tokenizer IDs, without batch-padding tokens.
20    pub input_ids: Vec<i64>,
21    /// Explicit token targets or the collator's selected ignore_index.
22    pub labels: Vec<i64>,
23}
24
25impl CausalExample {
26    /// Validate a nonempty encoded document without guessing prompt spans.
27    pub fn validate(&self, ignore_index: i64) -> Result<(), CausalDataError> {
28        if self.input_ids.is_empty() || self.input_ids.len() != self.labels.len() {
29            return Err(invalid("causal input_ids and labels must be equally sized nonempty sequences"));
30        }
31        if self.input_ids.iter().any(|&token| token < 0)
32            || self.labels.iter().any(|&label| label < 0 && label != ignore_index) {
33            return Err(invalid("causal token IDs must be nonnegative and negative labels must equal ignore_index"));
34        }
35        Ok(())
36    }
37
38    /// Actual encoded sequence length, excluding any future batch padding.
39    pub fn length(&self) -> usize { self.input_ids.len() }
40}
41
42/// Direction of rectangular batch padding, not a truncation policy.
43#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
44pub enum CausalPadding {
45    /// Place all real tokens before the padded suffix.
46    Right,
47    /// Place all real tokens after the padded prefix.
48    Left,
49}
50
51/// Device batch layout explicitly selected by the caller's model contract.
52#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
53pub enum CausalBatchLayout {
54    /// Concatenate documents, retaining cumulative boundaries and reset positions.
55    Packed,
56    /// Rectangular token rows with an explicit tokenizer padding ID/direction.
57    Padded {
58        /// Actual tokenizer's pad ID; never inferred from EOS or vocabulary size.
59        pad_token_id: i64,
60        /// Which side of each real document receives batch padding.
61        padding: CausalPadding,
62    },
63}
64
65/// Alignment used by the downstream causal loss.
66#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
67pub enum CausalTargetAlignment {
68    /// Hidden position t predicts label t+1; document-start targets are ignored.
69    NextToken,
70    /// Hidden position t predicts label t; supplied real labels remain unchanged.
71    SamePosition,
72}
73
74/// Explicit supervision/layout policy shared by every call of a data batcher.
75#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
76pub struct CausalBatcher {
77    /// Actual device token-axis geometry.
78    pub layout: CausalBatchLayout,
79    /// Label sentinel, distinct from attention visibility and padding token ID.
80    pub ignore_index: i64,
81    /// Must match the downstream loss's causal-shift setting.
82    pub target_alignment: CausalTargetAlignment,
83    /// Optional maximum actual sequence length; excess data is rejected, not cut.
84    pub maximum_sequence_length: Option<usize>,
85}
86
87/// Flat actual documents, ready for an explicit packed-attention model.
88#[derive(Clone, Debug)]
89pub struct PackedCausalBatch<B: Backend> {
90    /// Actual flat token IDs, without separator or filler tokens.
91    pub input_ids: Tensor<B, 1, Int>,
92    /// Supervised labels with ignored document starts for next-token training.
93    pub labels: Tensor<B, 1, Int>,
94    /// True for every real token; no physical padding in this representation.
95    pub attention_mask: Tensor<B, 1, Bool>,
96    /// Per-document positions reset to zero.
97    pub position_ids: Tensor<B, 1, Int>,
98    /// Cumulative physical token boundaries, beginning with zero.
99    pub boundaries: Vec<usize>,
100    /// Maximum actual document length.
101    pub maximum_sequence_length: usize,
102    /// Actual nonignored supervised targets under the selected alignment.
103    pub supervised_tokens: usize,
104    /// Actual label sentinel used during collation.
105    pub ignore_index: i64,
106    /// Target alignment recorded for the downstream loss contract.
107    pub target_alignment: CausalTargetAlignment,
108}
109
110impl<B: Backend> PackedCausalBatch<B> {
111    /// Actual document count, not token count or a padded batch dimension.
112    pub fn examples(&self) -> usize { self.boundaries.len() - 1 }
113
114    /// Move the actual tensors to the caller-selected backend device.
115    pub fn to_device(mut self, device: &B::Device) -> Self {
116        self.input_ids = self.input_ids.to_device(device);
117        self.labels = self.labels.to_device(device);
118        self.attention_mask = self.attention_mask.to_device(device);
119        self.position_ids = self.position_ids.to_device(device);
120        self
121    }
122}
123
124/// Rectangular actual documents with explicitly masked batch-padding cells.
125#[derive(Clone, Debug)]
126pub struct PaddedCausalBatch<B: Backend> {
127    /// Token IDs shaped [actual examples, maximum actual sequence length].
128    pub input_ids: Tensor<B, 2, Int>,
129    /// Same shape as input_ids; all padding cells contain ignore_index.
130    pub labels: Tensor<B, 2, Int>,
131    /// True exactly at real source tokens, never at padding cells.
132    pub attention_mask: Tensor<B, 2, Bool>,
133    /// Real positions reset per document; masked padding positions are zero.
134    pub position_ids: Tensor<B, 2, Int>,
135    /// Original encoded lengths in actual batch order.
136    pub lengths: Vec<usize>,
137    /// Actual nonignored supervised targets under the selected alignment.
138    pub supervised_tokens: usize,
139    /// Actual label sentinel used during collation.
140    pub ignore_index: i64,
141    /// Target alignment recorded for the downstream loss contract.
142    pub target_alignment: CausalTargetAlignment,
143}
144
145impl<B: Backend> PaddedCausalBatch<B> {
146    /// Actual source example count; no synthetic rows are created.
147    pub fn examples(&self) -> usize { self.lengths.len() }
148
149    /// Move the actual tensors while retaining the original length metadata.
150    pub fn to_device(mut self, device: &B::Device) -> Self {
151        self.input_ids = self.input_ids.to_device(device);
152        self.labels = self.labels.to_device(device);
153        self.attention_mask = self.attention_mask.to_device(device);
154        self.position_ids = self.position_ids.to_device(device);
155        self
156    }
157}
158
159/// A device batch carrying its actual geometry, not an inferred model family.
160#[derive(Clone, Debug)]
161pub enum CausalBatch<B: Backend> {
162    /// Flat documents for a boundary-aware packed model.
163    Packed(PackedCausalBatch<B>),
164    /// Masked rectangular rows for a padding-aware model.
165    Padded(PaddedCausalBatch<B>),
166}
167
168impl<B: Backend> CausalBatch<B> {
169    /// Actual supervised targets used for weighted microbatch accumulation.
170    pub fn supervised_tokens(&self) -> usize {
171        match self { Self::Packed(batch) => batch.supervised_tokens, Self::Padded(batch) => batch.supervised_tokens }
172    }
173
174    /// Actual encoded example count for committing the sample cursor.
175    pub fn examples(&self) -> usize {
176        match self { Self::Packed(batch) => batch.examples(), Self::Padded(batch) => batch.examples() }
177    }
178
179    /// Move a packed or rectangular batch without altering its layout.
180    pub fn to_device(self, device: &B::Device) -> Self {
181        match self { Self::Packed(batch) => Self::Packed(batch.to_device(device)), Self::Padded(batch) => Self::Padded(batch.to_device(device)) }
182    }
183}
184
185impl CausalBatcher {
186    /// Validate explicit policy and all actual source examples before uploading.
187    pub fn validate(&self, samples: &[CausalExample]) -> Result<(), CausalDataError> {
188        if samples.is_empty() || self.maximum_sequence_length == Some(0) {
189            return Err(invalid("causal batch must contain actual examples and a positive optional sequence limit"));
190        }
191        if matches!(self.layout, CausalBatchLayout::Padded { pad_token_id, .. } if pad_token_id < 0) {
192            return Err(invalid("supply a nonnegative actual padding token ID"));
193        }
194        for (index, sample) in samples.iter().enumerate() {
195            sample.validate(self.ignore_index).map_err(|error| invalid(format!("causal example {index}: {error}")))?;
196            if self.maximum_sequence_length.is_some_and(|limit| sample.length() > limit)
197                || i64::try_from(sample.length()).is_err() {
198                return Err(invalid(format!("causal example {index} exceeds the explicit sequence/integer limit")));
199            }
200        }
201        Ok(())
202    }
203
204    /// Collate explicit labels/IDs on the selected device, without tokenization.
205    ///
206    /// Next-token mode ignores each document's first target even for left
207    /// padding, so a padded hidden state cannot supervise the first real token.
208    /// Packed mode never inserts EOS, changes labels within a document, or
209    /// authorizes a model to attend across the retained boundaries.
210    pub fn collate<B: Backend>(&self, samples: Vec<CausalExample>, device: &B::Device) -> Result<CausalBatch<B>, CausalDataError> {
211        self.validate(&samples)?;
212        let maximum = samples.iter().map(CausalExample::length).max().unwrap();
213        let mut supervised_tokens = 0usize;
214        let label_at = |sample: &CausalExample, position: usize| {
215            if position == 0 && self.target_alignment == CausalTargetAlignment::NextToken { self.ignore_index }
216            else { sample.labels[position] }
217        };
218        match self.layout {
219            CausalBatchLayout::Packed => {
220                let total = samples.iter().try_fold(0usize, |total, sample| total.checked_add(sample.length()))
221                    .ok_or_else(|| invalid("packed causal token count overflow"))?;
222                let mut ids = Vec::with_capacity(total);
223                let mut labels = Vec::with_capacity(total);
224                let mut positions = Vec::with_capacity(total);
225                let mut boundaries = Vec::with_capacity(samples.len() + 1);
226                boundaries.push(0);
227                for sample in samples {
228                    for (position, &token) in sample.input_ids.iter().enumerate() {
229                        let label = label_at(&sample, position);
230                        supervised_tokens += usize::from(label != self.ignore_index);
231                        ids.push(token);
232                        labels.push(label);
233                        positions.push(position as i64);
234                    }
235                    boundaries.push(ids.len());
236                }
237                Ok(CausalBatch::Packed(PackedCausalBatch {
238                    input_ids: Tensor::from_data(TensorData::new(ids, [total]), (device, DType::I64)),
239                    labels: Tensor::from_data(TensorData::new(labels, [total]), (device, DType::I64)),
240                    position_ids: Tensor::from_data(TensorData::new(positions, [total]), (device, DType::I64)),
241                    attention_mask: Tensor::from_data(TensorData::new(vec![true; total], [total]), device),
242                    boundaries, maximum_sequence_length: maximum, supervised_tokens,
243                    ignore_index: self.ignore_index, target_alignment: self.target_alignment,
244                }))
245            }
246            CausalBatchLayout::Padded { pad_token_id, padding } => {
247                let lengths: Vec<_> = samples.iter().map(CausalExample::length).collect();
248                let slots = samples.len().checked_mul(maximum).ok_or_else(|| invalid("padded causal token count overflow"))?;
249                let shape = [samples.len(), maximum];
250                let mut ids = vec![pad_token_id; slots];
251                let mut labels = vec![self.ignore_index; slots];
252                let mut positions = vec![0i64; slots];
253                let mut mask = vec![false; slots];
254                for (row, sample) in samples.iter().enumerate() {
255                    let start = match padding { CausalPadding::Right => 0, CausalPadding::Left => maximum - sample.length() };
256                    for (position, &token) in sample.input_ids.iter().enumerate() {
257                        let cell = row * maximum + start + position;
258                        let label = label_at(sample, position);
259                        supervised_tokens += usize::from(label != self.ignore_index);
260                        ids[cell] = token;
261                        labels[cell] = label;
262                        positions[cell] = position as i64;
263                        mask[cell] = true;
264                    }
265                }
266                Ok(CausalBatch::Padded(PaddedCausalBatch {
267                    input_ids: Tensor::from_data(TensorData::new(ids, shape), (device, DType::I64)),
268                    labels: Tensor::from_data(TensorData::new(labels, shape), (device, DType::I64)),
269                    position_ids: Tensor::from_data(TensorData::new(positions, shape), (device, DType::I64)),
270                    attention_mask: Tensor::from_data(TensorData::new(mask, shape), device),
271                    lengths, supervised_tokens, ignore_index: self.ignore_index, target_alignment: self.target_alignment,
272                }))
273            }
274        }
275    }
276}
277
278#[cfg(feature = "dataset")]
279impl<B: Backend> super::dataloader::batcher::Batcher<B, CausalExample, Result<CausalBatch<B>, CausalDataError>> for CausalBatcher {
280    fn batch(&self, items: Vec<CausalExample>, device: &B::Device) -> Result<CausalBatch<B>, CausalDataError> {
281        self.collate(items, device)
282    }
283}