Skip to main content

strided_basic/
reduce_view.rs

1//! Reduce operations on dynamic-rank strided views.
2
3use 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
20/// Full reduction with map function: `reduce(init, op, map.(src))`.
21pub 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    // Parallel contiguous fast path: split into scheduler chunks with slice-based iteration.
83    // This enables LLVM auto-vectorization on each chunk, unlike the general threaded path
84    // which uses scalar pointer-offset loops.
85    #[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            // False sharing avoidance: space output slots by cache line size
131            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            // For complete reduction, strides_list has 2 entries:
139            // [0] = threadedout (stride 0 everywhere — broadcasting), [1] = src
140            // The spacing/taskindex mechanism addresses output slots.
141            let ndim = fused_dims.len();
142            let mut threaded_strides = Vec::with_capacity(ordered_strides.len() + 1);
143            threaded_strides.push(vec![0isize; ndim]); // threadedout: stride 0 (broadcast)
144            for s in &ordered_strides {
145                threaded_strides.push(s.clone());
146            }
147            let initial_offsets = vec![0isize; threaded_strides.len()];
148
149            // Mask costs for threadedout stride=0 dims (all dims, since it's fully broadcast)
150            // This means: do NOT split on dims where output stride is 0 — but for complete
151            // reduction, ALL output strides are 0, so costs would all be masked to 0.
152            // Julia handles this with the spacing mechanism: each task writes to its own slot.
153            // We keep costs unmasked so splitting still works.
154
155            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                    // offsets[0] = spacing * (taskindex - 1) for threadedout
166                    // offsets[1] = offset into src
167                    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            // Merge thread-local results
194            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
227/// Reduce along a single axis, returning a new StridedArray.
228pub 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    // Reject an element count beyond usize (a huge stride-0 broadcast) up
249    // front instead of allocating and looping over it.
250    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        // Reduce to scalar
264        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    // Build source strides for iteration over non-axis dimensions (same rank as out_dims)
282    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    // This remains a dedicated sequential traversal: each output owns a mutable
297    // reduction accumulator, which does not fit the current alias-safe strided
298    // fanout primitive without adding a separate reduction-output partitioner.
299    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            // Fast path: when both output and source have stride 1, swap to
310            // reduction-outer / output-inner with slices so LLVM can
311            // auto-vectorize the contiguous inner loop.
312            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                // First reduction element: fold it into `init` so a
317                // non-identity seed is kept (issue #231). Writing without
318                // reading the pre-seeded output keeps the original traffic.
319                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                // Remaining reduction elements: accumulate
324                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            // General path: output-outer, reduction-inner
339            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}