1use ruda_core::tensor::{host::HostTensor, element::Scalar, TensorMetadata};
2use num_traits::ToPrimitive;
3use super::{BinaryOp, binary_op, scalar_op};
4#[cfg(not(feature = "std"))]
5#[allow(unused_imports)]
6use num_traits::Float;
7
8pub fn float_add(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
9 binary_op(lhs, rhs, |a, b| a + b, |a, b| a + b, Some(BinaryOp::Add))
10}
11
12pub fn float_add_scalar(lhs: HostTensor, rhs: Scalar) -> HostTensor {
13 let rhs_val = rhs.to_f64().unwrap();
14 scalar_op(lhs, rhs_val, |a, b| a + b, |a, b| a + b)
15}
16
17pub fn float_sub(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
18 binary_op(lhs, rhs, |a, b| a - b, |a, b| a - b, Some(BinaryOp::Sub))
19}
20
21pub fn float_sub_scalar(lhs: HostTensor, rhs: Scalar) -> HostTensor {
22 let rhs_val = rhs.to_f64().unwrap();
23 scalar_op(lhs, rhs_val, |a, b| a - b, |a, b| a - b)
24}
25
26pub fn float_mul(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
27 binary_op(lhs, rhs, |a, b| a * b, |a, b| a * b, Some(BinaryOp::Mul))
28}
29
30pub fn float_mul_scalar(lhs: HostTensor, rhs: Scalar) -> HostTensor {
31 let rhs_val = rhs.to_f64().unwrap();
32 scalar_op(lhs, rhs_val, |a, b| a * b, |a, b| a * b)
33}
34
35pub fn float_div(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
36 binary_op(lhs, rhs, |a, b| a / b, |a, b| a / b, Some(BinaryOp::Div))
37}
38
39pub fn float_div_scalar(lhs: HostTensor, rhs: Scalar) -> HostTensor {
40 let rhs_val = rhs.to_f64().unwrap();
41 scalar_op(lhs, rhs_val, |a, b| a / b, |a, b| a / b)
42}
43
44pub fn float_remainder(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
45 binary_op(
47 lhs,
48 rhs,
49 |a, b| ((a % b) + b) % b,
50 |a, b| ((a % b) + b) % b,
51 None,
52 )
53}
54
55pub fn float_remainder_scalar(lhs: HostTensor, rhs: Scalar) -> HostTensor {
56 let rhs_val = rhs.to_f64().unwrap();
57 scalar_op(
59 lhs,
60 rhs_val,
61 |a, b| ((a % b) + b) % b,
62 |a, b| ((a % b) + b) % b,
63 )
64}
65
66pub fn float_powf(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
67 binary_op(lhs, rhs, |a: f32, b| a.powf(b), |a: f64, b| a.powf(b), None)
68}
69
70pub fn float_powf_scalar_impl(tensor: HostTensor, value: Scalar) -> HostTensor {
71 let exp = value.to_f64().unwrap();
72 scalar_op(tensor, exp, |a: f32, b| a.powf(b), |a: f64, b| a.powf(b))
73}
74
75pub fn float_atan2(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
76 binary_op(
77 lhs,
78 rhs,
79 |a: f32, b| a.atan2(b),
80 |a: f64, b| a.atan2(b),
81 None,
82 )
83}
84
85pub fn float_powi(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
86 super::integer_power::tensor(lhs, rhs)
87}
88
89pub fn float_powi_scalar(lhs: HostTensor, rhs: Scalar) -> HostTensor {
90 if let Scalar::UInt(exponent) = rhs
91 && exponent > i64::MAX as u64
92 {
93 return super::integer_power::scalar(lhs, false, exponent);
94 }
95 match rhs.to_i64().unwrap() {
96 0 => crate::fill::float_ones(lhs.shape(), lhs.dtype().into()),
97 1 => lhs,
98 2 => crate::binary::dispatch_float::float_mul(lhs.clone(), lhs),
99 -1 => crate::unary::recip(lhs),
100 -2 => crate::unary::recip(crate::binary::dispatch_float::float_mul(lhs.clone(), lhs)),
101 exponent => super::integer_power::scalar(lhs, exponent < 0, exponent.unsigned_abs()),
102 }
103}
104
105pub fn float_powf_scalar(tensor: HostTensor, value: Scalar) -> HostTensor {
106 if let Some(exp) = value.try_as_integer() {
107 crate::binary::dispatch_float::float_powi_scalar(tensor, exp)
108 } else {
109 crate::binary::dispatch_float::float_powf_scalar_impl(tensor, value)
110 }
111}
112