Skip to main content

luma_tensor/device/cpu/kernels/
reduce.rs

1//! Generic reduction kernels. Ported from luma-core `reduce.rs`.
2//!
3//! Reduces a single dimension at a time; multi-dim reductions are applied by the
4//! caller folding over dims (outermost-first, adjusting indices).
5
6use super::element::CpuNum;
7use super::iter::DimArray;
8use crate::{Layout, ReduceOp, Result, Shape};
9
10/// Which single-axis reduction to run.
11#[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
32/// Reduce over multiple dims. Reduces from the highest dim down (with keepdim to
33/// preserve dim indices), then optionally squeezes the reduced dims out.
34pub 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    // First reduction reads from the input storage + layout.
40    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    // reduce highest dim first so lower indices stay valid
49    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
90/// Reduce `x`/`layout` over a single `reduce_dim`. Returns the reduced buffer and
91/// its shape (with `reduce_dim` kept as size-1 if `keepdim`, else removed).
92pub 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
131/// argmin/argmax over a single dim. Returns `usize` indices (caller casts to the
132/// int storage dtype) and the result shape. Ties keep the first index.
133pub 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;