Skip to main content

ruprim_host/binary/
dispatch_float.rs

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    // Python/PyTorch-style remainder: result has same sign as divisor
46    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    // Python/PyTorch-style remainder: result has same sign as divisor
58    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