ruprim_host/comparison/
dispatch.rs1use ruda_core::tensor::{DType, host::HostTensor, element::Scalar};
2use num_traits::ToPrimitive;
3
4fn 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