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    index: usize,
657}
658
659impl<'a, T> Lane<'a, T> {
660    /// Return the remaining part of the lane as a slice, if it is contiguous.
661    pub fn as_slice(&self) -> Option<&'a [T]> {
662        self.view.data().map(|data| &data[self.index..])
663    }
664
665    /// Return the item at a given index in this lane.
666    pub fn get(&self, idx: usize) -> Option<&'a T> {
667        self.view.get([idx])
668    }
669
670    /// Return the entire lane as a 1D tensor view.
671    pub fn as_view(&self) -> NdTensorView<'a, T, 1> {
672        self.view
673    }
674}
675
676impl<'a, T> From<NdTensorView<'a, T, 1>> for Lane<'a, T> {
677    fn from(val: NdTensorView<'a, T, 1>) -> Self {
678        Lane {
679            view: val,
680            index: 0,
681        }
682    }
683}
684
685impl<'a, T> Iterator for Lane<'a, T> {
686    type Item = &'a T;
687
688    #[inline]
689    fn next(&mut self) -> Option<Self::Item> {
690        if self.index < self.view.len() {
691            let index = self.index;
692            self.index += 1;
693
694            // Safety: Index is in bounds for axis 0.
695            Some(unsafe { self.view.get_unchecked([index]) })
696        } else {
697            None
698        }
699    }
700
701    fn size_hint(&self) -> (usize, Option<usize>) {
702        let size = self.view.size(0);
703        (size, Some(size))
704    }
705}
706
707impl<T> ExactSizeIterator for Lane<'_, T> {}
708
709impl<T> FusedIterator for Lane<'_, T> {}
710
711impl<T: PartialEq> PartialEq<Lane<'_, T>> for Lane<'_, T> {
712    fn eq(&self, other: &Lane<'_, T>) -> bool {
713        self.view.slice(self.index..) == other.view.slice(other.index..)
714    }
715}
716
717impl<T: PartialEq> PartialEq<Lane<'_, T>> for LaneMut<'_, T> {
718    fn eq(&self, other: &Lane<'_, T>) -> bool {
719        self.view.slice(self.index..) == other.view.slice(other.index..)
720    }
721}
722
723impl<'a, T> Lanes<'a, T> {
724    /// Create an iterator which yields all possible slices over the `dim`
725    /// dimension of `tensor`.
726    pub(crate) fn new<L: Layout + RemoveDim + Clone>(
727        view: TensorBase<ViewData<'a, T>, L>,
728        dim: usize,
729    ) -> Lanes<'a, T> {
730        let size = view.size(dim);
731        let stride = view.stride(dim);
732        let lane_layout =
733            NdLayout::from_shape_and_strides([size], [stride], OverlapPolicy::AllowOverlap)
734                .unwrap();
735        Lanes {
736            data: view.storage(),
737            ranges: LaneRanges::new(view.layout(), dim),
738            lane_layout,
739        }
740    }
741}
742
743fn lane_for_offset_range<T>(
744    data: ViewData<T>,
745    layout: NdLayout<1>,
746    offsets: Range<usize>,
747) -> Lane<T> {
748    let view = NdTensorView::from_storage_and_layout(data.slice(offsets), layout);
749    Lane { view, index: 0 }
750}
751
752impl<'a, T> Iterator for Lanes<'a, T> {
753    type Item = Lane<'a, T>;
754
755    /// Yield the next slice over the target dimension.
756    #[inline]
757    fn next(&mut self) -> Option<Self::Item> {
758        self.ranges
759            .next()
760            .map(|range| lane_for_offset_range(self.data, self.lane_layout, range))
761    }
762
763    fn size_hint(&self) -> (usize, Option<usize>) {
764        self.ranges.size_hint()
765    }
766
767    fn fold<B, F>(self, init: B, mut f: F) -> B
768    where
769        Self: Sized,
770        F: FnMut(B, Self::Item) -> B,
771    {
772        self.ranges.fold(init, |acc, offsets| {
773            let lane = lane_for_offset_range(self.data, self.lane_layout, offsets);
774            f(acc, lane)
775        })
776    }
777}
778
779impl<T> DoubleEndedIterator for Lanes<'_, T> {
780    fn next_back(&mut self) -> Option<Self::Item> {
781        self.ranges
782            .next_back()
783            .map(|range| lane_for_offset_range(self.data, self.lane_layout, range))
784    }
785}
786
787impl<T> ExactSizeIterator for Lanes<'_, T> {}
788
789impl<T> FusedIterator for Lanes<'_, T> {}
790
791/// Mutable version of [`Lanes`].
792///
793/// Unlike [`Lanes`], this does not implement [`Iterator`] due to complications
794/// in implementing this for an iterator that returns mutable references, but
795/// it has a similar interface.
796pub struct LanesMut<'a, T> {
797    data: ViewMutData<'a, T>,
798    ranges: LaneRanges,
799    lane_layout: NdLayout<1>,
800}
801
802impl<'a, T> LanesMut<'a, T> {
803    /// Create an iterator which yields all possible slices over the `dim`
804    /// dimension of `view`.
805    pub(crate) fn new<L: Layout + RemoveDim + Clone>(
806        view: TensorBase<ViewMutData<'a, T>, L>,
807        dim: usize,
808    ) -> LanesMut<'a, T> {
809        // See notes in `Layout` about internal overlap.
810        assert!(
811            !view.is_broadcast(),
812            "Cannot mutably iterate over broadcasting view"
813        );
814
815        let size = view.size(dim);
816        let stride = view.stride(dim);
817
818        // We allow overlap here to handle the case where the stride is zero,
819        // but the tensor is empty. If the tensor was not empty, the assert above
820        // would have caught this.
821        let lane_layout =
822            NdLayout::from_shape_and_strides([size], [stride], OverlapPolicy::AllowOverlap)
823                .unwrap();
824
825        LanesMut {
826            ranges: LaneRanges::new(view.layout(), dim),
827            data: view.into_storage(),
828            lane_layout,
829        }
830    }
831}
832
833impl<'a, T> Iterator for LanesMut<'a, T> {
834    type Item = LaneMut<'a, T>;
835
836    #[inline]
837    fn next(&mut self) -> Option<LaneMut<'a, T>> {
838        self.ranges.next().map(|offsets| {
839            // Safety: Offsets range length is sufficient for layout, elements
840            // in each lane do not overlap.
841            unsafe {
842                LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
843            }
844        })
845    }
846
847    fn size_hint(&self) -> (usize, Option<usize>) {
848        self.ranges.size_hint()
849    }
850
851    fn fold<B, F>(mut self, init: B, mut f: F) -> B
852    where
853        Self: Sized,
854        F: FnMut(B, Self::Item) -> B,
855    {
856        self.ranges.fold(init, |acc, offsets| {
857            // Safety: Offsets range length is sufficient for layout, elements
858            // in each lane do not overlap.
859            let lane = unsafe {
860                LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
861            };
862            f(acc, lane)
863        })
864    }
865}
866
867impl<'a, T> ExactSizeIterator for LanesMut<'a, T> {}
868
869impl<'a, T> DoubleEndedIterator for LanesMut<'a, T> {
870    fn next_back(&mut self) -> Option<LaneMut<'a, T>> {
871        self.ranges.next_back().map(|offsets| {
872            // Safety: Offsets range length is sufficient for layout, elements
873            // in each lane do not overlap.
874            unsafe {
875                LaneMut::from_storage_layout(self.data.to_view_slice_mut(offsets), self.lane_layout)
876            }
877        })
878    }
879}
880
881/// Iterator over items in a 1D slice of a tensor.
882#[derive(Debug)]
883pub struct LaneMut<'a, T> {
884    view: NdTensorViewMut<'a, T, 1>,
885    index: usize,
886}
887
888impl<'a, T> LaneMut<'a, T> {
889    /// Create a new lane given the storage and layout.
890    ///
891    /// # Safety
892    ///
893    /// - Caller must ensure that no two lanes are created which overlap.
894    /// - Storage length must exceed `layout.min_data_len()`.
895    unsafe fn from_storage_layout(data: ViewMutData<'a, T>, layout: NdLayout<1>) -> Self {
896        let view = unsafe {
897            // Safety: Caller promises that each call uses the offset ranges for
898            // a different lane and that the range length is sufficient for the
899            // lane's size and stride.
900            NdTensorViewMut::from_storage_and_layout_unchecked(data, layout)
901        };
902        LaneMut { view, index: 0 }
903    }
904
905    /// Return the remaining part of the lane as a slice, if it is contiguous.
906    pub fn as_slice_mut(&mut self) -> Option<&mut [T]> {
907        self.view.data_mut().map(|data| &mut data[self.index..])
908    }
909
910    /// Return the entire lane as a mutable 1D tensor view.
911    pub fn into_view(self) -> NdTensorViewMut<'a, T, 1> {
912        self.view
913    }
914}
915
916impl<'a, T> Iterator for LaneMut<'a, T> {
917    type Item = &'a mut T;
918
919    #[inline]
920    fn next(&mut self) -> Option<Self::Item> {
921        if self.index < self.view.size(0) {
922            let index = self.index;
923            self.index += 1;
924            let item = unsafe { self.view.get_unchecked_mut([index]) };
925
926            // Transmute to preserve lifetime of data. This is safe as we
927            // yield each element only once.
928            Some(unsafe { transmute::<&mut T, Self::Item>(item) })
929        } else {
930            None
931        }
932    }
933
934    #[inline]
935    fn nth(&mut self, nth: usize) -> Option<Self::Item> {
936        self.index = (self.index + nth).min(self.view.size(0));
937        self.next()
938    }
939
940    fn size_hint(&self) -> (usize, Option<usize>) {
941        let size = self.view.size(0);
942        (size, Some(size))
943    }
944}
945
946impl<T> ExactSizeIterator for LaneMut<'_, T> {}
947
948impl<T: PartialEq> PartialEq<LaneMut<'_, T>> for LaneMut<'_, T> {
949    fn eq(&self, other: &LaneMut<'_, T>) -> bool {
950        self.view.slice(self.index..) == other.view.slice(other.index..)
951    }
952}
953
954/// Base for iterators over views of the inner dimensions of a tensor, where
955/// the inner dimensions have layout `L`.
956struct InnerIterBase<L: Layout> {
957    // Iterator over storage start offsets for each inner view. The storage
958    // range for each view is `offset..offset + inner_data_len`.
959    outer_offsets: Offsets,
960    inner_layout: L,
961    inner_data_len: usize,
962}
963
964impl<L: Layout + Clone> InnerIterBase<L> {
965    fn new_impl<PL: Layout, F: Fn(&[usize], &[usize]) -> L>(
966        parent_layout: &PL,
967        inner_dims: usize,
968        make_inner_layout: F,
969    ) -> InnerIterBase<L> {
970        assert!(parent_layout.ndim() >= inner_dims);
971        let outer_dims = parent_layout.ndim() - inner_dims;
972        let parent_shape = parent_layout.shape();
973        let parent_strides = parent_layout.strides();
974
975        let parent_dims: SmallVec<[usize; 5]> = parent_shape.iter().collect();
976        let (outer_shape, inner_shape) = parent_dims.as_ref().split_at(outer_dims);
977
978        let parent_strides: SmallVec<[usize; 5]> = parent_strides.iter().collect();
979        let (outer_strides, inner_strides) = parent_strides.as_ref().split_at(outer_dims);
980
981        let inner_layout = make_inner_layout(inner_shape, inner_strides);
982        let inner_data_len = inner_layout.min_data_len();
983
984        // If the inner views are empty, the tensor must have zero-length
985        // storage. Zero the outer strides so that `outer_offsets` always yields
986        // zero - the only valid storage offset.
987        let zero_strides: SmallVec<[usize; 5]>;
988        let outer_strides = if inner_data_len == 0 {
989            zero_strides = SmallVec::from_elem(0, outer_dims);
990            zero_strides.as_ref()
991        } else {
992            outer_strides
993        };
994
995        let outer_layout = DynLayout::from_shape_and_strides(
996            outer_shape,
997            outer_strides,
998            OverlapPolicy::AllowOverlap,
999        )
1000        .unwrap();
1001
1002        InnerIterBase {
1003            outer_offsets: Offsets::new(&outer_layout),
1004            inner_data_len,
1005            inner_layout,
1006        }
1007    }
1008}
1009
1010impl<const N: usize> InnerIterBase<NdLayout<N>> {
1011    pub(crate) fn new<L: Layout>(parent_layout: &L) -> Self {
1012        Self::new_impl(parent_layout, N, |inner_shape, inner_strides| {
1013            let inner_shape: [usize; N] = inner_shape.try_into().unwrap();
1014            let inner_strides: [usize; N] = inner_strides.try_into().unwrap();
1015            NdLayout::from_shape_and_strides(
1016                inner_shape,
1017                inner_strides,
1018                // We allow overlap here, but the view that owns `parent_layout`
1019                // will enforce there is no overlap if it is a mutable view.
1020                OverlapPolicy::AllowOverlap,
1021            )
1022            .expect("failed to create layout")
1023        })
1024    }
1025}
1026
1027impl InnerIterBase<DynLayout> {
1028    pub(crate) fn new_dyn<L: Layout>(parent_layout: &L, inner_dims: usize) -> Self {
1029        Self::new_impl(parent_layout, inner_dims, |inner_shape, inner_strides| {
1030            DynLayout::from_shape_and_strides(
1031                inner_shape,
1032                inner_strides,
1033                // We allow overlap here, but the view that owns `parent_layout`
1034                // will enforce there is no overlap if it is a mutable view.
1035                OverlapPolicy::AllowOverlap,
1036            )
1037            .expect("failed to create layout")
1038        })
1039    }
1040}
1041
1042impl<L: Layout> Iterator for InnerIterBase<L> {
1043    /// Storage offset range for next view
1044    type Item = Range<usize>;
1045
1046    fn next(&mut self) -> Option<Range<usize>> {
1047        self.outer_offsets
1048            .next()
1049            .map(|offset| offset..offset + self.inner_data_len)
1050    }
1051
1052    fn size_hint(&self) -> (usize, Option<usize>) {
1053        self.outer_offsets.size_hint()
1054    }
1055
1056    fn fold<B, F>(self, init: B, mut f: F) -> B
1057    where
1058        Self: Sized,
1059        F: FnMut(B, Self::Item) -> B,
1060    {
1061        self.outer_offsets.fold(init, |acc, offset| {
1062            f(acc, offset..offset + self.inner_data_len)
1063        })
1064    }
1065}
1066
1067impl<L: Layout> ExactSizeIterator for InnerIterBase<L> {}
1068
1069impl<L: Layout> DoubleEndedIterator for InnerIterBase<L> {
1070    fn next_back(&mut self) -> Option<Self::Item> {
1071        self.outer_offsets
1072            .next_back()
1073            .map(|offset| offset..offset + self.inner_data_len)
1074    }
1075}
1076
1077/// Iterator over views of the innermost dimensions of a tensor, where the
1078/// tensor has element type T and the inner dimensions have layout L.
1079pub struct InnerIter<'a, T, L: Layout> {
1080    base: InnerIterBase<L>,
1081    data: ViewData<'a, T>,
1082}
1083
1084impl<'a, T, const N: usize> InnerIter<'a, T, NdLayout<N>> {
1085    pub(crate) fn new<L: Layout + Clone>(view: TensorBase<ViewData<'a, T>, L>) -> Self {
1086        let base = InnerIterBase::new(&view);
1087        InnerIter {
1088            base,
1089            data: view.storage(),
1090        }
1091    }
1092}
1093
1094impl<'a, T> InnerIter<'a, T, DynLayout> {
1095    pub(crate) fn new_dyn<L: Layout + Clone>(
1096        view: TensorBase<ViewData<'a, T>, L>,
1097        inner_dims: usize,
1098    ) -> Self {
1099        let base = InnerIterBase::new_dyn(&view, inner_dims);
1100        InnerIter {
1101            base,
1102            data: view.storage(),
1103        }
1104    }
1105}
1106
1107impl<'a, T, L: Layout + Clone> Iterator for InnerIter<'a, T, L> {
1108    type Item = TensorBase<ViewData<'a, T>, L>;
1109
1110    fn next(&mut self) -> Option<Self::Item> {
1111        self.base.next().map(|offset_range| {
1112            TensorBase::from_storage_and_layout(
1113                self.data.slice(offset_range),
1114                self.base.inner_layout.clone(),
1115            )
1116        })
1117    }
1118
1119    fn size_hint(&self) -> (usize, Option<usize>) {
1120        self.base.size_hint()
1121    }
1122
1123    fn fold<B, F>(self, init: B, mut f: F) -> B
1124    where
1125        Self: Sized,
1126        F: FnMut(B, Self::Item) -> B,
1127    {
1128        let inner_layout = self.base.inner_layout.clone();
1129        self.base.fold(init, |acc, offset_range| {
1130            let item = TensorBase::from_storage_and_layout(
1131                self.data.slice(offset_range),
1132                inner_layout.clone(),
1133            );
1134            f(acc, item)
1135        })
1136    }
1137}
1138
1139impl<T, L: Layout + Clone> ExactSizeIterator for InnerIter<'_, T, L> {}
1140
1141impl<T, L: Layout + Clone> DoubleEndedIterator for InnerIter<'_, T, L> {
1142    fn next_back(&mut self) -> Option<Self::Item> {
1143        self.base.next_back().map(|offset_range| {
1144            TensorBase::from_storage_and_layout(
1145                self.data.slice(offset_range),
1146                self.base.inner_layout.clone(),
1147            )
1148        })
1149    }
1150}
1151
1152/// Iterator over mutable views of the innermost dimensions of a tensor, where
1153/// the tensor has element type T and the inner dimensions have layout L.
1154pub struct InnerIterMut<'a, T, L: Layout> {
1155    base: InnerIterBase<L>,
1156    data: ViewMutData<'a, T>,
1157}
1158
1159impl<'a, T, const N: usize> InnerIterMut<'a, T, NdLayout<N>> {
1160    pub(crate) fn new<L: Layout>(view: TensorBase<ViewMutData<'a, T>, L>) -> Self {
1161        let base = InnerIterBase::new(&view);
1162        InnerIterMut {
1163            base,
1164            data: view.into_storage(),
1165        }
1166    }
1167}
1168
1169impl<'a, T> InnerIterMut<'a, T, DynLayout> {
1170    pub(crate) fn new_dyn<L: Layout>(
1171        view: TensorBase<ViewMutData<'a, T>, L>,
1172        inner_dims: usize,
1173    ) -> Self {
1174        let base = InnerIterBase::new_dyn(&view, inner_dims);
1175        InnerIterMut {
1176            base,
1177            data: view.into_storage(),
1178        }
1179    }
1180}
1181
1182impl<'a, T, L: Layout + Clone> Iterator for InnerIterMut<'a, T, L> {
1183    type Item = TensorBase<ViewMutData<'a, T>, L>;
1184
1185    fn next(&mut self) -> Option<Self::Item> {
1186        self.base.next().map(|offset_range| {
1187            let storage = self.data.slice_mut(offset_range);
1188            let storage = unsafe {
1189                // Safety: The iterator was constructed from a tensor with a
1190                // non-overlapping layout, and no two views yielded by this
1191                // iterator overlap. Hence we can transmute the lifetime without
1192                // creating multiple mutable references to the same elements.
1193                std::mem::transmute::<ViewMutData<'_, T>, ViewMutData<'a, T>>(storage)
1194            };
1195            TensorBase::from_storage_and_layout(storage, self.base.inner_layout.clone())
1196        })
1197    }
1198
1199    fn size_hint(&self) -> (usize, Option<usize>) {
1200        self.base.size_hint()
1201    }
1202
1203    fn fold<B, F>(mut self, init: B, mut f: F) -> B
1204    where
1205        Self: Sized,
1206        F: FnMut(B, Self::Item) -> B,
1207    {
1208        let inner_layout = self.base.inner_layout.clone();
1209        self.base.fold(init, |acc, offset_range| {
1210            let storage = self.data.slice_mut(offset_range);
1211            let storage = unsafe {
1212                // Safety: The iterator was constructed from a tensor with a
1213                // non-overlapping layout, and no two views yielded by this
1214                // iterator overlap. Hence we can transmute the lifetime without
1215                // creating multiple mutable references to the same elements.
1216                std::mem::transmute::<ViewMutData<'_, T>, ViewMutData<'a, T>>(storage)
1217            };
1218            let item = TensorBase::from_storage_and_layout(storage, inner_layout.clone());
1219            f(acc, item)
1220        })
1221    }
1222}
1223
1224impl<T, L: Layout + Clone> ExactSizeIterator for InnerIterMut<'_, T, L> {}
1225
1226impl<'a, T, L: Layout + Clone> DoubleEndedIterator for InnerIterMut<'a, T, L> {
1227    fn next_back(&mut self) -> Option<Self::Item> {
1228        self.base.next_back().map(|offset_range| {
1229            let storage = self.data.slice_mut(offset_range);
1230            let storage = unsafe {
1231                // Safety: Outer view is non-broadcasting, and we increment the
1232                // outer index each time, so returned views will not overlap.
1233                std::mem::transmute::<ViewMutData<'_, T>, ViewMutData<'a, T>>(storage)
1234            };
1235            TensorBase::from_storage_and_layout(storage, self.base.inner_layout.clone())
1236        })
1237    }
1238}
1239
1240/// Iterator over slices of a tensor along an axis. See
1241/// [`TensorView::axis_iter`](crate::TensorView::axis_iter).
1242pub struct AxisIter<'a, T, L: Layout + RemoveDim> {
1243    view: TensorBase<ViewData<'a, T>, L>,
1244    axis: usize,
1245    index: usize,
1246    end: usize,
1247}
1248
1249impl<'a, T, L: MutLayout + RemoveDim> AxisIter<'a, T, L> {
1250    pub(crate) fn new(view: &TensorBase<ViewData<'a, T>, L>, axis: usize) -> AxisIter<'a, T, L> {
1251        assert!(axis < view.ndim());
1252        AxisIter {
1253            view: view.clone(),
1254            axis,
1255            index: 0,
1256            end: view.size(axis),
1257        }
1258    }
1259}
1260
1261impl<'a, T, L: MutLayout + RemoveDim> Iterator for AxisIter<'a, T, L> {
1262    type Item = TensorBase<ViewData<'a, T>, <L as RemoveDim>::Output>;
1263
1264    fn next(&mut self) -> Option<Self::Item> {
1265        if self.index >= self.end {
1266            None
1267        } else {
1268            let slice = self.view.index_axis(self.axis, self.index);
1269            self.index += 1;
1270            Some(slice)
1271        }
1272    }
1273
1274    fn size_hint(&self) -> (usize, Option<usize>) {
1275        let len = self.end - self.index;
1276        (len, Some(len))
1277    }
1278}
1279
1280impl<'a, T, L: MutLayout + RemoveDim> ExactSizeIterator for AxisIter<'a, T, L> {}
1281
1282impl<'a, T, L: MutLayout + RemoveDim> DoubleEndedIterator for AxisIter<'a, T, L> {
1283    fn next_back(&mut self) -> Option<Self::Item> {
1284        if self.index >= self.end {
1285            None
1286        } else {
1287            let slice = self.view.index_axis(self.axis, self.end - 1);
1288            self.end -= 1;
1289            Some(slice)
1290        }
1291    }
1292}
1293
1294/// Iterator over mutable slices of a tensor along an axis. See [`TensorViewMut::axis_iter_mut`].
1295pub struct AxisIterMut<'a, T, L: Layout + RemoveDim> {
1296    view: TensorBase<ViewMutData<'a, T>, L>,
1297    axis: usize,
1298    index: usize,
1299    end: usize,
1300}
1301
1302impl<'a, T, L: Layout + RemoveDim + Clone> AxisIterMut<'a, T, L> {
1303    pub(crate) fn new(
1304        view: TensorBase<ViewMutData<'a, T>, L>,
1305        axis: usize,
1306    ) -> AxisIterMut<'a, T, L> {
1307        // See notes in `Layout` about internal overlap.
1308        assert!(
1309            !view.layout().is_broadcast(),
1310            "Cannot mutably iterate over broadcasting view"
1311        );
1312        assert!(axis < view.ndim());
1313        AxisIterMut {
1314            axis,
1315            index: 0,
1316            end: view.size(axis),
1317            view,
1318        }
1319    }
1320}
1321
1322/// Mutable tensor view with one less dimension than `L`.
1323type SmallerMutView<'b, T, L> = TensorBase<ViewMutData<'b, T>, <L as RemoveDim>::Output>;
1324
1325impl<'a, T, L: MutLayout + RemoveDim> Iterator for AxisIterMut<'a, T, L> {
1326    type Item = TensorBase<ViewMutData<'a, T>, <L as RemoveDim>::Output>;
1327
1328    fn next(&mut self) -> Option<Self::Item> {
1329        if self.index >= self.end {
1330            None
1331        } else {
1332            let index = self.index;
1333            self.index += 1;
1334
1335            let slice = self.view.index_axis_mut(self.axis, index);
1336
1337            // Promote lifetime from self -> 'a.
1338            //
1339            // Safety: This is non-broadcasting view, and we increment the index
1340            // each time, so returned views will not overlap.
1341            let view = unsafe { transmute::<SmallerMutView<'_, T, L>, Self::Item>(slice) };
1342
1343            Some(view)
1344        }
1345    }
1346
1347    fn size_hint(&self) -> (usize, Option<usize>) {
1348        let len = self.end - self.index;
1349        (len, Some(len))
1350    }
1351}
1352
1353impl<'a, T, L: MutLayout + RemoveDim> ExactSizeIterator for AxisIterMut<'a, T, L> {}
1354
1355impl<'a, T, L: MutLayout + RemoveDim> DoubleEndedIterator for AxisIterMut<'a, T, L> {
1356    fn next_back(&mut self) -> Option<Self::Item> {
1357        if self.index >= self.end {
1358            None
1359        } else {
1360            let index = self.end - 1;
1361            self.end -= 1;
1362
1363            let slice = self.view.index_axis_mut(self.axis, index);
1364
1365            // Promote lifetime from self -> 'a.
1366            //
1367            // Safety: This is non-broadcasting view, and we increment the index
1368            // each time, so returned views will not overlap.
1369            let view = unsafe { transmute::<SmallerMutView<'_, T, L>, Self::Item>(slice) };
1370
1371            Some(view)
1372        }
1373    }
1374}
1375
1376/// Iterator over slices of a tensor along an axis. See
1377/// [`TensorView::axis_chunks`](crate::TensorView::axis_chunks).
1378pub struct AxisChunks<'a, T, L: MutLayout> {
1379    remainder: Option<TensorBase<ViewData<'a, T>, L>>,
1380    axis: usize,
1381    chunk_size: usize,
1382}
1383
1384impl<'a, T, L: MutLayout> AxisChunks<'a, T, L> {
1385    pub(crate) fn new(
1386        view: &TensorBase<ViewData<'a, T>, L>,
1387        axis: usize,
1388        chunk_size: usize,
1389    ) -> AxisChunks<'a, T, L> {
1390        assert!(chunk_size > 0, "chunk size must be > 0");
1391        AxisChunks {
1392            remainder: if view.size(axis) > 0 {
1393                Some(view.view())
1394            } else {
1395                None
1396            },
1397            axis,
1398            chunk_size,
1399        }
1400    }
1401}
1402
1403impl<'a, T, L: MutLayout> Iterator for AxisChunks<'a, T, L> {
1404    type Item = TensorBase<ViewData<'a, T>, L>;
1405
1406    fn next(&mut self) -> Option<Self::Item> {
1407        let remainder = self.remainder.take()?;
1408        let chunk_len = self.chunk_size.min(remainder.size(self.axis));
1409        let (current, next_remainder) = remainder.split_at(self.axis, chunk_len);
1410        self.remainder = if next_remainder.size(self.axis) > 0 {
1411            Some(next_remainder)
1412        } else {
1413            None
1414        };
1415        Some(current)
1416    }
1417
1418    fn size_hint(&self) -> (usize, Option<usize>) {
1419        let len = self
1420            .remainder
1421            .as_ref()
1422            .map(|r| r.size(self.axis))
1423            .unwrap_or(0)
1424            .div_ceil(self.chunk_size);
1425        (len, Some(len))
1426    }
1427}
1428
1429impl<'a, T, L: MutLayout> ExactSizeIterator for AxisChunks<'a, T, L> {}
1430
1431impl<'a, T, L: MutLayout> DoubleEndedIterator for AxisChunks<'a, T, L> {
1432    fn next_back(&mut self) -> Option<Self::Item> {
1433        let remainder = self.remainder.take()?;
1434        let chunk_len = self.chunk_size.min(remainder.size(self.axis));
1435        let (prev_remainder, current) =
1436            remainder.split_at(self.axis, remainder.size(self.axis) - chunk_len);
1437        self.remainder = if prev_remainder.size(self.axis) > 0 {
1438            Some(prev_remainder)
1439        } else {
1440            None
1441        };
1442        Some(current)
1443    }
1444}
1445
1446/// Iterator over mutable slices of a tensor along an axis. See [`TensorViewMut::axis_chunks_mut`].
1447pub struct AxisChunksMut<'a, T, L: MutLayout> {
1448    remainder: Option<TensorBase<ViewMutData<'a, T>, L>>,
1449    axis: usize,
1450    chunk_size: usize,
1451}
1452
1453impl<'a, T, L: MutLayout> AxisChunksMut<'a, T, L> {
1454    pub(crate) fn new(
1455        view: TensorBase<ViewMutData<'a, T>, L>,
1456        axis: usize,
1457        chunk_size: usize,
1458    ) -> AxisChunksMut<'a, T, L> {
1459        // See notes in `Layout` about internal overlap.
1460        assert!(
1461            !view.layout().is_broadcast(),
1462            "Cannot mutably iterate over broadcasting view"
1463        );
1464        assert!(chunk_size > 0, "chunk size must be > 0");
1465        AxisChunksMut {
1466            remainder: if view.size(axis) > 0 {
1467                Some(view)
1468            } else {
1469                None
1470            },
1471            axis,
1472            chunk_size,
1473        }
1474    }
1475}
1476
1477impl<'a, T, L: MutLayout> Iterator for AxisChunksMut<'a, T, L> {
1478    type Item = TensorBase<ViewMutData<'a, T>, L>;
1479
1480    fn next(&mut self) -> Option<Self::Item> {
1481        let remainder = self.remainder.take()?;
1482        let chunk_len = self.chunk_size.min(remainder.size(self.axis));
1483        let (current, next_remainder) = remainder.split_at_mut(self.axis, chunk_len);
1484        self.remainder = if next_remainder.size(self.axis) > 0 {
1485            Some(next_remainder)
1486        } else {
1487            None
1488        };
1489        Some(current)
1490    }
1491
1492    fn size_hint(&self) -> (usize, Option<usize>) {
1493        let len = self
1494            .remainder
1495            .as_ref()
1496            .map(|r| r.size(self.axis))
1497            .unwrap_or(0)
1498            .div_ceil(self.chunk_size);
1499        (len, Some(len))
1500    }
1501}
1502
1503impl<'a, T, L: MutLayout> ExactSizeIterator for AxisChunksMut<'a, T, L> {}
1504
1505impl<'a, T, L: MutLayout> DoubleEndedIterator for AxisChunksMut<'a, T, L> {
1506    fn next_back(&mut self) -> Option<Self::Item> {
1507        let remainder = self.remainder.take()?;
1508        let remainder_size = remainder.size(self.axis);
1509        let chunk_len = self.chunk_size.min(remainder_size);
1510        let (prev_remainder, current) =
1511            remainder.split_at_mut(self.axis, remainder_size - chunk_len);
1512        self.remainder = if prev_remainder.size(self.axis) > 0 {
1513            Some(prev_remainder)
1514        } else {
1515            None
1516        };
1517        Some(current)
1518    }
1519}
1520
1521/// Call `f` on each element of `view`.
1522pub(crate) fn for_each_mut<T, F: Fn(&mut T)>(mut view: TensorViewMut<T>, f: F) {
1523    while view.ndim() < 4 {
1524        view.insert_axis(0);
1525    }
1526
1527    // This could be improved by sorting dimensions of `view` in order of
1528    // decreasing stride. If the resulting view is contiguous, `f` can be
1529    // applied to the underlying data directly. Even if it isn't, this will
1530    // still make memory access as contiguous as possible.
1531
1532    view.inner_iter_mut::<4>().for_each(|mut src| {
1533        for i0 in 0..src.size(0) {
1534            for i1 in 0..src.size(1) {
1535                for i2 in 0..src.size(2) {
1536                    for i3 in 0..src.size(3) {
1537                        // Safety: i0..i3 are in `[0, src.size(i))`.
1538                        let x = unsafe { src.get_unchecked_mut([i0, i1, i2, i3]) };
1539                        f(x);
1540                    }
1541                }
1542            }
1543        }
1544    });
1545}
1546
1547// Tests for iterator internals. Most tests of iterators are currently done via
1548// tests on tensor methods.
1549#[cfg(test)]
1550mod tests {
1551    use super::{AxisChunks, AxisChunksMut, Lanes, LanesMut};
1552    use crate::{AsView, Layout, NdLayout, NdTensor, Tensor};
1553
1554    fn compare_reversed<T: PartialEq + std::fmt::Debug>(fwd: &[T], rev: &[T]) {
1555        assert_eq!(fwd.len(), rev.len());
1556        for (x, y) in fwd.iter().zip(rev.iter().rev()) {
1557            assert_eq!(x, y);
1558        }
1559    }
1560
1561    /// Apply a standard set of tests to an iterator.
1562    fn test_iterator<I: Iterator + ExactSizeIterator + DoubleEndedIterator>(
1563        create_iter: impl Fn() -> I,
1564        expected: &[I::Item],
1565    ) where
1566        I::Item: PartialEq + std::fmt::Debug,
1567    {
1568        let iter = create_iter();
1569
1570        let (min_len, max_len) = iter.size_hint();
1571        let items: Vec<_> = iter.collect();
1572
1573        assert_eq!(&items, expected);
1574
1575        // Test ExactSizeIterator via `size_hint`.
1576        assert_eq!(min_len, items.len(), "incorrect size lower bound");
1577        assert_eq!(max_len, Some(items.len()), "incorrect size upper bound");
1578
1579        // Test DoubleEndedIterator via `rev`.
1580        let rev_items: Vec<_> = create_iter().rev().collect();
1581        compare_reversed(&items, &rev_items);
1582
1583        // Test FusedIterator.
1584        let mut iter = create_iter();
1585        for _x in &mut iter { /* noop */ }
1586        assert_eq!(iter.next(), None);
1587
1588        // Test fold.
1589        let mut fold_items = Vec::new();
1590        let mut idx = 0;
1591        create_iter().fold(0, |acc, item| {
1592            assert_eq!(acc, idx);
1593            fold_items.push(item);
1594            idx += 1;
1595            idx
1596        });
1597        assert_eq!(items, fold_items);
1598    }
1599
1600    /// A collection that can be mutably iterated over multiple times.
1601    ///
1602    /// We use a different pattern for testing mutable iterators to avoid
1603    /// restrictions on values returned from `FnMut` closures.
1604    trait MutIterable {
1605        type Iter<'a>: Iterator + ExactSizeIterator + DoubleEndedIterator
1606        where
1607            Self: 'a;
1608
1609        fn iter_mut(&mut self) -> Self::Iter<'_>;
1610    }
1611
1612    /// Apply a standard set of tests to a mutable iterator.
1613    fn test_mut_iterator<M, T>(mut iterable: M, expected: &[T])
1614    where
1615        M: MutIterable,
1616        T: std::fmt::Debug,
1617        for<'a> <M::Iter<'a> as Iterator>::Item: std::fmt::Debug + PartialEq + PartialEq<T>,
1618    {
1619        // Test Iterator and ExactSizeIterator.
1620        {
1621            let iter = iterable.iter_mut();
1622            let (min_len, max_len) = iter.size_hint();
1623            let items: Vec<_> = iter.collect();
1624
1625            // Test `next`
1626            assert_eq!(items, expected);
1627
1628            // Test `size_hint`
1629            assert_eq!(min_len, items.len(), "incorrect size lower bound");
1630            assert_eq!(max_len, Some(items.len()), "incorrect size upper bound");
1631        }
1632
1633        // Test FusedIterator.
1634        {
1635            let mut iter = iterable.iter_mut();
1636            for _x in &mut iter { /* noop */ }
1637            assert!(iter.next().is_none());
1638        }
1639
1640        // Test DoubleEndedIterator via `rev`.
1641        //
1642        // We use `format!` here to convert mutable references into comparable
1643        // items that have no connection to the mutable references yielded by
1644        // the iterator. Ideally this should be replaced by a clone or something.
1645        {
1646            let items: Vec<_> = iterable.iter_mut().map(|x| format!("{:?}", x)).collect();
1647            let rev_items: Vec<_> = iterable
1648                .iter_mut()
1649                .rev()
1650                .map(|x| format!("{:?}", x))
1651                .collect();
1652            compare_reversed(&items, &rev_items);
1653        }
1654
1655        // Test fold.
1656        {
1657            let items: Vec<_> = iterable.iter_mut().map(|x| format!("{:?}", x)).collect();
1658            let mut fold_items = Vec::new();
1659            let mut idx = 0;
1660            iterable.iter_mut().fold(0, |acc, item| {
1661                assert_eq!(acc, idx);
1662                fold_items.push(format!("{:?}", item));
1663                idx += 1;
1664                idx
1665            });
1666            assert_eq!(items, fold_items);
1667        }
1668    }
1669
1670    #[test]
1671    fn test_axis_chunks() {
1672        let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1673        test_iterator(
1674            || tensor.axis_chunks(0, 1),
1675            &[tensor.slice(0..1), tensor.slice(1..2)],
1676        );
1677    }
1678
1679    #[test]
1680    fn test_axis_chunks_empty() {
1681        let x = Tensor::<i32>::zeros(&[5, 0]);
1682        assert!(AxisChunks::new(&x.view(), 1, 1).next().is_none());
1683    }
1684
1685    #[test]
1686    #[should_panic(expected = "chunk size must be > 0")]
1687    fn test_axis_chunks_zero_size() {
1688        let x = Tensor::<i32>::zeros(&[5, 0]);
1689        assert!(AxisChunks::new(&x.view(), 1, 0).next().is_none());
1690    }
1691
1692    #[test]
1693    fn test_axis_chunks_mut_empty() {
1694        let mut x = Tensor::<i32>::zeros(&[5, 0]);
1695        assert!(AxisChunksMut::new(x.view_mut(), 1, 1).next().is_none());
1696    }
1697
1698    #[test]
1699    fn test_axis_chunks_mut_rev() {
1700        let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1701        let fwd: Vec<_> = tensor
1702            .axis_chunks_mut(0, 1)
1703            .map(|view| view.to_vec())
1704            .collect();
1705        let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1706        let rev: Vec<_> = tensor
1707            .axis_chunks_mut(0, 1)
1708            .rev()
1709            .map(|view| view.to_vec())
1710            .collect();
1711        compare_reversed(&fwd, &rev);
1712    }
1713
1714    #[test]
1715    #[should_panic(expected = "chunk size must be > 0")]
1716    fn test_axis_chunks_mut_zero_size() {
1717        let mut x = Tensor::<i32>::zeros(&[5, 0]);
1718        assert!(AxisChunksMut::new(x.view_mut(), 1, 0).next().is_none());
1719    }
1720
1721    #[test]
1722    fn test_axis_iter() {
1723        let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1724        test_iterator(|| tensor.axis_iter(0), &[tensor.slice(0), tensor.slice(1)]);
1725    }
1726
1727    #[test]
1728    fn test_axis_iter_mut_rev() {
1729        let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1730        let fwd: Vec<_> = tensor.axis_iter_mut(0).map(|view| view.to_vec()).collect();
1731        let mut tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1732        let rev: Vec<_> = tensor
1733            .axis_iter_mut(0)
1734            .rev()
1735            .map(|view| view.to_vec())
1736            .collect();
1737        compare_reversed(&fwd, &rev);
1738    }
1739
1740    #[test]
1741    fn test_inner_iter() {
1742        let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1743        test_iterator(
1744            || tensor.inner_iter::<2>(),
1745            &[tensor.slice(0), tensor.slice(1)],
1746        );
1747    }
1748
1749    #[test]
1750    fn test_inner_iter_empty() {
1751        // Create a tensor view where the inner dimension has zero size and the
1752        // outer dimension has non-zero size and non-zero strides.
1753        let tensor = NdTensor::<i32, 2>::zeros([0, 3]);
1754        assert_eq!(tensor.strides(), [3, 1]);
1755        let view = tensor.permuted([1, 0]);
1756        assert_eq!(view.strides(), [1, 3]);
1757
1758        let mut count = 0;
1759        for lane in view.inner_iter::<1>() {
1760            assert_eq!(lane.shape(), [0]);
1761            count += 1;
1762        }
1763        assert_eq!(count, 3);
1764    }
1765
1766    #[test]
1767    fn test_inner_iter_mut() {
1768        struct InnerIterMutTest(NdTensor<i32, 3>);
1769
1770        impl MutIterable for InnerIterMutTest {
1771            type Iter<'a> = super::InnerIterMut<'a, i32, NdLayout<2>>;
1772
1773            fn iter_mut(&mut self) -> Self::Iter<'_> {
1774                self.0.inner_iter_mut::<2>()
1775            }
1776        }
1777
1778        let tensor = NdTensor::from([[[1, 2], [3, 4]], [[5, 6], [7, 8]]]);
1779        test_mut_iterator(
1780            InnerIterMutTest(tensor.clone()),
1781            &[tensor.slice(0), tensor.slice(1)],
1782        );
1783    }
1784
1785    #[test]
1786    fn test_lanes() {
1787        let x = NdTensor::from([[1, 2], [3, 4]]);
1788        test_iterator(
1789            || x.lanes(0),
1790            &[x.slice((.., 0)).into(), x.slice((.., 1)).into()],
1791        );
1792        test_iterator(|| x.lanes(1), &[x.slice(0).into(), x.slice(1).into()]);
1793    }
1794
1795    #[test]
1796    fn test_lanes_empty() {
1797        let x = Tensor::<i32>::zeros(&[5, 0]);
1798        assert!(Lanes::new(x.view().view_ref(), 0).next().is_none());
1799        assert!(Lanes::new(x.view().view_ref(), 1).next().is_none());
1800    }
1801
1802    #[test]
1803    fn test_lanes_mut() {
1804        use super::Lane;
1805
1806        struct LanesMutTest(NdTensor<i32, 2>);
1807
1808        impl MutIterable for LanesMutTest {
1809            type Iter<'a> = super::LanesMut<'a, i32>;
1810
1811            fn iter_mut(&mut self) -> Self::Iter<'_> {
1812                self.0.lanes_mut(0)
1813            }
1814        }
1815
1816        let tensor = NdTensor::from([[1, 2], [3, 4]]);
1817        test_mut_iterator::<_, Lane<i32>>(
1818            LanesMutTest(tensor.clone()),
1819            &[
1820                Lane::from(tensor.slice((.., 0))),
1821                Lane::from(tensor.slice((.., 1))),
1822            ],
1823        );
1824    }
1825
1826    #[test]
1827    fn test_lane_as_slice() {
1828        // Contiguous lane
1829        let x = NdTensor::from([0, 1, 2]);
1830        let mut lane = x.lanes(0).next().unwrap();
1831        assert_eq!(lane.as_slice(), Some([0, 1, 2].as_slice()));
1832        lane.next();
1833        assert_eq!(lane.as_slice(), Some([1, 2].as_slice()));
1834        lane.next();
1835        lane.next();
1836        assert_eq!(lane.as_slice(), Some([0i32; 0].as_slice()));
1837        lane.next();
1838        assert_eq!(lane.as_slice(), Some([0i32; 0].as_slice()));
1839
1840        // Non-contiguous lane
1841        let x = NdTensor::from([[1i32, 2], [3, 4]]);
1842        let lane = x.lanes(0).next().unwrap();
1843        assert_eq!(lane.as_slice(), None);
1844    }
1845
1846    #[test]
1847    fn test_lanes_mut_empty() {
1848        let mut x = Tensor::<i32>::zeros(&[5, 0]);
1849        assert!(LanesMut::new(x.mut_view_ref(), 0).next().is_none());
1850        assert!(LanesMut::new(x.mut_view_ref(), 1).next().is_none());
1851    }
1852
1853    #[test]
1854    fn test_iter_step_by() {
1855        let tensor = Tensor::<f32>::full(&[1, 3, 16, 8], 1.);
1856
1857        // Take a non-contiguous slice so we don't use the fast path for
1858        // contiguous tensors.
1859        let tensor = tensor.slice((.., .., 1.., ..));
1860
1861        let sum = tensor.iter().sum::<f32>();
1862        for n_skip in 0..tensor.len() {
1863            let sum_skip = tensor.iter().skip(n_skip).sum::<f32>();
1864            assert_eq!(
1865                sum_skip,
1866                sum - n_skip as f32,
1867                "wrong sum for n_skip={}",
1868                n_skip
1869            );
1870        }
1871    }
1872
1873    #[test]
1874    fn test_iter_broadcast() {
1875        let tensor = Tensor::<f32>::full(&[1], 1.);
1876        let broadcast = tensor.broadcast([1, 3, 16, 8]);
1877        assert_eq!(broadcast.iter().len(), broadcast.len());
1878        let count = broadcast.iter().count();
1879        assert_eq!(count, broadcast.len());
1880        let sum = broadcast.iter().sum::<f32>();
1881        assert_eq!(sum, broadcast.len() as f32);
1882    }
1883
1884    #[test]
1885    fn test_iter() {
1886        let tensor = NdTensor::from([[[1, 2], [3, 4]]]);
1887
1888        // Test iterator over contiguous tensor.
1889        test_iterator(|| tensor.iter().copied(), &[1, 2, 3, 4]);
1890
1891        // Test iterator over non-contiguous tensor.
1892        test_iterator(|| tensor.transposed().iter().copied(), &[1, 3, 2, 4]);
1893    }
1894
1895    #[test]
1896    fn test_iter_mut() {
1897        struct IterTest(NdTensor<i32, 3>);
1898
1899        impl MutIterable for IterTest {
1900            type Iter<'a> = super::IterMut<'a, i32>;
1901
1902            fn iter_mut(&mut self) -> Self::Iter<'_> {
1903                self.0.iter_mut()
1904            }
1905        }
1906
1907        let tensor = NdTensor::from([[[1, 2], [3, 4]]]);
1908        test_mut_iterator(IterTest(tensor), &[&1, &2, &3, &4]);
1909    }
1910
1911    #[test]
1912    #[ignore]
1913    fn bench_iter() {
1914        use crate::Layout;
1915        use rten_bench::run_bench;
1916
1917        type Elem = i32;
1918
1919        let tensor = std::hint::black_box(Tensor::<Elem>::full(&[1, 6, 768, 64], 1));
1920        let n_trials = 1000;
1921        let mut result = Elem::default();
1922
1923        fn reduce<'a>(iter: impl Iterator<Item = &'a Elem>) -> Elem {
1924            iter.fold(Elem::default(), |acc, x| acc.wrapping_add(*x))
1925        }
1926
1927        // Iterate directly over data slice.
1928        run_bench(n_trials, Some("slice iter"), || {
1929            result = reduce(tensor.data().unwrap().iter());
1930        });
1931        println!("sum {}", result);
1932
1933        // Use tensor iterator with contiguous tensor. This will use the fast
1934        // path which wraps a slice iterator.
1935        run_bench(n_trials, Some("contiguous iter"), || {
1936            result = reduce(tensor.iter());
1937        });
1938        println!("sum {}", result);
1939
1940        run_bench(n_trials, Some("contiguous reverse iter"), || {
1941            result = reduce(tensor.iter().rev());
1942        });
1943        println!("sum {}", result);
1944
1945        // Use tensor iterator with non-contiguous slice. This will fall back
1946        // to indexed iteration.
1947        let slice = tensor.slice((.., .., 1.., ..));
1948        assert!(!slice.is_contiguous());
1949        let n_trials = 1000;
1950        run_bench(n_trials, Some("non-contiguous iter"), || {
1951            result = reduce(slice.iter());
1952        });
1953        println!("sum {}", result);
1954
1955        // Reverse iteration with non-contiguous slice. This is much slower
1956        // because it translates linear indexes into offsets using division.
1957        let n_trials = 100;
1958        run_bench(n_trials, Some("non-contiguous reverse iter"), || {
1959            result = reduce(slice.iter().rev());
1960        });
1961        println!("sum {}", result);
1962    }
1963
1964    #[test]
1965    #[ignore]
1966    fn bench_inner_iter() {
1967        use crate::rng::XorShiftRng;
1968        use rten_bench::run_bench;
1969
1970        let n_trials = 100;
1971        let mut rng = XorShiftRng::new(1234);
1972
1973        // Tensor with many steps along the outer two dimensions relative to the
1974        // steps along the inner two dimensions. This emphasizes the overhead of
1975        // stepping `inner_iter`.
1976        let tensor = Tensor::<f32>::rand(&[512, 512, 12, 1], &mut rng);
1977
1978        let mut sum = 0.;
1979        run_bench(n_trials, Some("inner iter"), || {
1980            for inner in tensor.inner_iter::<2>() {
1981                for i0 in 0..inner.size(0) {
1982                    for i1 in 0..inner.size(1) {
1983                        sum += inner[[i0, i1]];
1984                    }
1985                }
1986            }
1987        });
1988        println!("sum {}", sum);
1989    }
1990}