torsh_data/collate/
builder.rs1use 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#[derive(Debug, Clone, Copy)]
17pub enum CollateStrategy {
18 Stack,
20 Optimized,
22 Padding,
24 Dynamic,
26 Cached,
28}
29
30pub 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 pub fn new() -> Self {
62 Self::default()
63 }
64
65 pub fn strategy(mut self, strategy: CollateStrategy) -> Self {
67 self.strategy = strategy;
68 self
69 }
70
71 pub fn with_padding(mut self, padding_value: T) -> Self {
73 self.padding_value = Some(padding_value);
74 self
75 }
76
77 pub fn max_length(mut self, max_length: usize) -> Self {
79 self.max_length = Some(max_length);
80 self
81 }
82
83 pub fn with_caching(mut self) -> Self {
85 self.use_caching = true;
86 self
87 }
88
89 pub fn batch_size_hint(mut self, size: usize) -> Self {
91 self.batch_size_hint = Some(size);
92 self
93 }
94
95 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 Box::new(OptimizedCollate)
117 }
118 }
119 }
120 }
121}