1use crate::fuse::{compress_dims, fuse_dims};
7use crate::{block, order, Result};
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 {
247 ($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 ($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
271macro_rules! block_loop {
277 ($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 ($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
311macro_rules! make_kernel {
317 ($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 ($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#[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 let mut idx = vec![0usize; rank];
453
454 loop {
455 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 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#[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 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
537pub 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 RowMajor,
567 ColMajor,
569}
570
571pub(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 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 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#[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#[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#[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;