1use crate::fuse::{compress_dims, fuse_dims};
7use crate::{block, order, Result, StridedError};
8
9pub const SMALL_TENSOR_THRESHOLD: usize = 1024;
13
14#[derive(Debug)]
16pub struct KernelPlan {
17 #[cfg_attr(not(test), allow(dead_code))]
18 pub(crate) order: Vec<usize>, pub block: Vec<usize>,
20}
21
22#[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
40pub(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 let order = order::compute_order(dims, strides_list, dest_index);
59
60 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 let fused_dims = fuse_dims(&ordered_dims, &ordered_strides_refs);
71
72 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 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
96pub(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 let fused = fuse_dims(dims, strides_list);
112 let (fused_dims, fused_strides) = compress_dims(&fused, &strides_owned);
113
114 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#[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 let ordered_dims: Vec<usize> = plan.order.iter().map(|&d| dims[d]).collect();
156 let ordered_blocks: Vec<usize> = plan.block.clone();
157
158 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 let mut offsets = vec![0isize; num_arrays];
168
169 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
237macro_rules! elem_loops {
256 ($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 ($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
292macro_rules! block_loop {
299 ($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 ($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
343macro_rules! make_kernel {
350 ($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 ($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#[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 let mut idx = vec![0usize; rank];
487
488 if dims.contains(&0) {
492 return Ok(());
493 }
494 loop {
495 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 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#[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 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
601pub 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 RowMajor,
631 ColMajor,
633}
634
635pub(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 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 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
679pub(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#[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#[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#[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;