1use torsh_core::error::{Result, TorshError};
7use torsh_tensor::Tensor;
8
9#[derive(Debug, Clone)]
15pub struct FunctionalConfig {
16 pub validate_inputs: bool,
18 pub eps: f32,
20 pub inplace: bool,
22 pub memory_opt: MemoryOptLevel,
24}
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28pub enum MemoryOptLevel {
29 None,
31 Balanced,
33 Maximum,
35}
36
37impl Default for FunctionalConfig {
38 fn default() -> Self {
39 Self {
40 validate_inputs: true,
41 eps: 1e-8,
42 inplace: false,
43 memory_opt: MemoryOptLevel::Balanced,
44 }
45 }
46}
47
48pub type FuncResult<T> = Result<T>;
50
51#[derive(Debug, Clone)]
53pub struct FunctionalBuilder {
54 config: FunctionalConfig,
55}
56
57impl FunctionalBuilder {
58 pub fn new() -> Self {
60 Self {
61 config: FunctionalConfig::default(),
62 }
63 }
64
65 pub fn validate(mut self, validate: bool) -> Self {
67 self.config.validate_inputs = validate;
68 self
69 }
70
71 pub fn eps(mut self, eps: f32) -> Self {
73 self.config.eps = eps;
74 self
75 }
76
77 pub fn inplace(mut self, inplace: bool) -> Self {
79 self.config.inplace = inplace;
80 self
81 }
82
83 pub fn memory_opt(mut self, level: MemoryOptLevel) -> Self {
85 self.config.memory_opt = level;
86 self
87 }
88
89 pub fn build(self) -> FunctionalConfig {
91 self.config
92 }
93}
94
95impl Default for FunctionalBuilder {
96 fn default() -> Self {
97 Self::new()
98 }
99}
100
101static DEFAULT_CONFIG: FunctionalConfig = FunctionalConfig {
103 validate_inputs: true,
104 eps: 1e-8,
105 inplace: false,
106 memory_opt: MemoryOptLevel::Balanced,
107};
108
109pub fn default_config() -> &'static FunctionalConfig {
111 &DEFAULT_CONFIG
112}
113
114pub fn optimized() -> FunctionalBuilder {
116 FunctionalBuilder::new()
117 .memory_opt(MemoryOptLevel::Maximum)
118 .inplace(true)
119}
120
121pub fn safe() -> FunctionalBuilder {
123 FunctionalBuilder::new()
124 .validate(true)
125 .memory_opt(MemoryOptLevel::None)
126}
127
128pub mod validation {
134 use super::*;
135 use torsh_core::TensorElement;
136
137 pub fn validate_not_empty<T: TensorElement>(tensor: &Tensor<T>, name: &str) -> FuncResult<()> {
139 if tensor.shape().numel() == 0 {
140 return Err(TorshError::InvalidArgument(format!(
141 "Tensor '{}' cannot be empty",
142 name
143 )));
144 }
145 Ok(())
146 }
147
148 pub fn validate_ndim(tensor: &Tensor, expected: usize, name: &str) -> FuncResult<()> {
150 let actual = tensor.shape().dims().len();
151 if actual != expected {
152 return Err(TorshError::InvalidArgument(format!(
153 "Tensor '{}' expected {}D, got {}D",
154 name, expected, actual
155 )));
156 }
157 Ok(())
158 }
159
160 pub fn validate_min_ndim(tensor: &Tensor, min_dims: usize, name: &str) -> FuncResult<()> {
162 let actual = tensor.shape().dims().len();
163 if actual < min_dims {
164 return Err(TorshError::InvalidArgument(format!(
165 "Tensor '{}' requires at least {}D, got {}D",
166 name, min_dims, actual
167 )));
168 }
169 Ok(())
170 }
171
172 pub fn validate_compatible_shapes(a: &Tensor, b: &Tensor, op_name: &str) -> FuncResult<()> {
174 let a_shape_binding = a.shape();
175 let b_shape_binding = b.shape();
176 let a_shape = a_shape_binding.dims();
177 let b_shape = b_shape_binding.dims();
178
179 if a_shape != b_shape && a.shape().numel() != 1 && b.shape().numel() != 1 {
181 if !are_broadcastable(a_shape, b_shape) {
183 return Err(TorshError::InvalidArgument(format!(
184 "Incompatible shapes for {}: {:?} vs {:?}",
185 op_name, a_shape, b_shape
186 )));
187 }
188 }
189 Ok(())
190 }
191
192 fn are_broadcastable(a: &[usize], b: &[usize]) -> bool {
194 let max_len = a.len().max(b.len());
195 let a_padded: Vec<usize> = (0..max_len)
196 .map(|i| {
197 if i < max_len - a.len() {
198 1
199 } else {
200 a[i - (max_len - a.len())]
201 }
202 })
203 .collect();
204 let b_padded: Vec<usize> = (0..max_len)
205 .map(|i| {
206 if i < max_len - b.len() {
207 1
208 } else {
209 b[i - (max_len - b.len())]
210 }
211 })
212 .collect();
213
214 a_padded
215 .iter()
216 .zip(b_padded.iter())
217 .all(|(a_dim, b_dim)| *a_dim == *b_dim || *a_dim == 1 || *b_dim == 1)
218 }
219
220 pub fn validate_range<T: PartialOrd + core::fmt::Display>(
222 value: T,
223 min: T,
224 max: T,
225 name: &str,
226 ) -> FuncResult<()> {
227 if value < min || value > max {
228 return Err(TorshError::InvalidArgument(format!(
229 "Parameter '{}' value {} not in range [{}, {}]",
230 name, value, min, max
231 )));
232 }
233 Ok(())
234 }
235
236 pub fn validate_positive<T: PartialOrd + Default + core::fmt::Display>(
238 value: T,
239 name: &str,
240 ) -> FuncResult<()> {
241 if value <= T::default() {
242 return Err(TorshError::InvalidArgument(format!(
243 "Parameter '{}' must be positive, got {}",
244 name, value
245 )));
246 }
247 Ok(())
248 }
249}
250
251pub mod numerics {
257 use super::*;
258
259 pub fn epsilon_tensor(like: &Tensor, eps: f32) -> FuncResult<Tensor> {
261 torsh_tensor::creation::full_like(like, eps)
262 }
263
264 pub fn safe_clamp(tensor: &Tensor, min_val: f32, max_val: f32) -> FuncResult<Tensor> {
266 let max_tensor = torsh_tensor::creation::full_like(tensor, max_val)?;
268 let min_tensor = torsh_tensor::creation::full_like(tensor, min_val)?;
269 let clamped_max = tensor.minimum(&max_tensor)?;
270 clamped_max.maximum(&min_tensor)
271 }
272
273 pub fn safe_div(numerator: &Tensor, denominator: &Tensor, eps: f32) -> FuncResult<Tensor> {
275 let eps_tensor = epsilon_tensor(denominator, eps)?;
276 let safe_denom = denominator.add(&eps_tensor)?;
277 numerator.div(&safe_denom)
278 }
279
280 pub fn safe_sqrt(tensor: &Tensor, eps: f32) -> FuncResult<Tensor> {
282 let eps_tensor = epsilon_tensor(tensor, eps)?;
283 let safe_tensor = tensor.add(&eps_tensor)?;
284 safe_tensor.sqrt()
285 }
286
287 pub fn safe_rsqrt(tensor: &Tensor, eps: f32) -> FuncResult<Tensor> {
289 let eps_tensor = epsilon_tensor(tensor, eps)?;
290 let safe_tensor = tensor.add(&eps_tensor)?;
291 safe_tensor.rsqrt()
292 }
293}
294
295pub mod performance {
301 use super::*;
302
303 pub fn with_memory_opt<F, T>(config: &FunctionalConfig, op: F) -> FuncResult<T>
305 where
306 F: FnOnce() -> FuncResult<T>,
307 {
308 match config.memory_opt {
310 MemoryOptLevel::None => op(),
311 MemoryOptLevel::Balanced => {
312 op()
314 }
315 MemoryOptLevel::Maximum => {
316 op()
318 }
319 }
320 }
321
322 pub fn should_use_inplace(config: &FunctionalConfig, tensor_size: usize) -> bool {
324 config.inplace && tensor_size > 1000 }
326}
327
328#[macro_export]
334macro_rules! validate_inputs {
335 ($config:expr, $($validation:expr),*) => {
336 if $config.validate_inputs {
337 $( $validation?; )*
338 }
339 };
340}
341
342#[macro_export]
344macro_rules! func_error {
345 ($op:expr, $context:expr) => {
346 $op.map_err(|e| torsh_core::error::TorshError::RuntimeError(format!("{}: {}", $context, e)))
347 };
348}
349
350#[derive(Debug, Clone)]
356pub struct ActivationConfig {
357 pub base: FunctionalConfig,
359 pub clamp_output: bool,
361 pub clamp_range: (f32, f32),
363}
364
365impl Default for ActivationConfig {
366 fn default() -> Self {
367 Self {
368 base: FunctionalConfig::default(),
369 clamp_output: false,
370 clamp_range: (-10.0, 10.0),
371 }
372 }
373}
374
375pub trait Activation {
377 fn apply(&self, input: &Tensor) -> FuncResult<Tensor>;
379
380 fn apply_with_config(&self, input: &Tensor, config: &ActivationConfig) -> FuncResult<Tensor> {
382 let result = self.apply(input)?;
383 if config.clamp_output {
384 let (min_val, max_val) = config.clamp_range;
385 numerics::safe_clamp(&result, min_val, max_val)
386 } else {
387 Ok(result)
388 }
389 }
390}
391
392#[cfg(test)]
397mod tests {
398 use super::*;
399 use approx::assert_relative_eq;
400
401 #[test]
406 fn test_functional_config_default() {
407 let config = FunctionalConfig::default();
408 assert!(config.validate_inputs);
409 assert_relative_eq!(config.eps, 1e-8);
410 assert!(!config.inplace);
411 assert_eq!(config.memory_opt, MemoryOptLevel::Balanced);
412 }
413
414 #[test]
415 fn test_functional_builder_basic() {
416 let config = FunctionalBuilder::new()
417 .validate(false)
418 .eps(1e-6)
419 .inplace(true)
420 .memory_opt(MemoryOptLevel::Maximum)
421 .build();
422
423 assert!(!config.validate_inputs);
424 assert_relative_eq!(config.eps, 1e-6);
425 assert!(config.inplace);
426 assert_eq!(config.memory_opt, MemoryOptLevel::Maximum);
427 }
428
429 #[test]
430 fn test_functional_builder_optimized() {
431 let config = optimized().build();
432 assert!(config.inplace);
433 assert_eq!(config.memory_opt, MemoryOptLevel::Maximum);
434 }
435
436 #[test]
437 fn test_functional_builder_safe() {
438 let config = safe().build();
439 assert!(config.validate_inputs);
440 assert_eq!(config.memory_opt, MemoryOptLevel::None);
441 }
442
443 #[test]
444 fn test_activation_config_default() {
445 let config = ActivationConfig::default();
446 assert!(!config.clamp_output);
447 assert_eq!(config.clamp_range, (-10.0, 10.0));
448 }
449
450 #[test]
455 fn test_validate_not_empty_valid() -> Result<()> {
456 let tensor = Tensor::from_vec(vec![1.0, 2.0, 3.0], &[3])?;
457 validation::validate_not_empty(&tensor, "test")?;
458 Ok(())
459 }
460
461 #[test]
462 fn test_validate_not_empty_invalid() {
463 let tensor: Tensor = Tensor::from_vec(vec![], &[0]).expect("Tensor should succeed");
464 let result = validation::validate_not_empty(&tensor, "test");
465 assert!(result.is_err());
466 }
467
468 #[test]
469 fn test_validate_ndim_valid() -> Result<()> {
470 let tensor = Tensor::from_vec(vec![1.0; 12], &[3, 4])?;
471 validation::validate_ndim(&tensor, 2, "test")?;
472 Ok(())
473 }
474
475 #[test]
476 fn test_validate_ndim_invalid() {
477 let tensor = Tensor::from_vec(vec![1.0; 12], &[3, 4]).expect("Tensor should succeed");
478 let result = validation::validate_ndim(&tensor, 3, "test");
479 assert!(result.is_err());
480 }
481
482 #[test]
483 fn test_validate_min_ndim_valid() -> Result<()> {
484 let tensor = Tensor::from_vec(vec![1.0; 24], &[2, 3, 4])?;
485 validation::validate_min_ndim(&tensor, 2, "test")?;
486 validation::validate_min_ndim(&tensor, 3, "test")?;
487 Ok(())
488 }
489
490 #[test]
491 fn test_validate_min_ndim_invalid() {
492 let tensor = Tensor::from_vec(vec![1.0; 12], &[3, 4]).expect("Tensor should succeed");
493 let result = validation::validate_min_ndim(&tensor, 3, "test");
494 assert!(result.is_err());
495 }
496
497 #[test]
498 fn test_validate_range_valid() -> Result<()> {
499 validation::validate_range(5.0, 0.0, 10.0, "test")?;
500 validation::validate_range(0.0, 0.0, 10.0, "test")?;
501 validation::validate_range(10.0, 0.0, 10.0, "test")?;
502 Ok(())
503 }
504
505 #[test]
506 fn test_validate_range_invalid() {
507 assert!(validation::validate_range(-1.0, 0.0, 10.0, "test").is_err());
508 assert!(validation::validate_range(11.0, 0.0, 10.0, "test").is_err());
509 }
510
511 #[test]
512 fn test_validate_positive_valid() -> Result<()> {
513 validation::validate_positive(1.0, "test")?;
514 validation::validate_positive(0.001, "test")?;
515 Ok(())
516 }
517
518 #[test]
519 fn test_validate_positive_invalid() {
520 assert!(validation::validate_positive(0.0, "test").is_err());
521 assert!(validation::validate_positive(-1.0, "test").is_err());
522 }
523
524 #[test]
525 fn test_validate_compatible_shapes_same() -> Result<()> {
526 let a = Tensor::from_vec(vec![1.0; 12], &[3, 4])?;
527 let b = Tensor::from_vec(vec![2.0; 12], &[3, 4])?;
528 validation::validate_compatible_shapes(&a, &b, "test")?;
529 Ok(())
530 }
531
532 #[test]
533 fn test_validate_compatible_shapes_scalar() -> Result<()> {
534 let a = Tensor::from_vec(vec![1.0; 12], &[3, 4])?;
535 let scalar = Tensor::from_vec(vec![2.0], &[1])?;
536 validation::validate_compatible_shapes(&a, &scalar, "test")?;
537 Ok(())
538 }
539
540 #[test]
541 fn test_validate_compatible_shapes_broadcastable() -> Result<()> {
542 let a = Tensor::from_vec(vec![1.0; 12], &[3, 4])?;
543 let b = Tensor::from_vec(vec![2.0; 4], &[1, 4])?;
544 validation::validate_compatible_shapes(&a, &b, "test")?;
545 Ok(())
546 }
547
548 #[test]
553 fn test_epsilon_tensor() -> Result<()> {
554 let tensor = Tensor::from_vec(vec![1.0, 2.0, 3.0], &[3])?;
555 let eps = numerics::epsilon_tensor(&tensor, 1e-5)?;
556
557 let eps_data = eps.to_vec()?;
558 assert_eq!(eps_data.len(), 3);
559 for &val in eps_data.iter() {
560 assert_relative_eq!(val, 1e-5, epsilon = 1e-10);
561 }
562 Ok(())
563 }
564
565 #[test]
566 fn test_safe_clamp_basic() -> Result<()> {
567 let tensor = Tensor::from_vec(vec![-5.0, 0.0, 5.0, 10.0, 15.0], &[5])?;
568 let clamped = numerics::safe_clamp(&tensor, 0.0, 10.0)?;
569
570 let clamped_data = clamped.to_vec()?;
571 assert_relative_eq!(clamped_data[0], 0.0, epsilon = 1e-6); assert_relative_eq!(clamped_data[1], 0.0, epsilon = 1e-6); assert_relative_eq!(clamped_data[2], 5.0, epsilon = 1e-6); assert_relative_eq!(clamped_data[3], 10.0, epsilon = 1e-6); assert_relative_eq!(clamped_data[4], 10.0, epsilon = 1e-6); Ok(())
577 }
578
579 #[test]
580 fn test_safe_clamp_negative_range() -> Result<()> {
581 let tensor = Tensor::from_vec(vec![-10.0, -5.0, 0.0, 5.0], &[4])?;
582 let clamped = numerics::safe_clamp(&tensor, -6.0, -2.0)?;
583
584 let clamped_data = clamped.to_vec()?;
585 assert_relative_eq!(clamped_data[0], -6.0, epsilon = 1e-6); assert_relative_eq!(clamped_data[1], -5.0, epsilon = 1e-6); assert_relative_eq!(clamped_data[2], -2.0, epsilon = 1e-6); assert_relative_eq!(clamped_data[3], -2.0, epsilon = 1e-6); Ok(())
590 }
591
592 #[test]
593 fn test_safe_div_basic() -> Result<()> {
594 let numerator = Tensor::from_vec(vec![10.0, 20.0, 30.0], &[3])?;
595 let denominator = Tensor::from_vec(vec![2.0, 4.0, 5.0], &[3])?;
596 let result = numerics::safe_div(&numerator, &denominator, 1e-8)?;
597
598 let result_data = result.to_vec()?;
599 assert_relative_eq!(result_data[0], 5.0, epsilon = 1e-5);
600 assert_relative_eq!(result_data[1], 5.0, epsilon = 1e-5);
601 assert_relative_eq!(result_data[2], 6.0, epsilon = 1e-5);
602 Ok(())
603 }
604
605 #[test]
606 fn test_safe_div_near_zero() -> Result<()> {
607 let numerator = Tensor::from_vec(vec![1.0], &[1])?;
608 let denominator = Tensor::from_vec(vec![0.0], &[1])?;
609 let result = numerics::safe_div(&numerator, &denominator, 1e-8)?;
610
611 let result_data = result.to_vec()?;
613 assert!(result_data[0].is_finite());
614 Ok(())
615 }
616
617 #[test]
618 fn test_safe_sqrt_basic() -> Result<()> {
619 let tensor = Tensor::from_vec(vec![4.0, 9.0, 16.0], &[3])?;
620 let result = numerics::safe_sqrt(&tensor, 1e-8)?;
621
622 let result_data = result.to_vec()?;
623 assert_relative_eq!(result_data[0], 2.0, epsilon = 1e-5);
624 assert_relative_eq!(result_data[1], 3.0, epsilon = 1e-5);
625 assert_relative_eq!(result_data[2], 4.0, epsilon = 1e-5);
626 Ok(())
627 }
628
629 #[test]
630 fn test_safe_sqrt_near_zero() -> Result<()> {
631 let tensor = Tensor::from_vec(vec![0.0], &[1])?;
632 let result = numerics::safe_sqrt(&tensor, 1e-8)?;
633
634 let result_data = result.to_vec()?;
636 assert!(result_data[0] > 0.0);
637 Ok(())
638 }
639
640 #[test]
641 fn test_safe_rsqrt_basic() -> Result<()> {
642 let tensor = Tensor::from_vec(vec![4.0, 9.0, 16.0], &[3])?;
643 let result = numerics::safe_rsqrt(&tensor, 1e-8)?;
644
645 let result_data = result.to_vec()?;
646 assert_relative_eq!(result_data[0], 0.5, epsilon = 1e-5); assert_relative_eq!(result_data[1], 1.0 / 3.0, epsilon = 1e-5); assert_relative_eq!(result_data[2], 0.25, epsilon = 1e-5); Ok(())
650 }
651
652 #[test]
653 fn test_safe_rsqrt_near_zero() -> Result<()> {
654 let tensor = Tensor::from_vec(vec![0.0], &[1])?;
655 let result = numerics::safe_rsqrt(&tensor, 1e-8)?;
656
657 let result_data = result.to_vec()?;
659 assert!(result_data[0].is_finite());
660 Ok(())
661 }
662
663 #[test]
668 fn test_with_memory_opt_none() -> Result<()> {
669 let config = FunctionalBuilder::new()
670 .memory_opt(MemoryOptLevel::None)
671 .build();
672
673 let result = performance::with_memory_opt(&config, || Ok(42))?;
674 assert_eq!(result, 42);
675 Ok(())
676 }
677
678 #[test]
679 fn test_with_memory_opt_balanced() -> Result<()> {
680 let config = FunctionalBuilder::new()
681 .memory_opt(MemoryOptLevel::Balanced)
682 .build();
683
684 let result = performance::with_memory_opt(&config, || Ok(42))?;
685 assert_eq!(result, 42);
686 Ok(())
687 }
688
689 #[test]
690 fn test_with_memory_opt_maximum() -> Result<()> {
691 let config = FunctionalBuilder::new()
692 .memory_opt(MemoryOptLevel::Maximum)
693 .build();
694
695 let result = performance::with_memory_opt(&config, || Ok(42))?;
696 assert_eq!(result, 42);
697 Ok(())
698 }
699
700 #[test]
701 fn test_should_use_inplace_small() {
702 let config = FunctionalBuilder::new().inplace(true).build();
703 assert!(!performance::should_use_inplace(&config, 100));
704 assert!(!performance::should_use_inplace(&config, 1000));
705 }
706
707 #[test]
708 fn test_should_use_inplace_large() {
709 let config = FunctionalBuilder::new().inplace(true).build();
710 assert!(performance::should_use_inplace(&config, 1001));
711 assert!(performance::should_use_inplace(&config, 10000));
712 }
713
714 #[test]
715 fn test_should_use_inplace_disabled() {
716 let config = FunctionalBuilder::new().inplace(false).build();
717 assert!(!performance::should_use_inplace(&config, 10000));
718 }
719}