1use crate::tensor::{Bool,DType,Int,Tensor,TensorData,backend::Backend};
3use super::causal::{CausalBatch,CausalBatchLayout,CausalBatcher,CausalExample,CausalPadding,CausalTargetAlignment,
4 PaddedCausalBatch,PackedCausalBatch};
5use serde::{Deserialize,Serialize};
6use std::{error::Error,fmt};
7
8#[derive(Clone,Debug,PartialEq,Eq,crate::record::Record)]
10pub struct Seq2SeqDataError(pub String);
11impl fmt::Display for Seq2SeqDataError {
12 fn fmt(&self,formatter: &mut fmt::Formatter<'_>) -> fmt::Result { formatter.write_str(&self.0) }
13}
14impl Error for Seq2SeqDataError {}
15fn invalid(message: impl Into<String>) -> Seq2SeqDataError { Seq2SeqDataError(message.into()) }
16
17#[derive(Clone,Debug,PartialEq,Eq,Serialize,Deserialize)]
19pub struct Seq2SeqExample {
20 pub encoder_input_ids: Vec<i64>,
22 pub decoder_input_ids: Vec<i64>,
24 pub labels: Vec<i64>,
26}
27
28#[derive(Clone,Copy,Debug,PartialEq,Eq,Serialize,Deserialize)]
30pub enum Seq2SeqBatchLayout {
31 Packed,
33 Padded {
35 encoder_pad_token_id: i64,
37 decoder_pad_token_id: i64,
39 encoder_padding: CausalPadding,
41 decoder_padding: CausalPadding,
43 },
44}
45
46#[derive(Clone,Debug,PartialEq,Eq,Serialize,Deserialize)]
48pub struct Seq2SeqBatcher {
49 pub layout: Seq2SeqBatchLayout,
51 pub ignore_index: i64,
53 pub maximum_encoder_length: Option<usize>,
55 pub maximum_decoder_length: Option<usize>,
57}
58
59#[derive(Clone,Debug)]
61pub struct PaddedTokenBatch<B: Backend> {
62 pub input_ids: Tensor<B,2,Int>,
64 pub attention_mask: Tensor<B,2,Bool>,
66 pub position_ids: Tensor<B,2,Int>,
68 pub lengths: Vec<usize>,
70}
71
72impl<B: Backend> PaddedTokenBatch<B> {
73 pub fn examples(&self) -> usize { self.lengths.len() }
75
76 pub fn validate(&self) -> Result<(),Seq2SeqDataError> {
78 let shape = self.input_ids.dims();
79 let device = self.input_ids.device();
80 if shape[0] != self.lengths.len() || self.lengths.iter().any(|&length|length > shape[1])
81 || self.attention_mask.dims() != shape || self.position_ids.dims() != shape
82 || self.attention_mask.device() != device || self.position_ids.device() != device {
83 return Err(invalid("padded input lengths, visibility, positions or device differ"));
84 }
85 Ok(())
86 }
87
88 pub fn to_device(mut self,device: &B::Device) -> Self {
90 self.input_ids = self.input_ids.to_device(device);
91 self.attention_mask = self.attention_mask.to_device(device);
92 self.position_ids = self.position_ids.to_device(device);
93 self
94 }
95}
96
97#[derive(Clone,Debug)]
99pub struct PackedTokenBatch<B: Backend> {
100 pub input_ids: Tensor<B,1,Int>,
102 pub attention_mask: Tensor<B,1,Bool>,
104 pub position_ids: Tensor<B,1,Int>,
106 pub boundaries: Vec<usize>,
108}
109
110impl<B: Backend> PackedTokenBatch<B> {
111 pub fn examples(&self) -> usize { self.boundaries.len().saturating_sub(1) }
113
114 pub fn validate(&self) -> Result<(),Seq2SeqDataError> {
116 let shape = self.input_ids.dims();
117 let device = self.input_ids.device();
118 if self.boundaries.is_empty() || self.boundaries[0] != 0 || self.boundaries.last() != Some(&shape[0])
119 || self.boundaries.windows(2).any(|pair|pair[0] > pair[1])
120 || self.attention_mask.dims() != shape || self.position_ids.dims() != shape
121 || self.attention_mask.device() != device || self.position_ids.device() != device {
122 return Err(invalid("packed input boundaries, visibility, positions or device differ"));
123 }
124 Ok(())
125 }
126
127 pub fn to_device(mut self,device: &B::Device) -> Self {
129 self.input_ids = self.input_ids.to_device(device);
130 self.attention_mask = self.attention_mask.to_device(device);
131 self.position_ids = self.position_ids.to_device(device);
132 self
133 }
134}
135
136impl<B: Backend> PaddedCausalBatch<B> {
137 pub fn token_inputs(&self) -> PaddedTokenBatch<B> {
139 PaddedTokenBatch {input_ids:self.input_ids.clone(),attention_mask:self.attention_mask.clone(),
140 position_ids:self.position_ids.clone(),lengths:self.lengths.clone()}
141 }
142}
143
144impl<B: Backend> PackedCausalBatch<B> {
145 pub fn token_inputs(&self) -> PackedTokenBatch<B> {
147 PackedTokenBatch {input_ids:self.input_ids.clone(),attention_mask:self.attention_mask.clone(),
148 position_ids:self.position_ids.clone(),boundaries:self.boundaries.clone()}
149 }
150}
151
152#[derive(Clone,Debug)]
154pub struct PaddedSeq2SeqBatch<B: Backend> {
155 pub encoder: PaddedTokenBatch<B>,
157 pub decoder: PaddedCausalBatch<B>,
159}
160
161#[derive(Clone,Debug)]
163pub struct PackedSeq2SeqBatch<B: Backend> {
164 pub encoder: PackedTokenBatch<B>,
166 pub decoder: PackedCausalBatch<B>,
168}
169
170#[derive(Clone,Debug)]
172pub enum Seq2SeqBatch<B: Backend> {
173 Padded(PaddedSeq2SeqBatch<B>),
175 Packed(PackedSeq2SeqBatch<B>),
177}
178
179impl<B: Backend> Seq2SeqBatch<B> {
180 pub fn supervised_tokens(&self) -> usize {
182 match self {Self::Padded(batch)=>batch.decoder.supervised_tokens,Self::Packed(batch)=>batch.decoder.supervised_tokens}
183 }
184
185 pub fn examples(&self) -> usize {
187 match self {Self::Padded(batch)=>batch.decoder.examples(),Self::Packed(batch)=>batch.decoder.examples()}
188 }
189
190 pub fn to_device(self,device: &B::Device) -> Self {
192 match self {
193 Self::Padded(batch)=>Self::Padded(PaddedSeq2SeqBatch {encoder:batch.encoder.to_device(device),decoder:batch.decoder.to_device(device)}),
194 Self::Packed(batch)=>Self::Packed(PackedSeq2SeqBatch {encoder:batch.encoder.to_device(device),decoder:batch.decoder.to_device(device)}),
195 }
196 }
197}
198
199impl Seq2SeqBatcher {
200 pub fn validate(&self,samples: &[Seq2SeqExample]) -> Result<(),Seq2SeqDataError> {
202 if samples.is_empty() || self.maximum_encoder_length == Some(0) || self.maximum_decoder_length == Some(0) {
203 return Err(invalid("paired batch needs actual examples and positive optional sequence limits"));
204 }
205 if matches!(self.layout,Seq2SeqBatchLayout::Padded {encoder_pad_token_id,decoder_pad_token_id,..}
206 if encoder_pad_token_id < 0 || decoder_pad_token_id < 0) {
207 return Err(invalid("supply nonnegative actual source/decoder padding IDs"));
208 }
209 let mut encoder_total = 0usize;
210 let mut decoder_total = 0usize;
211 for (index,sample) in samples.iter().enumerate() {
212 let encoder_length = sample.encoder_input_ids.len();
213 let decoder_length = sample.decoder_input_ids.len();
214 if decoder_length == 0 || decoder_length != sample.labels.len()
215 || sample.encoder_input_ids.iter().chain(&sample.decoder_input_ids).any(|&token|token < 0)
216 || sample.labels.iter().any(|&label|label < 0 && label != self.ignore_index) {
217 return Err(invalid(format!("paired example {index} has invalid token IDs or aligned decoder labels")));
218 }
219 if self.maximum_encoder_length.is_some_and(|limit|encoder_length > limit)
220 || self.maximum_decoder_length.is_some_and(|limit|decoder_length > limit)
221 || i64::try_from(encoder_length).is_err() || i64::try_from(decoder_length).is_err() {
222 return Err(invalid(format!("paired example {index} exceeds its explicit sequence/integer limit")));
223 }
224 encoder_total = encoder_total.checked_add(encoder_length).ok_or_else(||invalid("packed source token count overflow"))?;
225 decoder_total = decoder_total.checked_add(decoder_length).ok_or_else(||invalid("packed decoder token count overflow"))?;
226 }
227 if matches!(self.layout,Seq2SeqBatchLayout::Padded {..}) {
228 let encoder_max = samples.iter().map(|sample|sample.encoder_input_ids.len()).max().unwrap();
229 let decoder_max = samples.iter().map(|sample|sample.decoder_input_ids.len()).max().unwrap();
230 samples.len().checked_mul(encoder_max).ok_or_else(||invalid("padded source token count overflow"))?;
231 samples.len().checked_mul(decoder_max).ok_or_else(||invalid("padded decoder token count overflow"))?;
232 }
233 Ok(())
234 }
235
236 pub fn collate<B: Backend>(&self,samples: Vec<Seq2SeqExample>,device: &B::Device) -> Result<Seq2SeqBatch<B>,Seq2SeqDataError> {
239 self.validate(&samples)?;
240 let (encoder,decoder): (Vec<_>,Vec<_>) = samples.into_iter().map(|sample|
241 (sample.encoder_input_ids,CausalExample {input_ids:sample.decoder_input_ids,labels:sample.labels})).unzip();
242 let decoder_layout = match self.layout {
243 Seq2SeqBatchLayout::Packed=>CausalBatchLayout::Packed,
244 Seq2SeqBatchLayout::Padded {decoder_pad_token_id,decoder_padding,..}=>
245 CausalBatchLayout::Padded {pad_token_id:decoder_pad_token_id,padding:decoder_padding},
246 };
247 let decoder = CausalBatcher {layout:decoder_layout,ignore_index:self.ignore_index,
248 target_alignment:CausalTargetAlignment::SamePosition,maximum_sequence_length:self.maximum_decoder_length}
249 .collate(decoder,device).map_err(|error|invalid(error.0))?;
250 match decoder {
251 CausalBatch::Packed(decoder)=>{
252 let total: usize = encoder.iter().map(Vec::len).sum();
253 let mut ids = Vec::with_capacity(total);
254 let mut positions = Vec::with_capacity(total);
255 let mut boundaries = Vec::with_capacity(encoder.len()+1);
256 boundaries.push(0);
257 for sequence in encoder {
258 positions.extend((0..sequence.len()).map(|position|position as i64));
259 ids.extend(sequence);
260 boundaries.push(ids.len());
261 }
262 Ok(Seq2SeqBatch::Packed(PackedSeq2SeqBatch {decoder,encoder:PackedTokenBatch {
263 input_ids:Tensor::from_data(TensorData::new(ids,[total]),(device,DType::I64)),
264 position_ids:Tensor::from_data(TensorData::new(positions,[total]),(device,DType::I64)),
265 attention_mask:Tensor::from_data(TensorData::new(vec![true;total],[total]),device),boundaries,
266 }}))
267 },
268 CausalBatch::Padded(decoder)=>{
269 let Seq2SeqBatchLayout::Padded {encoder_pad_token_id,encoder_padding,..} = self.layout
270 else { return Err(invalid("actual decoder layout differs from paired collation policy")); };
271 let lengths: Vec<_> = encoder.iter().map(Vec::len).collect();
272 let maximum = *lengths.iter().max().unwrap();
273 let slots = encoder.len().checked_mul(maximum).ok_or_else(||invalid("source padded token count overflow"))?;
274 let shape = [encoder.len(),maximum];
275 let mut ids = vec![encoder_pad_token_id;slots];
276 let mut positions = vec![0i64;slots];
277 let mut visible = vec![false;slots];
278 for (row,sequence) in encoder.into_iter().enumerate() {
279 let start = match encoder_padding {CausalPadding::Right=>0,CausalPadding::Left=>maximum-sequence.len()};
280 for (position,token) in sequence.into_iter().enumerate() {
281 let cell = row*maximum+start+position;
282 ids[cell] = token;
283 positions[cell] = position as i64;
284 visible[cell] = true;
285 }
286 }
287 Ok(Seq2SeqBatch::Padded(PaddedSeq2SeqBatch {decoder,encoder:PaddedTokenBatch {
288 input_ids:Tensor::from_data(TensorData::new(ids,shape),(device,DType::I64)),
289 position_ids:Tensor::from_data(TensorData::new(positions,shape),(device,DType::I64)),
290 attention_mask:Tensor::from_data(TensorData::new(visible,shape),device),lengths,
291 }}))
292 },
293 }
294 }
295}
296
297#[cfg(feature = "dataset")]
298impl<B: Backend> super::dataloader::batcher::Batcher<B,Seq2SeqExample,Result<Seq2SeqBatch<B>,Seq2SeqDataError>> for Seq2SeqBatcher {
299 fn batch(&self,items: Vec<Seq2SeqExample>,device: &B::Device) -> Result<Seq2SeqBatch<B>,Seq2SeqDataError> {
300 self.collate(items,device)
301 }
302}