Skip to main content

candela/tensor/
iter.rs

1use std::iter::FusedIterator;
2
3use crate::tensor::mem_formats::layout::Layout;
4use crate::tensor::traits::StreamingIterator;
5
6pub struct ContiguousIter<'a, T: Clone> {
7    data: &'a [T],
8    offset: usize,
9    left_over: usize,
10}
11
12impl<'a, T: Clone> ContiguousIter<'a, T> {
13    pub fn new(data: &'a [T], offset: usize, len: usize) -> Self {
14        Self {
15            data,
16            offset,
17            left_over: len,
18        }
19    }
20}
21
22impl<'a, T: Clone> Iterator for ContiguousIter<'a, T> {
23    type Item = &'a T;
24
25    fn next(&mut self) -> Option<Self::Item> {
26        if self.left_over == 0 {
27            return None;
28        }
29
30        let item = &self.data[self.offset] as *const T;
31        self.offset += 1;
32        self.left_over -= 1;
33
34        Some(unsafe { &*item })
35    }
36
37    fn size_hint(&self) -> (usize, Option<usize>) {
38        (self.left_over, Some(self.left_over))
39    }
40}
41
42impl<'a, T: Clone> ExactSizeIterator for ContiguousIter<'a, T> {}
43
44impl<'a, T: Clone> FusedIterator for ContiguousIter<'a, T> {}
45
46///////////////////////////////////////////////////////////////
47
48// TODO: This impl is not correct anymore. If used please fix it like
49// the non-mut one.
50
51// pub struct MutContiguousIter<'a, T: Clone> {
52//     data: RwLockWriteGuard<'a, Vec<T>>,
53//     index: usize,
54// }
55//
56// impl<'a, T: Clone> MutContiguousIter<'a, T> {
57//     pub fn new(lock: &'a RwLock<Vec<T>>) -> Self {
58//         let data: RwLockWriteGuard<'_, Vec<T>> = lock.write();
59//         Self { data, index: 0 }
60//     }
61// }
62//
63// impl<'a, T: Clone> Iterator for MutContiguousIter<'a, T> {
64//     type Item = &'a mut T;
65//
66//     fn next(&mut self) -> Option<Self::Item> {
67//         if self.index >= self.data.len() {
68//             return None;
69//         }
70//
71//         let mut item = NonNull::new(&mut self.data[self.index] as *mut T).unwrap();
72//         self.index += 1;
73//
74//         return Some(unsafe { item.as_mut() });
75//     }
76//
77//     fn size_hint(&self) -> (usize, Option<usize>) {
78//         let len = self.data.len() - self.index;
79//
80//         (len, Some(len))
81//     }
82// }
83//
84// impl<'a, T: Clone> ExactSizeIterator for MutContiguousIter<'a, T> {}
85//
86// impl<'a, T: Clone> FusedIterator for MutContiguousIter<'a, T> {}
87//
88
89///////////////////////////////////////////////////////////////
90
91/// Iterator over a tensor's elements in logical (row-major) order.
92///
93/// Returned by [`Tensor::iter`](crate::Tensor::iter). It walks the backing
94/// buffer following the tensor's [`Layout`], so a sliced or transposed tensor
95/// yields its elements in the order its shape implies rather than in storage
96/// order.
97#[derive(Debug, Clone)]
98pub struct Iter<'a, T> {
99    data: &'a [T],
100    pos: isize,
101    counter: Box<[usize]>,
102    layout: &'a Layout,
103    left_over: usize,
104}
105
106impl<'a, T: Clone> Iter<'a, T> {
107    // TODO: data_len is used anywhere? Like at all? If not, maybe just remove it.
108    pub fn new(data: &'a [T], data_len: usize, layout: &'a Layout) -> Self {
109        let counter = vec![0; layout.shape().len()].into_boxed_slice();
110
111        Self {
112            data,
113            pos: layout.offset() as isize,
114            layout,
115            counter,
116            left_over: data_len,
117        }
118    }
119}
120
121impl<'a, T: Clone> Iterator for Iter<'a, T> {
122    type Item = &'a T;
123
124    fn next(&mut self) -> Option<Self::Item> {
125        if self.left_over == 0 {
126            return None;
127        }
128
129        let last = self.counter.len() - 1;
130        self.counter[last] += 1;
131        let mut step_dim = last;
132
133        for dim in (1..self.counter.len()).rev() {
134            if self.counter[dim] == self.layout.shape()[dim] {
135                self.counter[dim] = 0;
136                self.counter[dim - 1] += 1;
137
138                step_dim = dim - 1;
139                continue;
140            }
141            break;
142        }
143
144        let pos = self.pos as usize;
145
146        unsafe {
147            let item = &self.data[pos] as *const T;
148            self.pos += self.layout.adj_stride()[step_dim] as isize;
149            self.left_over -= 1;
150
151            Some(&*item)
152        }
153    }
154
155    fn size_hint(&self) -> (usize, Option<usize>) {
156        (self.left_over, Some(self.left_over))
157    }
158}
159
160impl<'a, T: Clone> ExactSizeIterator for Iter<'a, T> {}
161
162impl<'a, T: Clone> FusedIterator for Iter<'a, T> {}
163
164///////////////////////////////////////////////////////////////
165
166pub struct MutSliceIter<'a, T> {
167    data: &'a mut Vec<T>,
168    pos: isize,
169    counter: Box<[usize]>,
170    layout: &'a Layout,
171    left_over: usize,
172}
173
174impl<'a, T: Clone> MutSliceIter<'a, T> {
175    pub fn new(data: &'a mut Vec<T>, data_len: usize, layout: &'a Layout) -> Self {
176        let counter = vec![0; layout.shape().len()].into_boxed_slice();
177
178        Self {
179            data,
180            pos: layout.offset() as isize,
181            layout,
182            counter,
183            left_over: data_len,
184        }
185    }
186}
187
188impl<'a, T: Clone> Iterator for MutSliceIter<'a, T> {
189    type Item = &'a mut T;
190
191    fn next(&mut self) -> Option<Self::Item> {
192        if self.left_over == 0 {
193            return None;
194        }
195
196        let last = self.counter.len() - 1;
197        self.counter[last] += 1;
198        let mut step_dim = last;
199
200        for dim in (1..self.counter.len()).rev() {
201            if self.counter[dim] == self.layout.shape()[dim] {
202                self.counter[dim] = 0;
203                self.counter[dim - 1] += 1;
204
205                step_dim = dim - 1;
206                continue;
207            }
208            break;
209        }
210
211        let pos = self.pos as usize;
212
213        unsafe {
214            let item = &mut self.data[pos] as *mut T;
215            self.pos += self.layout.adj_stride()[step_dim] as isize;
216            self.left_over -= 1;
217
218            Some(&mut *item)
219        }
220    }
221
222    fn size_hint(&self) -> (usize, Option<usize>) {
223        (self.left_over, Some(self.left_over))
224    }
225}
226
227impl<'a, T: Clone> ExactSizeIterator for MutSliceIter<'a, T> {}
228
229impl<'a, T: Clone> FusedIterator for MutSliceIter<'a, T> {}
230
231///////////////////////////////////////////////////////////////
232
233/// A single event in a structural walk of a tensor, produced by
234/// [`Tensor::informed_iter`](crate::Tensor::informed_iter).
235///
236/// A walk is the nested loop it would take to visit every element - one loop per
237/// dimension, the innermost yielding values. Each variant marks a point in that
238/// loop: [`EnterDimension`] when a loop opens, [`Value`] for an element the
239/// innermost loop reads, [`ExitDimension`] when a loop closes, and [`End`] once
240/// every loop has finished.
241///
242/// The walk follows the tensor's logical layout, so a sliced or transposed
243/// tensor is visited in the order its shape implies.
244///
245/// [`EnterDimension`]: StepInfo::EnterDimension
246/// [`ExitDimension`]: StepInfo::ExitDimension
247/// [`Value`]: StepInfo::Value
248/// [`End`]: StepInfo::End
249///
250/// # Examples
251///
252/// ```
253/// use candela::{StepInfo, Tensor};
254///
255/// let t = Tensor::from_slice(&[1.0, 2.0], &[2]);
256/// let events: Vec<StepInfo<f64>> = t.informed_iter().collect();
257/// assert_eq!(events, vec![
258///     StepInfo::EnterDimension(0),
259///     StepInfo::Value(1.0),
260///     StepInfo::Value(2.0),
261///     StepInfo::ExitDimension(0),
262/// ]);
263/// ```
264#[derive(Debug, Clone, Copy, PartialEq)]
265pub enum StepInfo<T: Clone> {
266    /// A dimension's loop has opened; the payload is its index, `0` being the
267    /// outermost.
268    EnterDimension(usize),
269    /// A dimension's loop has closed; the payload is its index.
270    ExitDimension(usize),
271    /// An element, read by the innermost loop in logical order.
272    Value(T),
273    /// Every loop has finished. Iteration terminates by returning `None`, so the
274    /// iterator does not yield this variant.
275    End,
276}
277
278/// Structural walk over a tensor, yielding a [`StepInfo`] per element and per
279/// dimension boundary.
280///
281/// Returned by [`Tensor::informed_iter`](crate::Tensor::informed_iter). Unlike
282/// [`Iter`], which yields a flat stream of values, its events also mark where
283/// each sub-array opens and closes, which is what lets you reconstruct the
284/// tensor's nesting. See [`StepInfo`] for the event kinds and
285/// [`Tensor::informed_iter`](crate::Tensor::informed_iter) for a worked example.
286#[derive(Debug, Clone)]
287pub struct InformedIter<'a, T: Clone> {
288    buffer: &'a [T],
289    layout: &'a Layout,
290    next_state: StepInfo<T>,
291    pos: i64,
292    counter: Vec<usize>,
293}
294
295impl<'a, T: Clone> InformedIter<'a, T> {
296    pub fn new(data: &'a [T], layout: &'a Layout) -> Self {
297        let len = layout.shape().len();
298
299        Self {
300            buffer: data,
301            layout,
302            next_state: StepInfo::<T>::EnterDimension(0),
303            pos: layout.offset() as i64,
304            counter: vec![0; len],
305        }
306    }
307}
308
309impl<'a, T: Copy> Iterator for InformedIter<'a, T> {
310    type Item = StepInfo<T>;
311
312    fn next(&mut self) -> Option<Self::Item> {
313        match self.next_state {
314            StepInfo::EnterDimension(dim) => {
315                if dim == self.layout.shape().len() - 1 {
316                    self.next_state = StepInfo::Value(self.buffer[self.pos as usize]);
317
318                    return Some(StepInfo::EnterDimension(dim));
319                }
320
321                self.next_state = StepInfo::EnterDimension(dim + 1);
322
323                Some(StepInfo::EnterDimension(dim))
324            }
325            StepInfo::ExitDimension(dim) => {
326                if dim == 0 {
327                    self.next_state = StepInfo::End;
328                    return Some(StepInfo::ExitDimension(dim));
329                }
330
331                self.counter[dim] = 0;
332                self.counter[dim - 1] += 1;
333
334                if self.counter[dim - 1] == self.layout.shape()[dim - 1] {
335                    self.next_state = StepInfo::ExitDimension(dim - 1);
336                    return Some(StepInfo::ExitDimension(dim));
337                }
338
339                self.pos += self.layout.adj_stride()[dim - 1] as i64;
340                self.next_state = StepInfo::EnterDimension(dim);
341
342                Some(StepInfo::ExitDimension(dim))
343            }
344            StepInfo::Value(v) => {
345                let counter_last = self.counter.len() - 1;
346
347                if *self.counter.last().unwrap() == *self.layout.shape().last().unwrap() - 1 {
348                    self.next_state = StepInfo::ExitDimension(self.counter.len() - 1);
349                    self.counter[counter_last] = 0;
350
351                    return Some(StepInfo::Value(v));
352                }
353
354                self.pos += *self.layout.adj_stride().last().unwrap() as i64;
355                self.counter[counter_last] += 1;
356
357                self.next_state = StepInfo::Value(self.buffer[self.pos as usize]);
358
359                Some(StepInfo::Value(v))
360            }
361            StepInfo::End => None,
362        }
363    }
364
365    fn size_hint(&self) -> (usize, Option<usize>) {
366        let len = self.layout.len() - self.pos as usize;
367
368        (len, Some(len))
369    }
370}
371
372impl<'a, T: Copy> ExactSizeIterator for InformedIter<'a, T> {}
373
374impl<'a, T: Copy> FusedIterator for InformedIter<'a, T> {}
375
376/////////////////////////////////////////////////////////////
377pub struct PackedBuffer<'a, T: Clone> {
378    pub packing_buffer: &'a [T],
379    pub absolute_buffer_position: usize,
380}
381
382pub struct ChunkedSliceIter<I, T: Clone>
383where
384    I: IntoIterator<Item = T>,
385{
386    iter: I::IntoIter,
387    packing_buffer: Vec<T>,
388    absolute_buffer_position: usize,
389}
390
391impl<I, T: Clone + Default> ChunkedSliceIter<I, T>
392where
393    I: Iterator<Item = T>,
394{
395    pub fn new(iter: I, packing_buffer_size: usize) -> Self {
396        Self {
397            iter,
398            packing_buffer: vec![T::default(); packing_buffer_size],
399            absolute_buffer_position: 0,
400        }
401    }
402}
403
404impl<I, T: Clone> StreamingIterator for ChunkedSliceIter<I, T>
405where
406    I: IntoIterator<Item = T>,
407{
408    type Item<'a>
409        = PackedBuffer<'a, T>
410    where
411        Self: 'a;
412
413    fn next_stream<'a>(&'a mut self) -> Option<Self::Item<'a>> {
414        let mut len = 0;
415
416        for slot in &mut self.packing_buffer {
417            match self.iter.next() {
418                Some(v) => {
419                    *slot = v;
420                    len += 1;
421                }
422                None => break,
423            }
424        }
425
426        if len == 0 {
427            return None;
428        }
429
430        let pos = self.absolute_buffer_position;
431        self.absolute_buffer_position += len;
432
433        Some(PackedBuffer {
434            packing_buffer: &self.packing_buffer[..len],
435            absolute_buffer_position: pos,
436        })
437    }
438}
439
440/////////////////////////////////////////////////////////////
441pub struct ChunkedContiguousIter<'a, T: Clone> {
442    data: &'a [T],
443    packing_buffer_size: usize,
444    absolute_buffer_position: usize,
445}
446
447impl<'a, T: Clone> ChunkedContiguousIter<'a, T> {
448    pub fn new(data: &'a [T], packing_buffer_size: usize) -> Self {
449        Self {
450            data,
451            packing_buffer_size,
452            absolute_buffer_position: 0,
453        }
454    }
455}
456
457impl<'b, T: Clone> StreamingIterator for ChunkedContiguousIter<'b, T> {
458    type Item<'a>
459        = PackedBuffer<'a, T>
460    where
461        Self: 'a;
462
463    fn next_stream<'a>(&'a mut self) -> Option<Self::Item<'a>> {
464        if self.absolute_buffer_position >= self.data.len() {
465            return None;
466        }
467
468        let start = self.absolute_buffer_position;
469        let end = (self.absolute_buffer_position + self.packing_buffer_size).min(self.data.len());
470        self.absolute_buffer_position = end;
471
472        Some(PackedBuffer {
473            packing_buffer: &self.data[start..end],
474            absolute_buffer_position: start,
475        })
476    }
477}
478
479impl<'a, T: Clone> Iterator for ChunkedContiguousIter<'a, T> {
480    type Item = PackedBuffer<'a, T>;
481
482    fn next(&mut self) -> Option<Self::Item> {
483        if self.absolute_buffer_position >= self.data.len() {
484            return None;
485        }
486
487        let start = self.absolute_buffer_position;
488        let end = (self.absolute_buffer_position + self.packing_buffer_size).min(self.data.len());
489        self.absolute_buffer_position = end;
490
491        Some(PackedBuffer {
492            packing_buffer: &self.data[start..end],
493            absolute_buffer_position: start,
494        })
495    }
496}
497
498impl<'a, T: Clone> FusedIterator for ChunkedContiguousIter<'a, T> {}