ruprim_host/unary/
dispatch_float.rs1use ruda_core::tensor::{host::HostTensor, element::Scalar};
2use num_traits::ToPrimitive;
3use crate::unary;
4#[cfg(not(feature = "std"))]
5#[allow(unused_imports)]
6use num_traits::Float;
7
8pub fn float_neg(tensor: HostTensor) -> HostTensor {
9 unary::unary_op(tensor, |x: f32| -x, |x: f64| -x)
10}
11
12pub fn float_clamp(tensor: HostTensor, min: Scalar, max: Scalar) -> HostTensor {
13 let min32 = min.to_f32().unwrap();
14 let max32 = max.to_f32().unwrap();
15 let min64 = min.to_f64().unwrap();
16 let max64 = max.to_f64().unwrap();
17 unary::unary_op(
18 tensor,
19 move |x: f32| x.clamp(min32, max32),
20 move |x: f64| x.clamp(min64, max64),
21 )
22}
23
24pub fn float_clamp_min(tensor: HostTensor, min: Scalar) -> HostTensor {
25 let min32 = min.to_f32().unwrap();
26 let min64 = min.to_f64().unwrap();
27 unary::unary_op(
28 tensor,
29 move |x: f32| x.max(min32),
30 move |x: f64| x.max(min64),
31 )
32}
33
34pub fn float_clamp_max(tensor: HostTensor, max: Scalar) -> HostTensor {
35 let max32 = max.to_f32().unwrap();
36 let max64 = max.to_f64().unwrap();
37 unary::unary_op(
38 tensor,
39 move |x: f32| x.min(max32),
40 move |x: f64| x.min(max64),
41 )
42}
43
44pub fn float_sign(tensor: HostTensor) -> HostTensor {
45 unary::unary_op(
46 tensor,
47 |x: f32| {
48 if x.is_nan() {
49 x
50 } else if x > 0.0 {
51 1.0
52 } else if x < 0.0 {
53 -1.0
54 } else {
55 0.0
56 }
57 },
58 |x: f64| {
59 if x.is_nan() {
60 x
61 } else if x > 0.0 {
62 1.0
63 } else if x < 0.0 {
64 -1.0
65 } else {
66 0.0
67 }
68 },
69 )
70}
71
72pub fn float_is_nan(tensor: HostTensor, out_dtype: ruda_core::tensor::BoolDType) -> HostTensor {
73 unary::float_predicate(tensor, out_dtype, |x: f32| x.is_nan(), |x: f64| x.is_nan())
74}
75
76pub fn float_is_inf(tensor: HostTensor, out_dtype: ruda_core::tensor::BoolDType) -> HostTensor {
77 unary::float_predicate(
78 tensor,
79 out_dtype,
80 |x: f32| x.is_infinite(),
81 |x: f64| x.is_infinite(),
82 )
83}
84