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, StridedError};
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// Every loop below advances an offset only when another position along its
242// level follows, and each loop rewinds its own level on completion. Offsets
243// therefore only ever hold reachable positions of the validated layouts (plus
244// the caller's initial offsets), so no intermediate value can overflow even
245// when a layout ends within one stride of `isize::MAX` (issue #243
246// follow-up). The rewind distance is the difference of two reachable offsets,
247// so it is representable as well.
248
249/// Element-level nested loops for ranks >= 2.
250///
251/// Iterates over the element block, calling `f` at each innermost position.
252/// Levels are listed from outermost to innermost (excluding level 0 which is
253/// the callback level). For a rank-N kernel, element levels are N-1, N-2, ..., 1.
254/// Each level leaves the offsets where it found them.
255macro_rules! elem_loops {
256    // Base case: single level (innermost element level).
257    // Iterates blens[$lv] times, calling f and advancing by stride[$lv]
258    // between calls.
259    ($offsets:ident, $strides:ident, $f:ident, $blens:ident, $is:ident; $lv:literal) => {
260        for _i in 0..$blens[$lv] {
261            $f($offsets, $blens[0], &$is)?;
262            if _i + 1 < $blens[$lv] {
263                for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
264                    *o += s[$lv];
265                }
266            }
267        }
268        let _back = $blens[$lv].saturating_sub(1) as isize;
269        for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
270            *o -= _back * s[$lv];
271        }
272    };
273    // Recursive case: outermost level wraps inner levels, which rewind
274    // themselves; this level advances between inner passes.
275    ($offsets:ident, $strides:ident, $f:ident, $blens:ident, $is:ident;
276     $lv:literal, $next:literal $(, $rest:literal)*) => {
277        for _i in 0..$blens[$lv] {
278            elem_loops!($offsets, $strides, $f, $blens, $is; $next $(, $rest)*);
279            if _i + 1 < $blens[$lv] {
280                for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
281                    *o += s[$lv];
282                }
283            }
284        }
285        let _back = $blens[$lv].saturating_sub(1) as isize;
286        for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
287            *o -= _back * s[$lv];
288        }
289    };
290}
291
292/// Block-level nested while loops for ranks >= 2.
293///
294/// Each block level iterates over tiles of the corresponding dimension.
295/// At level 0 (innermost), element loops are executed. Each block level
296/// advances to the next tile only when one follows and rewinds itself on
297/// completion; inner levels rewind themselves.
298macro_rules! block_loop {
299    // Base case: single block level (the innermost).
300    ($dims:ident, $blocks:ident, $strides:ident, $offsets:ident, $f:ident,
301     $blens:ident, $is:ident; elem=[$($el:literal),+]; $lv0:literal; top=$top:literal) => {{
302        let mut _j = 0usize;
303        let mut _advanced = 0usize;
304        while _j < $dims[$lv0] {
305            $blens[$lv0] = $blocks[$lv0].max(1).min($dims[$lv0]).min($dims[$lv0] - _j);
306            elem_loops!($offsets, $strides, $f, $blens, $is; $($el),+);
307            _j += $blens[$lv0];
308            if _j < $dims[$lv0] {
309                for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
310                    *o += ($blens[$lv0] as isize) * s[$lv0];
311                }
312                _advanced = _j;
313            }
314        }
315        for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
316            *o -= (_advanced as isize) * s[$lv0];
317        }
318    }};
319    // Recursive case: outer block level wraps inner block levels.
320    ($dims:ident, $blocks:ident, $strides:ident, $offsets:ident, $f:ident,
321     $blens:ident, $is:ident; elem=[$($el:literal),+];
322     $lv:literal, $next:literal $(, $rest:literal)*; top=$top:literal) => {{
323        let mut _j = 0usize;
324        let mut _advanced = 0usize;
325        while _j < $dims[$lv] {
326            $blens[$lv] = $blocks[$lv].max(1).min($dims[$lv]).min($dims[$lv] - _j);
327            block_loop!($dims, $blocks, $strides, $offsets, $f, $blens, $is;
328                elem=[$($el),+]; $next $(, $rest)*; top=$top);
329            _j += $blens[$lv];
330            if _j < $dims[$lv] {
331                for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
332                    *o += ($blens[$lv] as isize) * s[$lv];
333                }
334                _advanced = _j;
335            }
336        }
337        for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
338            *o -= (_advanced as isize) * s[$lv];
339        }
340    }};
341}
342
343/// Generate a rank-specialized kernel function.
344///
345/// For rank 1 there are no element loops; only a single block loop on dim 0.
346/// For rank >= 2, block loops nest from outermost to innermost (dim 0), and
347/// element loops nest from dim N-1 down to dim 1 inside the innermost block.
348/// The offsets are back at their starting values on return.
349macro_rules! make_kernel {
350    // Rank 1: single block loop, no element nesting.
351    ($name:ident, rank=1) => {
352        #[inline]
353        fn $name<F>(
354            dims: &[usize],
355            blocks: &[usize],
356            strides: &[Vec<isize>],
357            offsets: &mut [isize],
358            f: &mut F,
359        ) -> Result<()>
360        where
361            F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
362        {
363            let d0 = dims[0];
364            let b0 = blocks[0].max(1).min(d0);
365            let inner_strides: Vec<isize> = strides.iter().map(|s| s[0]).collect();
366
367            let mut j0 = 0usize;
368            let mut advanced = 0usize;
369            while j0 < d0 {
370                let blen0 = b0.min(d0 - j0);
371                f(offsets, blen0, &inner_strides)?;
372                j0 += blen0;
373                if j0 < d0 {
374                    for (o, s) in offsets.iter_mut().zip(strides.iter()) {
375                        *o += (blen0 as isize) * s[0];
376                    }
377                    advanced = j0;
378                }
379            }
380            for (o, s) in offsets.iter_mut().zip(strides.iter()) {
381                *o -= (advanced as isize) * s[0];
382            }
383            Ok(())
384        }
385    };
386    // Rank >= 2: nested block loops + element loops.
387    //   block=[outermost, ..., 0]  - block loop levels from outermost to innermost
388    //   elem=[N-1, ..., 1]         - element loop levels from outermost to innermost
389    //   top=N-1                    - topmost element level (= rank - 1)
390    ($name:ident, rank=$rank:literal,
391     block=[$($blk:literal),+], elem=[$($el:literal),+], top=$top:literal) => {
392        #[inline]
393        fn $name<F>(
394            dims: &[usize],
395            blocks: &[usize],
396            strides: &[Vec<isize>],
397            offsets: &mut [isize],
398            f: &mut F,
399        ) -> Result<()>
400        where
401            F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
402        {
403            let inner_strides: Vec<isize> = strides.iter().map(|s| s[0]).collect();
404            let mut blens = [0usize; $rank];
405            block_loop!(dims, blocks, strides, offsets, f, blens, inner_strides;
406                elem=[$($el),+]; $($blk),+; top=$top);
407            Ok(())
408        }
409    };
410}
411
412make_kernel!(kernel_1d_inner, rank = 1);
413make_kernel!(
414    kernel_2d_inner,
415    rank = 2,
416    block = [1, 0],
417    elem = [1],
418    top = 1
419);
420make_kernel!(
421    kernel_3d_inner,
422    rank = 3,
423    block = [2, 1, 0],
424    elem = [2, 1],
425    top = 2
426);
427make_kernel!(
428    kernel_4d_inner,
429    rank = 4,
430    block = [3, 2, 1, 0],
431    elem = [3, 2, 1],
432    top = 3
433);
434make_kernel!(
435    kernel_5d_inner,
436    rank = 5,
437    block = [4, 3, 2, 1, 0],
438    elem = [4, 3, 2, 1],
439    top = 4
440);
441make_kernel!(
442    kernel_6d_inner,
443    rank = 6,
444    block = [5, 4, 3, 2, 1, 0],
445    elem = [5, 4, 3, 2, 1],
446    top = 5
447);
448make_kernel!(
449    kernel_7d_inner,
450    rank = 7,
451    block = [6, 5, 4, 3, 2, 1, 0],
452    elem = [6, 5, 4, 3, 2, 1],
453    top = 6
454);
455make_kernel!(
456    kernel_8d_inner,
457    rank = 8,
458    block = [7, 6, 5, 4, 3, 2, 1, 0],
459    elem = [7, 6, 5, 4, 3, 2, 1],
460    top = 7
461);
462
463/// N-dimensional kernel with inner block callback (iterative form).
464///
465/// Fallback for rank >= 9. Uses a carry-style increment over outer levels
466/// instead of compile-time unrolled nested loops.
467#[inline]
468fn kernel_nd_inner_iterative<F>(
469    dims: &[usize],
470    blocks: &[usize],
471    strides: &[Vec<isize>],
472    offsets: &mut [isize],
473    f: &mut F,
474) -> Result<()>
475where
476    F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
477{
478    let rank = dims.len();
479    debug_assert!(rank >= 9);
480
481    let d0 = dims[0];
482    let b0 = blocks[0].max(1).min(d0);
483    let inner_strides: Vec<isize> = strides.iter().map(|s| s[0]).collect();
484
485    // Current position for each outer level (1..rank-1). Level 0 uses block loop.
486    let mut idx = vec![0usize; rank];
487
488    // As in the rank-specialized kernels, offsets advance only when another
489    // position follows and each level rewinds from its last position, so they
490    // never leave the reachable span.
491    if dims.contains(&0) {
492        return Ok(());
493    }
494    loop {
495        // Level 0: callback over contiguous block fragments.
496        let mut j0 = 0usize;
497        let mut advanced = 0usize;
498        while j0 < d0 {
499            let blen0 = b0.min(d0 - j0);
500            f(offsets, blen0, &inner_strides)?;
501            j0 += blen0;
502            if j0 < d0 {
503                for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
504                    *offset += (blen0 as isize) * s[0];
505                }
506                advanced = j0;
507            }
508        }
509        for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
510            *offset -= (advanced as isize) * s[0];
511        }
512
513        // Carry-style increment for outer levels.
514        let mut level = 1usize;
515        loop {
516            if idx[level] + 1 < dims[level] {
517                idx[level] += 1;
518                for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
519                    *offset += s[level];
520                }
521                break;
522            }
523
524            let last = (dims[level] - 1) as isize;
525            idx[level] = 0;
526            for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
527                *offset -= last * s[level];
528            }
529            level += 1;
530            if level == rank {
531                return Ok(());
532            }
533        }
534    }
535}
536
537// ============================================================================
538// Pre-ordered iteration (for threaded leaf functions)
539// ============================================================================
540
541/// Iterate over blocks with pre-ordered dimensions and initial offsets.
542///
543/// Unlike `for_each_inner_block`, this function assumes that `dims`, `blocks`,
544/// and `strides` are **already in iteration order** (i.e., identity ordering).
545/// It also accepts `initial_offsets` which are added to the starting offsets
546/// before iteration begins.
547///
548/// This avoids the redundant re-ordering and per-callback `Vec` allocation
549/// that `for_each_inner_block_with_offsets` previously incurred.
550///
551/// The block-walking machinery is compiled once (`preordered_dyn`): a generic
552/// caller only instantiates its own callback, reached through one indirect call
553/// per inner block (a run of elements, never per element). Keeping the walk
554/// generic would duplicate it for every operation, dtype and element-op
555/// combination of every downstream crate.
556#[inline]
557pub(crate) fn for_each_inner_block_preordered<F>(
558    dims: &[usize],
559    blocks: &[usize],
560    strides: &[Vec<isize>],
561    initial_offsets: &[isize],
562    mut f: F,
563) -> Result<()>
564where
565    F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
566{
567    preordered_dyn(dims, blocks, strides, initial_offsets, &mut f)
568}
569
570type BlockFn<'a> = &'a mut dyn FnMut(&[isize], usize, &[isize]) -> Result<()>;
571
572#[inline(never)]
573fn preordered_dyn(
574    dims: &[usize],
575    blocks: &[usize],
576    strides: &[Vec<isize>],
577    initial_offsets: &[isize],
578    mut f: BlockFn<'_>,
579) -> Result<()> {
580    let rank = dims.len();
581    if rank == 0 {
582        return f(initial_offsets, 1, &[]);
583    }
584
585    // Start from initial_offsets (kernel functions reset to starting values at end)
586    let mut offsets = initial_offsets.to_vec();
587
588    match rank {
589        1 => kernel_1d_inner(dims, blocks, strides, &mut offsets, &mut f),
590        2 => kernel_2d_inner(dims, blocks, strides, &mut offsets, &mut f),
591        3 => kernel_3d_inner(dims, blocks, strides, &mut offsets, &mut f),
592        4 => kernel_4d_inner(dims, blocks, strides, &mut offsets, &mut f),
593        5 => kernel_5d_inner(dims, blocks, strides, &mut offsets, &mut f),
594        6 => kernel_6d_inner(dims, blocks, strides, &mut offsets, &mut f),
595        7 => kernel_7d_inner(dims, blocks, strides, &mut offsets, &mut f),
596        8 => kernel_8d_inner(dims, blocks, strides, &mut offsets, &mut f),
597        _ => kernel_nd_inner_iterative(dims, blocks, strides, &mut offsets, &mut f),
598    }
599}
600
601// ============================================================================
602// Utility functions
603// ============================================================================
604
605/// Check exact shape agreement before prevalidated replay.
606///
607/// # Examples
608///
609/// ```
610/// use strided_basic::execution::ensure_same_shape;
611/// ensure_same_shape(&[2, 3], &[2, 3]).unwrap();
612/// assert!(ensure_same_shape(&[2, 3], &[3, 2]).is_err());
613/// ```
614///
615/// # Errors
616/// Returns a rank or shape mismatch.
617pub fn ensure_same_shape(a: &[usize], b: &[usize]) -> Result<()> {
618    if a.len() != b.len() {
619        return Err(crate::StridedError::RankMismatch(a.len(), b.len()));
620    }
621    if a != b {
622        return Err(crate::StridedError::ShapeMismatch(a.to_vec(), b.to_vec()));
623    }
624    Ok(())
625}
626
627#[derive(Copy, Clone, Debug, Eq, PartialEq)]
628pub(crate) enum ContiguousLayout {
629    /// C-like layout: last axis varies fastest.
630    RowMajor,
631    /// Julia/Fortran-like layout: first axis varies fastest.
632    ColMajor,
633}
634
635/// Returns the contiguous memory layout kind for the given (dims, strides).
636///
637/// Notes:
638/// - Ignores axes with `dim <= 1` since they do not affect addressability.
639/// - Does not treat negative-stride views as contiguous for fast-path purposes.
640pub(crate) fn contiguous_layout(dims: &[usize], strides: &[isize]) -> Option<ContiguousLayout> {
641    if dims.len() != strides.len() {
642        return None;
643    }
644    if dims.is_empty() {
645        return Some(ContiguousLayout::RowMajor);
646    }
647
648    // Row-major: check from last to first.
649    let mut expected = 1isize;
650    let mut row_ok = true;
651    for (&dim, &stride) in dims.iter().rev().zip(strides.iter().rev()) {
652        if dim <= 1 {
653            continue;
654        }
655        if stride != expected {
656            row_ok = false;
657            break;
658        }
659        expected = expected.saturating_mul(dim as isize);
660    }
661    if row_ok {
662        return Some(ContiguousLayout::RowMajor);
663    }
664
665    // Col-major: check from first to last.
666    let mut expected = 1isize;
667    for (&dim, &stride) in dims.iter().zip(strides.iter()) {
668        if dim <= 1 {
669            continue;
670        }
671        if stride != expected {
672            return None;
673        }
674        expected = expected.saturating_mul(dim as isize);
675    }
676    Some(ContiguousLayout::ColMajor)
677}
678
679/// Checked element count of `dims`.
680///
681/// A shape containing a zero extent has zero elements even when the product
682/// of its other extents would overflow; such views pass
683/// `strided_view::validate_bounds`, so this must not form that product. A
684/// nonempty shape whose element count exceeds `usize::MAX` (for example a
685/// huge stride-0 broadcast) returns [`StridedError::OffsetOverflow`].
686pub(crate) fn total_len(dims: &[usize]) -> Result<usize> {
687    if dims.contains(&0) {
688        return Ok(0);
689    }
690    dims.iter()
691        .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
692        .ok_or(StridedError::OffsetOverflow)
693}
694
695/// Whether the sequential contiguous fast path should be used.
696///
697/// When the `parallel` feature is enabled and the total element count exceeds
698/// the threading threshold, we normally must *not* take the contiguous fast path
699/// so that the parallel kernel path can be reached. If the active Rayon pool has
700/// only one worker, however, there is no threaded path to reach and the
701/// sequential slice loop is the right kernel.
702#[inline]
703pub(crate) fn use_sequential_fast_path(total: usize) -> bool {
704    #[cfg(feature = "parallel")]
705    {
706        total <= crate::threading::MINTHREADLENGTH || crate::execution_policy::rayon_threads() <= 1
707    }
708    #[cfg(not(feature = "parallel"))]
709    {
710        let _ = total;
711        true
712    }
713}
714
715/// Returns the common contiguous layout if **all** provided stride arrays
716/// share the same contiguous layout for the given `dims`.
717///
718/// Returns `None` if `strides_list` is empty, any array is not contiguous,
719/// or any two arrays have different contiguous layouts.
720#[inline]
721pub(crate) fn same_contiguous_layout(
722    dims: &[usize],
723    strides_list: &[&[isize]],
724) -> Option<ContiguousLayout> {
725    let first = contiguous_layout(dims, strides_list.first()?)?;
726    for strides in &strides_list[1..] {
727        if contiguous_layout(dims, strides)? != first {
728            return None;
729        }
730    }
731    Some(first)
732}
733
734/// Returns the common contiguous layout only when the sequential fast path
735/// should be used (total elements <= threading threshold).
736///
737/// Returns [`StridedError::OffsetOverflow`] when the element count of `dims`
738/// does not fit in `usize`.
739#[inline]
740pub(crate) fn sequential_contiguous_layout(
741    dims: &[usize],
742    strides_list: &[&[isize]],
743) -> Result<Option<ContiguousLayout>> {
744    if !use_sequential_fast_path(total_len(dims)?) {
745        return Ok(None);
746    }
747    Ok(same_contiguous_layout(dims, strides_list))
748}
749
750#[cfg(test)]
751#[path = "kernel/tests/tests.rs"]
752mod tests;