Skip to main content

trustformers_core/tensor/
activations.rs

1//! Tensor activation functions.
2//!
3//! This module contains activation functions commonly used in neural networks.
4//!
5//! # Performance
6//!
7//! This module uses scirs2-core's SIMD-optimized activation functions for larger tensors:
8//! - `simd_gelu` - GELU activation (used in BERT, GPT, etc.)
9//! - `simd_swish` - Swish/SiLU activation (used in EfficientNet, GPT-NeoX)
10//! - `simd_sigmoid` - Sigmoid activation
11//! - `simd_tanh` - Tanh activation
12//!
13//! For tensors with <256 elements, uses scalar operations to avoid SIMD overhead.
14
15#![allow(deprecated)] // Using rand legacy API, will migrate to scirs2_core
16
17use super::Tensor;
18use crate::errors::{Result, TrustformersError};
19use scirs2_core::ndarray::{Axis, IxDyn};
20use scirs2_core::simd_ops::SimdUnifiedOps;
21
22/// Minimum tensor size to use SIMD operations (avoids overhead for small tensors)
23const MIN_SIZE_FOR_SIMD: usize = 256;
24
25/// Upcast a half-precision (F16/BF16) tensor to F32, run `op`, then downcast the
26/// F32 result back to the original half-precision dtype.
27///
28/// Half-precision floats lack the precision and dynamic range required for stable
29/// activation math (e.g. `exp` in softmax), so we compute in F32 and round the
30/// result back to F16/BF16 so the output dtype matches the input dtype.
31fn run_half_in_f32<F>(input: &Tensor, op: F) -> Result<Tensor>
32where
33    F: Fn(&Tensor) -> Result<Tensor>,
34{
35    match input {
36        Tensor::F16(a) => {
37            // Upcast F16 -> F32 for accurate computation, then downcast back to F16.
38            let upcast = Tensor::F32(a.mapv(|x| x.to_f32()));
39            match op(&upcast)? {
40                Tensor::F32(r) => Ok(Tensor::F16(r.mapv(half::f16::from_f32))),
41                other => other.to_dtype(crate::tensor::DType::F16),
42            }
43        },
44        Tensor::BF16(a) => {
45            // Upcast BF16 -> F32 for accurate computation, then downcast back to BF16.
46            let upcast = Tensor::F32(a.mapv(|x| x.to_f32()));
47            match op(&upcast)? {
48                Tensor::F32(r) => Ok(Tensor::BF16(r.mapv(half::bf16::from_f32))),
49                other => other.to_dtype(crate::tensor::DType::BF16),
50            }
51        },
52        _ => Err(TrustformersError::tensor_op_error(
53            "run_half_in_f32 called on a non-half-precision tensor",
54            "run_half_in_f32",
55        )),
56    }
57}
58
59impl Tensor {
60    /// ReLU activation function.
61    ///
62    /// # Returns
63    ///
64    /// A tensor with ReLU applied element-wise.
65    pub fn relu(&self) -> Result<Tensor> {
66        match self {
67            Tensor::F32(a) => {
68                let result = a.mapv(|x| x.max(0.0));
69                Ok(Tensor::F32(result))
70            },
71            Tensor::F64(a) => {
72                let result = a.mapv(|x| x.max(0.0));
73                Ok(Tensor::F64(result))
74            },
75            Tensor::F16(_) | Tensor::BF16(_) => run_half_in_f32(self, |t| t.relu()),
76            _ => Err(TrustformersError::tensor_op_error(
77                "ReLU not supported for this tensor type",
78                "relu",
79            )),
80        }
81    }
82
83    /// Sigmoid activation function.
84    ///
85    /// # Performance
86    ///
87    /// Uses scirs2-core's SIMD-accelerated sigmoid for tensors with ≥256 elements.
88    ///
89    /// # Returns
90    ///
91    /// A tensor with sigmoid applied element-wise.
92    pub fn sigmoid(&self) -> Result<Tensor> {
93        match self {
94            Tensor::F32(a) => {
95                let size = a.len();
96                if size >= MIN_SIZE_FOR_SIMD {
97                    // Use SIMD-accelerated sigmoid for larger tensors
98                    let shape = a.shape().to_vec();
99                    let flat = a.as_standard_layout();
100                    let flat_view = flat
101                        .view()
102                        .into_shape_with_order(size)
103                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
104                    let result_1d = f32::simd_sigmoid(&flat_view);
105                    let result = result_1d
106                        .into_shape_with_order(IxDyn(&shape))
107                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
108                    Ok(Tensor::F32(result))
109                } else {
110                    // Numerically stable sigmoid implementation for small tensors
111                    let result = a.mapv(|x| {
112                        if x >= 0.0 {
113                            let exp_neg_x = (-x).exp();
114                            1.0 / (1.0 + exp_neg_x)
115                        } else {
116                            let exp_x = x.exp();
117                            exp_x / (1.0 + exp_x)
118                        }
119                    });
120                    Ok(Tensor::F32(result))
121                }
122            },
123            Tensor::F64(a) => {
124                let size = a.len();
125                if size >= MIN_SIZE_FOR_SIMD {
126                    // Use SIMD-accelerated sigmoid for larger tensors
127                    let shape = a.shape().to_vec();
128                    let flat = a.as_standard_layout();
129                    let flat_view = flat
130                        .view()
131                        .into_shape_with_order(size)
132                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
133                    let result_1d = f64::simd_sigmoid(&flat_view);
134                    let result = result_1d
135                        .into_shape_with_order(IxDyn(&shape))
136                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
137                    Ok(Tensor::F64(result))
138                } else {
139                    // Numerically stable sigmoid implementation for small tensors
140                    let result = a.mapv(|x| {
141                        if x >= 0.0 {
142                            let exp_neg_x = (-x).exp();
143                            1.0 / (1.0 + exp_neg_x)
144                        } else {
145                            let exp_x = x.exp();
146                            exp_x / (1.0 + exp_x)
147                        }
148                    });
149                    Ok(Tensor::F64(result))
150                }
151            },
152            Tensor::F16(_) | Tensor::BF16(_) => run_half_in_f32(self, |t| t.sigmoid()),
153            _ => Err(TrustformersError::tensor_op_error(
154                "Sigmoid not supported for this tensor type",
155                "sigmoid",
156            )),
157        }
158    }
159
160    /// Tanh activation function.
161    ///
162    /// # Performance
163    ///
164    /// Uses scirs2-core's SIMD-accelerated tanh for tensors with ≥256 elements.
165    ///
166    /// # Returns
167    ///
168    /// A tensor with tanh applied element-wise.
169    pub fn tanh(&self) -> Result<Tensor> {
170        match self {
171            Tensor::F32(a) => {
172                let size = a.len();
173                if size >= MIN_SIZE_FOR_SIMD {
174                    // Use SIMD-accelerated tanh for larger tensors
175                    let shape = a.shape().to_vec();
176                    let flat = a.as_standard_layout();
177                    let flat_view = flat
178                        .view()
179                        .into_shape_with_order(size)
180                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
181                    let result_1d = f32::simd_tanh(&flat_view);
182                    let result = result_1d
183                        .into_shape_with_order(IxDyn(&shape))
184                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
185                    Ok(Tensor::F32(result))
186                } else {
187                    let result = a.mapv(|x| x.tanh());
188                    Ok(Tensor::F32(result))
189                }
190            },
191            Tensor::F64(a) => {
192                let size = a.len();
193                if size >= MIN_SIZE_FOR_SIMD {
194                    // Use SIMD-accelerated tanh for larger tensors
195                    let shape = a.shape().to_vec();
196                    let flat = a.as_standard_layout();
197                    let flat_view = flat
198                        .view()
199                        .into_shape_with_order(size)
200                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
201                    let result_1d = f64::simd_tanh(&flat_view);
202                    let result = result_1d
203                        .into_shape_with_order(IxDyn(&shape))
204                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
205                    Ok(Tensor::F64(result))
206                } else {
207                    let result = a.mapv(|x| x.tanh());
208                    Ok(Tensor::F64(result))
209                }
210            },
211            Tensor::F16(_) | Tensor::BF16(_) => run_half_in_f32(self, |t| t.tanh()),
212            _ => Err(TrustformersError::tensor_op_error(
213                "Tanh not supported for this tensor type",
214                "tanh",
215            )),
216        }
217    }
218
219    /// Softmax activation function.
220    ///
221    /// # Arguments
222    ///
223    /// * `axis` - The axis along which to apply softmax
224    ///
225    /// # Returns
226    ///
227    /// A tensor with softmax applied along the specified axis.
228    pub fn softmax(&self, axis: i32) -> Result<Tensor> {
229        match self {
230            Tensor::F32(a) => {
231                let ndim = a.ndim();
232                let axis = if axis < 0 { (ndim as i32 + axis) as usize } else { axis as usize };
233
234                if axis >= ndim {
235                    return Err(TrustformersError::shape_error(format!(
236                        "Axis {} is out of bounds for tensor with {} dimensions",
237                        axis, ndim
238                    )));
239                }
240
241                // Ensure contiguous input layout
242                let a_contiguous = a.as_standard_layout().to_owned();
243
244                // For numerical stability, subtract max before exp
245                let max_vals = a_contiguous.map_axis(Axis(axis), |lane| {
246                    lane.iter().fold(f32::NEG_INFINITY, |acc, &x| acc.max(x))
247                });
248
249                // Ensure contiguous max_vals and compute shifted values
250                let max_vals_contiguous = max_vals.as_standard_layout().to_owned();
251                let shifted = &a_contiguous - &max_vals_contiguous.insert_axis(Axis(axis));
252                let shifted_contiguous = shifted.as_standard_layout().to_owned();
253
254                // Compute exp and sum with contiguous layout
255                let exp_vals = shifted_contiguous.mapv(|x| x.exp());
256                let exp_vals_contiguous = exp_vals.as_standard_layout().to_owned();
257                let sum_exp = exp_vals_contiguous.sum_axis(Axis(axis));
258                let sum_exp_contiguous = sum_exp.as_standard_layout().to_owned();
259
260                // Protect against division by very small numbers
261                let protected_sum = sum_exp_contiguous.mapv(|x| {
262                    if x <= f32::MIN_POSITIVE {
263                        f32::MIN_POSITIVE
264                    } else {
265                        x
266                    }
267                });
268
269                // Final result with contiguous layout
270                let result = exp_vals_contiguous / protected_sum.insert_axis(Axis(axis));
271                let result_contiguous = result.as_standard_layout().to_owned();
272                Ok(Tensor::F32(result_contiguous))
273            },
274            Tensor::F64(a) => {
275                let ndim = a.ndim();
276                let axis = if axis < 0 { (ndim as i32 + axis) as usize } else { axis as usize };
277
278                if axis >= ndim {
279                    return Err(TrustformersError::shape_error(format!(
280                        "Axis {} is out of bounds for tensor with {} dimensions",
281                        axis, ndim
282                    )));
283                }
284
285                // Ensure contiguous input layout
286                let a_contiguous = a.as_standard_layout().to_owned();
287
288                let max_vals = a_contiguous.map_axis(Axis(axis), |lane| {
289                    lane.iter().fold(f64::NEG_INFINITY, |acc, &x| acc.max(x))
290                });
291
292                // Ensure contiguous layouts throughout computation
293                let max_vals_contiguous = max_vals.as_standard_layout().to_owned();
294                let shifted = &a_contiguous - &max_vals_contiguous.insert_axis(Axis(axis));
295                let shifted_contiguous = shifted.as_standard_layout().to_owned();
296
297                let exp_vals = shifted_contiguous.mapv(|x| x.exp());
298                let exp_vals_contiguous = exp_vals.as_standard_layout().to_owned();
299                let sum_exp = exp_vals_contiguous.sum_axis(Axis(axis));
300                let sum_exp_contiguous = sum_exp.as_standard_layout().to_owned();
301
302                // Protect against division by very small numbers
303                let protected_sum = sum_exp_contiguous.mapv(|x| {
304                    if x <= f64::MIN_POSITIVE {
305                        f64::MIN_POSITIVE
306                    } else {
307                        x
308                    }
309                });
310
311                let result = exp_vals_contiguous / protected_sum.insert_axis(Axis(axis));
312                let result_contiguous = result.as_standard_layout().to_owned();
313                Ok(Tensor::F64(result_contiguous))
314            },
315            Tensor::F16(_) | Tensor::BF16(_) => run_half_in_f32(self, |t| t.softmax(axis)),
316            _ => Err(TrustformersError::tensor_op_error(
317                "Softmax not supported for this tensor type",
318                "softmax",
319            )),
320        }
321    }
322
323    /// Dropout operation.
324    ///
325    /// # Arguments
326    ///
327    /// * `dropout_prob` - Probability of dropping each element
328    ///
329    /// # Returns
330    ///
331    /// A tensor with dropout applied.
332    pub fn dropout(&self, dropout_prob: f32) -> Result<Tensor> {
333        use scirs2_core::random::*;
334
335        if !(0.0..=1.0).contains(&dropout_prob) {
336            return Err(TrustformersError::tensor_op_error(
337                "Dropout probability must be between 0 and 1",
338                "dropout",
339            ));
340        }
341
342        if dropout_prob == 0.0 {
343            return Ok(self.clone());
344        }
345
346        match self {
347            Tensor::F32(a) => {
348                let mut rng = thread_rng();
349                let scale = 1.0 / (1.0 - dropout_prob);
350                let result =
351                    a.mapv(
352                        |x| {
353                            if rng.random::<f32>() < dropout_prob {
354                                0.0
355                            } else {
356                                x * scale
357                            }
358                        },
359                    );
360                Ok(Tensor::F32(result))
361            },
362            _ => Err(TrustformersError::tensor_op_error(
363                "Dropout not supported for this tensor type",
364                "dropout",
365            )),
366        }
367    }
368
369    /// GELU (Gaussian Error Linear Unit) activation function.
370    ///
371    /// # Performance
372    ///
373    /// Uses scirs2-core's SIMD-accelerated GELU for tensors with ≥256 elements.
374    /// GELU is widely used in Transformer models (BERT, GPT, etc.).
375    ///
376    /// # Returns
377    ///
378    /// A tensor with GELU applied element-wise.
379    pub fn gelu(&self) -> Result<Tensor> {
380        match self {
381            // Metal GPU path - stays on GPU!
382            #[cfg(all(target_os = "macos", feature = "metal"))]
383            Tensor::Metal(metal_data) => {
384                use crate::gpu_ops::metal::get_metal_backend;
385                use crate::tensor::MetalTensorData;
386
387                let backend = get_metal_backend()?;
388                let size = metal_data.shape.iter().product();
389
390                let output_buffer_id = backend.gelu_gpu_to_gpu(&metal_data.buffer_id, size)?;
391
392                Ok(Tensor::Metal(MetalTensorData {
393                    buffer_id: output_buffer_id,
394                    shape: metal_data.shape.clone(),
395                    dtype: metal_data.dtype,
396                }))
397            },
398            Tensor::F32(a) => {
399                let size = a.len();
400                if size >= MIN_SIZE_FOR_SIMD {
401                    // Use SIMD-accelerated GELU for larger tensors
402                    let shape = a.shape().to_vec();
403                    let flat = a.as_standard_layout();
404                    let flat_view = flat
405                        .view()
406                        .into_shape_with_order(size)
407                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
408                    let result_1d = f32::simd_gelu(&flat_view);
409                    let result = result_1d
410                        .into_shape_with_order(IxDyn(&shape))
411                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
412                    Ok(Tensor::F32(result))
413                } else {
414                    // Scalar path for small tensors
415                    let result = a.mapv(|x| {
416                        0.5 * x * (1.0 + (0.7978845608 * (x + 0.044715 * x.powi(3))).tanh())
417                    });
418                    Ok(Tensor::F32(result))
419                }
420            },
421            Tensor::F64(a) => {
422                let size = a.len();
423                if size >= MIN_SIZE_FOR_SIMD {
424                    // Use SIMD-accelerated GELU for larger tensors
425                    let shape = a.shape().to_vec();
426                    let flat = a.as_standard_layout();
427                    let flat_view = flat
428                        .view()
429                        .into_shape_with_order(size)
430                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
431                    let result_1d = f64::simd_gelu(&flat_view);
432                    let result = result_1d
433                        .into_shape_with_order(IxDyn(&shape))
434                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
435                    Ok(Tensor::F64(result))
436                } else {
437                    // Scalar path for small tensors
438                    let result = a.mapv(|x| {
439                        0.5 * x * (1.0 + (0.7978845608028654 * (x + 0.044715 * x.powi(3))).tanh())
440                    });
441                    Ok(Tensor::F64(result))
442                }
443            },
444            Tensor::F16(_) | Tensor::BF16(_) => run_half_in_f32(self, |t| t.gelu()),
445            _ => Err(TrustformersError::tensor_op_error(
446                "GELU not supported for this tensor type",
447                "gelu",
448            )),
449        }
450    }
451
452    /// Leaky ReLU activation function.
453    ///
454    /// # Arguments
455    ///
456    /// * `negative_slope` - The slope for negative values (default: 0.01)
457    ///
458    /// # Returns
459    ///
460    /// A tensor with Leaky ReLU applied element-wise.
461    pub fn leaky_relu(&self, negative_slope: f32) -> Result<Tensor> {
462        match self {
463            Tensor::F32(a) => {
464                let result = a.mapv(|x| if x > 0.0 { x } else { negative_slope * x });
465                Ok(Tensor::F32(result))
466            },
467            Tensor::F64(a) => {
468                let negative_slope = negative_slope as f64;
469                let result = a.mapv(|x| if x > 0.0 { x } else { negative_slope * x });
470                Ok(Tensor::F64(result))
471            },
472            Tensor::F16(_) | Tensor::BF16(_) => {
473                run_half_in_f32(self, |t| t.leaky_relu(negative_slope))
474            },
475            _ => Err(TrustformersError::tensor_op_error(
476                "Leaky ReLU not supported for this tensor type",
477                "leaky_relu",
478            )),
479        }
480    }
481
482    /// SiLU (Sigmoid-Linear Unit) activation function.
483    ///
484    /// Also known as Swish activation: f(x) = x * sigmoid(x)
485    ///
486    /// # Performance
487    ///
488    /// Uses scirs2-core's SIMD-accelerated Swish for tensors with ≥256 elements.
489    /// SiLU/Swish is used in EfficientNet, GPT-NeoX, and many modern architectures.
490    ///
491    /// # Returns
492    ///
493    /// A tensor with SiLU applied element-wise.
494    pub fn silu(&self) -> Result<Tensor> {
495        match self {
496            Tensor::F32(a) => {
497                let size = a.len();
498                if size >= MIN_SIZE_FOR_SIMD {
499                    // Use SIMD-accelerated Swish for larger tensors
500                    let shape = a.shape().to_vec();
501                    let flat = a.as_standard_layout();
502                    let flat_view = flat
503                        .view()
504                        .into_shape_with_order(size)
505                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
506                    let result_1d = f32::simd_swish(&flat_view);
507                    let result = result_1d
508                        .into_shape_with_order(IxDyn(&shape))
509                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
510                    Ok(Tensor::F32(result))
511                } else {
512                    // Scalar path for small tensors
513                    let result = a.mapv(|x| x * (1.0 / (1.0 + (-x).exp())));
514                    Ok(Tensor::F32(result))
515                }
516            },
517            Tensor::F64(a) => {
518                let size = a.len();
519                if size >= MIN_SIZE_FOR_SIMD {
520                    // Use SIMD-accelerated Swish for larger tensors
521                    let shape = a.shape().to_vec();
522                    let flat = a.as_standard_layout();
523                    let flat_view = flat
524                        .view()
525                        .into_shape_with_order(size)
526                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
527                    let result_1d = f64::simd_swish(&flat_view);
528                    let result = result_1d
529                        .into_shape_with_order(IxDyn(&shape))
530                        .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
531                    Ok(Tensor::F64(result))
532                } else {
533                    // Scalar path for small tensors
534                    let result = a.mapv(|x| x * (1.0 / (1.0 + (-x).exp())));
535                    Ok(Tensor::F64(result))
536                }
537            },
538            Tensor::F16(_) | Tensor::BF16(_) => run_half_in_f32(self, |t| t.silu()),
539            _ => Err(TrustformersError::tensor_op_error(
540                "SiLU not supported for this tensor type",
541                "silu",
542            )),
543        }
544    }
545
546    /// Swish activation function (alias for SiLU).
547    ///
548    /// Swish(x) = x * sigmoid(x) = SiLU(x)
549    pub fn swish(&self) -> Result<Tensor> {
550        self.silu()
551    }
552}
553
554#[cfg(test)]
555mod tests {
556    use crate::errors::Result;
557    use crate::tensor::Tensor;
558
559    #[test]
560    fn test_relu_positive() -> Result<()> {
561        let t = Tensor::from_data(vec![1.0, 2.0, 3.0], &[3])?;
562        let r = t.relu()?;
563        let data = r.data()?;
564        assert!((data[0] - 1.0).abs() < 1e-6);
565        assert!((data[1] - 2.0).abs() < 1e-6);
566        assert!((data[2] - 3.0).abs() < 1e-6);
567        Ok(())
568    }
569
570    #[test]
571    fn test_relu_negative() -> Result<()> {
572        let t = Tensor::from_data(vec![-1.0, -2.0, -3.0], &[3])?;
573        let r = t.relu()?;
574        let data = r.data()?;
575        for val in &data {
576            assert!(val.abs() < 1e-6);
577        }
578        Ok(())
579    }
580
581    #[test]
582    fn test_relu_mixed() -> Result<()> {
583        let t = Tensor::from_data(vec![-2.0, 0.0, 3.0], &[3])?;
584        let r = t.relu()?;
585        let data = r.data()?;
586        assert!(data[0].abs() < 1e-6);
587        assert!(data[1].abs() < 1e-6);
588        assert!((data[2] - 3.0).abs() < 1e-6);
589        Ok(())
590    }
591
592    #[test]
593    fn test_sigmoid_zero() -> Result<()> {
594        let t = Tensor::from_data(vec![0.0], &[1])?;
595        let r = t.sigmoid()?;
596        let data = r.data()?;
597        assert!((data[0] - 0.5).abs() < 1e-5);
598        Ok(())
599    }
600
601    #[test]
602    fn test_sigmoid_large_positive() -> Result<()> {
603        let t = Tensor::from_data(vec![10.0], &[1])?;
604        let r = t.sigmoid()?;
605        let data = r.data()?;
606        assert!((data[0] - 1.0).abs() < 1e-3);
607        Ok(())
608    }
609
610    #[test]
611    fn test_sigmoid_large_negative() -> Result<()> {
612        let t = Tensor::from_data(vec![-10.0], &[1])?;
613        let r = t.sigmoid()?;
614        let data = r.data()?;
615        assert!(data[0] < 1e-3);
616        Ok(())
617    }
618
619    #[test]
620    fn test_sigmoid_range() -> Result<()> {
621        let t = Tensor::from_data(vec![-5.0, -1.0, 0.0, 1.0, 5.0], &[5])?;
622        let r = t.sigmoid()?;
623        let data = r.data()?;
624        for val in &data {
625            assert!(*val >= 0.0 && *val <= 1.0);
626        }
627        Ok(())
628    }
629
630    #[test]
631    fn test_tanh_zero() -> Result<()> {
632        let t = Tensor::from_data(vec![0.0], &[1])?;
633        let r = t.tanh()?;
634        let data = r.data()?;
635        assert!(data[0].abs() < 1e-5);
636        Ok(())
637    }
638
639    #[test]
640    fn test_tanh_range() -> Result<()> {
641        let t = Tensor::from_data(vec![-10.0, -1.0, 0.0, 1.0, 10.0], &[5])?;
642        let r = t.tanh()?;
643        let data = r.data()?;
644        for val in &data {
645            assert!(*val >= -1.0 && *val <= 1.0);
646        }
647        Ok(())
648    }
649
650    #[test]
651    fn test_softmax_sums_to_one() -> Result<()> {
652        let t = Tensor::from_data(vec![1.0, 2.0, 3.0], &[3])?;
653        let r = t.softmax(0)?;
654        let data = r.data()?;
655        let sum: f32 = data.iter().sum();
656        assert!((sum - 1.0).abs() < 1e-5);
657        Ok(())
658    }
659
660    #[test]
661    fn test_softmax_positive() -> Result<()> {
662        let t = Tensor::from_data(vec![1.0, 2.0, 3.0], &[3])?;
663        let r = t.softmax(0)?;
664        let data = r.data()?;
665        for val in &data {
666            assert!(*val > 0.0);
667        }
668        Ok(())
669    }
670
671    #[test]
672    fn test_softmax_ordering() -> Result<()> {
673        let t = Tensor::from_data(vec![1.0, 2.0, 3.0], &[3])?;
674        let r = t.softmax(0)?;
675        let data = r.data()?;
676        assert!(data[0] < data[1]);
677        assert!(data[1] < data[2]);
678        Ok(())
679    }
680
681    #[test]
682    fn test_gelu() -> Result<()> {
683        let t = Tensor::from_data(vec![0.0, 1.0, -1.0], &[3])?;
684        let r = t.gelu()?;
685        let data = r.data()?;
686        // GELU(0) = 0
687        assert!(data[0].abs() < 1e-4);
688        // GELU(1) ~ 0.8413
689        assert!((data[1] - 0.8413).abs() < 0.02);
690        // GELU(-1) ~ -0.1587
691        assert!((data[2] - (-0.1587)).abs() < 0.02);
692        Ok(())
693    }
694
695    #[test]
696    fn test_leaky_relu_positive() -> Result<()> {
697        let t = Tensor::from_data(vec![1.0, 2.0], &[2])?;
698        let r = t.leaky_relu(0.01)?;
699        let data = r.data()?;
700        assert!((data[0] - 1.0).abs() < 1e-6);
701        assert!((data[1] - 2.0).abs() < 1e-6);
702        Ok(())
703    }
704
705    #[test]
706    fn test_leaky_relu_negative() -> Result<()> {
707        let t = Tensor::from_data(vec![-1.0, -2.0], &[2])?;
708        let r = t.leaky_relu(0.1)?;
709        let data = r.data()?;
710        assert!((data[0] - (-0.1)).abs() < 1e-5);
711        assert!((data[1] - (-0.2)).abs() < 1e-5);
712        Ok(())
713    }
714
715    #[test]
716    fn test_silu_zero() -> Result<()> {
717        let t = Tensor::from_data(vec![0.0], &[1])?;
718        let r = t.silu()?;
719        let data = r.data()?;
720        // SiLU(0) = 0 * sigmoid(0) = 0 * 0.5 = 0
721        assert!(data[0].abs() < 1e-5);
722        Ok(())
723    }
724
725    #[test]
726    fn test_silu_positive() -> Result<()> {
727        let t = Tensor::from_data(vec![2.0], &[1])?;
728        let r = t.silu()?;
729        let data = r.data()?;
730        // SiLU(2) = 2 * sigmoid(2) ~ 2 * 0.88 ~ 1.76
731        assert!(data[0] > 1.5 && data[0] < 2.0);
732        Ok(())
733    }
734
735    #[test]
736    fn test_swish_is_silu() -> Result<()> {
737        let t = Tensor::from_data(vec![1.0, 2.0, -1.0], &[3])?;
738        let silu = t.silu()?;
739        let swish = t.swish()?;
740        let silu_data = silu.data()?;
741        let swish_data = swish.data()?;
742        for i in 0..3 {
743            assert!((silu_data[i] - swish_data[i]).abs() < 1e-6);
744        }
745        Ok(())
746    }
747
748    #[test]
749    fn test_softmax_2d() -> Result<()> {
750        let t = Tensor::from_data(vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0], &[2, 3])?;
751        let r = t.softmax(-1)?;
752        assert_eq!(r.shape(), vec![2, 3]);
753        Ok(())
754    }
755
756    #[test]
757    fn test_relu_2d() -> Result<()> {
758        let t = Tensor::from_data(vec![-1.0, 2.0, -3.0, 4.0], &[2, 2])?;
759        let r = t.relu()?;
760        let data = r.data()?;
761        assert!(data[0].abs() < 1e-6);
762        assert!((data[1] - 2.0).abs() < 1e-6);
763        assert!(data[2].abs() < 1e-6);
764        assert!((data[3] - 4.0).abs() < 1e-6);
765        Ok(())
766    }
767
768    #[test]
769    fn test_dropout_zero_prob() -> Result<()> {
770        let t = Tensor::from_data(vec![1.0, 2.0, 3.0], &[3])?;
771        let r = t.dropout(0.0)?;
772        let data = r.data()?;
773        assert!((data[0] - 1.0).abs() < 1e-5);
774        assert!((data[1] - 2.0).abs() < 1e-5);
775        assert!((data[2] - 3.0).abs() < 1e-5);
776        Ok(())
777    }
778
779    #[test]
780    fn test_gelu_2d() -> Result<()> {
781        let t = Tensor::from_data(vec![0.0, 1.0, -1.0, 2.0], &[2, 2])?;
782        let r = t.gelu()?;
783        assert_eq!(r.shape(), vec![2, 2]);
784        Ok(())
785    }
786
787    // ---- Half-precision (F16 / BF16) upcast-path tests ----
788
789    use crate::tensor::DType;
790    use scirs2_core::ndarray::{ArrayD, IxDyn};
791
792    /// Build an F16 tensor from f32 values.
793    fn make_f16(data: &[f32], shape: &[usize]) -> Result<Tensor> {
794        let arr = ArrayD::from_shape_vec(
795            IxDyn(shape),
796            data.iter().map(|&x| half::f16::from_f32(x)).collect(),
797        )
798        .map_err(|e| crate::errors::TrustformersError::shape_error(e.to_string()))?;
799        Ok(Tensor::F16(arr))
800    }
801
802    /// Build a BF16 tensor from f32 values.
803    fn make_bf16(data: &[f32], shape: &[usize]) -> Result<Tensor> {
804        let arr = ArrayD::from_shape_vec(
805            IxDyn(shape),
806            data.iter().map(|&x| half::bf16::from_f32(x)).collect(),
807        )
808        .map_err(|e| crate::errors::TrustformersError::shape_error(e.to_string()))?;
809        Ok(Tensor::BF16(arr))
810    }
811
812    /// Read a half-precision tensor's values as f32 for assertions.
813    fn half_to_vec_f32(t: &Tensor) -> Vec<f32> {
814        match t {
815            Tensor::F16(a) => a.iter().map(|x| x.to_f32()).collect(),
816            Tensor::BF16(a) => a.iter().map(|x| x.to_f32()).collect(),
817            _ => panic!("expected a half-precision tensor"),
818        }
819    }
820
821    #[test]
822    fn test_relu_f16_bf16() -> Result<()> {
823        for (t, dt) in [
824            (make_f16(&[-1.0, 0.0, 2.5], &[3])?, DType::F16),
825            (make_bf16(&[-1.0, 0.0, 2.5], &[3])?, DType::BF16),
826        ] {
827            let r = t.relu()?;
828            assert_eq!(r.dtype(), dt);
829            assert_eq!(r.shape(), vec![3]);
830            let data = half_to_vec_f32(&r);
831            assert!(data.iter().all(|v| v.is_finite()));
832            assert!(data[0].abs() < 0.05);
833            assert!(data[1].abs() < 0.05);
834            assert!((data[2] - 2.5).abs() < 0.05);
835        }
836        Ok(())
837    }
838
839    #[test]
840    fn test_sigmoid_f16_bf16() -> Result<()> {
841        for (t, dt) in [
842            (make_f16(&[0.0, 4.0, -4.0], &[3])?, DType::F16),
843            (make_bf16(&[0.0, 4.0, -4.0], &[3])?, DType::BF16),
844        ] {
845            let r = t.sigmoid()?;
846            assert_eq!(r.dtype(), dt);
847            assert_eq!(r.shape(), vec![3]);
848            let data = half_to_vec_f32(&r);
849            assert!(data.iter().all(|v| v.is_finite() && *v >= 0.0 && *v <= 1.0));
850            assert!((data[0] - 0.5).abs() < 0.05);
851        }
852        Ok(())
853    }
854
855    #[test]
856    fn test_tanh_f16_bf16() -> Result<()> {
857        for (t, dt) in [
858            (make_f16(&[0.0, 2.0, -2.0], &[3])?, DType::F16),
859            (make_bf16(&[0.0, 2.0, -2.0], &[3])?, DType::BF16),
860        ] {
861            let r = t.tanh()?;
862            assert_eq!(r.dtype(), dt);
863            assert_eq!(r.shape(), vec![3]);
864            let data = half_to_vec_f32(&r);
865            assert!(data.iter().all(|v| v.is_finite() && *v >= -1.0 && *v <= 1.0));
866            assert!(data[0].abs() < 0.05);
867        }
868        Ok(())
869    }
870
871    #[test]
872    fn test_softmax_f16_bf16_rows_sum_to_one() -> Result<()> {
873        for (t, dt) in [
874            (
875                make_f16(&[1.0, 2.0, 3.0, 0.0, 1.0, 0.0], &[2, 3])?,
876                DType::F16,
877            ),
878            (
879                make_bf16(&[1.0, 2.0, 3.0, 0.0, 1.0, 0.0], &[2, 3])?,
880                DType::BF16,
881            ),
882        ] {
883            let r = t.softmax(-1)?;
884            assert_eq!(r.dtype(), dt);
885            assert_eq!(r.shape(), vec![2, 3]);
886            let data = half_to_vec_f32(&r);
887            assert!(data.iter().all(|v| v.is_finite()));
888            // Each row of 3 elements should sum to approximately 1.0.
889            let row0: f32 = data[0..3].iter().sum();
890            let row1: f32 = data[3..6].iter().sum();
891            assert!((row0 - 1.0).abs() < 0.05, "row0 sum = {}", row0);
892            assert!((row1 - 1.0).abs() < 0.05, "row1 sum = {}", row1);
893        }
894        Ok(())
895    }
896
897    #[test]
898    fn test_gelu_f16_bf16() -> Result<()> {
899        for (t, dt) in [
900            (make_f16(&[0.0, 1.0, -1.0], &[3])?, DType::F16),
901            (make_bf16(&[0.0, 1.0, -1.0], &[3])?, DType::BF16),
902        ] {
903            let r = t.gelu()?;
904            assert_eq!(r.dtype(), dt);
905            assert_eq!(r.shape(), vec![3]);
906            let data = half_to_vec_f32(&r);
907            assert!(data.iter().all(|v| v.is_finite()));
908            assert!(data[0].abs() < 0.05);
909            assert!((data[1] - 0.8413).abs() < 0.05);
910        }
911        Ok(())
912    }
913
914    #[test]
915    fn test_leaky_relu_f16_bf16() -> Result<()> {
916        for (t, dt) in [
917            (make_f16(&[-1.0, 2.0], &[2])?, DType::F16),
918            (make_bf16(&[-1.0, 2.0], &[2])?, DType::BF16),
919        ] {
920            let r = t.leaky_relu(0.1)?;
921            assert_eq!(r.dtype(), dt);
922            assert_eq!(r.shape(), vec![2]);
923            let data = half_to_vec_f32(&r);
924            assert!(data.iter().all(|v| v.is_finite()));
925            assert!((data[0] - (-0.1)).abs() < 0.05);
926            assert!((data[1] - 2.0).abs() < 0.05);
927        }
928        Ok(())
929    }
930
931    #[test]
932    fn test_silu_f16_bf16() -> Result<()> {
933        for (t, dt) in [
934            (make_f16(&[0.0, 2.0], &[2])?, DType::F16),
935            (make_bf16(&[0.0, 2.0], &[2])?, DType::BF16),
936        ] {
937            let r = t.silu()?;
938            assert_eq!(r.dtype(), dt);
939            assert_eq!(r.shape(), vec![2]);
940            let data = half_to_vec_f32(&r);
941            assert!(data.iter().all(|v| v.is_finite()));
942            assert!(data[0].abs() < 0.05);
943            assert!(data[1] > 1.5 && data[1] < 2.0);
944        }
945        Ok(())
946    }
947}