Skip to main content

ruprim_host/boolean/
indexing.rs

1use alloc::{vec, vec::Vec};
2use ruda_core::{bytes::Bytes, tensor::{DType, IntDType, Shape, host::{HostTensor, Layout}}};
3
4pub fn bool_select_or(
5    tensor: HostTensor,
6    dim: usize,
7    indices: HostTensor,
8    value: HostTensor,
9) -> HostTensor {
10    let mut result = crate::gather_scatter::select_add::<u8>(tensor, dim, indices, value);
11    // Clamp to 0/1: select_add sums u8 values, but bool OR saturates at 1
12    let storage: &mut [u8] = result.storage_mut();
13    for v in storage.iter_mut() {
14        if *v > 1 {
15            *v = 1;
16        }
17    }
18    result
19}
20
21pub async fn bool_argwhere(tensor: HostTensor, out_dtype: IntDType) -> HostTensor {
22    let tensor = tensor.to_contiguous();
23    let shape = tensor.layout().shape().clone();
24    let ndims = shape.num_dims();
25    let data: &[u8] = tensor.storage();
26    let n = shape.num_elements();
27
28    let count = data[..n].iter().filter(|&&v| v != 0).count();
29    let mut coords: Vec<isize> = Vec::with_capacity(count * ndims);
30    let strides = ruda_core::tensor::host::layout::contiguous_strides_usize(&shape);
31
32    for (flat_idx, &val) in data[..n].iter().enumerate() {
33        if val != 0 {
34            let mut remaining = flat_idx;
35            for &s in &strides {
36                coords.push((remaining / s) as isize);
37                remaining %= s;
38            }
39        }
40    }
41
42    let out_shape = Shape::from(vec![count, ndims]);
43    let result = HostTensor::new(
44        Bytes::from_elems(coords),
45        Layout::contiguous(out_shape),
46        ruda_core::tensor::host::dtype::INDEX_DTYPE,
47    );
48    if result.dtype() != DType::from(out_dtype) {
49        crate::cast::int_cast(result, out_dtype)
50    } else {
51        result
52    }
53}
54