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
657 index: usize,
659
660 end: usize,
662}
663
664impl<'a, T> Lane<'a, T> {
665 pub fn as_slice(&self) -> Option<&'a [T]> {
667 self.view.data().map(|data| &data[self.index..self.end])
668 }
669
670 pub fn get(&self, idx: usize) -> Option<&'a T> {
672 self.view.get([idx])
673 }
674
675 pub fn as_view(&self) -> NdTensorView<'a, T, 1> {
677 self.view
678 }
679}
680
681impl<'a, T> From<NdTensorView<'a, T, 1>> for Lane<'a, T> {
682 fn from(val: NdTensorView<'a, T, 1>) -> Self {
683 Lane {
684 index: 0,
685 end: val.size(0),
686 view: val,
687 }
688 }
689}
690
691impl<'a, T> Iterator for Lane<'a, T> {
692 type Item = &'a T;
693
694 #[inline]
695 fn next(&mut self) -> Option<Self::Item> {
696 if self.index < self.end {
697 let index = self.index;
698 self.index += 1;
699
700 Some(unsafe { self.view.get_unchecked([index]) })
702 } else {
703 None
704 }
705 }
706
707 fn size_hint(&self) -> (usize, Option<usize>) {
708 let len = self.end - self.index;
709 (len, Some(len))
710 }
711}
712
713impl<T> DoubleEndedIterator for Lane<'_, T> {
714 #[inline]
715 fn next_back(&mut self) -> Option<Self::Item> {
716 if self.index < self.end {
717 self.end -= 1;
718
719 Some(unsafe { self.view.get_unchecked([self.end]) })
721 } else {
722 None
723 }
724 }
725}
726
727impl<T> ExactSizeIterator for Lane<'_, T> {}
728
729impl<T> FusedIterator for Lane<'_, T> {}
730
731impl<T: PartialEq> PartialEq<Lane<'_, T>> for Lane<'_, T> {
732 fn eq(&self, other: &Lane<'_, T>) -> bool {
733 self.view.slice(self.index..self.end) == other.view.slice(other.index..other.end)
734 }
735}
736
737impl<T: PartialEq> PartialEq<Lane<'_, T>> for LaneMut<'_, T> {
738 fn eq(&self, other: &Lane<'_, T>) -> bool {
739 self.view.slice(self.index..self.end) == other.view.slice(other.index..other.end)
740 }
741}
742
743impl<'a, T> Lanes<'a, T> {
744 pub(crate) fn new<L: Layout + RemoveDim + Clone>(
747 view: TensorBase<ViewData<'a, T>, L>,
748 dim: usize,
749 ) -> Lanes<'a, T> {
750 let size = view.size(dim);
751 let stride = view.stride(dim);
752 let lane_layout =
753 NdLayout::from_shape_and_strides([size], [stride], OverlapPolicy::AllowOverlap)
754 .unwrap();
755 Lanes {
756 data: view.storage(),
757 ranges: LaneRanges::new(view.layout(), dim),
758 lane_layout,
759 }
760 }
761}
762
763fn lane_for_offset_range<T>(
764 data: ViewData<T>,
765 layout: NdLayout<1>,
766 offsets: Range<usize>,
767) -> Lane<T> {
768 let view = NdTensorView::from_storage_and_layout(data.slice(offsets), layout);
769 Lane {
770 index: 0,
771 end: view.size(0),
772 view,
773 }
774}
775
776impl<'a, T> Iterator for Lanes<'a, T> {
777 type Item = Lane<'a, T>;
778
779 #[inline]
781 fn next(&mut self) -> Option<Self::Item> {
782 self.ranges
783 .next()
784 .map(|range| lane_for_offset_range(self.data, self.lane_layout, range))
785 }
786
787 fn size_hint(&self) -> (usize, Option<usize>) {
788 self.ranges.size_hint()
789 }
790
791 fn fold<B, F>(self, init: B, mut f: F) -> B
792 where
793 Self: Sized,
794 F: FnMut(B, Self::Item) -> B,
795 {
796 self.ranges.fold(init, |acc, offsets| {
797 let lane = lane_for_offset_range(self.data, self.lane_layout, offsets);
798 f(acc, lane)
799 })
800 }
801}
802
803impl<T> DoubleEndedIterator for Lanes<'_, T> {
804 fn next_back(&mut self) -> Option<Self::Item> {
805 self.ranges
806 .next_back()
807 .map(|range| lane_for_offset_range(self.data, self.lane_layout, range))
808 }
809}
810
811impl<T> ExactSizeIterator for Lanes<'_, T> {}
812
813impl<T> FusedIterator for Lanes<'_, T> {}
814
815pub struct LanesMut<'a, T> {
821 data: ViewMutData<'a, T>,
822 ranges: LaneRanges,
823 lane_layout: NdLayout<1>,
824}
825
826impl<'a, T> LanesMut<'a, T> {
827 pub(crate) fn new<L: Layout + RemoveDim + Clone>(
830 view: TensorBase<ViewMutData<'a, T>, L>,
831 dim: usize,
832 ) -> LanesMut<'a, T> {
833 assert!(
835 !view.is_broadcast(),
836 "Cannot mutably iterate over broadcasting view"
837 );
838
839 let size = view.size(dim);
840 let stride = view.stride(dim);
841
842 let lane_layout =
846 NdLayout::from_shape_and_strides([size], [stride], OverlapPolicy::AllowOverlap)
847 .unwrap();
848
849 LanesMut {
850 ranges: LaneRanges::new(view.layout(), dim),
851 data: view.into_storage(),
852 lane_layout,
853 }
854 }
855}
856
857impl<'a, T> Iterator for LanesMut<'a, T> {
858 type Item = LaneMut<'a, T>;
859
860 #[inline]
861 fn next(&mut self) -> Option<LaneMut<'a, T>> {
862 self.ranges.next().map(|offsets| {
863 unsafe {
866 LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
867 }
868 })
869 }
870
871 fn size_hint(&self) -> (usize, Option<usize>) {
872 self.ranges.size_hint()
873 }
874
875 fn fold<B, F>(mut self, init: B, mut f: F) -> B
876 where
877 Self: Sized,
878 F: FnMut(B, Self::Item) -> B,
879 {
880 self.ranges.fold(init, |acc, offsets| {
881 let lane = unsafe {
884 LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
885 };
886 f(acc, lane)
887 })
888 }
889}
890
891impl<'a, T> ExactSizeIterator for LanesMut<'a, T> {}
892
893impl<'a, T> DoubleEndedIterator for LanesMut<'a, T> {
894 fn next_back(&mut self) -> Option<LaneMut<'a, T>> {
895 self.ranges.next_back().map(|offsets| {
896 unsafe {
899 LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
900 }
901 })
902 }
903}
904
905#[derive(Debug)]
907pub struct LaneMut<'a, T> {
908 view: NdTensorViewMut<'a, T, 1>,
909
910 index: usize,
912
913 end: usize,
915}
916
917impl<'a, T> LaneMut<'a, T> {
918 unsafe fn from_storage_layout(data: ViewMutData<'a, T>, layout: NdLayout<1>) -> Self {
925 let view = unsafe {
926 NdTensorViewMut::from_storage_and_layout_unchecked(data, layout)
930 };
931 LaneMut {
932 index: 0,
933 end: view.size(0),
934 view,
935 }
936 }
937
938 pub fn as_slice_mut(&mut self) -> Option<&mut [T]> {
940 let (index, end) = (self.index, self.end);
941 self.view.data_mut().map(|data| &mut data[index..end])
942 }
943
944 #[track_caller]
951 pub fn into_view(self) -> NdTensorViewMut<'a, T, 1> {
952 assert!(
953 self.index == 0 && self.end == self.view.size(0),
954 "lane has been stepped"
955 );
956 self.view
957 }
958}
959
960impl<'a, T> Iterator for LaneMut<'a, T> {
961 type Item = &'a mut T;
962
963 #[inline]
964 fn next(&mut self) -> Option<Self::Item> {
965 if self.index < self.end {
966 let index = self.index;
967 self.index += 1;
968 let item = unsafe { self.view.get_unchecked_mut([index]) };
969
970 Some(unsafe { transmute::<&mut T, Self::Item>(item) })
973 } else {
974 None
975 }
976 }
977
978 #[inline]
979 fn nth(&mut self, nth: usize) -> Option<Self::Item> {
980 self.index = self.index.saturating_add(nth).min(self.end);
981 self.next()
982 }
983
984 fn size_hint(&self) -> (usize, Option<usize>) {
985 let len = self.end - self.index;
986 (len, Some(len))
987 }
988}
989
990impl<T> DoubleEndedIterator for LaneMut<'_, T> {
991 #[inline]
992 fn next_back(&mut self) -> Option<Self::Item> {
993 if self.index < self.end {
994 self.end -= 1;
995 let item = unsafe { self.view.get_unchecked_mut([self.end]) };
996
997 Some(unsafe { transmute::<&mut T, Self::Item>(item) })
1000 } else {
1001 None
1002 }
1003 }
1004}
1005
1006impl<T> ExactSizeIterator for LaneMut<'_, T> {}
1007
1008impl<T: PartialEq> PartialEq<LaneMut<'_, T>> for LaneMut<'_, T> {
1009 fn eq(&self, other: &LaneMut<'_, T>) -> bool {
1010 self.view.slice(self.index..self.end) == other.view.slice(other.index..other.end)
1011 }
1012}
1013
1014struct InnerIterBase<L: Layout> {
1017 outer_offsets: Offsets,
1020 inner_layout: L,
1021 inner_data_len: usize,
1022}
1023
1024impl<L: Layout + Clone> InnerIterBase<L> {
1025 fn new_impl<PL: Layout, F: Fn(&[usize], &[usize]) -> L>(
1026 parent_layout: &PL,
1027 inner_dims: usize,
1028 make_inner_layout: F,
1029 ) -> InnerIterBase<L> {
1030 assert!(parent_layout.ndim() >= inner_dims);
1031 let outer_dims = parent_layout.ndim() - inner_dims;
1032 let parent_shape = parent_layout.shape();
1033 let parent_strides = parent_layout.strides();
1034
1035 let parent_dims: SmallVec<[usize; 5]> = parent_shape.iter().collect();
1036 let (outer_shape, inner_shape) = parent_dims.as_ref().split_at(outer_dims);
1037
1038 let parent_strides: SmallVec<[usize; 5]> = parent_strides.iter().collect();
1039 let (outer_strides, inner_strides) = parent_strides.as_ref().split_at(outer_dims);
1040
1041 let inner_layout = make_inner_layout(inner_shape, inner_strides);
1042 let inner_data_len = inner_layout.min_data_len();
1043
1044 let zero_strides: SmallVec<[usize; 5]>;
1048 let outer_strides = if inner_data_len == 0 {
1049 zero_strides = SmallVec::from_elem(0, outer_dims);
1050 zero_strides.as_ref()
1051 } else {
1052 outer_strides
1053 };
1054
1055 let outer_layout = DynLayout::from_shape_and_strides(
1056 outer_shape,
1057 outer_strides,
1058 OverlapPolicy::AllowOverlap,
1059 )
1060 .unwrap();
1061
1062 InnerIterBase {
1063 outer_offsets: Offsets::new(&outer_layout),
1064 inner_data_len,
1065 inner_layout,
1066 }
1067 }
1068}
1069
1070impl<const N: usize> InnerIterBase<NdLayout<N>> {
1071 pub(crate) fn new<L: Layout>(parent_layout: &L) -> Self {
1072 Self::new_impl(parent_layout, N, |inner_shape, inner_strides| {
1073 let inner_shape: [usize; N] = inner_shape.try_into().unwrap();
1074 let inner_strides: [usize; N] = inner_strides.try_into().unwrap();
1075 NdLayout::from_shape_and_strides(
1076 inner_shape,
1077 inner_strides,
1078 OverlapPolicy::AllowOverlap,
1081 )
1082 .expect("failed to create layout")
1083 })
1084 }
1085}
1086
1087impl InnerIterBase<DynLayout> {
1088 pub(crate) fn new_dyn<L: Layout>(parent_layout: &L, inner_dims: usize) -> Self {
1089 Self::new_impl(parent_layout, inner_dims, |inner_shape, inner_strides| {
1090 DynLayout::from_shape_and_strides(
1091 inner_shape,
1092 inner_strides,
1093 OverlapPolicy::AllowOverlap,
1096 )
1097 .expect("failed to create layout")
1098 })
1099 }
1100}
1101
1102impl<L: Layout> Iterator for InnerIterBase<L> {
1103 type Item = Range<usize>;
1105
1106 fn next(&mut self) -> Option<Range<usize>> {
1107 self.outer_offsets
1108 .next()
1109 .map(|offset| offset..offset + self.inner_data_len)
1110 }
1111
1112 fn size_hint(&self) -> (usize, Option<usize>) {
1113 self.outer_offsets.size_hint()
1114 }
1115
1116 fn fold<B, F>(self, init: B, mut f: F) -> B
1117 where
1118 Self: Sized,
1119 F: FnMut(B, Self::Item) -> B,
1120 {
1121 self.outer_offsets.fold(init, |acc, offset| {
1122 f(acc, offset..offset + self.inner_data_len)
1123 })
1124 }
1125}
1126
1127impl<L: Layout> ExactSizeIterator for InnerIterBase<L> {}
1128
1129impl<L: Layout> DoubleEndedIterator for InnerIterBase<L> {
1130 fn next_back(&mut self) -> Option<Self::Item> {
1131 self.outer_offsets
1132 .next_back()
1133 .map(|offset| offset..offset + self.inner_data_len)
1134 }
1135}
1136
1137pub struct InnerIter<'a, T, L: Layout> {
1140 base: InnerIterBase<L>,
1141 data: ViewData<'a, T>,
1142}
1143
1144impl<'a, T, const N: usize> InnerIter<'a, T, NdLayout<N>> {
1145 pub(crate) fn new<L: Layout + Clone>(view: TensorBase<ViewData<'a, T>, L>) -> Self {
1146 let base = InnerIterBase::new(&view);
1147 InnerIter {
1148 base,
1149 data: view.storage(),
1150 }
1151 }
1152}
1153
1154impl<'a, T> InnerIter<'a, T, DynLayout> {
1155 pub(crate) fn new_dyn<L: Layout + Clone>(
1156 view: TensorBase<ViewData<'a, T>, L>,
1157 inner_dims: usize,
1158 ) -> Self {
1159 let base = InnerIterBase::new_dyn(&view, inner_dims);
1160 InnerIter {
1161 base,
1162 data: view.storage(),
1163 }
1164 }
1165}
1166
1167impl<'a, T, L: Layout + Clone> Iterator for InnerIter<'a, T, L> {
1168 type Item = TensorBase<ViewData<'a, T>, L>;
1169
1170 fn next(&mut self) -> Option<Self::Item> {
1171 self.base.next().map(|offset_range| {
1172 TensorBase::from_storage_and_layout(
1173 self.data.slice(offset_range),
1174 self.base.inner_layout.clone(),
1175 )
1176 })
1177 }
1178
1179 fn size_hint(&self) -> (usize, Option<usize>) {
1180 self.base.size_hint()
1181 }
1182
1183 fn fold<B, F>(self, init: B, mut f: F) -> B
1184 where
1185 Self: Sized,
1186 F: FnMut(B, Self::Item) -> B,
1187 {
1188 let inner_layout = self.base.inner_layout.clone();
1189 self.base.fold(init, |acc, offset_range| {
1190 let item = TensorBase::from_storage_and_layout(
1191 self.data.slice(offset_range),
1192 inner_layout.clone(),
1193 );
1194 f(acc, item)
1195 })
1196 }
1197}
1198
1199impl<T, L: Layout + Clone> ExactSizeIterator for InnerIter<'_, T, L> {}
1200
1201impl<T, L: Layout + Clone> DoubleEndedIterator for InnerIter<'_, T, L> {
1202 fn next_back(&mut self) -> Option<Self::Item> {
1203 self.base.next_back().map(|offset_range| {
1204 TensorBase::from_storage_and_layout(
1205 self.data.slice(offset_range),
1206 self.base.inner_layout.clone(),
1207 )
1208 })
1209 }
1210}
1211
1212pub struct InnerIterMut<'a, T, L: Layout> {
1215 base: InnerIterBase<L>,
1216 data: ViewMutData<'a, T>,
1217}
1218
1219impl<'a, T, const N: usize> InnerIterMut<'a, T, NdLayout<N>> {
1220 pub(crate) fn new<L: Layout>(view: TensorBase<ViewMutData<'a, T>, L>) -> Self {
1221 let base = InnerIterBase::new(&view);
1222 InnerIterMut {
1223 base,
1224 data: view.into_storage(),
1225 }
1226 }
1227}
1228
1229impl<'a, T> InnerIterMut<'a, T, DynLayout> {
1230 pub(crate) fn new_dyn<L: Layout>(
1231 view: TensorBase<ViewMutData<'a, T>, L>,
1232 inner_dims: usize,
1233 ) -> Self {
1234 let base = InnerIterBase::new_dyn(&view, inner_dims);
1235 InnerIterMut {
1236 base,
1237 data: view.into_storage(),
1238 }
1239 }
1240}
1241
1242impl<'a, T, L: Layout + Clone> Iterator for InnerIterMut<'a, T, L> {
1243 type Item = TensorBase<ViewMutData<'a, T>, L>;
1244
1245 fn next(&mut self) -> Option<Self::Item> {
1246 self.base.next().map(|offset_range| {
1247 let storage = self.data.slice_mut(offset_range);
1248 let storage = unsafe {
1249 std::mem::transmute::<ViewMutData<'_, T>, ViewMutData<'a, T>>(storage)
1254 };
1255 TensorBase::from_storage_and_layout(storage, self.base.inner_layout.clone())
1256 })
1257 }
1258
1259 fn size_hint(&self) -> (usize, Option<usize>) {
1260 self.base.size_hint()
1261 }
1262
1263 fn fold<B, F>(mut self, init: B, mut f: F) -> B
1264 where
1265 Self: Sized,
1266 F: FnMut(B, Self::Item) -> B,
1267 {
1268 let inner_layout = self.base.inner_layout.clone();
1269 self.base.fold(init, |acc, offset_range| {
1270 let storage = self.data.slice_mut(offset_range);
1271 let storage = unsafe {
1272 std::mem::transmute::<ViewMutData<'_, T>, ViewMutData<'a, T>>(storage)
1277 };
1278 let item = TensorBase::from_storage_and_layout(storage, inner_layout.clone());
1279 f(acc, item)
1280 })
1281 }
1282}
1283
1284impl<T, L: Layout + Clone> ExactSizeIterator for InnerIterMut<'_, T, L> {}
1285
1286impl<'a, T, L: Layout + Clone> DoubleEndedIterator for InnerIterMut<'a, T, L> {
1287 fn next_back(&mut self) -> Option<Self::Item> {
1288 self.base.next_back().map(|offset_range| {
1289 let storage = self.data.slice_mut(offset_range);
1290 let storage = unsafe {
1291 std::mem::transmute::<ViewMutData<'_, T>, ViewMutData<'a, T>>(storage)
1294 };
1295 TensorBase::from_storage_and_layout(storage, self.base.inner_layout.clone())
1296 })
1297 }
1298}
1299
1300pub struct AxisIter<'a, T, L: Layout + RemoveDim> {
1303 view: TensorBase<ViewData<'a, T>, L>,
1304 axis: usize,
1305 index: usize,
1306 end: usize,
1307}
1308
1309impl<'a, T, L: MutLayout + RemoveDim> AxisIter<'a, T, L> {
1310 pub(crate) fn new(view: &TensorBase<ViewData<'a, T>, L>, axis: usize) -> AxisIter<'a, T, L> {
1311 assert!(axis < view.ndim());
1312 AxisIter {
1313 view: view.clone(),
1314 axis,
1315 index: 0,
1316 end: view.size(axis),
1317 }
1318 }
1319}
1320
1321impl<'a, T, L: MutLayout + RemoveDim> Iterator for AxisIter<'a, T, L> {
1322 type Item = TensorBase<ViewData<'a, T>, <L as RemoveDim>::Output>;
1323
1324 fn next(&mut self) -> Option<Self::Item> {
1325 if self.index >= self.end {
1326 None
1327 } else {
1328 let slice = self.view.index_axis(self.axis, self.index);
1329 self.index += 1;
1330 Some(slice)
1331 }
1332 }
1333
1334 fn size_hint(&self) -> (usize, Option<usize>) {
1335 let len = self.end - self.index;
1336 (len, Some(len))
1337 }
1338}
1339
1340impl<'a, T, L: MutLayout + RemoveDim> ExactSizeIterator for AxisIter<'a, T, L> {}
1341
1342impl<'a, T, L: MutLayout + RemoveDim> DoubleEndedIterator for AxisIter<'a, T, L> {
1343 fn next_back(&mut self) -> Option<Self::Item> {
1344 if self.index >= self.end {
1345 None
1346 } else {
1347 let slice = self.view.index_axis(self.axis, self.end - 1);
1348 self.end -= 1;
1349 Some(slice)
1350 }
1351 }
1352}
1353
1354pub struct AxisIterMut<'a, T, L: Layout + RemoveDim> {
1356 view: TensorBase<ViewMutData<'a, T>, L>,
1357 axis: usize,
1358 index: usize,
1359 end: usize,
1360}
1361
1362impl<'a, T, L: Layout + RemoveDim + Clone> AxisIterMut<'a, T, L> {
1363 pub(crate) fn new(
1364 view: TensorBase<ViewMutData<'a, T>, L>,
1365 axis: usize,
1366 ) -> AxisIterMut<'a, T, L> {
1367 assert!(
1369 !view.layout().is_broadcast(),
1370 "Cannot mutably iterate over broadcasting view"
1371 );
1372 assert!(axis < view.ndim());
1373 AxisIterMut {
1374 axis,
1375 index: 0,
1376 end: view.size(axis),
1377 view,
1378 }
1379 }
1380}
1381
1382type SmallerMutView<'b, T, L> = TensorBase<ViewMutData<'b, T>, <L as RemoveDim>::Output>;
1384
1385impl<'a, T, L: MutLayout + RemoveDim> Iterator for AxisIterMut<'a, T, L> {
1386 type Item = TensorBase<ViewMutData<'a, T>, <L as RemoveDim>::Output>;
1387
1388 fn next(&mut self) -> Option<Self::Item> {
1389 if self.index >= self.end {
1390 None
1391 } else {
1392 let index = self.index;
1393 self.index += 1;
1394
1395 let slice = self.view.index_axis_mut(self.axis, index);
1396
1397 let view = unsafe { transmute::<SmallerMutView<'_, T, L>, Self::Item>(slice) };
1402
1403 Some(view)
1404 }
1405 }
1406
1407 fn size_hint(&self) -> (usize, Option<usize>) {
1408 let len = self.end - self.index;
1409 (len, Some(len))
1410 }
1411}
1412
1413impl<'a, T, L: MutLayout + RemoveDim> ExactSizeIterator for AxisIterMut<'a, T, L> {}
1414
1415impl<'a, T, L: MutLayout + RemoveDim> DoubleEndedIterator for AxisIterMut<'a, T, L> {
1416 fn next_back(&mut self) -> Option<Self::Item> {
1417 if self.index >= self.end {
1418 None
1419 } else {
1420 let index = self.end - 1;
1421 self.end -= 1;
1422
1423 let slice = self.view.index_axis_mut(self.axis, index);
1424
1425 let view = unsafe { transmute::<SmallerMutView<'_, T, L>, Self::Item>(slice) };
1430
1431 Some(view)
1432 }
1433 }
1434}
1435
1436pub struct AxisChunks<'a, T, L: MutLayout> {
1439 remainder: Option<TensorBase<ViewData<'a, T>, L>>,
1440 axis: usize,
1441 chunk_size: usize,
1442}
1443
1444impl<'a, T, L: MutLayout> AxisChunks<'a, T, L> {
1445 pub(crate) fn new(
1446 view: &TensorBase<ViewData<'a, T>, L>,
1447 axis: usize,
1448 chunk_size: usize,
1449 ) -> AxisChunks<'a, T, L> {
1450 assert!(chunk_size > 0, "chunk size must be > 0");
1451 AxisChunks {
1452 remainder: if view.size(axis) > 0 {
1453 Some(view.view())
1454 } else {
1455 None
1456 },
1457 axis,
1458 chunk_size,
1459 }
1460 }
1461}
1462
1463impl<'a, T, L: MutLayout> Iterator for AxisChunks<'a, T, L> {
1464 type Item = TensorBase<ViewData<'a, T>, L>;
1465
1466 fn next(&mut self) -> Option<Self::Item> {
1467 let remainder = self.remainder.take()?;
1468 let chunk_len = self.chunk_size.min(remainder.size(self.axis));
1469 let (current, next_remainder) = remainder.split_at(self.axis, chunk_len);
1470 self.remainder = if next_remainder.size(self.axis) > 0 {
1471 Some(next_remainder)
1472 } else {
1473 None
1474 };
1475 Some(current)
1476 }
1477
1478 fn size_hint(&self) -> (usize, Option<usize>) {
1479 let len = self
1480 .remainder
1481 .as_ref()
1482 .map(|r| r.size(self.axis))
1483 .unwrap_or(0)
1484 .div_ceil(self.chunk_size);
1485 (len, Some(len))
1486 }
1487}
1488
1489impl<'a, T, L: MutLayout> ExactSizeIterator for AxisChunks<'a, T, L> {}
1490
1491impl<'a, T, L: MutLayout> DoubleEndedIterator for AxisChunks<'a, T, L> {
1492 fn next_back(&mut self) -> Option<Self::Item> {
1493 let remainder = self.remainder.take()?;
1494 let chunk_len = self.chunk_size.min(remainder.size(self.axis));
1495 let (prev_remainder, current) =
1496 remainder.split_at(self.axis, remainder.size(self.axis) - chunk_len);
1497 self.remainder = if prev_remainder.size(self.axis) > 0 {
1498 Some(prev_remainder)
1499 } else {
1500 None
1501 };
1502 Some(current)
1503 }
1504}
1505
1506pub struct AxisChunksMut<'a, T, L: MutLayout> {
1508 remainder: Option<TensorBase<ViewMutData<'a, T>, L>>,
1509 axis: usize,
1510 chunk_size: usize,
1511}
1512
1513impl<'a, T, L: MutLayout> AxisChunksMut<'a, T, L> {
1514 pub(crate) fn new(
1515 view: TensorBase<ViewMutData<'a, T>, L>,
1516 axis: usize,
1517 chunk_size: usize,
1518 ) -> AxisChunksMut<'a, T, L> {
1519 assert!(
1521 !view.layout().is_broadcast(),
1522 "Cannot mutably iterate over broadcasting view"
1523 );
1524 assert!(chunk_size > 0, "chunk size must be > 0");
1525 AxisChunksMut {
1526 remainder: if view.size(axis) > 0 {
1527 Some(view)
1528 } else {
1529 None
1530 },
1531 axis,
1532 chunk_size,
1533 }
1534 }
1535}
1536
1537impl<'a, T, L: MutLayout> Iterator for AxisChunksMut<'a, T, L> {
1538 type Item = TensorBase<ViewMutData<'a, T>, L>;
1539
1540 fn next(&mut self) -> Option<Self::Item> {
1541 let remainder = self.remainder.take()?;
1542 let chunk_len = self.chunk_size.min(remainder.size(self.axis));
1543 let (current, next_remainder) = remainder.split_at_mut(self.axis, chunk_len);
1544 self.remainder = if next_remainder.size(self.axis) > 0 {
1545 Some(next_remainder)
1546 } else {
1547 None
1548 };
1549 Some(current)
1550 }
1551
1552 fn size_hint(&self) -> (usize, Option<usize>) {
1553 let len = self
1554 .remainder
1555 .as_ref()
1556 .map(|r| r.size(self.axis))
1557 .unwrap_or(0)
1558 .div_ceil(self.chunk_size);
1559 (len, Some(len))
1560 }
1561}
1562
1563impl<'a, T, L: MutLayout> ExactSizeIterator for AxisChunksMut<'a, T, L> {}
1564
1565impl<'a, T, L: MutLayout> DoubleEndedIterator for AxisChunksMut<'a, T, L> {
1566 fn next_back(&mut self) -> Option<Self::Item> {
1567 let remainder = self.remainder.take()?;
1568 let remainder_size = remainder.size(self.axis);
1569 let chunk_len = self.chunk_size.min(remainder_size);
1570 let (prev_remainder, current) =
1571 remainder.split_at_mut(self.axis, remainder_size - chunk_len);
1572 self.remainder = if prev_remainder.size(self.axis) > 0 {
1573 Some(prev_remainder)
1574 } else {
1575 None
1576 };
1577 Some(current)
1578 }
1579}
1580
1581pub(crate) fn for_each_mut<T, F: Fn(&mut T)>(mut view: TensorViewMut<T>, f: F) {
1583 while view.ndim() < 4 {
1584 view.insert_axis(0);
1585 }
1586
1587 view.inner_iter_mut::<4>().for_each(|mut src| {
1593 for i0 in 0..src.size(0) {
1594 for i1 in 0..src.size(1) {
1595 for i2 in 0..src.size(2) {
1596 for i3 in 0..src.size(3) {
1597 let x = unsafe { src.get_unchecked_mut([i0, i1, i2, i3]) };
1599 f(x);
1600 }
1601 }
1602 }
1603 }
1604 });
1605}
1606
1607#[cfg(test)]
1610mod tests {
1611 use super::{AxisChunks, AxisChunksMut, Lanes, LanesMut};
1612 use crate::{AsView, Layout, NdLayout, NdTensor, Tensor};
1613
1614 fn compare_reversed<T: PartialEq + std::fmt::Debug>(fwd: &[T], rev: &[T]) {
1615 assert_eq!(fwd.len(), rev.len());
1616 for (x, y) in fwd.iter().zip(rev.iter().rev()) {
1617 assert_eq!(x, y);
1618 }
1619 }
1620
1621 fn test_iterator<I: Iterator + ExactSizeIterator + DoubleEndedIterator>(
1623 create_iter: impl Fn() -> I,
1624 expected: &[I::Item],
1625 ) where
1626 I::Item: PartialEq + std::fmt::Debug,
1627 {
1628 let iter = create_iter();
1629
1630 let (min_len, max_len) = iter.size_hint();
1631 let items: Vec<_> = iter.collect();
1632
1633 assert_eq!(&items, expected);
1634
1635 assert_eq!(min_len, items.len(), "incorrect size lower bound");
1637 assert_eq!(max_len, Some(items.len()), "incorrect size upper bound");
1638
1639 let rev_items: Vec<_> = create_iter().rev().collect();
1641 compare_reversed(&items, &rev_items);
1642
1643 let mut iter = create_iter();
1645 for _x in &mut iter { }
1646 assert_eq!(iter.next(), None);
1647
1648 let mut fold_items = Vec::new();
1650 let mut idx = 0;
1651 create_iter().fold(0, |acc, item| {
1652 assert_eq!(acc, idx);
1653 fold_items.push(item);
1654 idx += 1;
1655 idx
1656 });
1657 assert_eq!(items, fold_items);
1658 }
1659
1660 trait MutIterable {
1665 type Iter<'a>: Iterator + ExactSizeIterator + DoubleEndedIterator
1666 where
1667 Self: 'a;
1668
1669 fn iter_mut(&mut self) -> Self::Iter<'_>;
1670 }
1671
1672 fn test_mut_iterator<M, T>(mut iterable: M, expected: &[T])
1674 where
1675 M: MutIterable,
1676 T: std::fmt::Debug,
1677 for<'a> <M::Iter<'a> as Iterator>::Item: std::fmt::Debug + PartialEq + PartialEq<T>,
1678 {
1679 {
1681 let iter = iterable.iter_mut();
1682 let (min_len, max_len) = iter.size_hint();
1683 let items: Vec<_> = iter.collect();
1684
1685 assert_eq!(items, expected);
1687
1688 assert_eq!(min_len, items.len(), "incorrect size lower bound");
1690 assert_eq!(max_len, Some(items.len()), "incorrect size upper bound");
1691 }
1692
1693 {
1695 let mut iter = iterable.iter_mut();
1696 for _x in &mut iter { }
1697 assert!(iter.next().is_none());
1698 }
1699
1700 {
1706 let items: Vec<_> = iterable.iter_mut().map(|x| format!("{:?}", x)).collect();
1707 let rev_items: Vec<_> = iterable
1708 .iter_mut()
1709 .rev()
1710 .map(|x| format!("{:?}", x))
1711 .collect();
1712 compare_reversed(&items, &rev_items);
1713 }
1714
1715 {
1717 let items: Vec<_> = iterable.iter_mut().map(|x| format!("{:?}", x)).collect();
1718 let mut fold_items = Vec::new();
1719 let mut idx = 0;
1720 iterable.iter_mut().fold(0, |acc, item| {
1721 assert_eq!(acc, idx);
1722 fold_items.push(format!("{:?}", item));
1723 idx += 1;
1724 idx
1725 });
1726 assert_eq!(items, fold_items);
1727 }
1728 }
1729
1730 #[test]
1731 fn test_axis_chunks() {
1732 let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1733 test_iterator(
1734 || tensor.axis_chunks(0, 1),
1735 &[tensor.slice(0..1), tensor.slice(1..2)],
1736 );
1737 }
1738
1739 #[test]
1740 fn test_axis_chunks_empty() {
1741 let x = Tensor::<i32>::zeros(&[5, 0]);
1742 assert!(AxisChunks::new(&x.view(), 1, 1).next().is_none());
1743 }
1744
1745 #[test]
1746 #[should_panic(expected = "chunk size must be > 0")]
1747 fn test_axis_chunks_zero_size() {
1748 let x = Tensor::<i32>::zeros(&[5, 0]);
1749 assert!(AxisChunks::new(&x.view(), 1, 0).next().is_none());
1750 }
1751
1752 #[test]
1753 fn test_axis_chunks_mut_empty() {
1754 let mut x = Tensor::<i32>::zeros(&[5, 0]);
1755 assert!(AxisChunksMut::new(x.view_mut(), 1, 1).next().is_none());
1756 }
1757
1758 #[test]
1759 fn test_axis_chunks_mut_rev() {
1760 let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1761 let fwd: Vec<_> = tensor
1762 .axis_chunks_mut(0, 1)
1763 .map(|view| view.to_vec())
1764 .collect();
1765 let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1766 let rev: Vec<_> = tensor
1767 .axis_chunks_mut(0, 1)
1768 .rev()
1769 .map(|view| view.to_vec())
1770 .collect();
1771 compare_reversed(&fwd, &rev);
1772 }
1773
1774 #[test]
1775 #[should_panic(expected = "chunk size must be > 0")]
1776 fn test_axis_chunks_mut_zero_size() {
1777 let mut x = Tensor::<i32>::zeros(&[5, 0]);
1778 assert!(AxisChunksMut::new(x.view_mut(), 1, 0).next().is_none());
1779 }
1780
1781 #[test]
1782 fn test_axis_iter() {
1783 let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1784 test_iterator(|| tensor.axis_iter(0), &[tensor.slice(0), tensor.slice(1)]);
1785 }
1786
1787 #[test]
1788 fn test_axis_iter_mut_rev() {
1789 let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1790 let fwd: Vec<_> = tensor.axis_iter_mut(0).map(|view| view.to_vec()).collect();
1791 let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1792 let rev: Vec<_> = tensor
1793 .axis_iter_mut(0)
1794 .rev()
1795 .map(|view| view.to_vec())
1796 .collect();
1797 compare_reversed(&fwd, &rev);
1798 }
1799
1800 #[test]
1801 fn test_inner_iter() {
1802 let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1803 test_iterator(
1804 || tensor.inner_iter::<2>(),
1805 &[tensor.slice(0), tensor.slice(1)],
1806 );
1807 }
1808
1809 #[test]
1810 fn test_inner_iter_empty() {
1811 let tensor = NdTensor::<i32, 2>::zeros([0, 3]);
1814 assert_eq!(tensor.strides(), [3, 1]);
1815 let view = tensor.permuted([1, 0]);
1816 assert_eq!(view.strides(), [1, 3]);
1817
1818 let mut count = 0;
1819 for lane in view.inner_iter::<1>() {
1820 assert_eq!(lane.shape(), [0]);
1821 count += 1;
1822 }
1823 assert_eq!(count, 3);
1824 }
1825
1826 #[test]
1827 fn test_inner_iter_mut() {
1828 struct InnerIterMutTest(NdTensor<i32, 3>);
1829
1830 impl MutIterable for InnerIterMutTest {
1831 type Iter<'a> = super::InnerIterMut<'a, i32, NdLayout<2>>;
1832
1833 fn iter_mut(&mut self) -> Self::Iter<'_> {
1834 self.0.inner_iter_mut::<2>()
1835 }
1836 }
1837
1838 let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1839 test_mut_iterator(
1840 InnerIterMutTest(tensor.clone()),
1841 &[tensor.slice(0), tensor.slice(1)],
1842 );
1843 }
1844
1845 #[test]
1846 fn test_lanes() {
1847 let x = NdTensor::from([[1, 2], [3, 4]]);
1848 test_iterator(
1849 || x.lanes(0),
1850 &[x.slice((.., 0)).into(), x.slice((.., 1)).into()],
1851 );
1852 test_iterator(|| x.lanes(1), &[x.slice(0).into(), x.slice(1).into()]);
1853 }
1854
1855 #[test]
1856 fn test_lanes_empty() {
1857 let x = Tensor::<i32>::zeros(&[5, 0]);
1858 assert!(Lanes::new(x.view().view_ref(), 0).next().is_none());
1859 assert!(Lanes::new(x.view().view_ref(), 1).next().is_none());
1860 }
1861
1862 #[test]
1863 fn test_lanes_mut() {
1864 use super::Lane;
1865
1866 struct LanesMutTest(NdTensor<i32, 2>);
1867
1868 impl MutIterable for LanesMutTest {
1869 type Iter<'a> = super::LanesMut<'a, i32>;
1870
1871 fn iter_mut(&mut self) -> Self::Iter<'_> {
1872 self.0.lanes_mut(0)
1873 }
1874 }
1875
1876 let tensor = NdTensor::from([[1, 2], [3, 4]]);
1877 test_mut_iterator::<_, Lane<i32>>(
1878 LanesMutTest(tensor.clone()),
1879 &[
1880 Lane::from(tensor.slice((.., 0))),
1881 Lane::from(tensor.slice((.., 1))),
1882 ],
1883 );
1884 }
1885
1886 #[test]
1887 fn test_lane() {
1888 let x = NdTensor::from([[1, 2], [3, 4]]);
1889 test_iterator(|| x.lanes(0).next().unwrap(), &[&1, &3]);
1890 test_iterator(|| x.lanes(1).next().unwrap(), &[&1, &2]);
1891 }
1892
1893 #[test]
1894 fn test_lane_mut() {
1895 struct LaneMutTest(NdTensor<i32, 2>);
1896
1897 impl MutIterable for LaneMutTest {
1898 type Iter<'a> = super::LaneMut<'a, i32>;
1899
1900 fn iter_mut(&mut self) -> Self::Iter<'_> {
1901 self.0.lanes_mut(0).next().unwrap()
1902 }
1903 }
1904
1905 let tensor = NdTensor::from([[1, 2], [3, 4]]);
1906 test_mut_iterator(LaneMutTest(tensor), &[&1, &3]);
1907 }
1908
1909 #[test]
1910 fn test_lane_mut_nth() {
1911 let mut x = NdTensor::from([1, 2, 3]);
1912
1913 let mut lane = x.lanes_mut(0).next().unwrap();
1914 assert_eq!(lane.nth(1), Some(&mut 2));
1915 assert_eq!(lane.next(), Some(&mut 3));
1916
1917 let mut lane = x.lanes_mut(0).next().unwrap();
1919 assert_eq!(lane.next(), Some(&mut 1));
1920 assert_eq!(lane.nth(usize::MAX), None);
1921 assert_eq!(lane.next(), None);
1922 }
1923
1924 #[test]
1925 fn test_lane_mut_into_view() {
1926 let mut x = NdTensor::from([1, 2, 3, 4]);
1927 let lane = x.lanes_mut(0).next().unwrap();
1928 assert_eq!(lane.into_view(), NdTensor::from([1, 2, 3, 4]));
1929 }
1930
1931 #[test]
1932 #[should_panic(expected = "lane has been stepped")]
1933 fn test_lane_mut_into_view_after_next() {
1934 let mut x = NdTensor::from([1, 2, 3, 4]);
1935 let mut lane = x.lanes_mut(0).next().unwrap();
1936 lane.next();
1937 lane.into_view();
1938 }
1939
1940 #[test]
1941 #[should_panic(expected = "lane has been stepped")]
1942 fn test_lane_mut_into_view_after_next_back() {
1943 let mut x = NdTensor::from([1, 2, 3, 4]);
1944 let mut lane = x.lanes_mut(0).next().unwrap();
1945 lane.next_back();
1946 lane.into_view();
1947 }
1948
1949 #[test]
1950 fn test_lane_as_slice() {
1951 let x = NdTensor::from([0, 1, 2]);
1953 let mut lane = x.lanes(0).next().unwrap();
1954 assert_eq!(lane.as_slice(), Some([0, 1, 2].as_slice()));
1955 lane.next();
1956 assert_eq!(lane.as_slice(), Some([1, 2].as_slice()));
1957 lane.next();
1958 lane.next();
1959 assert_eq!(lane.as_slice(), Some([0i32; 0].as_slice()));
1960 lane.next();
1961 assert_eq!(lane.as_slice(), Some([0i32; 0].as_slice()));
1962
1963 let x = NdTensor::from([[1i32, 2], [3, 4]]);
1965 let lane = x.lanes(0).next().unwrap();
1966 assert_eq!(lane.as_slice(), None);
1967 }
1968
1969 #[test]
1970 fn test_lanes_mut_empty() {
1971 let mut x = Tensor::<i32>::zeros(&[5, 0]);
1972 assert!(LanesMut::new(x.mut_view_ref(), 0).next().is_none());
1973 assert!(LanesMut::new(x.mut_view_ref(), 1).next().is_none());
1974 }
1975
1976 #[test]
1977 fn test_iter_step_by() {
1978 let tensor = Tensor::<f32>::full(&[1, 3, 16, 8], 1.);
1979
1980 let tensor = tensor.slice((.., .., 1.., ..));
1983
1984 let sum = tensor.iter().sum::<f32>();
1985 for n_skip in 0..tensor.len() {
1986 let sum_skip = tensor.iter().skip(n_skip).sum::<f32>();
1987 assert_eq!(
1988 sum_skip,
1989 sum - n_skip as f32,
1990 "wrong sum for n_skip={}",
1991 n_skip
1992 );
1993 }
1994 }
1995
1996 #[test]
1997 fn test_iter_broadcast() {
1998 let tensor = Tensor::<f32>::full(&[1], 1.);
1999 let broadcast = tensor.broadcast([1, 3, 16, 8]);
2000 assert_eq!(broadcast.iter().len(), broadcast.len());
2001 let count = broadcast.iter().count();
2002 assert_eq!(count, broadcast.len());
2003 let sum = broadcast.iter().sum::<f32>();
2004 assert_eq!(sum, broadcast.len() as f32);
2005 }
2006
2007 #[test]
2008 fn test_iter() {
2009 let tensor = NdTensor::from([[[1, 2], [3, 4]]]);
2010
2011 test_iterator(|| tensor.iter().copied(), &[1, 2, 3, 4]);
2013
2014 test_iterator(|| tensor.transposed().iter().copied(), &[1, 3, 2, 4]);
2016 }
2017
2018 #[test]
2019 fn test_iter_mut() {
2020 struct IterTest(NdTensor<i32, 3>);
2021
2022 impl MutIterable for IterTest {
2023 type Iter<'a> = super::IterMut<'a, i32>;
2024
2025 fn iter_mut(&mut self) -> Self::Iter<'_> {
2026 self.0.iter_mut()
2027 }
2028 }
2029
2030 let tensor = NdTensor::from([[[1, 2], [3, 4]]]);
2031 test_mut_iterator(IterTest(tensor), &[&1, &2, &3, &4]);
2032 }
2033
2034 #[test]
2035 #[ignore]
2036 fn bench_iter() {
2037 use crate::Layout;
2038 use rten_bench::run_bench;
2039
2040 type Elem = i32;
2041
2042 let tensor = std::hint::black_box(Tensor::<Elem>::full(&[1, 6, 768, 64], 1));
2043 let n_trials = 1000;
2044 let mut result = Elem::default();
2045
2046 fn reduce<'a>(iter: impl Iterator<Item = &'a Elem>) -> Elem {
2047 iter.fold(Elem::default(), |acc, x| acc.wrapping_add(*x))
2048 }
2049
2050 run_bench(n_trials, Some("slice iter"), || {
2052 result = reduce(tensor.data().unwrap().iter());
2053 });
2054 println!("sum {}", result);
2055
2056 run_bench(n_trials, Some("contiguous iter"), || {
2059 result = reduce(tensor.iter());
2060 });
2061 println!("sum {}", result);
2062
2063 run_bench(n_trials, Some("contiguous reverse iter"), || {
2064 result = reduce(tensor.iter().rev());
2065 });
2066 println!("sum {}", result);
2067
2068 let slice = tensor.slice((.., .., 1.., ..));
2071 assert!(!slice.is_contiguous());
2072 let n_trials = 1000;
2073 run_bench(n_trials, Some("non-contiguous iter"), || {
2074 result = reduce(slice.iter());
2075 });
2076 println!("sum {}", result);
2077
2078 let n_trials = 100;
2081 run_bench(n_trials, Some("non-contiguous reverse iter"), || {
2082 result = reduce(slice.iter().rev());
2083 });
2084 println!("sum {}", result);
2085 }
2086
2087 #[test]
2088 #[ignore]
2089 fn bench_inner_iter() {
2090 use crate::rng::XorShiftRng;
2091 use rten_bench::run_bench;
2092
2093 let n_trials = 100;
2094 let mut rng = XorShiftRng::new(1234);
2095
2096 let tensor = Tensor::<f32>::rand(&[512, 512, 12, 1], &mut rng);
2100
2101 let mut sum = 0.;
2102 run_bench(n_trials, Some("inner iter"), || {
2103 for inner in tensor.inner_iter::<2>() {
2104 for i0 in 0..inner.size(0) {
2105 for i1 in 0..inner.size(1) {
2106 sum += inner[[i0, i1]];
2107 }
2108 }
2109 }
2110 });
2111 println!("sum {}", sum);
2112 }
2113}