Skip to main content

strided_basic/
reduce_view.rs

1//! Reduce operations on dynamic-rank strided views.
2
3#[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
22/// Number of independent accumulators of [`fold_contiguous`].
23const FOLD_LANES: usize = 16;
24
25/// Minimum outputs along a unit-stride kept axis for the column sweep of
26/// [`reduce_axis`].
27const AXIS_SWEEP_MIN_LEN: usize = 16;
28
29/// Byte budget of one output block of the column sweep of [`reduce_axis`].
30const AXIS_SWEEP_BLOCK_BYTES: usize = 64 * 1024;
31
32/// Folds a contiguous run with [`FOLD_LANES`] independent accumulators.
33///
34/// Lane `i` starts from element `i` and folds elements `i + FOLD_LANES`,
35/// `i + 2 * FOLD_LANES`, ... in order; the remainder is folded into the first
36/// lanes. `init` is combined exactly once, before lane 0, and the lanes are
37/// then combined left to right. Independent lanes let the compiler keep
38/// several vector accumulators in flight instead of one serial dependency
39/// chain, which is what makes a closure-based sum or max approach the speed
40/// of a hand written lane loop.
41#[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/// Folds `len` elements starting at `ptr` with step `stride`.
72///
73/// # Safety
74/// `ptr.offset(i * stride)` must be a readable element for every `i < len`.
75#[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        // SAFETY: the caller proves `len` readable unit-stride elements.
93        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        // SAFETY: the caller proves every element of the run is readable.
100        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
108/// Full reduction with map function: `reduce(init, op, map.(src))`.
109///
110/// # Evaluation order
111///
112/// `reduce_fn` must be associative and commutative, and `init` must be an
113/// identity of `reduce_fn` whenever the reduction runs on several threads.
114/// A contiguous source is folded with sixteen independent lanes (see the
115/// crate source for the exact lane assignment), non-contiguous sources are
116/// traversed in a cache-friendly loop order, and parallel execution folds
117/// chunks independently and combines them in chunk order. Floating point
118/// results can therefore differ from a strict left to right fold in the last
119/// bits; for a fixed shape, layout and thread count the result is repeatable.
120///
121/// Because `map_fn` and `reduce_fn` are opaque closures this entry cannot
122/// select the dtype-specific SIMD kernels of [`crate::ErasedReducePlan`];
123/// callers that reduce with a known operation (sum, product, max, min) on a
124/// supported dtype can use that plan for the fastest path.
125pub 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    // Parallel contiguous fast path: split into scheduler chunks with slice-based iteration.
164    // This enables LLVM auto-vectorization on each chunk, unlike the general threaded path
165    // which uses scalar pointer-offset loops.
166    #[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            // False sharing avoidance: space output slots by cache line size
205            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            // For complete reduction, strides_list has 2 entries:
213            // [0] = threadedout (stride 0 everywhere — broadcasting), [1] = src
214            // The spacing/taskindex mechanism addresses output slots.
215            let ndim = fused_dims.len();
216            let mut threaded_strides = Vec::with_capacity(ordered_strides.len() + 1);
217            threaded_strides.push(vec![0isize; ndim]); // threadedout: stride 0 (broadcast)
218            for s in &ordered_strides {
219                threaded_strides.push(s.clone());
220            }
221            let initial_offsets = vec![0isize; threaded_strides.len()];
222
223            // Mask costs for threadedout stride=0 dims (all dims, since it's fully broadcast)
224            // This means: do NOT split on dims where output stride is 0 — but for complete
225            // reduction, ALL output strides are 0, so costs would all be masked to 0.
226            // Julia handles this with the spacing mechanism: each task writes to its own slot.
227            // We keep costs unmasked so splitting still works.
228
229            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                    // offsets[0] = spacing * (taskindex - 1) for threadedout
240                    // offsets[1] = offset into src
241                    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            // Merge thread-local results
268            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
301/// Reduce along a single axis, returning a new StridedArray.
302///
303/// Every output folds `init` once and then its reduced elements. When the
304/// reduced axis has unit stride those elements are folded with the lane
305/// scheme of [`reduce`], so `reduce_fn` must be associative and commutative.
306/// When the leading kept axis has unit stride, a block of adjacent outputs
307/// is accumulated one contiguous source column at a time, which keeps the
308/// left to right order per output. With the `parallel` feature, outputs are
309/// split across threads once the number of source elements read exceeds the
310/// threading threshold; each output is computed by one thread, so the result
311/// does not depend on the thread count.
312pub 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    // Reject an element count beyond usize (a huge stride-0 broadcast) up
333    // front instead of allocating and looping over it.
334    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        // SAFETY: the view proves `axis_len` readable elements along `axis`.
350        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                // SAFETY: worker ranges are disjoint output ranges of the
385                // freshly allocated column-major output, and the view proves
386                // every source offset reachable from its kept coordinates.
387                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    // SAFETY: as for the parallel ranges, with one range covering all outputs.
401    unsafe { reduce_axis_range::<T, Op, M, R, U>(src_ptr, out_ptr, 0..total_out, &parts) };
402    Ok(out)
403}
404
405/// Shared inputs of [`reduce_axis_range`].
406struct AxisParts<'a, M, R, U> {
407    /// `(extent, source stride)` of every kept axis, in output order.
408    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
416/// Computes outputs `range` (column-major output indices) of [`reduce_axis`].
417///
418/// # Safety
419/// `out` must be the column-major output of `parts.kept` extents, `range` a
420/// subrange of it that no other thread writes, and every source offset
421/// reachable from a kept coordinate plus `k * parts.axis_stride` for
422/// `k < parts.axis_len` must be a readable element of `src`.
423unsafe 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    // Decode the first output once, then advance incrementally.
441    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            // SAFETY: the `len` outputs are inside the caller's range and
459            // their unit-stride source columns are readable.
460            unsafe { sweep_block::<T, Op, M, R, U>(src, src_off, out.add(output), len, parts) };
461        } else {
462            // SAFETY: one reachable output and its reduced run.
463            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        // Advance the kept coordinates by `len` along the leading axis; a
480        // block never crosses the end of the leading axis.
481        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/// Accumulates `len` adjacent outputs, one contiguous source column per
495/// reduced index, keeping the left to right order per output.
496///
497/// # Safety
498/// `out` must be `len` writable initialized outputs and `src + src_off +
499/// k * parts.axis_stride` must start `len` readable elements for every
500/// `k < parts.axis_len`.
501#[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    // SAFETY: forwarded caller contract.
518    let out = unsafe { std::slice::from_raw_parts_mut(out, len) };
519    let column = |k: usize| {
520        // SAFETY: forwarded caller contract for reduced index `k`.
521        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}