1use crate::tensor::{Bool, DType, Int, Tensor, TensorData, backend::Backend};
3use serde::{Deserialize, Serialize};
4use std::{error::Error, fmt};
5
6#[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#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
18pub struct CausalExample {
19 pub input_ids: Vec<i64>,
21 pub labels: Vec<i64>,
23}
24
25impl CausalExample {
26 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 pub fn length(&self) -> usize { self.input_ids.len() }
40}
41
42#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
44pub enum CausalPadding {
45 Right,
47 Left,
49}
50
51#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
53pub enum CausalBatchLayout {
54 Packed,
56 Padded {
58 pad_token_id: i64,
60 padding: CausalPadding,
62 },
63}
64
65#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
67pub enum CausalTargetAlignment {
68 NextToken,
70 SamePosition,
72}
73
74#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
76pub struct CausalBatcher {
77 pub layout: CausalBatchLayout,
79 pub ignore_index: i64,
81 pub target_alignment: CausalTargetAlignment,
83 pub maximum_sequence_length: Option<usize>,
85}
86
87#[derive(Clone, Debug)]
89pub struct PackedCausalBatch<B: Backend> {
90 pub input_ids: Tensor<B, 1, Int>,
92 pub labels: Tensor<B, 1, Int>,
94 pub attention_mask: Tensor<B, 1, Bool>,
96 pub position_ids: Tensor<B, 1, Int>,
98 pub boundaries: Vec<usize>,
100 pub maximum_sequence_length: usize,
102 pub supervised_tokens: usize,
104 pub ignore_index: i64,
106 pub target_alignment: CausalTargetAlignment,
108}
109
110impl<B: Backend> PackedCausalBatch<B> {
111 pub fn examples(&self) -> usize { self.boundaries.len() - 1 }
113
114 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#[derive(Clone, Debug)]
126pub struct PaddedCausalBatch<B: Backend> {
127 pub input_ids: Tensor<B, 2, Int>,
129 pub labels: Tensor<B, 2, Int>,
131 pub attention_mask: Tensor<B, 2, Bool>,
133 pub position_ids: Tensor<B, 2, Int>,
135 pub lengths: Vec<usize>,
137 pub supervised_tokens: usize,
139 pub ignore_index: i64,
141 pub target_alignment: CausalTargetAlignment,
143}
144
145impl<B: Backend> PaddedCausalBatch<B> {
146 pub fn examples(&self) -> usize { self.lengths.len() }
148
149 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#[derive(Clone, Debug)]
161pub enum CausalBatch<B: Backend> {
162 Packed(PackedCausalBatch<B>),
164 Padded(PaddedCausalBatch<B>),
166}
167
168impl<B: Backend> CausalBatch<B> {
169 pub fn supervised_tokens(&self) -> usize {
171 match self { Self::Packed(batch) => batch.supervised_tokens, Self::Padded(batch) => batch.supervised_tokens }
172 }
173
174 pub fn examples(&self) -> usize {
176 match self { Self::Packed(batch) => batch.examples(), Self::Padded(batch) => batch.examples() }
177 }
178
179 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 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 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}