Skip to main content

diskann_utils/views/rowmajor/
iter.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6use std::{marker::PhantomData, num::NonZeroUsize, ptr::NonNull};
7
8use crate::views::rowmajor::{Layout, Matrix, Mut, Ref};
9
10//------//
11// Rows //
12//------//
13
14/// An iterator over rows in a matrix. See: [`Matrix::rows`].
15#[derive(Debug)]
16pub struct Rows<'a, T> {
17    ptr: NonNull<T>,
18    remaining: usize,
19    ncols: usize,
20    _lifetime: PhantomData<&'a [T]>,
21}
22
23impl<'a, T> Rows<'a, T> {
24    pub(super) fn new(m: Ref<'a, T>) -> Self {
25        let layout = m.layout();
26        Self {
27            ptr: m.as_nonnull(),
28            remaining: layout.nrows(),
29            ncols: layout.ncols(),
30            _lifetime: PhantomData,
31        }
32    }
33}
34
35// SAFETY: `Rows<'_, T>` owns a shared slice borrow, so sending it requires `T: Sync`.
36unsafe impl<T> Send for Rows<'_, T> where T: Sync {}
37// SAFETY: Shared access to `Rows<'_, T>` exposes only shared access to `T`.
38unsafe impl<T> Sync for Rows<'_, T> where T: Sync {}
39
40impl<'a, T> Iterator for Rows<'a, T> {
41    type Item = &'a [T];
42    fn next(&mut self) -> Option<&'a [T]> {
43        self.remaining.checked_sub(1).map(|remaining| {
44            // SAFETY: Construction from a valid `Ref` guarantees that each remaining row
45            // contains `ncols` initialized elements beginning at `self.ptr`.
46            let item =
47                unsafe { std::slice::from_raw_parts(self.ptr.as_ptr().cast_const(), self.ncols) };
48            self.remaining = remaining;
49
50            // SAFETY: Advancing by one row remains within or one past the original matrix
51            // span. The validated parent layout guarantees that the offset is representable.
52            self.ptr = unsafe { self.ptr.add(self.ncols) };
53            item
54        })
55    }
56
57    fn size_hint(&self) -> (usize, Option<usize>) {
58        (self.remaining, Some(self.remaining))
59    }
60}
61
62impl<T> ExactSizeIterator for Rows<'_, T> {}
63impl<T> std::iter::FusedIterator for Rows<'_, T> {}
64
65//---------//
66// RowsMut //
67//---------//
68
69/// An iterator over mutable rows in a matrix. See: [`crate::views::rowmajor::MatrixMut::rows_mut`].
70#[derive(Debug)]
71pub struct RowsMut<'a, T> {
72    ptr: NonNull<T>,
73    remaining: usize,
74    ncols: usize,
75    _lifetime: PhantomData<&'a mut [T]>,
76}
77
78impl<'a, T> RowsMut<'a, T> {
79    pub(super) fn new(m: Mut<'a, T>) -> Self {
80        let layout = m.layout();
81        Self {
82            ptr: m.as_nonnull(),
83            remaining: layout.nrows(),
84            ncols: layout.ncols(),
85            _lifetime: PhantomData,
86        }
87    }
88}
89
90// SAFETY: `RowsMut<'_, T>` owns an exclusive slice borrow, so sending it requires `T: Send`.
91unsafe impl<T> Send for RowsMut<'_, T> where T: Send {}
92// SAFETY: Shared access to `RowsMut<'_, T>` exposes only shared access to `T`.
93unsafe impl<T> Sync for RowsMut<'_, T> where T: Sync {}
94
95impl<'a, T> Iterator for RowsMut<'a, T> {
96    type Item = &'a mut [T];
97    fn next(&mut self) -> Option<&'a mut [T]> {
98        self.remaining.checked_sub(1).map(|remaining| {
99            // SAFETY: Construction from a valid `Mut` guarantees that each remaining row
100            // contains `ncols` initialized elements beginning at `self.ptr`. Advancing the
101            // pointer after every yield makes nonempty returned rows disjoint; zero-length
102            // rows do not access memory and may share an address.
103            let item = unsafe { std::slice::from_raw_parts_mut(self.ptr.as_ptr(), self.ncols) };
104            self.remaining = remaining;
105
106            // SAFETY: Advancing by one row remains within or one past the original matrix
107            // span. The validated parent layout guarantees that the offset is representable.
108            self.ptr = unsafe { self.ptr.add(self.ncols) };
109            item
110        })
111    }
112
113    fn size_hint(&self) -> (usize, Option<usize>) {
114        (self.remaining, Some(self.remaining))
115    }
116}
117
118impl<T> ExactSizeIterator for RowsMut<'_, T> {}
119impl<T> std::iter::FusedIterator for RowsMut<'_, T> {}
120
121//---------//
122// Windows //
123//---------//
124
125/// An iterator over rows in a matrix. See: [`Matrix::window_iter`].
126#[derive(Debug)]
127pub struct Windows<'a, T> {
128    ptr: NonNull<T>,
129    remaining: usize,
130    batchsize: NonZeroUsize,
131    ncols: usize,
132    _lifetime: PhantomData<&'a [T]>,
133}
134
135impl<'a, T> Windows<'a, T> {
136    pub(super) fn new(m: Ref<'a, T>, batchsize: NonZeroUsize) -> Self {
137        let layout = m.layout();
138        Self {
139            ptr: m.as_nonnull(),
140            remaining: layout.nrows(),
141            batchsize,
142            ncols: layout.ncols(),
143            _lifetime: PhantomData,
144        }
145    }
146}
147
148// SAFETY: `Windows<'_, T>` owns a shared slice borrow, so sending it requires `T: Sync`.
149unsafe impl<T> Send for Windows<'_, T> where T: Sync {}
150// SAFETY: Shared access to `Windows<'_, T>` exposes only shared access to `T`.
151unsafe impl<T> Sync for Windows<'_, T> where T: Sync {}
152
153impl<'a, T> Iterator for Windows<'a, T> {
154    type Item = Ref<'a, T>;
155    fn next(&mut self) -> Option<Ref<'a, T>> {
156        if self.remaining == 0 {
157            None
158        } else {
159            let next_remaining = self.remaining.saturating_sub(self.batchsize.get());
160            let nrows = self.remaining - next_remaining;
161
162            // SAFETY: `self.ptr` starts the remaining suffix of a valid `Ref`, and `nrows`
163            // does not exceed that suffix. Keeping the parent's column count therefore
164            // produces a valid subview and a layout no larger than the parent layout.
165            let window = unsafe {
166                Ref {
167                    ptr: self.ptr,
168                    layout: Layout::new_unchecked(nrows, self.ncols),
169                    _lifetime: PhantomData,
170                }
171            };
172
173            // SAFETY: Advancing by the yielded window remains within or one past the
174            // original matrix span. The validated parent layout guarantees that the
175            // multiplication and pointer offset are representable.
176            self.ptr = unsafe { self.ptr.add(nrows * self.ncols) };
177            self.remaining = next_remaining;
178            Some(window)
179        }
180    }
181
182    fn size_hint(&self) -> (usize, Option<usize>) {
183        let remaining = self.remaining.div_ceil(self.batchsize.get());
184        (remaining, Some(remaining))
185    }
186}
187
188impl<T> ExactSizeIterator for Windows<'_, T> {}
189impl<T> std::iter::FusedIterator for Windows<'_, T> {}