Skip to main content

ruprim_host/comparison/
dispatch.rs

1use ruda_core::tensor::{DType, host::HostTensor, element::Scalar};
2use num_traits::ToPrimitive;
3
4/// Convert a Scalar to (i64, u64) pair for the given dtype.
5/// Only the matching type's conversion is validated; the other gets a dummy 0.
6fn scalar_to_int_pair(dtype: DType, rhs: &Scalar) -> (i64, u64) {
7    if dtype == DType::U64 {
8        (0, rhs.to_u64().unwrap())
9    } else {
10        (rhs.to_i64().unwrap(), 0)
11    }
12}
13
14pub fn float_equal_elem(
15    lhs: HostTensor,
16    rhs: Scalar,
17    out_dtype: ruda_core::tensor::BoolDType,
18) -> HostTensor {
19    crate::comparison::equal_elem(lhs, rhs.to_f64().unwrap(), out_dtype)
20}
21
22pub fn float_greater_elem(
23    lhs: HostTensor,
24    rhs: Scalar,
25    out_dtype: ruda_core::tensor::BoolDType,
26) -> HostTensor {
27    crate::comparison::greater_elem(lhs, rhs.to_f64().unwrap(), out_dtype)
28}
29
30pub fn float_greater_equal_elem(
31    lhs: HostTensor,
32    rhs: Scalar,
33    out_dtype: ruda_core::tensor::BoolDType,
34) -> HostTensor {
35    crate::comparison::greater_equal_elem(lhs, rhs.to_f64().unwrap(), out_dtype)
36}
37
38pub fn float_lower_elem(
39    lhs: HostTensor,
40    rhs: Scalar,
41    out_dtype: ruda_core::tensor::BoolDType,
42) -> HostTensor {
43    crate::comparison::lower_elem(lhs, rhs.to_f64().unwrap(), out_dtype)
44}
45
46pub fn float_lower_equal_elem(
47    lhs: HostTensor,
48    rhs: Scalar,
49    out_dtype: ruda_core::tensor::BoolDType,
50) -> HostTensor {
51    crate::comparison::lower_equal_elem(lhs, rhs.to_f64().unwrap(), out_dtype)
52}
53
54pub fn float_not_equal_elem(
55    lhs: HostTensor,
56    rhs: Scalar,
57    out_dtype: ruda_core::tensor::BoolDType,
58) -> HostTensor {
59    crate::comparison::not_equal_elem(lhs, rhs.to_f64().unwrap(), out_dtype)
60}
61
62pub fn int_equal_elem(
63    lhs: HostTensor,
64    rhs: Scalar,
65    out_dtype: ruda_core::tensor::BoolDType,
66) -> HostTensor {
67    let (i, u) = scalar_to_int_pair(lhs.dtype(), &rhs);
68    crate::comparison::int_equal_elem(lhs, i, u, out_dtype)
69}
70
71pub fn int_greater_elem(
72    lhs: HostTensor,
73    rhs: Scalar,
74    out_dtype: ruda_core::tensor::BoolDType,
75) -> HostTensor {
76    let (i, u) = scalar_to_int_pair(lhs.dtype(), &rhs);
77    crate::comparison::int_greater_elem(lhs, i, u, out_dtype)
78}
79
80pub fn int_greater_equal_elem(
81    lhs: HostTensor,
82    rhs: Scalar,
83    out_dtype: ruda_core::tensor::BoolDType,
84) -> HostTensor {
85    let (i, u) = scalar_to_int_pair(lhs.dtype(), &rhs);
86    crate::comparison::int_greater_equal_elem(lhs, i, u, out_dtype)
87}
88
89pub fn int_lower_elem(
90    lhs: HostTensor,
91    rhs: Scalar,
92    out_dtype: ruda_core::tensor::BoolDType,
93) -> HostTensor {
94    let (i, u) = scalar_to_int_pair(lhs.dtype(), &rhs);
95    crate::comparison::int_lower_elem(lhs, i, u, out_dtype)
96}
97
98pub fn int_lower_equal_elem(
99    lhs: HostTensor,
100    rhs: Scalar,
101    out_dtype: ruda_core::tensor::BoolDType,
102) -> HostTensor {
103    let (i, u) = scalar_to_int_pair(lhs.dtype(), &rhs);
104    crate::comparison::int_lower_equal_elem(lhs, i, u, out_dtype)
105}
106
107pub fn int_not_equal_elem(
108    lhs: HostTensor,
109    rhs: ruda_core::tensor::element::Scalar,
110    out_dtype: ruda_core::tensor::BoolDType,
111) -> HostTensor {
112    let (i, u) = scalar_to_int_pair(lhs.dtype(), &rhs);
113    crate::comparison::int_not_equal_elem(lhs, i, u, out_dtype)
114}
115