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: usize = fused_dims.iter().product();
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
249 let out_dims: Vec<usize> = src_dims
250 .iter()
251 .enumerate()
252 .filter(|(i, _)| *i != axis)
253 .map(|(_, &d)| d)
254 .collect();
255
256 let axis_len = src_dims[axis];
257 let axis_stride = src_strides[axis];
258
259 if out_dims.is_empty() {
260 let mut acc = init;
262 let mut offset = 0isize;
263 for _ in 0..axis_len {
264 let val = Op::apply(unsafe { *src_ptr.offset(offset) });
265 let mapped = map_fn(val);
266 acc = reduce_fn(acc, mapped);
267 offset += axis_stride;
268 }
269 let strides = col_major_strides(&[1]);
270 return StridedArray::from_parts(vec![acc], &[1], &strides, 0);
271 }
272
273 let total_out: usize = out_dims.iter().product();
274 let out_strides = col_major_strides(&out_dims);
275 let mut out =
276 StridedArray::from_parts(vec![init.clone(); total_out], &out_dims, &out_strides, 0)?;
277
278 let src_kept_strides: Vec<isize> = src_strides
280 .iter()
281 .enumerate()
282 .filter(|(i, _)| *i != axis)
283 .map(|(_, &s)| s)
284 .collect();
285
286 let elem_size = std::mem::size_of::<T>().max(std::mem::size_of::<U>());
287 let strides_list: [&[isize]; 2] = [&out_strides, &src_kept_strides];
288 let (fused_dims, ordered_strides, plan) =
289 build_plan_fused(&out_dims, &strides_list, Some(0), elem_size);
290
291 let out_ptr = out.view_mut().as_mut_ptr();
292
293 let initial_offsets = vec![0isize; ordered_strides.len()];
297 for_each_inner_block_preordered(
298 &fused_dims,
299 &plan.block,
300 &ordered_strides,
301 &initial_offsets,
302 |offsets, len, strides| {
303 let out_step = strides[0];
304 let src_step = strides[1];
305
306 if out_step == 1 && src_step == 1 && axis_len > 1 {
310 let n = len as usize;
311 let out_slice =
312 unsafe { std::slice::from_raw_parts_mut(out_ptr.offset(offsets[0]), n) };
313 let src0 = unsafe { std::slice::from_raw_parts(src_ptr.offset(offsets[1]), n) };
315 for i in 0..n {
316 out_slice[i] = map_fn(Op::apply(src0[i]));
317 }
318 for k in 1..axis_len {
320 let src_k = unsafe {
321 std::slice::from_raw_parts(
322 src_ptr.offset(offsets[1] + k as isize * axis_stride),
323 n,
324 )
325 };
326 for i in 0..n {
327 out_slice[i] = reduce_fn(out_slice[i].clone(), map_fn(Op::apply(src_k[i])));
328 }
329 }
330 return Ok(());
331 }
332
333 let mut out_off = offsets[0];
335 let mut src_off = offsets[1];
336 for _ in 0..len {
337 let mut acc = init.clone();
338 let mut ptr = unsafe { src_ptr.offset(src_off) };
339 for _ in 0..axis_len {
340 let val = Op::apply(unsafe { *ptr });
341 let mapped = map_fn(val);
342 acc = reduce_fn(acc, mapped);
343 unsafe {
344 ptr = ptr.offset(axis_stride);
345 }
346 }
347 unsafe {
348 *out_ptr.offset(out_off) = acc;
349 }
350 out_off += out_step;
351 src_off += src_step;
352 }
353 Ok(())
354 },
355 )?;
356
357 Ok(out)
358}