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
59impl<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
87impl<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 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
195impl<D: Device> Tensor<D, Float> {
198 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 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}