1#![allow(deprecated)] use super::Tensor;
18use crate::errors::{Result, TrustformersError};
19use scirs2_core::ndarray::{Axis, IxDyn};
20use scirs2_core::simd_ops::SimdUnifiedOps;
21
22const MIN_SIZE_FOR_SIMD: usize = 256;
24
25fn 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 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 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 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 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 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 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 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 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 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 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 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 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 let a_contiguous = a.as_standard_layout().to_owned();
243
244 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 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 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 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 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 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 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 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 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 pub fn gelu(&self) -> Result<Tensor> {
380 match self {
381 #[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 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 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 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 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 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 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 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 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 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 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 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 assert!(data[0].abs() < 1e-4);
688 assert!((data[1] - 0.8413).abs() < 0.02);
690 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 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 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 use crate::tensor::DType;
790 use scirs2_core::ndarray::{ArrayD, IxDyn};
791
792 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 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 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 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}