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#[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 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
164pub 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#[derive(Debug, Clone, Copy, PartialEq)]
265pub enum StepInfo<T: Clone> {
266 EnterDimension(usize),
269 ExitDimension(usize),
271 Value(T),
273 End,
276}
277
278#[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
376pub 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
440pub 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> {}