Skip to main content

torsh_data/collate/
builder.rs

1//! Collate builder and strategy definitions
2
3use super::{
4    advanced::{CachedCollate, DynamicBatchCollateWrapper, PadCollate},
5    core::DefaultCollate,
6    optimized::OptimizedCollate,
7};
8use crate::collate::Collate;
9use torsh_core::dtype::TensorElement;
10use torsh_tensor::Tensor;
11
12#[cfg(not(feature = "std"))]
13use alloc::boxed::Box;
14
15/// Different collation strategies
16#[derive(Debug, Clone, Copy)]
17pub enum CollateStrategy {
18    /// Simple stacking (default)
19    Stack,
20    /// Optimized for performance
21    Optimized,
22    /// Variable-length sequences with padding
23    Padding,
24    /// Dynamic batching
25    Dynamic,
26    /// Cached collation for repeated use
27    Cached,
28}
29
30/// Unified collate builder for creating collate functions with different strategies
31pub struct CollateBuilder<T> {
32    strategy: CollateStrategy,
33    padding_value: Option<T>,
34    max_length: Option<usize>,
35    use_caching: bool,
36    batch_size_hint: Option<usize>,
37}
38
39impl<T: TensorElement> Default for CollateBuilder<T> {
40    fn default() -> Self {
41        Self {
42            strategy: CollateStrategy::Stack,
43            padding_value: None,
44            max_length: None,
45            use_caching: false,
46            batch_size_hint: None,
47        }
48    }
49}
50
51impl<
52        T: TensorElement
53            + std::ops::Add<Output = T>
54            + std::ops::Sub<Output = T>
55            + std::ops::Mul<Output = T>
56            + std::ops::Div<Output = T>
57            + Default,
58    > CollateBuilder<T>
59{
60    /// Create a new collate builder
61    pub fn new() -> Self {
62        Self::default()
63    }
64
65    /// Set the collation strategy
66    pub fn strategy(mut self, strategy: CollateStrategy) -> Self {
67        self.strategy = strategy;
68        self
69    }
70
71    /// Set padding value for variable-length sequences
72    pub fn with_padding(mut self, padding_value: T) -> Self {
73        self.padding_value = Some(padding_value);
74        self
75    }
76
77    /// Set maximum sequence length
78    pub fn max_length(mut self, max_length: usize) -> Self {
79        self.max_length = Some(max_length);
80        self
81    }
82
83    /// Enable caching for better performance
84    pub fn with_caching(mut self) -> Self {
85        self.use_caching = true;
86        self
87    }
88
89    /// Provide batch size hint for optimization
90    pub fn batch_size_hint(mut self, size: usize) -> Self {
91        self.batch_size_hint = Some(size);
92        self
93    }
94
95    /// Build the collate function
96    pub fn build(self) -> Box<dyn Collate<Tensor<T>, Output = Tensor<T>> + Send + Sync>
97    where
98        T: Copy + 'static,
99    {
100        match self.strategy {
101            CollateStrategy::Stack => Box::new(DefaultCollate),
102            CollateStrategy::Optimized => Box::new(OptimizedCollate),
103            CollateStrategy::Padding => {
104                let padding_value = self.padding_value.unwrap_or_default();
105                Box::new(PadCollate::new(padding_value))
106            }
107            CollateStrategy::Dynamic => {
108                let padding_value = self.padding_value.unwrap_or_default();
109                Box::new(DynamicBatchCollateWrapper::new(padding_value))
110            }
111            CollateStrategy::Cached => {
112                if cfg!(feature = "std") {
113                    Box::new(CachedCollate::new(1000))
114                } else {
115                    // Fallback to optimized for no_std
116                    Box::new(OptimizedCollate)
117                }
118            }
119        }
120    }
121}