Skip to main content

ruda_model/data/
seq2seq.rs

1//! Paired encoder/decoder collation with explicit already-aligned teacher-forcing targets.
2use 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/// Invalid actual paired input/label geometry or explicit layout configuration.
9#[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/// One actual encoded source and its explicit decoder inputs and same-position labels.
18#[derive(Clone,Debug,PartialEq,Eq,Serialize,Deserialize)]
19pub struct Seq2SeqExample {
20    /// Actual source IDs; no implicit tokenization, BOS, EOS or separator insertion.
21    pub encoder_input_ids: Vec<i64>,
22    /// Actual nonempty teacher-forcing inputs, prepared by the caller's tokenizer/model contract.
23    pub decoder_input_ids: Vec<i64>,
24    /// Label t is supervised by decoder hidden position t; no second shift is performed.
25    pub labels: Vec<i64>,
26}
27
28/// Paired document layout, with independently supplied source/decoder padding policies.
29#[derive(Clone,Copy,Debug,PartialEq,Eq,Serialize,Deserialize)]
30pub enum Seq2SeqBatchLayout {
31    /// Independent source/decoder concatenations; corresponding document indices remain paired.
32    Packed,
33    /// Independent rectangular source/decoder lengths and explicit tokenizer padding IDs.
34    Padded {
35        /// Actual encoder padding ID, not inferred from decoder or EOS.
36        encoder_pad_token_id: i64,
37        /// Actual decoder padding ID, not inferred from encoder or EOS.
38        decoder_pad_token_id: i64,
39        /// Source padding direction.
40        encoder_padding: CausalPadding,
41        /// Decoder padding direction.
42        decoder_padding: CausalPadding,
43    },
44}
45
46/// Native paired collation without model-specific target construction or truncation.
47#[derive(Clone,Debug,PartialEq,Eq,Serialize,Deserialize)]
48pub struct Seq2SeqBatcher {
49    /// Actual source/decoder tensor layout.
50    pub layout: Seq2SeqBatchLayout,
51    /// Explicit decoder target sentinel; independent of either attention mask.
52    pub ignore_index: i64,
53    /// Optional source limit; excess documents are rejected, never truncated.
54    pub maximum_encoder_length: Option<usize>,
55    /// Optional decoder limit, independently enforced before uploading tensors.
56    pub maximum_decoder_length: Option<usize>,
57}
58
59/// Actual padded token inputs without labels, for a model's encoder or decoder.
60#[derive(Clone,Debug)]
61pub struct PaddedTokenBatch<B: Backend> {
62    /// [examples,maximum actual length], including explicitly masked padding cells.
63    pub input_ids: Tensor<B,2,Int>,
64    /// True exactly at actual input tokens, not a supervision/prompt mask.
65    pub attention_mask: Tensor<B,2,Bool>,
66    /// Actual per-document reset positions, with zero in masked padding cells.
67    pub position_ids: Tensor<B,2,Int>,
68    /// Original lengths in actual paired example order.
69    pub lengths: Vec<usize>,
70}
71
72impl<B: Backend> PaddedTokenBatch<B> {
73    /// Actual document count, independent of rectangular token storage size.
74    pub fn examples(&self) -> usize { self.lengths.len() }
75
76    /// Check metadata geometry/device without downloading token or mask contents.
77    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    /// Move actual input tensors without changing their geometry or host length metadata.
89    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/// Actual flat document inputs without labels or synthetic separator tokens.
98#[derive(Clone,Debug)]
99pub struct PackedTokenBatch<B: Backend> {
100    /// Flat actual tokenizer IDs.
101    pub input_ids: Tensor<B,1,Int>,
102    /// Actual input visibility; ordinary packed collation marks every real token visible.
103    pub attention_mask: Tensor<B,1,Bool>,
104    /// Actual per-document positions reset to zero.
105    pub position_ids: Tensor<B,1,Int>,
106    /// Cumulative source/decoder boundaries, independently beginning at zero.
107    pub boundaries: Vec<usize>,
108}
109
110impl<B: Backend> PackedTokenBatch<B> {
111    /// Actual retained document count, including any explicitly empty sources.
112    pub fn examples(&self) -> usize { self.boundaries.len().saturating_sub(1) }
113
114    /// Check exact physical boundaries/geometry/device without tensor readback.
115    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    /// Move actual flat inputs while retaining their exact document boundaries.
128    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    /// Label-free input view; cloning tensor handles does not duplicate token values.
138    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    /// Label-free flat input view with the exact original document boundaries.
146    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/// Native paired padded inputs and unchanged aligned decoder supervision.
153#[derive(Clone,Debug)]
154pub struct PaddedSeq2SeqBatch<B: Backend> {
155    /// Actual encoder inputs; no unused encoder-label tensor is allocated.
156    pub encoder: PaddedTokenBatch<B>,
157    /// Actual decoder inputs/labels; target_alignment is SamePosition.
158    pub decoder: PaddedCausalBatch<B>,
159}
160
161/// Native paired packed inputs with independent actual source/decoder boundaries.
162#[derive(Clone,Debug)]
163pub struct PackedSeq2SeqBatch<B: Backend> {
164    /// Actual encoder documents, including explicitly empty sources if supplied.
165    pub encoder: PackedTokenBatch<B>,
166    /// Actual aligned decoder supervision; no label shift crosses documents.
167    pub decoder: PackedCausalBatch<B>,
168}
169
170/// Explicit native paired device batch for padded or packed encoder-decoder models.
171#[derive(Clone,Debug)]
172pub enum Seq2SeqBatch<B: Backend> {
173    /// Independently padded source and decoder input axes.
174    Padded(PaddedSeq2SeqBatch<B>),
175    /// Independently packed paired document axes.
176    Packed(PackedSeq2SeqBatch<B>),
177}
178
179impl<B: Backend> Seq2SeqBatch<B> {
180    /// Actual decoder supervised labels for weighted gradient accumulation.
181    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    /// Actual paired example count for committing the source sampler cursor.
186    pub fn examples(&self) -> usize {
187        match self {Self::Padded(batch)=>batch.decoder.examples(),Self::Packed(batch)=>batch.decoder.examples()}
188    }
189
190    /// Move both source and decoder tensors to the same caller-selected device.
191    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    /// Validate both actual sequence axes and all explicit policies before tensor upload.
201    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    /// Collate exact supplied pairs; reuse native aligned decoder label/count handling.
237    /// No implicit teacher-forcing shift, BOS/EOS insertion, truncation or prompt masking.
238    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}