Skip to main content

ruprim_host/comparison/
predicate_reduce.rs

1use super::*;
2
3// ============================================================================
4// any / all operations
5// ============================================================================
6
7/// Check if any element is non-zero (float tensors).
8pub fn any_float(tensor: HostTensor, out_dtype: BoolDType) -> HostTensor {
9    let has_any = match tensor.dtype() {
10        DType::F32 => iter_elements::<f32>(&tensor).any(|x| x != 0.0),
11        DType::F64 => iter_elements::<f64>(&tensor).any(|x| x != 0.0),
12        DType::F16 => iter_elements::<f16>(&tensor).any(|x: f16| x.to_f32() != 0.0),
13        DType::BF16 => iter_elements::<bf16>(&tensor).any(|x: bf16| x.to_f32() != 0.0),
14        _ => panic!("any_float: unsupported dtype {:?}", tensor.dtype()),
15    };
16    bool_scalar(has_any, out_dtype)
17}
18
19/// Check if any element along a dimension is non-zero (float tensors).
20pub fn any_float_dim(tensor: HostTensor, dim: usize, out_dtype: BoolDType) -> HostTensor {
21    reduce_bool_dim(&tensor, dim, false, |a, b| a || b, out_dtype)
22}
23
24/// Check if all elements are non-zero (float tensors).
25pub fn all_float(tensor: HostTensor, out_dtype: BoolDType) -> HostTensor {
26    let all = match tensor.dtype() {
27        DType::F32 => iter_elements::<f32>(&tensor).all(|x| x != 0.0),
28        DType::F64 => iter_elements::<f64>(&tensor).all(|x| x != 0.0),
29        DType::F16 => iter_elements::<f16>(&tensor).all(|x: f16| x.to_f32() != 0.0),
30        DType::BF16 => iter_elements::<bf16>(&tensor).all(|x: bf16| x.to_f32() != 0.0),
31        _ => panic!("all_float: unsupported dtype {:?}", tensor.dtype()),
32    };
33    bool_scalar(all, out_dtype)
34}
35
36/// Check if all elements along a dimension are non-zero (float tensors).
37pub fn all_float_dim(tensor: HostTensor, dim: usize, out_dtype: BoolDType) -> HostTensor {
38    reduce_bool_dim(&tensor, dim, true, |a, b| a && b, out_dtype)
39}
40
41/// Check if any element is non-zero (int tensors).
42pub fn any_int(tensor: HostTensor, out_dtype: BoolDType) -> HostTensor {
43    let has_any = match tensor.dtype() {
44        DType::I64 => iter_elements::<i64>(&tensor).any(|x| x != 0),
45        DType::I32 => iter_elements::<i32>(&tensor).any(|x| x != 0),
46        DType::I16 => iter_elements::<i16>(&tensor).any(|x| x != 0),
47        DType::I8 => iter_elements::<i8>(&tensor).any(|x| x != 0),
48        DType::U64 => iter_elements::<u64>(&tensor).any(|x| x != 0),
49        DType::U32 => iter_elements::<u32>(&tensor).any(|x| x != 0),
50        DType::U16 => iter_elements::<u16>(&tensor).any(|x| x != 0),
51        DType::U8 => iter_elements::<u8>(&tensor).any(|x| x != 0),
52        _ => panic!("any_int: unsupported dtype {:?}", tensor.dtype()),
53    };
54    bool_scalar(has_any, out_dtype)
55}
56
57/// Check if any element along a dimension is non-zero (int tensors).
58pub fn any_int_dim(tensor: HostTensor, dim: usize, out_dtype: BoolDType) -> HostTensor {
59    reduce_bool_dim_int(&tensor, dim, false, |a, b| a || b, out_dtype)
60}
61
62/// Check if all elements are non-zero (int tensors).
63pub fn all_int(tensor: HostTensor, out_dtype: BoolDType) -> HostTensor {
64    let all = match tensor.dtype() {
65        DType::I64 => iter_elements::<i64>(&tensor).all(|x| x != 0),
66        DType::I32 => iter_elements::<i32>(&tensor).all(|x| x != 0),
67        DType::I16 => iter_elements::<i16>(&tensor).all(|x| x != 0),
68        DType::I8 => iter_elements::<i8>(&tensor).all(|x| x != 0),
69        DType::U64 => iter_elements::<u64>(&tensor).all(|x| x != 0),
70        DType::U32 => iter_elements::<u32>(&tensor).all(|x| x != 0),
71        DType::U16 => iter_elements::<u16>(&tensor).all(|x| x != 0),
72        DType::U8 => iter_elements::<u8>(&tensor).all(|x| x != 0),
73        _ => panic!("all_int: unsupported dtype {:?}", tensor.dtype()),
74    };
75    bool_scalar(all, out_dtype)
76}
77
78/// Check if all elements along a dimension are non-zero (int tensors).
79pub fn all_int_dim(tensor: HostTensor, dim: usize, out_dtype: BoolDType) -> HostTensor {
80    reduce_bool_dim_int(&tensor, dim, true, |a, b| a && b, out_dtype)
81}
82
83/// Check if any bool element is true.
84pub fn any_bool(tensor: HostTensor, out_dtype: BoolDType) -> HostTensor {
85    let tensor = tensor.to_contiguous();
86    let data: &[u8] = tensor.bytes();
87    bool_scalar(data.iter().any(|&x| x != 0), out_dtype)
88}
89
90/// Check if any bool element along a dimension is true.
91pub fn any_bool_dim(tensor: HostTensor, dim: usize, out_dtype: BoolDType) -> HostTensor {
92    reduce_bool_dim_raw(&tensor, dim, false, |a, b| a || b, out_dtype)
93}
94
95/// Check if all bool elements are true.
96pub fn all_bool(tensor: HostTensor, out_dtype: BoolDType) -> HostTensor {
97    let tensor = tensor.to_contiguous();
98    let data: &[u8] = tensor.bytes();
99    bool_scalar(data.iter().all(|&x| x != 0), out_dtype)
100}
101
102/// Check if all bool elements along a dimension are true.
103pub fn all_bool_dim(tensor: HostTensor, dim: usize, out_dtype: BoolDType) -> HostTensor {
104    reduce_bool_dim_raw(&tensor, dim, true, |a, b| a && b, out_dtype)
105}
106