ruprim_host/boolean/
mod.rs1use alloc::{vec, vec::Vec};
2use ruda_core::tensor::{DType, Shape, host::HostTensor};
3
4mod binary;
5use binary::{BoolBinaryOp, bool_binary_op_simd};
6mod indexing;
7pub use indexing::{bool_argwhere, bool_select_or};
8
9pub fn bool_equal(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
10 use ruda_core::tensor::host::strided_index::StridedIter;
11
12 let (lhs, rhs) = crate::expand::broadcast_binary(lhs, rhs);
18
19 let out_dtype = ruda_core::tensor::BoolDType::from(lhs.dtype());
20 let shape = lhs.layout().shape().clone();
21 let lhs_storage: &[u8] = lhs.bytes();
22 let rhs_storage: &[u8] = rhs.bytes();
23
24 let result: Vec<u8> = match (
25 lhs.layout().contiguous_offsets(),
26 rhs.layout().contiguous_offsets(),
27 ) {
28 (Some((l_start, l_end)), Some((r_start, r_end))) => {
29 let l_slice = &lhs_storage[l_start..l_end];
30 let r_slice = &rhs_storage[r_start..r_end];
31 l_slice
32 .iter()
33 .zip(r_slice)
34 .map(|(&a, &b)| (a == b) as u8)
35 .collect()
36 }
37 _ => {
38 let lhs_iter = StridedIter::new(lhs.layout());
39 let rhs_iter = StridedIter::new(rhs.layout());
40 lhs_iter
41 .zip(rhs_iter)
42 .map(|(li, ri)| (lhs_storage[li] == rhs_storage[ri]) as u8)
43 .collect()
44 }
45 };
46
47 crate::comparison::make_bool_tensor(result, shape, out_dtype)
48}
49
50pub fn bool_not(mut tensor: HostTensor) -> HostTensor {
51 use ruda_core::tensor::host::strided_index::StridedIter;
52
53 debug_assert!(
54 matches!(
55 tensor.dtype(),
56 DType::Bool(ruda_core::tensor::BoolStore::Native | ruda_core::tensor::BoolStore::U8)
57 ),
58 "bool_not: only Bool(Native) and Bool(U8) are supported, got {:?}",
59 tensor.dtype()
60 );
61
62 if tensor.is_unique()
66 && tensor.layout().is_contiguous()
67 && tensor.layout().start_offset() == 0
68 {
69 let storage = tensor.storage_mut::<u8>();
70 crate::simd::bool_not_inplace_u8(storage);
71 return tensor;
72 }
73
74 let out_dtype = ruda_core::tensor::BoolDType::from(tensor.dtype());
77 let shape = tensor.layout().shape().clone();
78 let storage: &[u8] = tensor.bytes();
79
80 let result: Vec<u8> = match tensor.layout().contiguous_offsets() {
81 Some((start, end)) => {
82 let slice = &storage[start..end];
83 let mut out = vec![0u8; slice.len()];
84 crate::simd::bool_not_u8(slice, &mut out);
85 out
86 }
87 None => StridedIter::new(tensor.layout())
88 .map(|idx| (storage[idx] == 0) as u8)
89 .collect(),
90 };
91
92 crate::comparison::make_bool_tensor(result, shape, out_dtype)
93}
94
95pub fn bool_and(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
96 bool_binary_op_simd(lhs, rhs, BoolBinaryOp::And)
97}
98
99pub fn bool_or(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
100 bool_binary_op_simd(lhs, rhs, BoolBinaryOp::Or)
101}
102
103pub fn bool_xor(lhs: HostTensor, rhs: HostTensor) -> HostTensor {
104 bool_binary_op_simd(lhs, rhs, BoolBinaryOp::Xor)
105}
106
107pub fn bool_ones(
108 shape: Shape,
109 dtype: ruda_core::tensor::BoolDType,
110) -> HostTensor {
111 let num_elements = shape.num_elements();
112 let data = vec![1u8; num_elements];
113 crate::comparison::make_bool_tensor(data, shape, dtype)
114}
115
116pub fn bool_equal_elem(lhs: HostTensor, rhs: ruda_core::tensor::element::Scalar) -> HostTensor {
117 use ruda_core::tensor::host::strided_index::StridedIter;
118
119 let out_dtype = ruda_core::tensor::BoolDType::from(lhs.dtype());
120 let shape = lhs.layout().shape().clone();
121 let storage: &[u8] = lhs.bytes();
122 let rhs_bool: bool = rhs.elem();
123 let rhs_val = rhs_bool as u8;
124
125 let result: Vec<u8> = match lhs.layout().contiguous_offsets() {
126 Some((start, end)) => storage[start..end]
127 .iter()
128 .map(|&v| (v == rhs_val) as u8)
129 .collect(),
130 None => StridedIter::new(lhs.layout())
131 .map(|idx| (storage[idx] == rhs_val) as u8)
132 .collect(),
133 };
134
135 crate::comparison::make_bool_tensor(result, shape, out_dtype)
136}