Skip to main content

ruprim_host/reduce/
dispatch.rs

1use ruda_core::{bytes::Bytes, tensor::{DType, Shape, host::{HostTensor, Layout}}};
2use crate::unary;
3
4pub fn float_max_dim_with_indices(
5    tensor: HostTensor,
6    dim: usize,
7    indices_dtype: ruda_core::tensor::IntDType,
8) -> (HostTensor, HostTensor) {
9    let (values, indices) = crate::reduce::max_dim_with_indices(tensor, dim);
10    if indices.dtype() != DType::from(indices_dtype) {
11        (values, crate::cast::int_cast(indices, indices_dtype))
12    } else {
13        (values, indices)
14    }
15}
16
17pub fn float_min_dim_with_indices(
18    tensor: HostTensor,
19    dim: usize,
20    indices_dtype: ruda_core::tensor::IntDType,
21) -> (HostTensor, HostTensor) {
22    let (values, indices) = crate::reduce::min_dim_with_indices(tensor, dim);
23    if indices.dtype() != DType::from(indices_dtype) {
24        (values, crate::cast::int_cast(indices, indices_dtype))
25    } else {
26        (values, indices)
27    }
28}
29
30pub fn float_argmax(
31    tensor: HostTensor,
32    dim: usize,
33    out_dtype: ruda_core::tensor::IntDType,
34) -> HostTensor {
35    let result = crate::reduce::argmax(tensor, dim);
36    if result.dtype() != DType::from(out_dtype) {
37        crate::cast::int_cast(result, out_dtype)
38    } else {
39        result
40    }
41}
42
43pub fn float_argmin(
44    tensor: HostTensor,
45    dim: usize,
46    out_dtype: ruda_core::tensor::IntDType,
47) -> HostTensor {
48    let result = crate::reduce::argmin(tensor, dim);
49    if result.dtype() != DType::from(out_dtype) {
50        crate::cast::int_cast(result, out_dtype)
51    } else {
52        result
53    }
54}
55
56pub fn float_max_abs(tensor: HostTensor) -> HostTensor {
57    let abs = unary::abs(tensor);
58    crate::reduce::max(abs)
59}
60
61pub fn float_max_abs_dim(tensor: HostTensor, dim: usize) -> HostTensor {
62    let abs = unary::abs(tensor);
63    crate::reduce::max_dim(abs, dim)
64}
65
66pub fn int_mean(tensor: HostTensor) -> HostTensor {
67    let n = tensor.layout().num_elements();
68    assert!(n > 0, "int_mean: cannot take mean of empty tensor");
69    let dtype = tensor.dtype();
70    let sum_result = crate::reduce::sum(tensor);
71    // Compute in i64 to avoid truncation of n for small int types
72    macro_rules! compute_mean {
73        ($ty:ty) => {{
74            let data: &[$ty] = sum_result.storage();
75            let mean_val = (data[0] as i64 / n as i64) as $ty;
76            HostTensor::new(
77                Bytes::from_elems(alloc::vec![mean_val]),
78                Layout::contiguous(Shape::from(alloc::vec![1])),
79                dtype,
80            )
81        }};
82    }
83    match dtype {
84        DType::I64 => compute_mean!(i64),
85        DType::I32 => compute_mean!(i32),
86        DType::I16 => compute_mean!(i16),
87        DType::I8 => compute_mean!(i8),
88        other => panic!("int_mean: unsupported dtype {:?}", other),
89    }
90}
91
92pub fn int_max_abs(tensor: HostTensor) -> HostTensor {
93    let abs = crate::unary::int_abs(tensor);
94    crate::reduce::max(abs)
95}
96
97pub fn int_max_abs_dim(tensor: HostTensor, dim: usize) -> HostTensor {
98    let abs = crate::unary::int_abs(tensor);
99    crate::reduce::max_dim(abs, dim)
100}
101