1use crate::kernel::{
4 build_plan_fused, for_each_inner_block_preordered, same_contiguous_layout,
5 sequential_contiguous_layout, total_len,
6};
7use crate::maybe_sync::{MaybeSendSync, MaybeSync};
8use crate::simd;
9use crate::view::{col_major_strides, StridedArray, StridedView};
10use crate::{Result, StridedError};
11use strided_view::ElementOp;
12
13#[cfg(feature = "parallel")]
14use crate::fuse::compute_costs;
15#[cfg(feature = "parallel")]
16use crate::threading::{
17 for_each_inner_block_with_offsets, mapreduce_threaded, SendPtr, MINTHREADLENGTH,
18};
19
20pub fn reduce<T: Copy + MaybeSendSync, Op: ElementOp<T>, M, R, U>(
22 src: &StridedView<T, Op>,
23 map_fn: M,
24 reduce_fn: R,
25 init: U,
26) -> Result<U>
27where
28 M: Fn(T) -> U + MaybeSync,
29 R: Fn(U, U) -> U + MaybeSync,
30 U: Clone + MaybeSendSync,
31{
32 reduce_impl(src, map_fn, reduce_fn, init, true)
33}
34
35pub(crate) fn reduce_serial<T: Copy + MaybeSendSync, Op: ElementOp<T>, M, R, U>(
36 src: &StridedView<T, Op>,
37 map_fn: M,
38 reduce_fn: R,
39 init: U,
40) -> Result<U>
41where
42 M: Fn(T) -> U + MaybeSync,
43 R: Fn(U, U) -> U + MaybeSync,
44 U: Clone + MaybeSendSync,
45{
46 reduce_impl(src, map_fn, reduce_fn, init, false)
47}
48
49fn reduce_impl<T: Copy + MaybeSendSync, Op: ElementOp<T>, M, R, U>(
50 src: &StridedView<T, Op>,
51 map_fn: M,
52 reduce_fn: R,
53 init: U,
54 allow_ambient_parallel: bool,
55) -> Result<U>
56where
57 M: Fn(T) -> U + MaybeSync,
58 R: Fn(U, U) -> U + MaybeSync,
59 U: Clone + MaybeSendSync,
60{
61 let src_ptr = src.ptr();
62 let src_dims = src.dims();
63 let src_strides = src.strides();
64
65 let contiguous = if allow_ambient_parallel {
66 sequential_contiguous_layout(src_dims, &[src_strides])?
67 } else {
68 same_contiguous_layout(src_dims, &[src_strides])
69 };
70 if contiguous.is_some() {
71 let len = total_len(src_dims)?;
72 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
73 return Ok(simd::dispatch_if_large(len, || {
74 let mut acc = init;
75 for &val in src.iter() {
76 acc = reduce_fn(acc, map_fn(Op::apply(val)));
77 }
78 acc
79 }));
80 }
81
82 #[cfg(feature = "parallel")]
86 {
87 let total = total_len(src_dims)?;
88 let nthreads = if allow_ambient_parallel {
89 crate::execution_policy::rayon_threads()
90 } else {
91 1
92 };
93 if total > MINTHREADLENGTH
94 && nthreads > 1
95 && same_contiguous_layout(src_dims, &[src_strides]).is_some()
96 {
97 let src_slice = unsafe { std::slice::from_raw_parts(src_ptr, total) };
98 let result = crate::threading::parallel_map_reduce(
99 0..total,
100 nthreads,
101 &|range| {
102 simd::dispatch_if_large(range.len(), || {
103 let mut acc = init.clone();
104 for &val in &src_slice[range] {
105 acc = reduce_fn(acc, map_fn(Op::apply(val)));
106 }
107 acc
108 })
109 },
110 &|a, b| reduce_fn(a, b),
111 );
112 return Ok(result);
113 }
114 }
115
116 let strides_list: [&[isize]; 1] = [src_strides];
117
118 let (fused_dims, ordered_strides, plan) =
119 build_plan_fused(src_dims, &strides_list, None, std::mem::size_of::<T>());
120
121 #[cfg(feature = "parallel")]
122 {
123 let total = total_len(&fused_dims)?;
124 let nthreads = if allow_ambient_parallel {
125 crate::execution_policy::rayon_threads()
126 } else {
127 1
128 };
129 if total > MINTHREADLENGTH && nthreads > 1 {
130 let spacing = (64 / std::mem::size_of::<U>()).max(1);
132 let mut threadedout = vec![init.clone(); spacing * nthreads];
133 let threadedout_ptr = SendPtr(threadedout.as_mut_ptr());
134 let src_send = SendPtr(src_ptr as *mut T);
135
136 let costs = compute_costs(&ordered_strides);
137
138 let ndim = fused_dims.len();
142 let mut threaded_strides = Vec::with_capacity(ordered_strides.len() + 1);
143 threaded_strides.push(vec![0isize; ndim]); for s in &ordered_strides {
145 threaded_strides.push(s.clone());
146 }
147 let initial_offsets = vec![0isize; threaded_strides.len()];
148
149 mapreduce_threaded(
156 &fused_dims,
157 &plan.block,
158 &threaded_strides,
159 &initial_offsets,
160 &costs,
161 nthreads,
162 spacing as isize,
163 1,
164 &|dims, blocks, strides_list, offsets| {
165 let out_offset = offsets[0] as usize;
168 let src_offsets = &offsets[1..];
169
170 for_each_inner_block_with_offsets(
171 dims,
172 blocks,
173 &strides_list[1..],
174 src_offsets,
175 |offsets, len, strides| {
176 let mut ptr = unsafe { src_send.as_const().offset(offsets[0]) };
177 let stride = strides[0];
178 let slot = unsafe { &mut *threadedout_ptr.as_ptr().add(out_offset) };
179 for _ in 0..len {
180 let val = Op::apply(unsafe { *ptr });
181 let mapped = map_fn(val);
182 *slot = reduce_fn(slot.clone(), mapped);
183 unsafe {
184 ptr = ptr.offset(stride);
185 }
186 }
187 Ok(())
188 },
189 )
190 },
191 )?;
192
193 let mut result = init;
195 for i in 0..nthreads {
196 result = reduce_fn(result, threadedout[i * spacing].clone());
197 }
198 return Ok(result);
199 }
200 }
201
202 let mut acc = init;
203 let initial_offsets = vec![0isize; ordered_strides.len()];
204 for_each_inner_block_preordered(
205 &fused_dims,
206 &plan.block,
207 &ordered_strides,
208 &initial_offsets,
209 |offsets, len, strides| {
210 let mut ptr = unsafe { src_ptr.offset(offsets[0]) };
211 let stride = strides[0];
212 for _ in 0..len {
213 let val = Op::apply(unsafe { *ptr });
214 let mapped = map_fn(val);
215 acc = reduce_fn(acc.clone(), mapped);
216 unsafe {
217 ptr = ptr.offset(stride);
218 }
219 }
220 Ok(())
221 },
222 )?;
223
224 Ok(acc)
225}
226
227pub fn reduce_axis<T: Copy + MaybeSendSync, Op: ElementOp<T>, M, R, U>(
229 src: &StridedView<T, Op>,
230 axis: usize,
231 map_fn: M,
232 reduce_fn: R,
233 init: U,
234) -> Result<StridedArray<U>>
235where
236 M: Fn(T) -> U + MaybeSync,
237 R: Fn(U, U) -> U + MaybeSync,
238 U: Clone + MaybeSendSync,
239{
240 let rank = src.ndim();
241 if axis >= rank {
242 return Err(StridedError::InvalidAxis { axis, rank });
243 }
244
245 let src_dims = src.dims();
246 let src_strides = src.strides();
247 let src_ptr = src.ptr();
248 total_len(src_dims)?;
251
252 let out_dims: Vec<usize> = src_dims
253 .iter()
254 .enumerate()
255 .filter(|(i, _)| *i != axis)
256 .map(|(_, &d)| d)
257 .collect();
258
259 let axis_len = src_dims[axis];
260 let axis_stride = src_strides[axis];
261
262 if out_dims.is_empty() {
263 let mut acc = init;
265 let mut offset = 0isize;
266 for _ in 0..axis_len {
267 let val = Op::apply(unsafe { *src_ptr.offset(offset) });
268 let mapped = map_fn(val);
269 acc = reduce_fn(acc, mapped);
270 offset += axis_stride;
271 }
272 let strides = col_major_strides(&[1]);
273 return StridedArray::from_parts(vec![acc], &[1], &strides, 0);
274 }
275
276 let total_out = total_len(&out_dims)?;
277 let out_strides = col_major_strides(&out_dims);
278 let mut out =
279 StridedArray::from_parts(vec![init.clone(); total_out], &out_dims, &out_strides, 0)?;
280
281 let src_kept_strides: Vec<isize> = src_strides
283 .iter()
284 .enumerate()
285 .filter(|(i, _)| *i != axis)
286 .map(|(_, &s)| s)
287 .collect();
288
289 let elem_size = std::mem::size_of::<T>().max(std::mem::size_of::<U>());
290 let strides_list: [&[isize]; 2] = [&out_strides, &src_kept_strides];
291 let (fused_dims, ordered_strides, plan) =
292 build_plan_fused(&out_dims, &strides_list, Some(0), elem_size);
293
294 let out_ptr = out.view_mut().as_mut_ptr();
295
296 let initial_offsets = vec![0isize; ordered_strides.len()];
300 for_each_inner_block_preordered(
301 &fused_dims,
302 &plan.block,
303 &ordered_strides,
304 &initial_offsets,
305 |offsets, len, strides| {
306 let out_step = strides[0];
307 let src_step = strides[1];
308
309 if out_step == 1 && src_step == 1 && axis_len > 1 {
313 let n = len as usize;
314 let out_slice =
315 unsafe { std::slice::from_raw_parts_mut(out_ptr.offset(offsets[0]), n) };
316 let src0 = unsafe { std::slice::from_raw_parts(src_ptr.offset(offsets[1]), n) };
320 for i in 0..n {
321 out_slice[i] = reduce_fn(init.clone(), map_fn(Op::apply(src0[i])));
322 }
323 for k in 1..axis_len {
325 let src_k = unsafe {
326 std::slice::from_raw_parts(
327 src_ptr.offset(offsets[1] + k as isize * axis_stride),
328 n,
329 )
330 };
331 for i in 0..n {
332 out_slice[i] = reduce_fn(out_slice[i].clone(), map_fn(Op::apply(src_k[i])));
333 }
334 }
335 return Ok(());
336 }
337
338 let mut out_off = offsets[0];
340 let mut src_off = offsets[1];
341 for _ in 0..len {
342 let mut acc = init.clone();
343 let mut ptr = unsafe { src_ptr.offset(src_off) };
344 for _ in 0..axis_len {
345 let val = Op::apply(unsafe { *ptr });
346 let mapped = map_fn(val);
347 acc = reduce_fn(acc, mapped);
348 unsafe {
349 ptr = ptr.offset(axis_stride);
350 }
351 }
352 unsafe {
353 *out_ptr.offset(out_off) = acc;
354 }
355 out_off += out_step;
356 src_off += src_step;
357 }
358 Ok(())
359 },
360 )?;
361
362 Ok(out)
363}