Skip to main content

strided_basic/
kernel.rs

1//! Kernel iteration engine ported from Julia's Strided.jl/src/mapreduce.jl
2//!
3//! This module implements the core iteration engine that follows Julia's
4//! `_mapreduce_kernel!` pattern for cache-optimized strided array operations.
5
6use crate::fuse::{compress_dims, fuse_dims};
7use crate::{block, order, Result};
8
9/// Maximum total elements for the small tensor fast path.
10/// Tensors at or below this size skip `compute_order` and `compute_block_sizes`
11/// since they fit in L1 cache and blocking provides no benefit.
12pub const SMALL_TENSOR_THRESHOLD: usize = 1024;
13
14/// Blocking metadata produced by the execution planners; not a layout-validity proof.
15#[derive(Debug)]
16pub struct KernelPlan {
17    #[cfg_attr(not(test), allow(dead_code))]
18    pub(crate) order: Vec<usize>, // outer -> inner
19    pub block: Vec<usize>,
20}
21
22/// Build an execution plan for strided iteration (used only in tests).
23///
24/// This follows Julia's `_mapreduce_fuse!` -> `_mapreduce_order!` -> `_mapreduce_block!` pipeline:
25/// 1. Fuse contiguous dimensions
26/// 2. Compute optimal iteration order
27/// 3. Compute block sizes for cache efficiency
28#[cfg(test)]
29pub(crate) fn build_plan(
30    dims: &[usize],
31    strides_list: &[&[isize]],
32    dest_index: Option<usize>,
33    elem_size: usize,
34) -> KernelPlan {
35    let order = order::compute_order(dims, strides_list, dest_index);
36    let block = block::compute_block_sizes(dims, &order, strides_list, elem_size);
37    KernelPlan { order, block }
38}
39
40/// Build an execution plan with dimension fusion.
41///
42/// Pipeline: order -> reorder -> fuse -> block.
43///
44/// Ordering first ensures that dimensions are sorted by stride importance
45/// (smallest stride innermost). Fusing *after* ordering catches contiguous
46/// dimensions regardless of the original memory layout (column-major,
47/// row-major, or any permutation).
48///
49/// Returns `(fused_dims, ordered_strides, KernelPlan)` where dimensions
50/// and strides are already in iteration order (plan.order is identity).
51pub(crate) fn build_plan_fused(
52    dims: &[usize],
53    strides_list: &[&[isize]],
54    dest_index: Option<usize>,
55    elem_size: usize,
56) -> (Vec<usize>, Vec<Vec<isize>>, KernelPlan) {
57    // 1. Compute optimal iteration order on original dims
58    let order = order::compute_order(dims, strides_list, dest_index);
59
60    // 2. Reorder dims and strides
61    let ordered_dims: Vec<usize> = order.iter().map(|&d| dims[d]).collect();
62    let ordered_strides: Vec<Vec<isize>> = strides_list
63        .iter()
64        .map(|strides| order.iter().map(|&d| strides[d]).collect())
65        .collect();
66    let ordered_strides_refs: Vec<&[isize]> =
67        ordered_strides.iter().map(|s| s.as_slice()).collect();
68
69    // 3. Fuse contiguous dimensions in ordered space
70    let fused_dims = fuse_dims(&ordered_dims, &ordered_strides_refs);
71
72    // 4. Compress: remove size-1 dimensions to reduce loop depth
73    let (compressed_dims, compressed_strides) = compress_dims(&fused_dims, &ordered_strides);
74    let compressed_strides_refs: Vec<&[isize]> =
75        compressed_strides.iter().map(|s| s.as_slice()).collect();
76
77    // 5. Compute blocks with identity ordering (already ordered)
78    let identity: Vec<usize> = (0..compressed_dims.len()).collect();
79    let block = block::compute_block_sizes(
80        &compressed_dims,
81        &identity,
82        &compressed_strides_refs,
83        elem_size,
84    );
85
86    (
87        compressed_dims,
88        compressed_strides,
89        KernelPlan {
90            order: identity,
91            block,
92        },
93    )
94}
95
96/// Simplified plan for small tensors that fit in L1 cache.
97///
98/// Skips `compute_order` and `compute_block_sizes` since blocking is
99/// unnecessary for small data. Only fuses and compresses dimensions
100/// to reduce loop depth.
101///
102/// Returns `(fused_dims, fused_strides, KernelPlan)` where the plan's
103/// order is identity and block sizes equal the full dimension sizes.
104pub(crate) fn build_plan_fused_small(
105    dims: &[usize],
106    strides_list: &[&[isize]],
107) -> (Vec<usize>, Vec<Vec<isize>>, KernelPlan) {
108    let strides_owned: Vec<Vec<isize>> = strides_list.iter().map(|s| s.to_vec()).collect();
109
110    // Single fuse + compress pass (no ordering needed)
111    let fused = fuse_dims(dims, strides_list);
112    let (fused_dims, fused_strides) = compress_dims(&fused, &strides_owned);
113
114    // Block = full dimension (no tiling)
115    let block = fused_dims.clone();
116    let identity: Vec<usize> = (0..fused_dims.len()).collect();
117
118    (
119        fused_dims,
120        fused_strides,
121        KernelPlan {
122            order: identity,
123            block,
124        },
125    )
126}
127
128// ============================================================================
129// Block-based iteration with inner stride callback
130// ============================================================================
131
132/// Iterate over blocks, calling f with (offsets, block_len, inner_strides).
133///
134/// The callback receives the current byte offsets for each array, the number
135/// of elements in the innermost block, and the innermost strides for each array.
136/// This allows the caller to implement vectorized inner loops.
137#[cfg(test)]
138#[inline]
139pub(crate) fn for_each_inner_block<F>(
140    dims: &[usize],
141    plan: &KernelPlan,
142    strides_list: &[&[isize]],
143    mut f: F,
144) -> Result<()>
145where
146    F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
147{
148    let rank = dims.len();
149    if rank == 0 {
150        let offsets = vec![0isize; strides_list.len()];
151        return f(&offsets, 1, &[]);
152    }
153
154    // Reorder dimensions and strides according to plan
155    let ordered_dims: Vec<usize> = plan.order.iter().map(|&d| dims[d]).collect();
156    let ordered_blocks: Vec<usize> = plan.block.clone();
157
158    // Build stride vectors for each array, ordered by plan
159    let num_arrays = strides_list.len();
160    let mut ordered_strides: Vec<Vec<isize>> = Vec::with_capacity(num_arrays);
161    for strides in strides_list {
162        let s: Vec<isize> = plan.order.iter().map(|&d| strides[d]).collect();
163        ordered_strides.push(s);
164    }
165
166    // Initial offsets (all zero)
167    let mut offsets = vec![0isize; num_arrays];
168
169    // Call the specialized kernel based on rank
170    match rank {
171        1 => kernel_1d_inner(
172            &ordered_dims,
173            &ordered_blocks,
174            &ordered_strides,
175            &mut offsets,
176            &mut f,
177        ),
178        2 => kernel_2d_inner(
179            &ordered_dims,
180            &ordered_blocks,
181            &ordered_strides,
182            &mut offsets,
183            &mut f,
184        ),
185        3 => kernel_3d_inner(
186            &ordered_dims,
187            &ordered_blocks,
188            &ordered_strides,
189            &mut offsets,
190            &mut f,
191        ),
192        4 => kernel_4d_inner(
193            &ordered_dims,
194            &ordered_blocks,
195            &ordered_strides,
196            &mut offsets,
197            &mut f,
198        ),
199        5 => kernel_5d_inner(
200            &ordered_dims,
201            &ordered_blocks,
202            &ordered_strides,
203            &mut offsets,
204            &mut f,
205        ),
206        6 => kernel_6d_inner(
207            &ordered_dims,
208            &ordered_blocks,
209            &ordered_strides,
210            &mut offsets,
211            &mut f,
212        ),
213        7 => kernel_7d_inner(
214            &ordered_dims,
215            &ordered_blocks,
216            &ordered_strides,
217            &mut offsets,
218            &mut f,
219        ),
220        8 => kernel_8d_inner(
221            &ordered_dims,
222            &ordered_blocks,
223            &ordered_strides,
224            &mut offsets,
225            &mut f,
226        ),
227        _ => kernel_nd_inner_iterative(
228            &ordered_dims,
229            &ordered_blocks,
230            &ordered_strides,
231            &mut offsets,
232            &mut f,
233        ),
234    }
235}
236
237// ============================================================================
238// Macro-generated rank-specialized kernels (inner-block callback)
239// ============================================================================
240
241/// Element-level nested loops for ranks >= 2.
242///
243/// Iterates over the element block, calling `f` at each innermost position.
244/// Levels are listed from outermost to innermost (excluding level 0 which is
245/// the callback level). For a rank-N kernel, element levels are N-1, N-2, ..., 1.
246macro_rules! elem_loops {
247    // Base case: single level (innermost element level).
248    // Iterates blens[$lv] times, calling f then advancing by stride[$lv].
249    ($offsets:ident, $strides:ident, $f:ident, $blens:ident, $is:ident; $lv:literal) => {
250        for _ in 0..$blens[$lv] {
251            $f($offsets, $blens[0], &$is)?;
252            for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
253                *o += s[$lv];
254            }
255        }
256    };
257    // Recursive case: outermost level wraps inner levels.
258    // After the inner loop, resets the inner level's offset and advances this level.
259    ($offsets:ident, $strides:ident, $f:ident, $blens:ident, $is:ident;
260     $lv:literal, $next:literal $(, $rest:literal)*) => {
261        for _ in 0..$blens[$lv] {
262            elem_loops!($offsets, $strides, $f, $blens, $is; $next $(, $rest)*);
263            for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
264                *o -= ($blens[$next] as isize) * s[$next];
265                *o += s[$lv];
266            }
267        }
268    };
269}
270
271/// Block-level nested while loops for ranks >= 2.
272///
273/// Each block level iterates over tiles of the corresponding dimension.
274/// At level 0 (innermost), element loops are executed. After each inner body,
275/// offsets are adjusted to reset the inner dimension and advance the current level.
276macro_rules! block_loop {
277    // Base case: single block level (the innermost).
278    // Runs element loops, then resets top element level and advances this block level.
279    ($dims:ident, $blocks:ident, $strides:ident, $offsets:ident, $f:ident,
280     $blens:ident, $is:ident; elem=[$($el:literal),+]; $lv0:literal; top=$top:literal) => {{
281        let mut _j = 0usize;
282        while _j < $dims[$lv0] {
283            $blens[$lv0] = $blocks[$lv0].max(1).min($dims[$lv0]).min($dims[$lv0] - _j);
284            elem_loops!($offsets, $strides, $f, $blens, $is; $($el),+);
285            for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
286                *o -= ($blens[$top] as isize) * s[$top];
287                *o += ($blens[$lv0] as isize) * s[$lv0];
288            }
289            _j += $blens[$lv0];
290        }
291    }};
292    // Recursive case: outer block level wraps inner block levels.
293    // After the inner body, resets the next-inner dimension and advances this level.
294    ($dims:ident, $blocks:ident, $strides:ident, $offsets:ident, $f:ident,
295     $blens:ident, $is:ident; elem=[$($el:literal),+];
296     $lv:literal, $next:literal $(, $rest:literal)*; top=$top:literal) => {{
297        let mut _j = 0usize;
298        while _j < $dims[$lv] {
299            $blens[$lv] = $blocks[$lv].max(1).min($dims[$lv]).min($dims[$lv] - _j);
300            block_loop!($dims, $blocks, $strides, $offsets, $f, $blens, $is;
301                elem=[$($el),+]; $next $(, $rest)*; top=$top);
302            for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
303                *o -= ($dims[$next] as isize) * s[$next];
304                *o += ($blens[$lv] as isize) * s[$lv];
305            }
306            _j += $blens[$lv];
307        }
308    }};
309}
310
311/// Generate a rank-specialized kernel function.
312///
313/// For rank 1 there are no element loops; only a single block loop on dim 0.
314/// For rank >= 2, block loops nest from outermost to innermost (dim 0), and
315/// element loops nest from dim N-1 down to dim 1 inside the innermost block.
316macro_rules! make_kernel {
317    // Rank 1: single block loop, no element nesting.
318    ($name:ident, rank=1) => {
319        #[inline]
320        fn $name<F>(
321            dims: &[usize],
322            blocks: &[usize],
323            strides: &[Vec<isize>],
324            offsets: &mut [isize],
325            f: &mut F,
326        ) -> Result<()>
327        where
328            F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
329        {
330            let d0 = dims[0];
331            let b0 = blocks[0].max(1).min(d0);
332            let inner_strides: Vec<isize> = strides.iter().map(|s| s[0]).collect();
333
334            let mut j0 = 0usize;
335            while j0 < d0 {
336                let blen0 = b0.min(d0 - j0);
337                f(offsets, blen0, &inner_strides)?;
338                for (o, s) in offsets.iter_mut().zip(strides.iter()) {
339                    *o += (blen0 as isize) * s[0];
340                }
341                j0 += blen0;
342            }
343            for (o, s) in offsets.iter_mut().zip(strides.iter()) {
344                *o -= (d0 as isize) * s[0];
345            }
346            Ok(())
347        }
348    };
349    // Rank >= 2: nested block loops + element loops.
350    //   block=[outermost, ..., 0]  - block loop levels from outermost to innermost
351    //   elem=[N-1, ..., 1]         - element loop levels from outermost to innermost
352    //   top=N-1                    - topmost element level (= rank - 1)
353    ($name:ident, rank=$rank:literal,
354     block=[$($blk:literal),+], elem=[$($el:literal),+], top=$top:literal) => {
355        #[inline]
356        fn $name<F>(
357            dims: &[usize],
358            blocks: &[usize],
359            strides: &[Vec<isize>],
360            offsets: &mut [isize],
361            f: &mut F,
362        ) -> Result<()>
363        where
364            F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
365        {
366            let inner_strides: Vec<isize> = strides.iter().map(|s| s[0]).collect();
367            let mut blens = [0usize; $rank];
368            block_loop!(dims, blocks, strides, offsets, f, blens, inner_strides;
369                elem=[$($el),+]; $($blk),+; top=$top);
370            for (o, s) in offsets.iter_mut().zip(strides.iter()) {
371                *o -= (dims[$top] as isize) * s[$top];
372            }
373            Ok(())
374        }
375    };
376}
377
378make_kernel!(kernel_1d_inner, rank = 1);
379make_kernel!(
380    kernel_2d_inner,
381    rank = 2,
382    block = [1, 0],
383    elem = [1],
384    top = 1
385);
386make_kernel!(
387    kernel_3d_inner,
388    rank = 3,
389    block = [2, 1, 0],
390    elem = [2, 1],
391    top = 2
392);
393make_kernel!(
394    kernel_4d_inner,
395    rank = 4,
396    block = [3, 2, 1, 0],
397    elem = [3, 2, 1],
398    top = 3
399);
400make_kernel!(
401    kernel_5d_inner,
402    rank = 5,
403    block = [4, 3, 2, 1, 0],
404    elem = [4, 3, 2, 1],
405    top = 4
406);
407make_kernel!(
408    kernel_6d_inner,
409    rank = 6,
410    block = [5, 4, 3, 2, 1, 0],
411    elem = [5, 4, 3, 2, 1],
412    top = 5
413);
414make_kernel!(
415    kernel_7d_inner,
416    rank = 7,
417    block = [6, 5, 4, 3, 2, 1, 0],
418    elem = [6, 5, 4, 3, 2, 1],
419    top = 6
420);
421make_kernel!(
422    kernel_8d_inner,
423    rank = 8,
424    block = [7, 6, 5, 4, 3, 2, 1, 0],
425    elem = [7, 6, 5, 4, 3, 2, 1],
426    top = 7
427);
428
429/// N-dimensional kernel with inner block callback (iterative form).
430///
431/// Fallback for rank >= 9. Uses a carry-style increment over outer levels
432/// instead of compile-time unrolled nested loops.
433#[inline]
434fn kernel_nd_inner_iterative<F>(
435    dims: &[usize],
436    blocks: &[usize],
437    strides: &[Vec<isize>],
438    offsets: &mut [isize],
439    f: &mut F,
440) -> Result<()>
441where
442    F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
443{
444    let rank = dims.len();
445    debug_assert!(rank >= 9);
446
447    let d0 = dims[0];
448    let b0 = blocks[0].max(1).min(d0);
449    let inner_strides: Vec<isize> = strides.iter().map(|s| s[0]).collect();
450
451    // Current position for each outer level (1..rank-1). Level 0 uses block loop.
452    let mut idx = vec![0usize; rank];
453
454    loop {
455        // Level 0: callback over contiguous block fragments.
456        let mut j0 = 0usize;
457        while j0 < d0 {
458            let blen0 = b0.min(d0 - j0);
459            f(offsets, blen0, &inner_strides)?;
460            for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
461                *offset += (blen0 as isize) * s[0];
462            }
463            j0 += blen0;
464        }
465        for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
466            *offset -= (d0 as isize) * s[0];
467        }
468
469        // Carry-style increment for outer levels.
470        let mut level = 1usize;
471        loop {
472            for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
473                *offset += s[level];
474            }
475            idx[level] += 1;
476            if idx[level] < dims[level] {
477                break;
478            }
479
480            idx[level] = 0;
481            for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
482                *offset -= (dims[level] as isize) * s[level];
483            }
484            level += 1;
485            if level == rank {
486                return Ok(());
487            }
488        }
489    }
490}
491
492// ============================================================================
493// Pre-ordered iteration (for threaded leaf functions)
494// ============================================================================
495
496/// Iterate over blocks with pre-ordered dimensions and initial offsets.
497///
498/// Unlike `for_each_inner_block`, this function assumes that `dims`, `blocks`,
499/// and `strides` are **already in iteration order** (i.e., identity ordering).
500/// It also accepts `initial_offsets` which are added to the starting offsets
501/// before iteration begins.
502///
503/// This avoids the redundant re-ordering and per-callback `Vec` allocation
504/// that `for_each_inner_block_with_offsets` previously incurred.
505#[inline]
506pub(crate) fn for_each_inner_block_preordered<F>(
507    dims: &[usize],
508    blocks: &[usize],
509    strides: &[Vec<isize>],
510    initial_offsets: &[isize],
511    mut f: F,
512) -> Result<()>
513where
514    F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
515{
516    let rank = dims.len();
517    if rank == 0 {
518        return f(initial_offsets, 1, &[]);
519    }
520
521    // Start from initial_offsets (kernel functions reset to starting values at end)
522    let mut offsets = initial_offsets.to_vec();
523
524    match rank {
525        1 => kernel_1d_inner(dims, blocks, strides, &mut offsets, &mut f),
526        2 => kernel_2d_inner(dims, blocks, strides, &mut offsets, &mut f),
527        3 => kernel_3d_inner(dims, blocks, strides, &mut offsets, &mut f),
528        4 => kernel_4d_inner(dims, blocks, strides, &mut offsets, &mut f),
529        5 => kernel_5d_inner(dims, blocks, strides, &mut offsets, &mut f),
530        6 => kernel_6d_inner(dims, blocks, strides, &mut offsets, &mut f),
531        7 => kernel_7d_inner(dims, blocks, strides, &mut offsets, &mut f),
532        8 => kernel_8d_inner(dims, blocks, strides, &mut offsets, &mut f),
533        _ => kernel_nd_inner_iterative(dims, blocks, strides, &mut offsets, &mut f),
534    }
535}
536
537// ============================================================================
538// Utility functions
539// ============================================================================
540
541/// Check exact shape agreement before prevalidated replay.
542///
543/// # Examples
544///
545/// ```
546/// use strided_basic::execution::ensure_same_shape;
547/// ensure_same_shape(&[2, 3], &[2, 3]).unwrap();
548/// assert!(ensure_same_shape(&[2, 3], &[3, 2]).is_err());
549/// ```
550///
551/// # Errors
552/// Returns a rank or shape mismatch.
553pub fn ensure_same_shape(a: &[usize], b: &[usize]) -> Result<()> {
554    if a.len() != b.len() {
555        return Err(crate::StridedError::RankMismatch(a.len(), b.len()));
556    }
557    if a != b {
558        return Err(crate::StridedError::ShapeMismatch(a.to_vec(), b.to_vec()));
559    }
560    Ok(())
561}
562
563#[derive(Copy, Clone, Debug, Eq, PartialEq)]
564pub(crate) enum ContiguousLayout {
565    /// C-like layout: last axis varies fastest.
566    RowMajor,
567    /// Julia/Fortran-like layout: first axis varies fastest.
568    ColMajor,
569}
570
571/// Returns the contiguous memory layout kind for the given (dims, strides).
572///
573/// Notes:
574/// - Ignores axes with `dim <= 1` since they do not affect addressability.
575/// - Does not treat negative-stride views as contiguous for fast-path purposes.
576pub(crate) fn contiguous_layout(dims: &[usize], strides: &[isize]) -> Option<ContiguousLayout> {
577    if dims.len() != strides.len() {
578        return None;
579    }
580    if dims.is_empty() {
581        return Some(ContiguousLayout::RowMajor);
582    }
583
584    // Row-major: check from last to first.
585    let mut expected = 1isize;
586    let mut row_ok = true;
587    for (&dim, &stride) in dims.iter().rev().zip(strides.iter().rev()) {
588        if dim <= 1 {
589            continue;
590        }
591        if stride != expected {
592            row_ok = false;
593            break;
594        }
595        expected = expected.saturating_mul(dim as isize);
596    }
597    if row_ok {
598        return Some(ContiguousLayout::RowMajor);
599    }
600
601    // Col-major: check from first to last.
602    let mut expected = 1isize;
603    for (&dim, &stride) in dims.iter().zip(strides.iter()) {
604        if dim <= 1 {
605            continue;
606        }
607        if stride != expected {
608            return None;
609        }
610        expected = expected.saturating_mul(dim as isize);
611    }
612    Some(ContiguousLayout::ColMajor)
613}
614
615pub(crate) fn total_len(dims: &[usize]) -> usize {
616    if dims.is_empty() {
617        return 1;
618    }
619    dims.iter().product()
620}
621
622/// Whether the sequential contiguous fast path should be used.
623///
624/// When the `parallel` feature is enabled and the total element count exceeds
625/// the threading threshold, we normally must *not* take the contiguous fast path
626/// so that the parallel kernel path can be reached. If the active Rayon pool has
627/// only one worker, however, there is no threaded path to reach and the
628/// sequential slice loop is the right kernel.
629#[inline]
630pub(crate) fn use_sequential_fast_path(total: usize) -> bool {
631    #[cfg(feature = "parallel")]
632    {
633        total <= crate::threading::MINTHREADLENGTH || crate::execution_policy::rayon_threads() <= 1
634    }
635    #[cfg(not(feature = "parallel"))]
636    {
637        let _ = total;
638        true
639    }
640}
641
642/// Returns the common contiguous layout if **all** provided stride arrays
643/// share the same contiguous layout for the given `dims`.
644///
645/// Returns `None` if `strides_list` is empty, any array is not contiguous,
646/// or any two arrays have different contiguous layouts.
647#[inline]
648pub(crate) fn same_contiguous_layout(
649    dims: &[usize],
650    strides_list: &[&[isize]],
651) -> Option<ContiguousLayout> {
652    let first = contiguous_layout(dims, strides_list.first()?)?;
653    for strides in &strides_list[1..] {
654        if contiguous_layout(dims, strides)? != first {
655            return None;
656        }
657    }
658    Some(first)
659}
660
661/// Returns the common contiguous layout only when the sequential fast path
662/// should be used (total elements <= threading threshold).
663#[inline]
664pub(crate) fn sequential_contiguous_layout(
665    dims: &[usize],
666    strides_list: &[&[isize]],
667) -> Option<ContiguousLayout> {
668    if !use_sequential_fast_path(total_len(dims)) {
669        return None;
670    }
671    same_contiguous_layout(dims, strides_list)
672}
673
674#[cfg(test)]
675#[path = "kernel/tests/tests.rs"]
676mod tests;