1#[cfg(feature = "parallel")]
4use crate::kernel::same_contiguous_layout;
5use crate::kernel::{
6 build_plan_fused, for_each_inner_block_preordered, sequential_contiguous_layout, total_len,
7};
8use crate::maybe_sync::{MaybeSendSync, MaybeSync};
9use crate::simd;
10use crate::view::{col_major_strides, StridedArray, StridedView};
11use crate::{Result, StridedError};
12use std::ops::Range;
13use strided_view::ElementOp;
14
15#[cfg(feature = "parallel")]
16use crate::fuse::compute_costs;
17#[cfg(feature = "parallel")]
18use crate::threading::{
19 for_each_inner_block_with_offsets, mapreduce_threaded, SendPtr, MINTHREADLENGTH,
20};
21
22const FOLD_LANES: usize = 16;
24
25const AXIS_SWEEP_MIN_LEN: usize = 16;
28
29const AXIS_SWEEP_BLOCK_BYTES: usize = 64 * 1024;
31
32#[inline(always)]
42fn fold_contiguous<T, Op, M, R, U>(src: &[T], init: U, map_fn: &M, reduce_fn: &R) -> U
43where
44 T: Copy,
45 Op: ElementOp<T>,
46 M: Fn(T) -> U,
47 R: Fn(U, U) -> U,
48 U: Clone,
49{
50 if src.len() < 2 * FOLD_LANES {
51 let mut acc = init;
52 for &value in src {
53 acc = reduce_fn(acc, map_fn(Op::apply(value)));
54 }
55 return acc;
56 }
57 let (head, rest) = src.split_at(FOLD_LANES);
58 let mut lanes: [U; FOLD_LANES] = core::array::from_fn(|lane| map_fn(Op::apply(head[lane])));
59 let mut chunks = rest.chunks_exact(FOLD_LANES);
60 for chunk in chunks.by_ref() {
61 for (lane, &value) in lanes.iter_mut().zip(chunk) {
62 *lane = reduce_fn(lane.clone(), map_fn(Op::apply(value)));
63 }
64 }
65 for (lane, &value) in lanes.iter_mut().zip(chunks.remainder()) {
66 *lane = reduce_fn(lane.clone(), map_fn(Op::apply(value)));
67 }
68 lanes.into_iter().fold(init, reduce_fn)
69}
70
71#[inline(always)]
76unsafe fn fold_run<T, Op, M, R, U>(
77 ptr: *const T,
78 stride: isize,
79 len: usize,
80 init: U,
81 map_fn: &M,
82 reduce_fn: &R,
83) -> U
84where
85 T: Copy,
86 Op: ElementOp<T>,
87 M: Fn(T) -> U,
88 R: Fn(U, U) -> U,
89 U: Clone,
90{
91 if stride == 1 {
92 let run = unsafe { std::slice::from_raw_parts(ptr, len) };
94 return fold_contiguous::<T, Op, M, R, U>(run, init, map_fn, reduce_fn);
95 }
96 let mut acc = init;
97 let mut cursor = ptr;
98 for index in 0..len {
99 acc = reduce_fn(acc, map_fn(Op::apply(unsafe { *cursor })));
101 if index + 1 < len {
102 cursor = cursor.wrapping_offset(stride);
103 }
104 }
105 acc
106}
107
108pub fn reduce<T: Copy + MaybeSendSync, Op: ElementOp<T>, M, R, U>(
126 src: &StridedView<T, Op>,
127 map_fn: M,
128 reduce_fn: R,
129 init: U,
130) -> Result<U>
131where
132 M: Fn(T) -> U + MaybeSync,
133 R: Fn(U, U) -> U + MaybeSync,
134 U: Clone + MaybeSendSync,
135{
136 reduce_impl(src, map_fn, reduce_fn, init)
137}
138
139fn reduce_impl<T: Copy + MaybeSendSync, Op: ElementOp<T>, M, R, U>(
140 src: &StridedView<T, Op>,
141 map_fn: M,
142 reduce_fn: R,
143 init: U,
144) -> Result<U>
145where
146 M: Fn(T) -> U + MaybeSync,
147 R: Fn(U, U) -> U + MaybeSync,
148 U: Clone + MaybeSendSync,
149{
150 let src_ptr = src.ptr();
151 let src_dims = src.dims();
152 let src_strides = src.strides();
153
154 let contiguous = sequential_contiguous_layout(src_dims, &[src_strides])?;
155 if contiguous.is_some() {
156 let len = total_len(src_dims)?;
157 let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
158 return Ok(simd::dispatch_if_large(len, || {
159 fold_contiguous::<T, Op, M, R, U>(src, init, &map_fn, &reduce_fn)
160 }));
161 }
162
163 #[cfg(feature = "parallel")]
167 {
168 let total = total_len(src_dims)?;
169 let nthreads = crate::execution_policy::rayon_threads();
170 if total > MINTHREADLENGTH
171 && nthreads > 1
172 && same_contiguous_layout(src_dims, &[src_strides]).is_some()
173 {
174 let src_slice = unsafe { std::slice::from_raw_parts(src_ptr, total) };
175 let result = crate::threading::parallel_map_reduce(
176 0..total,
177 nthreads,
178 &|range| {
179 simd::dispatch_if_large(range.len(), || {
180 fold_contiguous::<T, Op, M, R, U>(
181 &src_slice[range],
182 init.clone(),
183 &map_fn,
184 &reduce_fn,
185 )
186 })
187 },
188 &|a, b| reduce_fn(a, b),
189 );
190 return Ok(result);
191 }
192 }
193
194 let strides_list: [&[isize]; 1] = [src_strides];
195
196 let (fused_dims, ordered_strides, plan) =
197 build_plan_fused(src_dims, &strides_list, None, std::mem::size_of::<T>());
198
199 #[cfg(feature = "parallel")]
200 {
201 let total = total_len(&fused_dims)?;
202 let nthreads = crate::execution_policy::rayon_threads();
203 if total > MINTHREADLENGTH && nthreads > 1 {
204 let spacing = (64 / std::mem::size_of::<U>()).max(1);
206 let mut threadedout = vec![init.clone(); spacing * nthreads];
207 let threadedout_ptr = SendPtr(threadedout.as_mut_ptr());
208 let src_send = SendPtr(src_ptr as *mut T);
209
210 let costs = compute_costs(&ordered_strides);
211
212 let ndim = fused_dims.len();
216 let mut threaded_strides = Vec::with_capacity(ordered_strides.len() + 1);
217 threaded_strides.push(vec![0isize; ndim]); for s in &ordered_strides {
219 threaded_strides.push(s.clone());
220 }
221 let initial_offsets = vec![0isize; threaded_strides.len()];
222
223 mapreduce_threaded(
230 &fused_dims,
231 &plan.block,
232 &threaded_strides,
233 &initial_offsets,
234 &costs,
235 nthreads,
236 spacing as isize,
237 1,
238 &|dims, blocks, strides_list, offsets| {
239 let out_offset = offsets[0] as usize;
242 let src_offsets = &offsets[1..];
243
244 for_each_inner_block_with_offsets(
245 dims,
246 blocks,
247 &strides_list[1..],
248 src_offsets,
249 |offsets, len, strides| {
250 let mut ptr = unsafe { src_send.as_const().offset(offsets[0]) };
251 let stride = strides[0];
252 let slot = unsafe { &mut *threadedout_ptr.as_ptr().add(out_offset) };
253 for _ in 0..len {
254 let val = Op::apply(unsafe { *ptr });
255 let mapped = map_fn(val);
256 *slot = reduce_fn(slot.clone(), mapped);
257 unsafe {
258 ptr = ptr.offset(stride);
259 }
260 }
261 Ok(())
262 },
263 )
264 },
265 )?;
266
267 let mut result = init;
269 for i in 0..nthreads {
270 result = reduce_fn(result, threadedout[i * spacing].clone());
271 }
272 return Ok(result);
273 }
274 }
275
276 let mut acc = init;
277 let initial_offsets = vec![0isize; ordered_strides.len()];
278 for_each_inner_block_preordered(
279 &fused_dims,
280 &plan.block,
281 &ordered_strides,
282 &initial_offsets,
283 |offsets, len, strides| {
284 let mut ptr = unsafe { src_ptr.offset(offsets[0]) };
285 let stride = strides[0];
286 for _ in 0..len {
287 let val = Op::apply(unsafe { *ptr });
288 let mapped = map_fn(val);
289 acc = reduce_fn(acc.clone(), mapped);
290 unsafe {
291 ptr = ptr.offset(stride);
292 }
293 }
294 Ok(())
295 },
296 )?;
297
298 Ok(acc)
299}
300
301pub fn reduce_axis<T: Copy + MaybeSendSync, Op: ElementOp<T>, M, R, U>(
313 src: &StridedView<T, Op>,
314 axis: usize,
315 map_fn: M,
316 reduce_fn: R,
317 init: U,
318) -> Result<StridedArray<U>>
319where
320 M: Fn(T) -> U + MaybeSync,
321 R: Fn(U, U) -> U + MaybeSync,
322 U: Clone + MaybeSendSync,
323{
324 let rank = src.ndim();
325 if axis >= rank {
326 return Err(StridedError::InvalidAxis { axis, rank });
327 }
328
329 let src_dims = src.dims();
330 let src_strides = src.strides();
331 let src_ptr = src.ptr();
332 total_len(src_dims)?;
335
336 let axis_len = src_dims[axis];
337 let axis_stride = src_strides[axis];
338
339 let kept: Vec<(usize, isize)> = src_dims
340 .iter()
341 .zip(src_strides)
342 .enumerate()
343 .filter(|(i, _)| *i != axis)
344 .map(|(_, (&d, &s))| (d, s))
345 .collect();
346 let out_dims: Vec<usize> = kept.iter().map(|&(d, _)| d).collect();
347
348 if out_dims.is_empty() {
349 let acc = unsafe {
351 fold_run::<T, Op, M, R, U>(src_ptr, axis_stride, axis_len, init, &map_fn, &reduce_fn)
352 };
353 let strides = col_major_strides(&[1]);
354 return StridedArray::from_parts(vec![acc], &[1], &strides, 0);
355 }
356
357 let total_out = total_len(&out_dims)?;
358 let out_strides = col_major_strides(&out_dims);
359 let mut out =
360 StridedArray::from_parts(vec![init.clone(); total_out], &out_dims, &out_strides, 0)?;
361 if total_out == 0 || axis_len == 0 {
362 return Ok(out);
363 }
364 let out_ptr = out.view_mut().as_mut_ptr();
365
366 let parts = AxisParts {
367 kept: &kept,
368 axis_len,
369 axis_stride,
370 map_fn: &map_fn,
371 reduce_fn: &reduce_fn,
372 init: &init,
373 };
374
375 #[cfg(feature = "parallel")]
376 {
377 let work = total_out.saturating_mul(axis_len);
378 let nthreads = crate::threading::parallel_threads_for_len(work).min(total_out);
379 if nthreads > 1 {
380 let src_send = SendPtr(src_ptr as *mut T);
381 let out_send = SendPtr(out_ptr);
382 let parts = &parts;
383 crate::threading::parallel_for_each(0..total_out, nthreads, &|range| {
384 unsafe {
388 reduce_axis_range::<T, Op, M, R, U>(
389 src_send.as_const(),
390 out_send.as_ptr(),
391 range,
392 parts,
393 )
394 }
395 });
396 return Ok(out);
397 }
398 }
399
400 unsafe { reduce_axis_range::<T, Op, M, R, U>(src_ptr, out_ptr, 0..total_out, &parts) };
402 Ok(out)
403}
404
405struct AxisParts<'a, M, R, U> {
407 kept: &'a [(usize, isize)],
409 axis_len: usize,
410 axis_stride: isize,
411 map_fn: &'a M,
412 reduce_fn: &'a R,
413 init: &'a U,
414}
415
416unsafe fn reduce_axis_range<T, Op, M, R, U>(
424 src: *const T,
425 out: *mut U,
426 range: Range<usize>,
427 parts: &AxisParts<'_, M, R, U>,
428) where
429 T: Copy,
430 Op: ElementOp<T>,
431 M: Fn(T) -> U,
432 R: Fn(U, U) -> U,
433 U: Clone,
434{
435 let kept = parts.kept;
436 let (lead_extent, lead_stride) = kept[0];
437 let sweep = lead_stride == 1 && lead_extent >= AXIS_SWEEP_MIN_LEN && parts.axis_len > 1;
438 let block = (AXIS_SWEEP_BLOCK_BYTES / std::mem::size_of::<U>().max(1)).max(AXIS_SWEEP_MIN_LEN);
439
440 let mut coords = vec![0usize; kept.len()];
442 let mut rest = range.start;
443 let mut src_off = 0isize;
444 for (coord, &(extent, stride)) in coords.iter_mut().zip(kept) {
445 *coord = rest % extent;
446 rest /= extent;
447 src_off += *coord as isize * stride;
448 }
449
450 let mut output = range.start;
451 while output < range.end {
452 let len = if sweep {
453 (lead_extent - coords[0]).min(range.end - output).min(block)
454 } else {
455 1
456 };
457 if sweep {
458 unsafe { sweep_block::<T, Op, M, R, U>(src, src_off, out.add(output), len, parts) };
461 } else {
462 unsafe {
464 let acc = fold_run::<T, Op, M, R, U>(
465 src.offset(src_off),
466 parts.axis_stride,
467 parts.axis_len,
468 parts.init.clone(),
469 parts.map_fn,
470 parts.reduce_fn,
471 );
472 *out.add(output) = acc;
473 }
474 }
475 output += len;
476 if output >= range.end {
477 break;
478 }
479 coords[0] += len;
482 src_off += len as isize * lead_stride;
483 let mut axis = 0;
484 while coords[axis] == kept[axis].0 {
485 src_off -= kept[axis].0 as isize * kept[axis].1;
486 coords[axis] = 0;
487 axis += 1;
488 coords[axis] += 1;
489 src_off += kept[axis].1;
490 }
491 }
492}
493
494#[inline(always)]
502unsafe fn sweep_block<T, Op, M, R, U>(
503 src: *const T,
504 src_off: isize,
505 out: *mut U,
506 len: usize,
507 parts: &AxisParts<'_, M, R, U>,
508) where
509 T: Copy,
510 Op: ElementOp<T>,
511 M: Fn(T) -> U,
512 R: Fn(U, U) -> U,
513 U: Clone,
514{
515 let map_fn = parts.map_fn;
516 let reduce_fn = parts.reduce_fn;
517 let out = unsafe { std::slice::from_raw_parts_mut(out, len) };
519 let column = |k: usize| {
520 unsafe {
522 std::slice::from_raw_parts(src.offset(src_off + k as isize * parts.axis_stride), len)
523 }
524 };
525 let first = column(0);
526 for (slot, &value) in out.iter_mut().zip(first) {
527 *slot = reduce_fn(parts.init.clone(), map_fn(Op::apply(value)));
528 }
529 let mut k = 1;
530 while k + 4 <= parts.axis_len {
531 let (c0, c1, c2, c3) = (column(k), column(k + 1), column(k + 2), column(k + 3));
532 for (index, slot) in out.iter_mut().enumerate() {
533 let mut acc = reduce_fn(slot.clone(), map_fn(Op::apply(c0[index])));
534 acc = reduce_fn(acc, map_fn(Op::apply(c1[index])));
535 acc = reduce_fn(acc, map_fn(Op::apply(c2[index])));
536 acc = reduce_fn(acc, map_fn(Op::apply(c3[index])));
537 *slot = acc;
538 }
539 k += 4;
540 }
541 while k < parts.axis_len {
542 let values = column(k);
543 for (slot, &value) in out.iter_mut().zip(values) {
544 *slot = reduce_fn(slot.clone(), map_fn(Op::apply(value)));
545 }
546 k += 1;
547 }
548}