ruprim_host/reduce/
dispatch.rs1use 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 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