1use std::iter::FusedIterator;
4use std::mem::transmute;
5use std::ops::Range;
6
7use rten_base::iter::SplitIterator;
8use smallvec::SmallVec;
9
10use super::{AsView, DynLayout, NdTensorView, NdTensorViewMut, TensorBase, TensorViewMut};
11use crate::layout::{Layout, MutLayout, NdLayout, OverlapPolicy, RemoveDim, SizeArray, merge_axes};
12use crate::storage::{StorageMut, ViewData, ViewMutData};
13
14mod parallel;
15
16#[derive(Copy, Clone, Debug, Default)]
18struct IterPos {
19 remaining: usize,
21
22 offset: usize,
24
25 stride: usize,
27
28 max_remaining: usize,
30}
31
32impl IterPos {
33 fn from_size_stride(size: usize, stride: usize) -> Self {
34 let remaining = size.saturating_sub(1);
35 IterPos {
36 remaining,
37 offset: 0,
38 stride,
39 max_remaining: remaining,
40 }
41 }
42
43 #[inline(always)]
44 fn step(&mut self) -> bool {
45 if self.remaining != 0 {
46 self.remaining -= 1;
47 self.offset += self.stride;
48 true
49 } else {
50 self.remaining = self.max_remaining;
51 self.offset = 0;
52 false
53 }
54 }
55
56 fn size(&self) -> usize {
58 self.max_remaining + 1
61 }
62
63 fn index(&self) -> usize {
65 self.max_remaining - self.remaining
66 }
67
68 fn set_index(&mut self, index: usize) {
70 self.remaining = self.max_remaining - index;
71 self.offset = index * self.stride;
72 }
73}
74
75const INNER_NDIM: usize = 2;
76
77#[derive(Clone, Debug)]
79struct OffsetsBase {
80 len: usize,
85
86 inner_offset: usize,
88
89 inner_pos: [IterPos; INNER_NDIM],
91
92 outer_offset: usize,
94
95 outer_pos: Vec<IterPos>,
102}
103
104impl OffsetsBase {
105 fn new<L: Layout>(layout: &L) -> OffsetsBase {
107 let merged = merge_axes(&layout.shape(), &layout.strides());
110
111 let inner_pos_pad = INNER_NDIM.saturating_sub(merged.len());
112 let n_outer = merged.len().saturating_sub(INNER_NDIM);
113
114 let inner_pos = std::array::from_fn(|dim| {
115 let (size, stride) = if dim < inner_pos_pad {
116 (1, 0)
117 } else {
118 merged[n_outer + dim - inner_pos_pad]
119 };
120 IterPos::from_size_stride(size, stride)
121 });
122
123 let outer_pos = (0..n_outer)
124 .map(|i| {
125 let (size, stride) = merged[i];
126 IterPos::from_size_stride(size, stride)
127 })
128 .collect();
129
130 OffsetsBase {
131 len: merged.iter().map(|dim| dim.0).product(),
132 inner_pos,
133 inner_offset: 0,
134 outer_pos,
135 outer_offset: 0,
136 }
137 }
138
139 fn step_outer_pos(&mut self) -> bool {
144 let mut done = self.outer_pos.is_empty();
145 for (i, dim) in self.outer_pos.iter_mut().enumerate().rev() {
146 if dim.step() {
147 break;
148 } else if i == 0 {
149 done = true;
150 }
151 }
152 self.outer_offset = self.outer_pos.iter().map(|p| p.offset).sum();
153 !done
154 }
155
156 fn pos(&self, dim: usize) -> IterPos {
157 let outer_ndim = self.outer_pos.len();
158 if dim >= outer_ndim {
159 self.inner_pos[dim - outer_ndim]
160 } else {
161 self.outer_pos[dim]
162 }
163 }
164
165 fn pos_mut(&mut self, dim: usize) -> &mut IterPos {
166 let outer_ndim = self.outer_pos.len();
167 if dim >= outer_ndim {
168 &mut self.inner_pos[dim - outer_ndim]
169 } else {
170 &mut self.outer_pos[dim]
171 }
172 }
173
174 fn step_by(&mut self, n: usize) {
176 let mut remaining = n.min(self.len);
177 self.len -= remaining;
178
179 for dim in (0..self.ndim()).rev() {
180 if remaining == 0 {
181 break;
182 }
183
184 let pos = self.pos_mut(dim);
185 let size = pos.size();
186 let new_index = pos.index() + remaining;
187 pos.set_index(new_index % size);
188 remaining = new_index / size;
189 }
190
191 self.inner_offset = self.inner_pos.iter().map(|p| p.offset).sum();
193 self.outer_offset = self.outer_pos.iter().map(|p| p.offset).sum();
194 }
195
196 fn ndim(&self) -> usize {
197 self.outer_pos.len() + self.inner_pos.len()
198 }
199
200 fn offset_from_linear_index(&self, index: usize) -> usize {
203 let mut offset = 0;
204 let mut shape_product = 1;
205 for dim in (0..self.ndim()).rev() {
206 let pos = self.pos(dim);
207 let dim_index = (index / shape_product) % pos.size();
208 shape_product *= pos.size();
209 offset += dim_index * pos.stride;
210 }
211 offset
212 }
213
214 fn truncate(&mut self, len: usize) {
216 self.len = self.len.min(len);
220 }
221}
222
223impl Iterator for OffsetsBase {
224 type Item = usize;
225
226 #[inline(always)]
227 fn next(&mut self) -> Option<usize> {
228 if self.len == 0 {
229 return None;
230 }
231 let offset = self.outer_offset + self.inner_offset;
232
233 self.len -= 1;
234
235 self.inner_offset += self.inner_pos[1].stride;
238
239 if !self.inner_pos[1].step() {
242 if !self.inner_pos[0].step() {
243 self.step_outer_pos();
244 }
245
246 self.inner_offset = self.inner_pos[0].offset;
251 }
252
253 Some(offset)
254 }
255
256 fn size_hint(&self) -> (usize, Option<usize>) {
257 (self.len, Some(self.len))
258 }
259
260 fn fold<B, F>(mut self, init: B, mut f: F) -> B
261 where
262 Self: Sized,
263 F: FnMut(B, usize) -> B,
264 {
265 if self.len == 0 {
267 return init;
268 }
269
270 let mut accum = init;
271 'outer: loop {
272 for i0 in self.inner_pos[0].index()..self.inner_pos[0].size() {
273 for i1 in self.inner_pos[1].index()..self.inner_pos[1].size() {
274 let inner_offset =
275 i0 * self.inner_pos[0].stride + i1 * self.inner_pos[1].stride;
276 accum = f(accum, self.outer_offset + inner_offset);
277
278 self.len -= 1;
279 if self.len == 0 {
280 break 'outer;
281 }
282 }
283 self.inner_pos[1].set_index(0);
284 }
285 self.inner_pos[0].set_index(0);
286
287 if !self.step_outer_pos() {
288 break;
289 }
290 }
291
292 accum
293 }
294}
295
296impl ExactSizeIterator for OffsetsBase {}
297
298impl DoubleEndedIterator for OffsetsBase {
299 fn next_back(&mut self) -> Option<usize> {
300 if self.len == 0 {
301 return None;
302 }
303
304 let index = self.len - 1;
307 let offset = self.offset_from_linear_index(index);
308 self.len -= 1;
309
310 Some(offset)
311 }
312}
313
314impl SplitIterator for OffsetsBase {
315 fn split_at(mut self, index: usize) -> (Self, Self) {
318 assert!(self.len >= index);
319
320 let mut right = self.clone();
321 OffsetsBase::step_by(&mut right, index);
322
323 self.truncate(index);
324
325 (self, right)
326 }
327}
328
329pub struct Iter<'a, T> {
331 offsets: Offsets,
332 data: ViewData<'a, T>,
333}
334
335impl<'a, T> Iter<'a, T> {
336 pub(super) fn new<L: Layout + Clone>(view: TensorBase<ViewData<'a, T>, L>) -> Iter<'a, T> {
337 Iter {
338 offsets: Offsets::new(view.layout()),
339 data: view.storage(),
340 }
341 }
342}
343
344impl<T> Clone for Iter<'_, T> {
345 fn clone(&self) -> Self {
346 Iter {
347 offsets: self.offsets.clone(),
348 data: self.data,
349 }
350 }
351}
352
353impl<'a, T> Iterator for Iter<'a, T> {
354 type Item = &'a T;
355
356 #[inline(always)]
357 fn next(&mut self) -> Option<Self::Item> {
358 let offset = self.offsets.next()?;
359
360 Some(unsafe { self.data.get_unchecked(offset) })
362 }
363
364 fn size_hint(&self) -> (usize, Option<usize>) {
365 self.offsets.size_hint()
366 }
367
368 fn nth(&mut self, n: usize) -> Option<Self::Item> {
369 let offset = self.offsets.nth(n)?;
370
371 Some(unsafe { self.data.get_unchecked(offset) })
373 }
374
375 fn fold<B, F>(self, init: B, mut f: F) -> B
376 where
377 Self: Sized,
378 F: FnMut(B, Self::Item) -> B,
379 {
380 self.offsets.fold(init, |acc, offset| {
381 let item = unsafe { self.data.get_unchecked(offset) };
383 f(acc, item)
384 })
385 }
386}
387
388impl<'a, T> DoubleEndedIterator for Iter<'a, T> {
389 fn next_back(&mut self) -> Option<Self::Item> {
390 let offset = self.offsets.next_back()?;
391
392 Some(unsafe { self.data.get_unchecked(offset) })
394 }
395}
396
397impl<T> ExactSizeIterator for Iter<'_, T> {}
398
399impl<T> FusedIterator for Iter<'_, T> {}
400
401unsafe fn transmute_lifetime_mut<'a, 'b, T>(x: &'a mut T) -> &'b mut T {
404 unsafe { transmute::<&'a mut T, &'b mut T>(x) }
405}
406
407pub struct IterMut<'a, T> {
409 offsets: Offsets,
410 data: ViewMutData<'a, T>,
411}
412
413impl<'a, T> IterMut<'a, T> {
414 pub(super) fn new<L: Layout + Clone>(
415 view: TensorBase<ViewMutData<'a, T>, L>,
416 ) -> IterMut<'a, T> {
417 IterMut {
418 offsets: Offsets::new(view.layout()),
419 data: view.into_storage(),
420 }
421 }
422}
423
424impl<'a, T> Iterator for IterMut<'a, T> {
425 type Item = &'a mut T;
426
427 #[inline]
428 fn next(&mut self) -> Option<Self::Item> {
429 let offset = self.offsets.next()?;
430
431 Some(unsafe { transmute_lifetime_mut(self.data.get_unchecked_mut(offset)) })
434 }
435
436 #[inline]
437 fn size_hint(&self) -> (usize, Option<usize>) {
438 self.offsets.size_hint()
439 }
440
441 fn nth(&mut self, n: usize) -> Option<Self::Item> {
442 let offset = self.offsets.nth(n)?;
443
444 Some(unsafe { transmute_lifetime_mut(self.data.get_unchecked_mut(offset)) })
447 }
448
449 fn fold<B, F>(mut self, init: B, mut f: F) -> B
450 where
451 Self: Sized,
452 F: FnMut(B, Self::Item) -> B,
453 {
454 self.offsets.fold(init, |acc, offset| {
455 let item = unsafe { transmute_lifetime_mut(self.data.get_unchecked_mut(offset)) };
458 f(acc, item)
459 })
460 }
461}
462
463impl<T> DoubleEndedIterator for IterMut<'_, T> {
464 fn next_back(&mut self) -> Option<Self::Item> {
465 let offset = self.offsets.next_back()?;
466
467 Some(unsafe { transmute_lifetime_mut(self.data.get_unchecked_mut(offset)) })
470 }
471}
472
473impl<T> ExactSizeIterator for IterMut<'_, T> {}
474
475impl<T> FusedIterator for IterMut<'_, T> {}
476
477#[derive(Clone)]
478enum OffsetsKind {
479 Range(Range<usize>),
480 Indexing(OffsetsBase),
481}
482
483#[derive(Clone)]
490struct Offsets {
491 base: OffsetsKind,
492}
493
494impl Offsets {
495 fn new<L: Layout>(layout: &L) -> Offsets {
496 Offsets {
497 base: if layout.is_contiguous() {
498 OffsetsKind::Range(0..layout.min_data_len())
499 } else {
500 OffsetsKind::Indexing(OffsetsBase::new(layout))
501 },
502 }
503 }
504}
505
506impl Iterator for Offsets {
507 type Item = usize;
508
509 #[inline]
510 fn next(&mut self) -> Option<Self::Item> {
511 match &mut self.base {
512 OffsetsKind::Range(r) => r.next(),
513 OffsetsKind::Indexing(base) => base.next(),
514 }
515 }
516
517 fn size_hint(&self) -> (usize, Option<usize>) {
518 match &self.base {
519 OffsetsKind::Range(r) => r.size_hint(),
520 OffsetsKind::Indexing(base) => (base.len, Some(base.len)),
521 }
522 }
523
524 fn nth(&mut self, n: usize) -> Option<Self::Item> {
525 match &mut self.base {
526 OffsetsKind::Range(r) => r.nth(n),
527 OffsetsKind::Indexing(base) => {
528 base.step_by(n);
529 self.next()
530 }
531 }
532 }
533
534 fn fold<B, F>(self, init: B, f: F) -> B
535 where
536 Self: Sized,
537 F: FnMut(B, Self::Item) -> B,
538 {
539 match self.base {
540 OffsetsKind::Range(r) => r.fold(init, f),
541 OffsetsKind::Indexing(base) => base.fold(init, f),
542 }
543 }
544}
545
546impl DoubleEndedIterator for Offsets {
547 fn next_back(&mut self) -> Option<Self::Item> {
548 match &mut self.base {
549 OffsetsKind::Range(r) => r.next_back(),
550 OffsetsKind::Indexing(base) => base.next_back(),
551 }
552 }
553}
554
555impl ExactSizeIterator for Offsets {}
556
557impl FusedIterator for Offsets {}
558
559struct LaneRanges {
562 offsets: Offsets,
564
565 dim_size: usize,
567 dim_stride: usize,
568}
569
570impl LaneRanges {
571 fn new<L: Layout + RemoveDim>(layout: &L, dim: usize) -> LaneRanges {
572 let offsets = if layout.is_empty() {
575 Offsets::new(layout)
576 } else {
577 let other_dims = layout.remove_dim(dim);
578 Offsets::new(&other_dims)
579 };
580
581 LaneRanges {
582 offsets,
583 dim_size: layout.size(dim),
584 dim_stride: layout.stride(dim),
585 }
586 }
587
588 fn lane_offset_range(&self, start_offset: usize) -> Range<usize> {
591 lane_offsets(start_offset, self.dim_size, self.dim_stride)
592 }
593}
594
595fn lane_offsets(start_offset: usize, size: usize, stride: usize) -> Range<usize> {
596 start_offset..start_offset + (size - 1) * stride + 1
597}
598
599impl Iterator for LaneRanges {
600 type Item = Range<usize>;
601
602 #[inline]
603 fn next(&mut self) -> Option<Range<usize>> {
604 self.offsets
605 .next()
606 .map(|offset| self.lane_offset_range(offset))
607 }
608
609 fn size_hint(&self) -> (usize, Option<usize>) {
610 self.offsets.size_hint()
611 }
612
613 fn fold<B, F>(self, init: B, mut f: F) -> B
614 where
615 Self: Sized,
616 F: FnMut(B, Self::Item) -> B,
617 {
618 let Self {
619 offsets,
620 dim_size,
621 dim_stride,
622 } = self;
623
624 offsets.fold(init, |acc, offset| {
625 f(acc, lane_offsets(offset, dim_size, dim_stride))
626 })
627 }
628}
629
630impl DoubleEndedIterator for LaneRanges {
631 fn next_back(&mut self) -> Option<Range<usize>> {
632 self.offsets
633 .next_back()
634 .map(|offset| self.lane_offset_range(offset))
635 }
636}
637
638impl ExactSizeIterator for LaneRanges {}
639
640impl FusedIterator for LaneRanges {}
641
642pub struct Lanes<'a, T> {
647 data: ViewData<'a, T>,
648 ranges: LaneRanges,
649 lane_layout: NdLayout<1>,
650}
651
652#[derive(Clone, Debug)]
654pub struct Lane<'a, T> {
655 view: NdTensorView<'a, T, 1>,
656 index: usize,
657}
658
659impl<'a, T> Lane<'a, T> {
660 pub fn as_slice(&self) -> Option<&'a [T]> {
662 self.view.data().map(|data| &data[self.index..])
663 }
664
665 pub fn get(&self, idx: usize) -> Option<&'a T> {
667 self.view.get([idx])
668 }
669
670 pub fn as_view(&self) -> NdTensorView<'a, T, 1> {
672 self.view
673 }
674}
675
676impl<'a, T> From<NdTensorView<'a, T, 1>> for Lane<'a, T> {
677 fn from(val: NdTensorView<'a, T, 1>) -> Self {
678 Lane {
679 view: val,
680 index: 0,
681 }
682 }
683}
684
685impl<'a, T> Iterator for Lane<'a, T> {
686 type Item = &'a T;
687
688 #[inline]
689 fn next(&mut self) -> Option<Self::Item> {
690 if self.index < self.view.len() {
691 let index = self.index;
692 self.index += 1;
693
694 Some(unsafe { self.view.get_unchecked([index]) })
696 } else {
697 None
698 }
699 }
700
701 fn size_hint(&self) -> (usize, Option<usize>) {
702 let size = self.view.size(0);
703 (size, Some(size))
704 }
705}
706
707impl<T> ExactSizeIterator for Lane<'_, T> {}
708
709impl<T> FusedIterator for Lane<'_, T> {}
710
711impl<T: PartialEq> PartialEq<Lane<'_, T>> for Lane<'_, T> {
712 fn eq(&self, other: &Lane<'_, T>) -> bool {
713 self.view.slice(self.index..) == other.view.slice(other.index..)
714 }
715}
716
717impl<T: PartialEq> PartialEq<Lane<'_, T>> for LaneMut<'_, T> {
718 fn eq(&self, other: &Lane<'_, T>) -> bool {
719 self.view.slice(self.index..) == other.view.slice(other.index..)
720 }
721}
722
723impl<'a, T> Lanes<'a, T> {
724 pub(crate) fn new<L: Layout + RemoveDim + Clone>(
727 view: TensorBase<ViewData<'a, T>, L>,
728 dim: usize,
729 ) -> Lanes<'a, T> {
730 let size = view.size(dim);
731 let stride = view.stride(dim);
732 let lane_layout =
733 NdLayout::from_shape_and_strides([size], [stride], OverlapPolicy::AllowOverlap)
734 .unwrap();
735 Lanes {
736 data: view.storage(),
737 ranges: LaneRanges::new(view.layout(), dim),
738 lane_layout,
739 }
740 }
741}
742
743fn lane_for_offset_range<T>(
744 data: ViewData<T>,
745 layout: NdLayout<1>,
746 offsets: Range<usize>,
747) -> Lane<T> {
748 let view = NdTensorView::from_storage_and_layout(data.slice(offsets), layout);
749 Lane { view, index: 0 }
750}
751
752impl<'a, T> Iterator for Lanes<'a, T> {
753 type Item = Lane<'a, T>;
754
755 #[inline]
757 fn next(&mut self) -> Option<Self::Item> {
758 self.ranges
759 .next()
760 .map(|range| lane_for_offset_range(self.data, self.lane_layout, range))
761 }
762
763 fn size_hint(&self) -> (usize, Option<usize>) {
764 self.ranges.size_hint()
765 }
766
767 fn fold<B, F>(self, init: B, mut f: F) -> B
768 where
769 Self: Sized,
770 F: FnMut(B, Self::Item) -> B,
771 {
772 self.ranges.fold(init, |acc, offsets| {
773 let lane = lane_for_offset_range(self.data, self.lane_layout, offsets);
774 f(acc, lane)
775 })
776 }
777}
778
779impl<T> DoubleEndedIterator for Lanes<'_, T> {
780 fn next_back(&mut self) -> Option<Self::Item> {
781 self.ranges
782 .next_back()
783 .map(|range| lane_for_offset_range(self.data, self.lane_layout, range))
784 }
785}
786
787impl<T> ExactSizeIterator for Lanes<'_, T> {}
788
789impl<T> FusedIterator for Lanes<'_, T> {}
790
791pub struct LanesMut<'a, T> {
797 data: ViewMutData<'a, T>,
798 ranges: LaneRanges,
799 lane_layout: NdLayout<1>,
800}
801
802impl<'a, T> LanesMut<'a, T> {
803 pub(crate) fn new<L: Layout + RemoveDim + Clone>(
806 view: TensorBase<ViewMutData<'a, T>, L>,
807 dim: usize,
808 ) -> LanesMut<'a, T> {
809 assert!(
811 !view.is_broadcast(),
812 "Cannot mutably iterate over broadcasting view"
813 );
814
815 let size = view.size(dim);
816 let stride = view.stride(dim);
817
818 let lane_layout =
822 NdLayout::from_shape_and_strides([size], [stride], OverlapPolicy::AllowOverlap)
823 .unwrap();
824
825 LanesMut {
826 ranges: LaneRanges::new(view.layout(), dim),
827 data: view.into_storage(),
828 lane_layout,
829 }
830 }
831}
832
833impl<'a, T> Iterator for LanesMut<'a, T> {
834 type Item = LaneMut<'a, T>;
835
836 #[inline]
837 fn next(&mut self) -> Option<LaneMut<'a, T>> {
838 self.ranges.next().map(|offsets| {
839 unsafe {
842 LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
843 }
844 })
845 }
846
847 fn size_hint(&self) -> (usize, Option<usize>) {
848 self.ranges.size_hint()
849 }
850
851 fn fold<B, F>(mut self, init: B, mut f: F) -> B
852 where
853 Self: Sized,
854 F: FnMut(B, Self::Item) -> B,
855 {
856 self.ranges.fold(init, |acc, offsets| {
857 let lane = unsafe {
860 LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
861 };
862 f(acc, lane)
863 })
864 }
865}
866
867impl<'a, T> ExactSizeIterator for LanesMut<'a, T> {}
868
869impl<'a, T> DoubleEndedIterator for LanesMut<'a, T> {
870 fn next_back(&mut self) -> Option<LaneMut<'a, T>> {
871 self.ranges.next_back().map(|offsets| {
872 unsafe {
875 LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
876 }
877 })
878 }
879}
880
881#[derive(Debug)]
883pub struct LaneMut<'a, T> {
884 view: NdTensorViewMut<'a, T, 1>,
885 index: usize,
886}
887
888impl<'a, T> LaneMut<'a, T> {
889 unsafe fn from_storage_layout(data: ViewMutData<'a, T>, layout: NdLayout<1>) -> Self {
896 let view = unsafe {
897 NdTensorViewMut::from_storage_and_layout_unchecked(data, layout)
901 };
902 LaneMut { view, index: 0 }
903 }
904
905 pub fn as_slice_mut(&mut self) -> Option<&mut [T]> {
907 self.view.data_mut().map(|data| &mut data[self.index..])
908 }
909
910 pub fn into_view(self) -> NdTensorViewMut<'a, T, 1> {
912 self.view
913 }
914}
915
916impl<'a, T> Iterator for LaneMut<'a, T> {
917 type Item = &'a mut T;
918
919 #[inline]
920 fn next(&mut self) -> Option<Self::Item> {
921 if self.index < self.view.size(0) {
922 let index = self.index;
923 self.index += 1;
924 let item = unsafe { self.view.get_unchecked_mut([index]) };
925
926 Some(unsafe { transmute::<&mut T, Self::Item>(item) })
929 } else {
930 None
931 }
932 }
933
934 #[inline]
935 fn nth(&mut self, nth: usize) -> Option<Self::Item> {
936 self.index = (self.index + nth).min(self.view.size(0));
937 self.next()
938 }
939
940 fn size_hint(&self) -> (usize, Option<usize>) {
941 let size = self.view.size(0);
942 (size, Some(size))
943 }
944}
945
946impl<T> ExactSizeIterator for LaneMut<'_, T> {}
947
948impl<T: PartialEq> PartialEq<LaneMut<'_, T>> for LaneMut<'_, T> {
949 fn eq(&self, other: &LaneMut<'_, T>) -> bool {
950 self.view.slice(self.index..) == other.view.slice(other.index..)
951 }
952}
953
954struct InnerIterBase<L: Layout> {
957 outer_offsets: Offsets,
960 inner_layout: L,
961 inner_data_len: usize,
962}
963
964impl<L: Layout + Clone> InnerIterBase<L> {
965 fn new_impl<PL: Layout, F: Fn(&[usize], &[usize]) -> L>(
966 parent_layout: &PL,
967 inner_dims: usize,
968 make_inner_layout: F,
969 ) -> InnerIterBase<L> {
970 assert!(parent_layout.ndim() >= inner_dims);
971 let outer_dims = parent_layout.ndim() - inner_dims;
972 let parent_shape = parent_layout.shape();
973 let parent_strides = parent_layout.strides();
974
975 let parent_dims: SmallVec<[usize; 5]> = parent_shape.iter().collect();
976 let (outer_shape, inner_shape) = parent_dims.as_ref().split_at(outer_dims);
977
978 let parent_strides: SmallVec<[usize; 5]> = parent_strides.iter().collect();
979 let (outer_strides, inner_strides) = parent_strides.as_ref().split_at(outer_dims);
980
981 let inner_layout = make_inner_layout(inner_shape, inner_strides);
982 let inner_data_len = inner_layout.min_data_len();
983
984 let zero_strides: SmallVec<[usize; 5]>;
988 let outer_strides = if inner_data_len == 0 {
989 zero_strides = SmallVec::from_elem(0, outer_dims);
990 zero_strides.as_ref()
991 } else {
992 outer_strides
993 };
994
995 let outer_layout = DynLayout::from_shape_and_strides(
996 outer_shape,
997 outer_strides,
998 OverlapPolicy::AllowOverlap,
999 )
1000 .unwrap();
1001
1002 InnerIterBase {
1003 outer_offsets: Offsets::new(&outer_layout),
1004 inner_data_len,
1005 inner_layout,
1006 }
1007 }
1008}
1009
1010impl<const N: usize> InnerIterBase<NdLayout<N>> {
1011 pub(crate) fn new<L: Layout>(parent_layout: &L) -> Self {
1012 Self::new_impl(parent_layout, N, |inner_shape, inner_strides| {
1013 let inner_shape: [usize; N] = inner_shape.try_into().unwrap();
1014 let inner_strides: [usize; N] = inner_strides.try_into().unwrap();
1015 NdLayout::from_shape_and_strides(
1016 inner_shape,
1017 inner_strides,
1018 OverlapPolicy::AllowOverlap,
1021 )
1022 .expect("failed to create layout")
1023 })
1024 }
1025}
1026
1027impl InnerIterBase<DynLayout> {
1028 pub(crate) fn new_dyn<L: Layout>(parent_layout: &L, inner_dims: usize) -> Self {
1029 Self::new_impl(parent_layout, inner_dims, |inner_shape, inner_strides| {
1030 DynLayout::from_shape_and_strides(
1031 inner_shape,
1032 inner_strides,
1033 OverlapPolicy::AllowOverlap,
1036 )
1037 .expect("failed to create layout")
1038 })
1039 }
1040}
1041
1042impl<L: Layout> Iterator for InnerIterBase<L> {
1043 type Item = Range<usize>;
1045
1046 fn next(&mut self) -> Option<Range<usize>> {
1047 self.outer_offsets
1048 .next()
1049 .map(|offset| offset..offset + self.inner_data_len)
1050 }
1051
1052 fn size_hint(&self) -> (usize, Option<usize>) {
1053 self.outer_offsets.size_hint()
1054 }
1055
1056 fn fold<B, F>(self, init: B, mut f: F) -> B
1057 where
1058 Self: Sized,
1059 F: FnMut(B, Self::Item) -> B,
1060 {
1061 self.outer_offsets.fold(init, |acc, offset| {
1062 f(acc, offset..offset + self.inner_data_len)
1063 })
1064 }
1065}
1066
1067impl<L: Layout> ExactSizeIterator for InnerIterBase<L> {}
1068
1069impl<L: Layout> DoubleEndedIterator for InnerIterBase<L> {
1070 fn next_back(&mut self) -> Option<Self::Item> {
1071 self.outer_offsets
1072 .next_back()
1073 .map(|offset| offset..offset + self.inner_data_len)
1074 }
1075}
1076
1077pub struct InnerIter<'a, T, L: Layout> {
1080 base: InnerIterBase<L>,
1081 data: ViewData<'a, T>,
1082}
1083
1084impl<'a, T, const N: usize> InnerIter<'a, T, NdLayout<N>> {
1085 pub(crate) fn new<L: Layout + Clone>(view: TensorBase<ViewData<'a, T>, L>) -> Self {
1086 let base = InnerIterBase::new(&view);
1087 InnerIter {
1088 base,
1089 data: view.storage(),
1090 }
1091 }
1092}
1093
1094impl<'a, T> InnerIter<'a, T, DynLayout> {
1095 pub(crate) fn new_dyn<L: Layout + Clone>(
1096 view: TensorBase<ViewData<'a, T>, L>,
1097 inner_dims: usize,
1098 ) -> Self {
1099 let base = InnerIterBase::new_dyn(&view, inner_dims);
1100 InnerIter {
1101 base,
1102 data: view.storage(),
1103 }
1104 }
1105}
1106
1107impl<'a, T, L: Layout + Clone> Iterator for InnerIter<'a, T, L> {
1108 type Item = TensorBase<ViewData<'a, T>, L>;
1109
1110 fn next(&mut self) -> Option<Self::Item> {
1111 self.base.next().map(|offset_range| {
1112 TensorBase::from_storage_and_layout(
1113 self.data.slice(offset_range),
1114 self.base.inner_layout.clone(),
1115 )
1116 })
1117 }
1118
1119 fn size_hint(&self) -> (usize, Option<usize>) {
1120 self.base.size_hint()
1121 }
1122
1123 fn fold<B, F>(self, init: B, mut f: F) -> B
1124 where
1125 Self: Sized,
1126 F: FnMut(B, Self::Item) -> B,
1127 {
1128 let inner_layout = self.base.inner_layout.clone();
1129 self.base.fold(init, |acc, offset_range| {
1130 let item = TensorBase::from_storage_and_layout(
1131 self.data.slice(offset_range),
1132 inner_layout.clone(),
1133 );
1134 f(acc, item)
1135 })
1136 }
1137}
1138
1139impl<T, L: Layout + Clone> ExactSizeIterator for InnerIter<'_, T, L> {}
1140
1141impl<T, L: Layout + Clone> DoubleEndedIterator for InnerIter<'_, T, L> {
1142 fn next_back(&mut self) -> Option<Self::Item> {
1143 self.base.next_back().map(|offset_range| {
1144 TensorBase::from_storage_and_layout(
1145 self.data.slice(offset_range),
1146 self.base.inner_layout.clone(),
1147 )
1148 })
1149 }
1150}
1151
1152pub struct InnerIterMut<'a, T, L: Layout> {
1155 base: InnerIterBase<L>,
1156 data: ViewMutData<'a, T>,
1157}
1158
1159impl<'a, T, const N: usize> InnerIterMut<'a, T, NdLayout<N>> {
1160 pub(crate) fn new<L: Layout>(view: TensorBase<ViewMutData<'a, T>, L>) -> Self {
1161 let base = InnerIterBase::new(&view);
1162 InnerIterMut {
1163 base,
1164 data: view.into_storage(),
1165 }
1166 }
1167}
1168
1169impl<'a, T> InnerIterMut<'a, T, DynLayout> {
1170 pub(crate) fn new_dyn<L: Layout>(
1171 view: TensorBase<ViewMutData<'a, T>, L>,
1172 inner_dims: usize,
1173 ) -> Self {
1174 let base = InnerIterBase::new_dyn(&view, inner_dims);
1175 InnerIterMut {
1176 base,
1177 data: view.into_storage(),
1178 }
1179 }
1180}
1181
1182impl<'a, T, L: Layout + Clone> Iterator for InnerIterMut<'a, T, L> {
1183 type Item = TensorBase<ViewMutData<'a, T>, L>;
1184
1185 fn next(&mut self) -> Option<Self::Item> {
1186 self.base.next().map(|offset_range| {
1187 let storage = self.data.slice_mut(offset_range);
1188 let storage = unsafe {
1189 std::mem::transmute::<ViewMutData<'_, T>, ViewMutData<'a, T>>(storage)
1194 };
1195 TensorBase::from_storage_and_layout(storage, self.base.inner_layout.clone())
1196 })
1197 }
1198
1199 fn size_hint(&self) -> (usize, Option<usize>) {
1200 self.base.size_hint()
1201 }
1202
1203 fn fold<B, F>(mut self, init: B, mut f: F) -> B
1204 where
1205 Self: Sized,
1206 F: FnMut(B, Self::Item) -> B,
1207 {
1208 let inner_layout = self.base.inner_layout.clone();
1209 self.base.fold(init, |acc, offset_range| {
1210 let storage = self.data.slice_mut(offset_range);
1211 let storage = unsafe {
1212 std::mem::transmute::<ViewMutData<'_, T>, ViewMutData<'a, T>>(storage)
1217 };
1218 let item = TensorBase::from_storage_and_layout(storage, inner_layout.clone());
1219 f(acc, item)
1220 })
1221 }
1222}
1223
1224impl<T, L: Layout + Clone> ExactSizeIterator for InnerIterMut<'_, T, L> {}
1225
1226impl<'a, T, L: Layout + Clone> DoubleEndedIterator for InnerIterMut<'a, T, L> {
1227 fn next_back(&mut self) -> Option<Self::Item> {
1228 self.base.next_back().map(|offset_range| {
1229 let storage = self.data.slice_mut(offset_range);
1230 let storage = unsafe {
1231 std::mem::transmute::<ViewMutData<'_, T>, ViewMutData<'a, T>>(storage)
1234 };
1235 TensorBase::from_storage_and_layout(storage, self.base.inner_layout.clone())
1236 })
1237 }
1238}
1239
1240pub struct AxisIter<'a, T, L: Layout + RemoveDim> {
1243 view: TensorBase<ViewData<'a, T>, L>,
1244 axis: usize,
1245 index: usize,
1246 end: usize,
1247}
1248
1249impl<'a, T, L: MutLayout + RemoveDim> AxisIter<'a, T, L> {
1250 pub(crate) fn new(view: &TensorBase<ViewData<'a, T>, L>, axis: usize) -> AxisIter<'a, T, L> {
1251 assert!(axis < view.ndim());
1252 AxisIter {
1253 view: view.clone(),
1254 axis,
1255 index: 0,
1256 end: view.size(axis),
1257 }
1258 }
1259}
1260
1261impl<'a, T, L: MutLayout + RemoveDim> Iterator for AxisIter<'a, T, L> {
1262 type Item = TensorBase<ViewData<'a, T>, <L as RemoveDim>::Output>;
1263
1264 fn next(&mut self) -> Option<Self::Item> {
1265 if self.index >= self.end {
1266 None
1267 } else {
1268 let slice = self.view.index_axis(self.axis, self.index);
1269 self.index += 1;
1270 Some(slice)
1271 }
1272 }
1273
1274 fn size_hint(&self) -> (usize, Option<usize>) {
1275 let len = self.end - self.index;
1276 (len, Some(len))
1277 }
1278}
1279
1280impl<'a, T, L: MutLayout + RemoveDim> ExactSizeIterator for AxisIter<'a, T, L> {}
1281
1282impl<'a, T, L: MutLayout + RemoveDim> DoubleEndedIterator for AxisIter<'a, T, L> {
1283 fn next_back(&mut self) -> Option<Self::Item> {
1284 if self.index >= self.end {
1285 None
1286 } else {
1287 let slice = self.view.index_axis(self.axis, self.end - 1);
1288 self.end -= 1;
1289 Some(slice)
1290 }
1291 }
1292}
1293
1294pub struct AxisIterMut<'a, T, L: Layout + RemoveDim> {
1296 view: TensorBase<ViewMutData<'a, T>, L>,
1297 axis: usize,
1298 index: usize,
1299 end: usize,
1300}
1301
1302impl<'a, T, L: Layout + RemoveDim + Clone> AxisIterMut<'a, T, L> {
1303 pub(crate) fn new(
1304 view: TensorBase<ViewMutData<'a, T>, L>,
1305 axis: usize,
1306 ) -> AxisIterMut<'a, T, L> {
1307 assert!(
1309 !view.layout().is_broadcast(),
1310 "Cannot mutably iterate over broadcasting view"
1311 );
1312 assert!(axis < view.ndim());
1313 AxisIterMut {
1314 axis,
1315 index: 0,
1316 end: view.size(axis),
1317 view,
1318 }
1319 }
1320}
1321
1322type SmallerMutView<'b, T, L> = TensorBase<ViewMutData<'b, T>, <L as RemoveDim>::Output>;
1324
1325impl<'a, T, L: MutLayout + RemoveDim> Iterator for AxisIterMut<'a, T, L> {
1326 type Item = TensorBase<ViewMutData<'a, T>, <L as RemoveDim>::Output>;
1327
1328 fn next(&mut self) -> Option<Self::Item> {
1329 if self.index >= self.end {
1330 None
1331 } else {
1332 let index = self.index;
1333 self.index += 1;
1334
1335 let slice = self.view.index_axis_mut(self.axis, index);
1336
1337 let view = unsafe { transmute::<SmallerMutView<'_, T, L>, Self::Item>(slice) };
1342
1343 Some(view)
1344 }
1345 }
1346
1347 fn size_hint(&self) -> (usize, Option<usize>) {
1348 let len = self.end - self.index;
1349 (len, Some(len))
1350 }
1351}
1352
1353impl<'a, T, L: MutLayout + RemoveDim> ExactSizeIterator for AxisIterMut<'a, T, L> {}
1354
1355impl<'a, T, L: MutLayout + RemoveDim> DoubleEndedIterator for AxisIterMut<'a, T, L> {
1356 fn next_back(&mut self) -> Option<Self::Item> {
1357 if self.index >= self.end {
1358 None
1359 } else {
1360 let index = self.end - 1;
1361 self.end -= 1;
1362
1363 let slice = self.view.index_axis_mut(self.axis, index);
1364
1365 let view = unsafe { transmute::<SmallerMutView<'_, T, L>, Self::Item>(slice) };
1370
1371 Some(view)
1372 }
1373 }
1374}
1375
1376pub struct AxisChunks<'a, T, L: MutLayout> {
1379 remainder: Option<TensorBase<ViewData<'a, T>, L>>,
1380 axis: usize,
1381 chunk_size: usize,
1382}
1383
1384impl<'a, T, L: MutLayout> AxisChunks<'a, T, L> {
1385 pub(crate) fn new(
1386 view: &TensorBase<ViewData<'a, T>, L>,
1387 axis: usize,
1388 chunk_size: usize,
1389 ) -> AxisChunks<'a, T, L> {
1390 assert!(chunk_size > 0, "chunk size must be > 0");
1391 AxisChunks {
1392 remainder: if view.size(axis) > 0 {
1393 Some(view.view())
1394 } else {
1395 None
1396 },
1397 axis,
1398 chunk_size,
1399 }
1400 }
1401}
1402
1403impl<'a, T, L: MutLayout> Iterator for AxisChunks<'a, T, L> {
1404 type Item = TensorBase<ViewData<'a, T>, L>;
1405
1406 fn next(&mut self) -> Option<Self::Item> {
1407 let remainder = self.remainder.take()?;
1408 let chunk_len = self.chunk_size.min(remainder.size(self.axis));
1409 let (current, next_remainder) = remainder.split_at(self.axis, chunk_len);
1410 self.remainder = if next_remainder.size(self.axis) > 0 {
1411 Some(next_remainder)
1412 } else {
1413 None
1414 };
1415 Some(current)
1416 }
1417
1418 fn size_hint(&self) -> (usize, Option<usize>) {
1419 let len = self
1420 .remainder
1421 .as_ref()
1422 .map(|r| r.size(self.axis))
1423 .unwrap_or(0)
1424 .div_ceil(self.chunk_size);
1425 (len, Some(len))
1426 }
1427}
1428
1429impl<'a, T, L: MutLayout> ExactSizeIterator for AxisChunks<'a, T, L> {}
1430
1431impl<'a, T, L: MutLayout> DoubleEndedIterator for AxisChunks<'a, T, L> {
1432 fn next_back(&mut self) -> Option<Self::Item> {
1433 let remainder = self.remainder.take()?;
1434 let chunk_len = self.chunk_size.min(remainder.size(self.axis));
1435 let (prev_remainder, current) =
1436 remainder.split_at(self.axis, remainder.size(self.axis) - chunk_len);
1437 self.remainder = if prev_remainder.size(self.axis) > 0 {
1438 Some(prev_remainder)
1439 } else {
1440 None
1441 };
1442 Some(current)
1443 }
1444}
1445
1446pub struct AxisChunksMut<'a, T, L: MutLayout> {
1448 remainder: Option<TensorBase<ViewMutData<'a, T>, L>>,
1449 axis: usize,
1450 chunk_size: usize,
1451}
1452
1453impl<'a, T, L: MutLayout> AxisChunksMut<'a, T, L> {
1454 pub(crate) fn new(
1455 view: TensorBase<ViewMutData<'a, T>, L>,
1456 axis: usize,
1457 chunk_size: usize,
1458 ) -> AxisChunksMut<'a, T, L> {
1459 assert!(
1461 !view.layout().is_broadcast(),
1462 "Cannot mutably iterate over broadcasting view"
1463 );
1464 assert!(chunk_size > 0, "chunk size must be > 0");
1465 AxisChunksMut {
1466 remainder: if view.size(axis) > 0 {
1467 Some(view)
1468 } else {
1469 None
1470 },
1471 axis,
1472 chunk_size,
1473 }
1474 }
1475}
1476
1477impl<'a, T, L: MutLayout> Iterator for AxisChunksMut<'a, T, L> {
1478 type Item = TensorBase<ViewMutData<'a, T>, L>;
1479
1480 fn next(&mut self) -> Option<Self::Item> {
1481 let remainder = self.remainder.take()?;
1482 let chunk_len = self.chunk_size.min(remainder.size(self.axis));
1483 let (current, next_remainder) = remainder.split_at_mut(self.axis, chunk_len);
1484 self.remainder = if next_remainder.size(self.axis) > 0 {
1485 Some(next_remainder)
1486 } else {
1487 None
1488 };
1489 Some(current)
1490 }
1491
1492 fn size_hint(&self) -> (usize, Option<usize>) {
1493 let len = self
1494 .remainder
1495 .as_ref()
1496 .map(|r| r.size(self.axis))
1497 .unwrap_or(0)
1498 .div_ceil(self.chunk_size);
1499 (len, Some(len))
1500 }
1501}
1502
1503impl<'a, T, L: MutLayout> ExactSizeIterator for AxisChunksMut<'a, T, L> {}
1504
1505impl<'a, T, L: MutLayout> DoubleEndedIterator for AxisChunksMut<'a, T, L> {
1506 fn next_back(&mut self) -> Option<Self::Item> {
1507 let remainder = self.remainder.take()?;
1508 let remainder_size = remainder.size(self.axis);
1509 let chunk_len = self.chunk_size.min(remainder_size);
1510 let (prev_remainder, current) =
1511 remainder.split_at_mut(self.axis, remainder_size - chunk_len);
1512 self.remainder = if prev_remainder.size(self.axis) > 0 {
1513 Some(prev_remainder)
1514 } else {
1515 None
1516 };
1517 Some(current)
1518 }
1519}
1520
1521pub(crate) fn for_each_mut<T, F: Fn(&mut T)>(mut view: TensorViewMut<T>, f: F) {
1523 while view.ndim() < 4 {
1524 view.insert_axis(0);
1525 }
1526
1527 view.inner_iter_mut::<4>().for_each(|mut src| {
1533 for i0 in 0..src.size(0) {
1534 for i1 in 0..src.size(1) {
1535 for i2 in 0..src.size(2) {
1536 for i3 in 0..src.size(3) {
1537 let x = unsafe { src.get_unchecked_mut([i0, i1, i2, i3]) };
1539 f(x);
1540 }
1541 }
1542 }
1543 }
1544 });
1545}
1546
1547#[cfg(test)]
1550mod tests {
1551 use super::{AxisChunks, AxisChunksMut, Lanes, LanesMut};
1552 use crate::{AsView, Layout, NdLayout, NdTensor, Tensor};
1553
1554 fn compare_reversed<T: PartialEq + std::fmt::Debug>(fwd: &[T], rev: &[T]) {
1555 assert_eq!(fwd.len(), rev.len());
1556 for (x, y) in fwd.iter().zip(rev.iter().rev()) {
1557 assert_eq!(x, y);
1558 }
1559 }
1560
1561 fn test_iterator<I: Iterator + ExactSizeIterator + DoubleEndedIterator>(
1563 create_iter: impl Fn() -> I,
1564 expected: &[I::Item],
1565 ) where
1566 I::Item: PartialEq + std::fmt::Debug,
1567 {
1568 let iter = create_iter();
1569
1570 let (min_len, max_len) = iter.size_hint();
1571 let items: Vec<_> = iter.collect();
1572
1573 assert_eq!(&items, expected);
1574
1575 assert_eq!(min_len, items.len(), "incorrect size lower bound");
1577 assert_eq!(max_len, Some(items.len()), "incorrect size upper bound");
1578
1579 let rev_items: Vec<_> = create_iter().rev().collect();
1581 compare_reversed(&items, &rev_items);
1582
1583 let mut iter = create_iter();
1585 for _x in &mut iter { }
1586 assert_eq!(iter.next(), None);
1587
1588 let mut fold_items = Vec::new();
1590 let mut idx = 0;
1591 create_iter().fold(0, |acc, item| {
1592 assert_eq!(acc, idx);
1593 fold_items.push(item);
1594 idx += 1;
1595 idx
1596 });
1597 assert_eq!(items, fold_items);
1598 }
1599
1600 trait MutIterable {
1605 type Iter<'a>: Iterator + ExactSizeIterator + DoubleEndedIterator
1606 where
1607 Self: 'a;
1608
1609 fn iter_mut(&mut self) -> Self::Iter<'_>;
1610 }
1611
1612 fn test_mut_iterator<M, T>(mut iterable: M, expected: &[T])
1614 where
1615 M: MutIterable,
1616 T: std::fmt::Debug,
1617 for<'a> <M::Iter<'a> as Iterator>::Item: std::fmt::Debug + PartialEq + PartialEq<T>,
1618 {
1619 {
1621 let iter = iterable.iter_mut();
1622 let (min_len, max_len) = iter.size_hint();
1623 let items: Vec<_> = iter.collect();
1624
1625 assert_eq!(items, expected);
1627
1628 assert_eq!(min_len, items.len(), "incorrect size lower bound");
1630 assert_eq!(max_len, Some(items.len()), "incorrect size upper bound");
1631 }
1632
1633 {
1635 let mut iter = iterable.iter_mut();
1636 for _x in &mut iter { }
1637 assert!(iter.next().is_none());
1638 }
1639
1640 {
1646 let items: Vec<_> = iterable.iter_mut().map(|x| format!("{:?}", x)).collect();
1647 let rev_items: Vec<_> = iterable
1648 .iter_mut()
1649 .rev()
1650 .map(|x| format!("{:?}", x))
1651 .collect();
1652 compare_reversed(&items, &rev_items);
1653 }
1654
1655 {
1657 let items: Vec<_> = iterable.iter_mut().map(|x| format!("{:?}", x)).collect();
1658 let mut fold_items = Vec::new();
1659 let mut idx = 0;
1660 iterable.iter_mut().fold(0, |acc, item| {
1661 assert_eq!(acc, idx);
1662 fold_items.push(format!("{:?}", item));
1663 idx += 1;
1664 idx
1665 });
1666 assert_eq!(items, fold_items);
1667 }
1668 }
1669
1670 #[test]
1671 fn test_axis_chunks() {
1672 let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1673 test_iterator(
1674 || tensor.axis_chunks(0, 1),
1675 &[tensor.slice(0..1), tensor.slice(1..2)],
1676 );
1677 }
1678
1679 #[test]
1680 fn test_axis_chunks_empty() {
1681 let x = Tensor::<i32>::zeros(&[5, 0]);
1682 assert!(AxisChunks::new(&x.view(), 1, 1).next().is_none());
1683 }
1684
1685 #[test]
1686 #[should_panic(expected = "chunk size must be > 0")]
1687 fn test_axis_chunks_zero_size() {
1688 let x = Tensor::<i32>::zeros(&[5, 0]);
1689 assert!(AxisChunks::new(&x.view(), 1, 0).next().is_none());
1690 }
1691
1692 #[test]
1693 fn test_axis_chunks_mut_empty() {
1694 let mut x = Tensor::<i32>::zeros(&[5, 0]);
1695 assert!(AxisChunksMut::new(x.view_mut(), 1, 1).next().is_none());
1696 }
1697
1698 #[test]
1699 fn test_axis_chunks_mut_rev() {
1700 let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1701 let fwd: Vec<_> = tensor
1702 .axis_chunks_mut(0, 1)
1703 .map(|view| view.to_vec())
1704 .collect();
1705 let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1706 let rev: Vec<_> = tensor
1707 .axis_chunks_mut(0, 1)
1708 .rev()
1709 .map(|view| view.to_vec())
1710 .collect();
1711 compare_reversed(&fwd, &rev);
1712 }
1713
1714 #[test]
1715 #[should_panic(expected = "chunk size must be > 0")]
1716 fn test_axis_chunks_mut_zero_size() {
1717 let mut x = Tensor::<i32>::zeros(&[5, 0]);
1718 assert!(AxisChunksMut::new(x.view_mut(), 1, 0).next().is_none());
1719 }
1720
1721 #[test]
1722 fn test_axis_iter() {
1723 let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1724 test_iterator(|| tensor.axis_iter(0), &[tensor.slice(0), tensor.slice(1)]);
1725 }
1726
1727 #[test]
1728 fn test_axis_iter_mut_rev() {
1729 let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1730 let fwd: Vec<_> = tensor.axis_iter_mut(0).map(|view| view.to_vec()).collect();
1731 let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1732 let rev: Vec<_> = tensor
1733 .axis_iter_mut(0)
1734 .rev()
1735 .map(|view| view.to_vec())
1736 .collect();
1737 compare_reversed(&fwd, &rev);
1738 }
1739
1740 #[test]
1741 fn test_inner_iter() {
1742 let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1743 test_iterator(
1744 || tensor.inner_iter::<2>(),
1745 &[tensor.slice(0), tensor.slice(1)],
1746 );
1747 }
1748
1749 #[test]
1750 fn test_inner_iter_empty() {
1751 let tensor = NdTensor::<i32, 2>::zeros([0, 3]);
1754 assert_eq!(tensor.strides(), [3, 1]);
1755 let view = tensor.permuted([1, 0]);
1756 assert_eq!(view.strides(), [1, 3]);
1757
1758 let mut count = 0;
1759 for lane in view.inner_iter::<1>() {
1760 assert_eq!(lane.shape(), [0]);
1761 count += 1;
1762 }
1763 assert_eq!(count, 3);
1764 }
1765
1766 #[test]
1767 fn test_inner_iter_mut() {
1768 struct InnerIterMutTest(NdTensor<i32, 3>);
1769
1770 impl MutIterable for InnerIterMutTest {
1771 type Iter<'a> = super::InnerIterMut<'a, i32, NdLayout<2>>;
1772
1773 fn iter_mut(&mut self) -> Self::Iter<'_> {
1774 self.0.inner_iter_mut::<2>()
1775 }
1776 }
1777
1778 let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1779 test_mut_iterator(
1780 InnerIterMutTest(tensor.clone()),
1781 &[tensor.slice(0), tensor.slice(1)],
1782 );
1783 }
1784
1785 #[test]
1786 fn test_lanes() {
1787 let x = NdTensor::from([[1, 2], [3, 4]]);
1788 test_iterator(
1789 || x.lanes(0),
1790 &[x.slice((.., 0)).into(), x.slice((.., 1)).into()],
1791 );
1792 test_iterator(|| x.lanes(1), &[x.slice(0).into(), x.slice(1).into()]);
1793 }
1794
1795 #[test]
1796 fn test_lanes_empty() {
1797 let x = Tensor::<i32>::zeros(&[5, 0]);
1798 assert!(Lanes::new(x.view().view_ref(), 0).next().is_none());
1799 assert!(Lanes::new(x.view().view_ref(), 1).next().is_none());
1800 }
1801
1802 #[test]
1803 fn test_lanes_mut() {
1804 use super::Lane;
1805
1806 struct LanesMutTest(NdTensor<i32, 2>);
1807
1808 impl MutIterable for LanesMutTest {
1809 type Iter<'a> = super::LanesMut<'a, i32>;
1810
1811 fn iter_mut(&mut self) -> Self::Iter<'_> {
1812 self.0.lanes_mut(0)
1813 }
1814 }
1815
1816 let tensor = NdTensor::from([[1, 2], [3, 4]]);
1817 test_mut_iterator::<_, Lane<i32>>(
1818 LanesMutTest(tensor.clone()),
1819 &[
1820 Lane::from(tensor.slice((.., 0))),
1821 Lane::from(tensor.slice((.., 1))),
1822 ],
1823 );
1824 }
1825
1826 #[test]
1827 fn test_lane_as_slice() {
1828 let x = NdTensor::from([0, 1, 2]);
1830 let mut lane = x.lanes(0).next().unwrap();
1831 assert_eq!(lane.as_slice(), Some([0, 1, 2].as_slice()));
1832 lane.next();
1833 assert_eq!(lane.as_slice(), Some([1, 2].as_slice()));
1834 lane.next();
1835 lane.next();
1836 assert_eq!(lane.as_slice(), Some([0i32; 0].as_slice()));
1837 lane.next();
1838 assert_eq!(lane.as_slice(), Some([0i32; 0].as_slice()));
1839
1840 let x = NdTensor::from([[1i32, 2], [3, 4]]);
1842 let lane = x.lanes(0).next().unwrap();
1843 assert_eq!(lane.as_slice(), None);
1844 }
1845
1846 #[test]
1847 fn test_lanes_mut_empty() {
1848 let mut x = Tensor::<i32>::zeros(&[5, 0]);
1849 assert!(LanesMut::new(x.mut_view_ref(), 0).next().is_none());
1850 assert!(LanesMut::new(x.mut_view_ref(), 1).next().is_none());
1851 }
1852
1853 #[test]
1854 fn test_iter_step_by() {
1855 let tensor = Tensor::<f32>::full(&[1, 3, 16, 8], 1.);
1856
1857 let tensor = tensor.slice((.., .., 1.., ..));
1860
1861 let sum = tensor.iter().sum::<f32>();
1862 for n_skip in 0..tensor.len() {
1863 let sum_skip = tensor.iter().skip(n_skip).sum::<f32>();
1864 assert_eq!(
1865 sum_skip,
1866 sum - n_skip as f32,
1867 "wrong sum for n_skip={}",
1868 n_skip
1869 );
1870 }
1871 }
1872
1873 #[test]
1874 fn test_iter_broadcast() {
1875 let tensor = Tensor::<f32>::full(&[1], 1.);
1876 let broadcast = tensor.broadcast([1, 3, 16, 8]);
1877 assert_eq!(broadcast.iter().len(), broadcast.len());
1878 let count = broadcast.iter().count();
1879 assert_eq!(count, broadcast.len());
1880 let sum = broadcast.iter().sum::<f32>();
1881 assert_eq!(sum, broadcast.len() as f32);
1882 }
1883
1884 #[test]
1885 fn test_iter() {
1886 let tensor = NdTensor::from([[[1, 2], [3, 4]]]);
1887
1888 test_iterator(|| tensor.iter().copied(), &[1, 2, 3, 4]);
1890
1891 test_iterator(|| tensor.transposed().iter().copied(), &[1, 3, 2, 4]);
1893 }
1894
1895 #[test]
1896 fn test_iter_mut() {
1897 struct IterTest(NdTensor<i32, 3>);
1898
1899 impl MutIterable for IterTest {
1900 type Iter<'a> = super::IterMut<'a, i32>;
1901
1902 fn iter_mut(&mut self) -> Self::Iter<'_> {
1903 self.0.iter_mut()
1904 }
1905 }
1906
1907 let tensor = NdTensor::from([[[1, 2], [3, 4]]]);
1908 test_mut_iterator(IterTest(tensor), &[&1, &2, &3, &4]);
1909 }
1910
1911 #[test]
1912 #[ignore]
1913 fn bench_iter() {
1914 use crate::Layout;
1915 use rten_bench::run_bench;
1916
1917 type Elem = i32;
1918
1919 let tensor = std::hint::black_box(Tensor::<Elem>::full(&[1, 6, 768, 64], 1));
1920 let n_trials = 1000;
1921 let mut result = Elem::default();
1922
1923 fn reduce<'a>(iter: impl Iterator<Item = &'a Elem>) -> Elem {
1924 iter.fold(Elem::default(), |acc, x| acc.wrapping_add(*x))
1925 }
1926
1927 run_bench(n_trials, Some("slice iter"), || {
1929 result = reduce(tensor.data().unwrap().iter());
1930 });
1931 println!("sum {}", result);
1932
1933 run_bench(n_trials, Some("contiguous iter"), || {
1936 result = reduce(tensor.iter());
1937 });
1938 println!("sum {}", result);
1939
1940 run_bench(n_trials, Some("contiguous reverse iter"), || {
1941 result = reduce(tensor.iter().rev());
1942 });
1943 println!("sum {}", result);
1944
1945 let slice = tensor.slice((.., .., 1.., ..));
1948 assert!(!slice.is_contiguous());
1949 let n_trials = 1000;
1950 run_bench(n_trials, Some("non-contiguous iter"), || {
1951 result = reduce(slice.iter());
1952 });
1953 println!("sum {}", result);
1954
1955 let n_trials = 100;
1958 run_bench(n_trials, Some("non-contiguous reverse iter"), || {
1959 result = reduce(slice.iter().rev());
1960 });
1961 println!("sum {}", result);
1962 }
1963
1964 #[test]
1965 #[ignore]
1966 fn bench_inner_iter() {
1967 use crate::rng::XorShiftRng;
1968 use rten_bench::run_bench;
1969
1970 let n_trials = 100;
1971 let mut rng = XorShiftRng::new(1234);
1972
1973 let tensor = Tensor::<f32>::rand(&[512, 512, 12, 1], &mut rng);
1977
1978 let mut sum = 0.;
1979 run_bench(n_trials, Some("inner iter"), || {
1980 for inner in tensor.inner_iter::<2>() {
1981 for i0 in 0..inner.size(0) {
1982 for i1 in 0..inner.size(1) {
1983 sum += inner[[i0, i1]];
1984 }
1985 }
1986 }
1987 });
1988 println!("sum {}", sum);
1989 }
1990}