ruprim_host/boolean/
indexing.rs1use 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 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