Skip to main content

luma_tensor/ops/
reduce.rs

1use crate::{DTypeKind, Device, Dim, Dims, Float, Int, Layout, ReduceOp, Shape, Tensor, TensorMeta};
2
3pub trait ReduceDTypeKind<D: Device>: DTypeKind<D> {
4    fn arg_reduce_dispatch(
5        x: &Self::Storage,
6        layout: &Layout,
7        dim: usize,
8        keepdim: bool,
9        take_max: bool,
10    ) -> crate::Result<(D::IntStorage, Shape)>;
11    fn reduce_dispatch(x: &Self::Storage, l: &Layout, dims: &[usize], keepdim: bool, op: ReduceOp)
12    -> crate::Result<(Self::Storage, Shape)>;
13}
14
15impl<D: Device> ReduceDTypeKind<D> for Float {
16    fn arg_reduce_dispatch(
17        x: &Self::Storage,
18        layout: &crate::Layout,
19        dim: usize,
20        keepdim: bool,
21        take_max: bool,
22    ) -> crate::Result<(D::IntStorage, crate::Shape)> {
23        D::f_arg_reduce(x, layout, dim, keepdim, take_max)
24    }
25
26    fn reduce_dispatch(
27        x: &Self::Storage,
28        l: &Layout,
29        dims: &[usize],
30        keepdim: bool,
31        op: ReduceOp,
32    ) -> crate::Result<(Self::Storage, Shape)> {
33        D::f_reduce(x, l, dims, keepdim, op)
34    }
35}
36
37impl<D: Device> ReduceDTypeKind<D> for Int {
38    fn arg_reduce_dispatch(
39        x: &Self::Storage,
40        layout: &crate::Layout,
41        dim: usize,
42        keepdim: bool,
43        take_max: bool,
44    ) -> crate::Result<(D::IntStorage, crate::Shape)> {
45        D::i_arg_reduce(x, layout, dim, keepdim, take_max)
46    }
47
48    fn reduce_dispatch(
49        x: &Self::Storage,
50        l: &Layout,
51        dims: &[usize],
52        keepdim: bool,
53        op: ReduceOp,
54    ) -> crate::Result<(Self::Storage, Shape)> {
55        D::i_reduce(x, l, dims, keepdim, op)
56    }
57}
58
59// ============================================================================
60// Generic argmin/argmax for Float and Int
61// ============================================================================
62
63impl<D: Device, K: ReduceDTypeKind<D>> Tensor<D, K> {
64    pub fn argmin<Dm: Dim>(&self, dim: Dm) -> crate::Result<Tensor<D, Int>> {
65        self.arg_reduce_dispatch(dim, false, false)
66    }
67
68    pub fn argmin_keepdim<Dm: Dim>(&self, dim: Dm) -> crate::Result<Tensor<D, Int>> {
69        self.arg_reduce_dispatch(dim, true, false)
70    }
71
72    pub fn argmax<Dm: Dim>(&self, dim: Dm) -> crate::Result<Tensor<D, Int>> {
73        self.arg_reduce_dispatch(dim, false, true)
74    }
75
76    pub fn argmax_keepdim<Dm: Dim>(&self, dim: Dm) -> crate::Result<Tensor<D, Int>> {
77        self.arg_reduce_dispatch(dim, true, true)
78    }
79
80    fn arg_reduce_dispatch<Dm: Dim>(&self, dim: Dm, keepdim: bool, take_max: bool) -> crate::Result<Tensor<D, Int>> {
81        let d = dim.to_index(self.shape(), "argmin/argmax")?;
82        let (storage, shape) = K::arg_reduce_dispatch(&*self.storage_read()?, self.layout(), d, keepdim, take_max)?;
83        Ok(Tensor::<D, Int>::from_storage(storage, shape, ()))
84    }
85}
86
87// ============================================================================
88// Float-specific: mean and variance
89// ============================================================================
90
91impl<D: Device> Tensor<D, Float> {
92    pub fn mean<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
93        let d = dim.to_index(self.shape(), "mean")?;
94        self.f_reduce(&[d], false, ReduceOp::Mean)
95    }
96
97    pub fn mean_keepdim<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
98        let d = dim.to_index(self.shape(), "mean_keepdim")?;
99        self.f_reduce(&[d], true, ReduceOp::Mean)
100    }
101
102    pub fn mean_all(&self) -> crate::Result<Self> {
103        let dims: Vec<usize> = (0..self.rank()).collect();
104        self.f_reduce(&dims, false, ReduceOp::Mean)
105    }
106
107    pub fn var<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
108        let d = dim.to_index(self.shape(), "var")?;
109        self.var_impl(d, false, false)
110    }
111
112    pub fn var_keepdim<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
113        let d = dim.to_index(self.shape(), "var_keepdim")?;
114        self.var_impl(d, true, false)
115    }
116
117    pub fn var_all(&self) -> crate::Result<Self> {
118        self.flatten_all()?.var(0)
119    }
120
121    pub fn var_unbiased<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
122        let d = dim.to_index(self.shape(), "var_unbiased")?;
123        self.var_impl(d, false, true)
124    }
125
126    pub fn var_unbiased_keepdim<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
127        let d = dim.to_index(self.shape(), "var_unbiased_keepdim")?;
128        self.var_impl(d, true, true)
129    }
130
131    pub fn var_unbiased_all(&self) -> crate::Result<Self> {
132        self.flatten_all()?.var_unbiased(0)
133    }
134
135    /// `var = mean((x - mean(x))^2)`, optionally with Bessel correction.
136    fn var_impl(&self, dim: usize, keepdim: bool, unbiased: bool) -> crate::Result<Self> {
137        let mean = self.mean_keepdim(dim)?;
138        let diff = self.sub(&mean.broadcast_as(self.shape().clone())?)?;
139        let sq = diff.sqr()?;
140        let v = sq.mean_keepdim(dim)?;
141        let result = if keepdim { v.clone() } else { v.squeeze(dim)? };
142        if unbiased {
143            let n = self.dims()[dim] as f64;
144            result.mul_scalar(n / (n - 1.0))
145        } else {
146            Ok(result)
147        }
148    }
149
150    fn f_reduce(&self, dims: &[usize], keepdim: bool, op: ReduceOp) -> crate::Result<Self> {
151        let (s, shape) = D::f_reduce(&*self.storage_read()?, self.layout(), dims, keepdim, op)?;
152        let meta = <Float as DTypeKind<D>>::Meta::on_reduce(self, dims, op);
153        Ok(Self::from_storage(s, shape, meta))
154    }
155}
156
157impl<D: Device, K: ReduceDTypeKind<D>> Tensor<D, K> {
158    fn reduce_impl(&self, dims: &[usize], keepdim: bool, op: ReduceOp) -> crate::Result<Self> {
159        let (s, shape) = K::reduce_dispatch(&*self.storage_read()?, self.layout(), dims, keepdim, op)?;
160        Ok(Self::from_storage(s, shape, K::Meta::on_reduce(self, dims, op)))
161    }
162}
163
164macro_rules! reduce_dispatch {
165    ($name:ident, $keep:ident, $all:ident, $variant:ident) => {
166        pub fn $name<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
167            let d = dim.to_index(self.shape(), stringify!($name))?;
168            self.reduce_impl(&[d], false, ReduceOp::$variant)
169        }
170
171        pub fn $keep<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
172            let d = dim.to_index(self.shape(), stringify!($keep))?;
173            self.reduce_impl(&[d], true, ReduceOp::$variant)
174        }
175
176        pub fn $all(&self) -> crate::Result<Self> {
177            let dims: Vec<usize> = (0..self.rank()).collect();
178            self.reduce_impl(&dims, false, ReduceOp::$variant)
179        }
180    };
181}
182
183impl<D: Device, K: ReduceDTypeKind<D>> Tensor<D, K> {
184    reduce_dispatch!(sum, sum_keepdim, sum_all, Sum);
185    reduce_dispatch!(max, max_keepdim, max_all, Max);
186    reduce_dispatch!(min, min_keepdim, min_all, Min);
187    reduce_dispatch!(prod, prod_keepdim, prod_all, Prod);
188
189    pub fn sum_dims<Ds: Dims>(&self, dims: Ds, keepdim: bool) -> crate::Result<Self> {
190        let dims = dims.to_indexes(self.shape(), "sum_dims")?;
191        self.reduce_impl(&dims, keepdim, ReduceOp::Sum)
192    }
193}
194
195// ---- Float-only: std, logsumexp ----
196
197impl<D: Device> Tensor<D, Float> {
198    /// Standard deviation along `dim`.
199    pub fn std<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
200        self.var(dim)?.sqrt()
201    }
202
203    pub fn std_keepdim<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
204        self.var_keepdim(dim)?.sqrt()
205    }
206
207    pub fn std_all(&self) -> crate::Result<Self> {
208        self.var_all()?.sqrt()
209    }
210
211    /// Log-sum-exp along `dim`, numerically stable.
212    pub fn logsumexp<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
213        let m = self.max_keepdim(dim)?;
214        let e = self.sub(&m.broadcast_as(self.shape().clone())?)?.exp()?;
215        let s = e.sum_keepdim(dim)?;
216        s.ln()?.add(&m)?.squeeze(dim)
217    }
218
219    pub fn logsumexp_keepdim<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
220        let m = self.max_keepdim(dim)?;
221        let e = self.sub(&m.broadcast_as(self.shape().clone())?)?.exp()?;
222        e.sum_keepdim(dim)?.ln()?.add(&m)
223    }
224
225    pub fn logsumexp_all(&self) -> crate::Result<Self> {
226        self.flatten_all()?.logsumexp(0)
227    }
228}