Skip to main content

rten_tensor/
iterators.rs

1//! Iterators over tensor elements and sub-views.
2
3use 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/// Tracks the iteration position within a single dimension.
17#[derive(Copy, Clone, Debug, Default)]
18struct IterPos {
19    /// Remaining steps along this dimension before it needs to be reset.
20    remaining: usize,
21
22    /// Current index in this dimension pre-multiplied by stride.
23    offset: usize,
24
25    /// Update to `offset` for each step.
26    stride: usize,
27
28    /// Maximum value of `self.remaining`. Used when resetting position.
29    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    /// Return the size of this dimension.
57    fn size(&self) -> usize {
58        // nb. The size is always > 0 since if any dim has zero size, the
59        // iterator will have a length of zero.
60        self.max_remaining + 1
61    }
62
63    /// Return the current index along this dimension.
64    fn index(&self) -> usize {
65        self.max_remaining - self.remaining
66    }
67
68    /// Set the current index along this dimension.
69    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/// Iterator over offsets of a tensor's elements.
78#[derive(Clone, Debug)]
79struct OffsetsBase {
80    /// Remaining number of elements this iterator will yield.
81    ///
82    /// The offsets and positions in other fields are only valid if this is
83    /// non-zero.
84    len: usize,
85
86    /// Component of next element offset from innermost (fastest-changing) dims.
87    inner_offset: usize,
88
89    /// Current position in innermost dims.
90    inner_pos: [IterPos; INNER_NDIM],
91
92    /// Component of next element offset from outermost (slowest-changing) dims.
93    outer_offset: usize,
94
95    /// Current position in outermost dims.
96    ///
97    /// Optimization note: The number of outermost dims will usually be small,
98    /// so you might be tempted to use `SmallVec`. However this resulted in
99    /// worse performance for `IndexingIterBase::step`, as the compiler was
100    /// less likely/able to unroll iteration loops.
101    outer_pos: Vec<IterPos>,
102}
103
104impl OffsetsBase {
105    /// Create an iterator over element offsets in `tensor`.
106    fn new<L: Layout>(layout: &L) -> OffsetsBase {
107        // Merge axes to maximize the number of iterations that use the fast
108        // path for stepping over the inner dimensions.
109        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    /// Step in the outer dimensions.
140    ///
141    /// Returns `true` if the position was advanced or `false` if the end was
142    /// reached.
143    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    /// Advance iterator by up to `n` indices.
175    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        // Update offset of next element.
192        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    /// Compute the storage offset of an element given a linear index into a
201    /// tensor's element sequence.
202    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    /// Truncate this iterator so that it yields at most `len` elements.
215    fn truncate(&mut self, len: usize) {
216        // We adjust `self.len` here but not any of the iteration positions.
217        // This means that methods like `next` and `fold` must always check
218        // `self.len` before each step.
219        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        // Optimistically update offset, assuming we haven't reached the
236        // end of the last dimension.
237        self.inner_offset += self.inner_pos[1].stride;
238
239        // Use a fast path to step inner dimensions and fall back to the slower
240        // path to step the outer dimensions only when we reach the end.
241        if !self.inner_pos[1].step() {
242            if !self.inner_pos[0].step() {
243                self.step_outer_pos();
244            }
245
246            // `inner_offset` is the sum of `inner_pos[i].offset`. It only
247            // contains two entries, and we know `inner_pos[1].offset` is zero
248            // since `inner_pos[1].step()` returned false. Hence we can use
249            // an assignment.
250            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        // Iter positions are only valid if `self.len > 0`.
266        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        // This is inefficient compared to forward iteration, but that's OK
305        // because reverse iteration is not performance critical.
306        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    /// Split this iterator into two. The left result visits indices before
316    /// `index`, the right result visits indices from `index` onwards.
317    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
329/// Iterator over elements of a tensor, in their logical order.
330pub 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        // Safety: Offset is valid for data length.
361        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        // Safety: Offset is valid for data length.
372        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            // Safety: Offset is valid for data length.
382            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        // Safety: Offset is valid for data length.
393        Some(unsafe { self.data.get_unchecked(offset) })
394    }
395}
396
397impl<T> ExactSizeIterator for Iter<'_, T> {}
398
399impl<T> FusedIterator for Iter<'_, T> {}
400
401/// Wrapper around [`transmute`] which allows transmuting only the lifetime,
402/// not the type, of a reference.
403unsafe 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
407/// Mutable iterator over elements of a tensor.
408pub 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        // Safety: Offset is valid for data length, `offsets.next` yields each
432        // offset only once.
433        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        // Safety: Offset is valid for data length, `offsets.next` yields each
445        // offset only once.
446        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            // Safety: Offset is valid for data length, `offsets.fold` yields
456            // each offset only once.
457            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        // Safety: Offset is valid for data length, `offsets.next` yields each
468        // offset only once.
469        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/// Iterator over element offsets of a tensor.
484///
485/// `Offsets` does not hold a reference to the tensor, allowing the tensor to
486/// be modified during iteration. It is the caller's responsibilty not to modify
487/// the tensor in ways that invalidate the offset sequence returned by this
488/// iterator.
489#[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
559/// Iterator over the ranges of a tensor's data that correspond to 1D lanes
560/// along a particular dimension.
561struct LaneRanges {
562    /// Start offsets of each lane.
563    offsets: Offsets,
564
565    // Number of elements in each lane and gap between them.
566    dim_size: usize,
567    dim_stride: usize,
568}
569
570impl LaneRanges {
571    fn new<L: Layout + RemoveDim>(layout: &L, dim: usize) -> LaneRanges {
572        // If the layout is empty (has any zero-sized dims), we need to make
573        // sure that `offsets` is as well.
574        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    /// Return the range of storage offsets for a 1D lane where the first
589    /// element is at `start_offset`.
590    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
642/// Iterator over 1D slices of a tensor along a target dimension of size N.
643///
644/// Conceptually this iterator steps through every distinct slice of a tensor
645/// where a target dim is varied from 0..N and other indices are held fixed.
646pub struct Lanes<'a, T> {
647    data: ViewData<'a, T>,
648    ranges: LaneRanges,
649    lane_layout: NdLayout<1>,
650}
651
652/// Iterator over items in a 1D slice of a tensor.
653#[derive(Clone, Debug)]
654pub struct Lane<'a, T> {
655    view: NdTensorView<'a, T, 1>,
656
657    /// Index of the next item yielded from the front of the lane.
658    index: usize,
659
660    /// Index one past the next item yielded from the back of the lane.
661    end: usize,
662}
663
664impl<'a, T> Lane<'a, T> {
665    /// Return the remaining part of the lane as a slice, if it is contiguous.
666    pub fn as_slice(&self) -> Option<&'a [T]> {
667        self.view.data().map(|data| &data[self.index..self.end])
668    }
669
670    /// Return the item at a given index in this lane.
671    pub fn get(&self, idx: usize) -> Option<&'a T> {
672        self.view.get([idx])
673    }
674
675    /// Return the entire lane as a 1D tensor view.
676    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            // Safety: Index is in bounds for axis 0.
701            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            // Safety: Index is in bounds for axis 0.
720            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    /// Create an iterator which yields all possible slices over the `dim`
745    /// dimension of `tensor`.
746    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    /// Yield the next slice over the target dimension.
780    #[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
815/// Mutable version of [`Lanes`].
816///
817/// Unlike [`Lanes`], this does not implement [`Iterator`] due to complications
818/// in implementing this for an iterator that returns mutable references, but
819/// it has a similar interface.
820pub 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    /// Create an iterator which yields all possible slices over the `dim`
828    /// dimension of `view`.
829    pub(crate) fn new<L: Layout + RemoveDim + Clone>(
830        view: TensorBase<ViewMutData<'a, T>, L>,
831        dim: usize,
832    ) -> LanesMut<'a, T> {
833        // See notes in `Layout` about internal overlap.
834        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        // We allow overlap here to handle the case where the stride is zero,
843        // but the tensor is empty. If the tensor was not empty, the assert above
844        // would have caught this.
845        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            // Safety: Offsets range length is sufficient for layout, elements
864            // in each lane do not overlap.
865            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            // Safety: Offsets range length is sufficient for layout, elements
882            // in each lane do not overlap.
883            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            // Safety: Offsets range length is sufficient for layout, elements
897            // in each lane do not overlap.
898            unsafe {
899                LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
900            }
901        })
902    }
903}
904
905/// Iterator over items in a 1D slice of a tensor.
906#[derive(Debug)]
907pub struct LaneMut<'a, T> {
908    view: NdTensorViewMut<'a, T, 1>,
909
910    /// Index of the next item yielded from the front of the lane.
911    index: usize,
912
913    /// Index one past the next item yielded from the back of the lane.
914    end: usize,
915}
916
917impl<'a, T> LaneMut<'a, T> {
918    /// Create a new lane given the storage and layout.
919    ///
920    /// # Safety
921    ///
922    /// - Caller must ensure that no two lanes are created which overlap.
923    /// - Storage length must exceed `layout.min_data_len()`.
924    unsafe fn from_storage_layout(data: ViewMutData<'a, T>, layout: NdLayout<1>) -> Self {
925        let view = unsafe {
926            // Safety: Caller promises that each call uses the offset ranges for
927            // a different lane and that the range length is sufficient for the
928            // lane's size and stride.
929            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    /// Return the remaining part of the lane as a slice, if it is contiguous.
939    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    /// Return the entire lane as a mutable 1D tensor view.
945    ///
946    /// # Panics
947    ///
948    /// Panics if the lane has been stepped, as the view would then alias
949    /// references which the iterator has already yielded.
950    #[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            // Transmute to preserve lifetime of data. This is safe as we
971            // yield each element only once.
972            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            // Transmute to preserve lifetime of data. This is safe as we
998            // yield each element only once.
999            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
1014/// Base for iterators over views of the inner dimensions of a tensor, where
1015/// the inner dimensions have layout `L`.
1016struct InnerIterBase<L: Layout> {
1017    // Iterator over storage start offsets for each inner view. The storage
1018    // range for each view is `offset..offset + inner_data_len`.
1019    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        // If the inner views are empty, the tensor must have zero-length
1045        // storage. Zero the outer strides so that `outer_offsets` always yields
1046        // zero - the only valid storage offset.
1047        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                // We allow overlap here, but the view that owns `parent_layout`
1079                // will enforce there is no overlap if it is a mutable view.
1080                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                // We allow overlap here, but the view that owns `parent_layout`
1094                // will enforce there is no overlap if it is a mutable view.
1095                OverlapPolicy::AllowOverlap,
1096            )
1097            .expect("failed to create layout")
1098        })
1099    }
1100}
1101
1102impl<L: Layout> Iterator for InnerIterBase<L> {
1103    /// Storage offset range for next view
1104    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
1137/// Iterator over views of the innermost dimensions of a tensor, where the
1138/// tensor has element type T and the inner dimensions have layout L.
1139pub 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
1212/// Iterator over mutable views of the innermost dimensions of a tensor, where
1213/// the tensor has element type T and the inner dimensions have layout L.
1214pub 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                // Safety: The iterator was constructed from a tensor with a
1250                // non-overlapping layout, and no two views yielded by this
1251                // iterator overlap. Hence we can transmute the lifetime without
1252                // creating multiple mutable references to the same elements.
1253                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                // Safety: The iterator was constructed from a tensor with a
1273                // non-overlapping layout, and no two views yielded by this
1274                // iterator overlap. Hence we can transmute the lifetime without
1275                // creating multiple mutable references to the same elements.
1276                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                // Safety: Outer view is non-broadcasting, and we increment the
1292                // outer index each time, so returned views will not overlap.
1293                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
1300/// Iterator over slices of a tensor along an axis. See
1301/// [`TensorView::axis_iter`](crate::TensorView::axis_iter).
1302pub 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
1354/// Iterator over mutable slices of a tensor along an axis. See [`TensorViewMut::axis_iter_mut`].
1355pub 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        // See notes in `Layout` about internal overlap.
1368        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
1382/// Mutable tensor view with one less dimension than `L`.
1383type 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            // Promote lifetime from self -> 'a.
1398            //
1399            // Safety: This is non-broadcasting view, and we increment the index
1400            // each time, so returned views will not overlap.
1401            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            // Promote lifetime from self -> 'a.
1426            //
1427            // Safety: This is non-broadcasting view, and we increment the index
1428            // each time, so returned views will not overlap.
1429            let view = unsafe { transmute::<SmallerMutView<'_, T, L>, Self::Item>(slice) };
1430
1431            Some(view)
1432        }
1433    }
1434}
1435
1436/// Iterator over slices of a tensor along an axis. See
1437/// [`TensorView::axis_chunks`](crate::TensorView::axis_chunks).
1438pub 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
1506/// Iterator over mutable slices of a tensor along an axis. See [`TensorViewMut::axis_chunks_mut`].
1507pub 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        // See notes in `Layout` about internal overlap.
1520        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
1581/// Call `f` on each element of `view`.
1582pub(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    // This could be improved by sorting dimensions of `view` in order of
1588    // decreasing stride. If the resulting view is contiguous, `f` can be
1589    // applied to the underlying data directly. Even if it isn't, this will
1590    // still make memory access as contiguous as possible.
1591
1592    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                        // Safety: i0..i3 are in `[0, src.size(i))`.
1598                        let x = unsafe { src.get_unchecked_mut([i0, i1, i2, i3]) };
1599                        f(x);
1600                    }
1601                }
1602            }
1603        }
1604    });
1605}
1606
1607// Tests for iterator internals. Most tests of iterators are currently done via
1608// tests on tensor methods.
1609#[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    /// Apply a standard set of tests to an iterator.
1622    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        // Test ExactSizeIterator via `size_hint`.
1636        assert_eq!(min_len, items.len(), "incorrect size lower bound");
1637        assert_eq!(max_len, Some(items.len()), "incorrect size upper bound");
1638
1639        // Test DoubleEndedIterator via `rev`.
1640        let rev_items: Vec<_> = create_iter().rev().collect();
1641        compare_reversed(&items, &rev_items);
1642
1643        // Test FusedIterator.
1644        let mut iter = create_iter();
1645        for _x in &mut iter { /* noop */ }
1646        assert_eq!(iter.next(), None);
1647
1648        // Test fold.
1649        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    /// A collection that can be mutably iterated over multiple times.
1661    ///
1662    /// We use a different pattern for testing mutable iterators to avoid
1663    /// restrictions on values returned from `FnMut` closures.
1664    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    /// Apply a standard set of tests to a mutable iterator.
1673    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        // Test Iterator and ExactSizeIterator.
1680        {
1681            let iter = iterable.iter_mut();
1682            let (min_len, max_len) = iter.size_hint();
1683            let items: Vec<_> = iter.collect();
1684
1685            // Test `next`
1686            assert_eq!(items, expected);
1687
1688            // Test `size_hint`
1689            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        // Test FusedIterator.
1694        {
1695            let mut iter = iterable.iter_mut();
1696            for _x in &mut iter { /* noop */ }
1697            assert!(iter.next().is_none());
1698        }
1699
1700        // Test DoubleEndedIterator via `rev`.
1701        //
1702        // We use `format!` here to convert mutable references into comparable
1703        // items that have no connection to the mutable references yielded by
1704        // the iterator. Ideally this should be replaced by a clone or something.
1705        {
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        // Test fold.
1716        {
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        // Create a tensor view where the inner dimension has zero size and the
1812        // outer dimension has non-zero size and non-zero strides.
1813        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        // Skipping past the end must not wrap the cursor around.
1918        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        // Contiguous lane
1952        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        // Non-contiguous lane
1964        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        // Take a non-contiguous slice so we don't use the fast path for
1981        // contiguous tensors.
1982        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 over contiguous tensor.
2012        test_iterator(|| tensor.iter().copied(), &[1, 2, 3, 4]);
2013
2014        // Test iterator over non-contiguous tensor.
2015        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        // Iterate directly over data slice.
2051        run_bench(n_trials, Some("slice iter"), || {
2052            result = reduce(tensor.data().unwrap().iter());
2053        });
2054        println!("sum {}", result);
2055
2056        // Use tensor iterator with contiguous tensor. This will use the fast
2057        // path which wraps a slice iterator.
2058        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        // Use tensor iterator with non-contiguous slice. This will fall back
2069        // to indexed iteration.
2070        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        // Reverse iteration with non-contiguous slice. This is much slower
2079        // because it translates linear indexes into offsets using division.
2080        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        // Tensor with many steps along the outer two dimensions relative to the
2097        // steps along the inner two dimensions. This emphasizes the overhead of
2098        // stepping `inner_iter`.
2099        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}