Skip to main content

ruprim_host/boolean/
mod.rs

1use 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    // Broadcast to a common shape before comparing. The contiguous fast
13    // path below uses `zip`, which silently truncates to the shorter
14    // operand; and the output shape is taken from lhs, so mismatched
15    // operands would otherwise produce a result vec shorter than the
16    // output layout claims.
17    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    // Fast path: in-place for unique, contiguous tensors at offset 0. This
63    // preserves the input tensor's dtype tag implicitly (the in-place SIMD
64    // ops flip bytes without touching the dtype tag).
65    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    // Allocating path for shared, non-contiguous, or offset tensors:
75    // preserve the input's bool dtype for the new tensor.
76    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}