luma_tensor/device/cpu/kernels/
reduce.rs1use super::element::CpuNum;
7use super::iter::DimArray;
8use crate::{Layout, ReduceOp, Result, Shape};
9
10#[derive(Clone, Copy)]
12pub enum Reducer {
13 Sum,
14 Mean,
15 Min,
16 Max,
17 Product,
18}
19
20impl From<ReduceOp> for Reducer {
21 fn from(op: ReduceOp) -> Self {
22 match op {
23 ReduceOp::Sum => Reducer::Sum,
24 ReduceOp::Mean => Reducer::Mean,
25 ReduceOp::Min => Reducer::Min,
26 ReduceOp::Max => Reducer::Max,
27 ReduceOp::Prod => Reducer::Product,
28 }
29 }
30}
31
32pub fn reduce_dims<T: CpuNum>(x: &[T], layout: &Layout, dims: &[usize], keepdim: bool, reducer: Reducer) -> Result<(Vec<T>, Shape)> {
35 let mut sorted: Vec<usize> = dims.to_vec();
36 sorted.sort_unstable();
37 sorted.dedup();
38
39 let mut cur_layout = layout.clone();
41 let mut cur_data: Vec<T>;
42 let (mut data, mut shape): (Vec<T>, Shape);
43
44 if sorted.is_empty() {
45 return Ok((crate::device::cpu::kernels::iter::gather(x, layout), layout.shape().clone()));
46 }
47
48 let mut iter = sorted.iter().rev();
50 let first = *iter.next().unwrap();
51 let (d, s) = reduce_dim(x, &cur_layout, first, true, reducer)?;
52 data = d;
53 shape = s;
54 for &dim in iter {
55 cur_layout = Layout::contiguous(shape.clone());
56 cur_data = data;
57 let (d, s) = reduce_dim(&cur_data, &cur_layout, dim, true, reducer)?;
58 data = d;
59 shape = s;
60 }
61
62 if !keepdim {
63 let mut out_dims: Vec<usize> = Vec::new();
64 for (i, &d) in shape.dims().iter().enumerate() {
65 if !sorted.contains(&i) {
66 out_dims.push(d);
67 }
68 }
69 shape = Shape::from(out_dims);
70 }
71 Ok((data, shape))
72}
73
74impl Reducer {
75 fn apply<T: CpuNum, I: Iterator<Item = T> + ExactSizeIterator>(&self, iter: I) -> T {
76 match self {
77 Reducer::Sum => iter.sum(),
78 Reducer::Mean => {
79 let len = iter.len();
80 let s: T = iter.sum();
81 s / T::from_usize(len.max(1))
82 }
83 Reducer::Min => iter.reduce(|a, b| T::minimum(a, b)).unwrap_or(T::ZERO),
84 Reducer::Max => iter.reduce(|a, b| T::maximum(a, b)).unwrap_or(T::ZERO),
85 Reducer::Product => iter.product(),
86 }
87 }
88}
89
90pub fn reduce_dim<T: CpuNum>(x: &[T], layout: &Layout, reduce_dim: usize, keepdim: bool, reducer: Reducer) -> Result<(Vec<T>, Shape)> {
93 let reduce_dim_stride = layout.stride()[reduce_dim];
94 let reduce_dim_size = layout.dims()[reduce_dim];
95
96 let dst: Vec<T> = if layout.is_contiguous() && reduce_dim_stride == 1 {
97 let x = &x[layout.start_offset()..];
98 (0..layout.element_count() / reduce_dim_size)
99 .map(|i| {
100 let chunk = &x[i * reduce_dim_size..i * reduce_dim_size + reduce_dim_size];
101 reducer.apply(chunk.iter().copied())
102 })
103 .collect()
104 } else {
105 let dst_len = layout.element_count() / reduce_dim_size;
106 let mut dst: Vec<T> = Vec::with_capacity(dst_len);
107 let collapsed = layout.narrow(reduce_dim, 0, 1)?;
108 if reduce_dim_stride == 1 {
109 for src_index in collapsed.storage_indices() {
110 let chunk = &x[src_index..src_index + reduce_dim_size];
111 dst.push(reducer.apply(chunk.iter().copied()));
112 }
113 } else {
114 for src_index in collapsed.storage_indices() {
115 let arr = DimArray::new(&x[src_index..], reduce_dim_size, reduce_dim_stride);
116 dst.push(reducer.apply(arr.into_iter()));
117 }
118 }
119 dst
120 };
121
122 let mut shape = layout.dims().to_vec();
123 if keepdim {
124 shape[reduce_dim] = 1;
125 } else {
126 shape.remove(reduce_dim);
127 }
128 Ok((dst, Shape::from(shape)))
129}
130
131pub fn arg_reduce<T: CpuNum>(x: &[T], layout: &Layout, dim: usize, keepdim: bool, take_max: bool) -> Result<(Vec<usize>, Shape)> {
134 let reduce_dim_stride = layout.stride()[dim];
135 let reduce_dim_size = layout.dims()[dim];
136
137 let arg = |iter: DimArrayIter<T>| -> usize {
138 iter.enumerate()
139 .reduce(|(ia, a), (ib, b)| {
140 let keep_a = if take_max { a >= b } else { a <= b };
141 if keep_a { (ia, a) } else { (ib, b) }
142 })
143 .map(|(i, _)| i)
144 .unwrap_or(0)
145 };
146
147 let dst_len = layout.element_count() / reduce_dim_size;
148 let mut dst: Vec<usize> = Vec::with_capacity(dst_len);
149 let collapsed = layout.narrow(dim, 0, 1)?;
150 for src_index in collapsed.storage_indices() {
151 let arr = DimArray::new(&x[src_index..], reduce_dim_size, reduce_dim_stride);
152 dst.push(arg(arr.into_iter()));
153 }
154
155 let mut shape = layout.dims().to_vec();
156 if keepdim {
157 shape[dim] = 1;
158 } else {
159 shape.remove(dim);
160 }
161 Ok((dst, Shape::from(shape)))
162}
163
164use super::iter::DimArrayIter;