1use alloc::vec::Vec;
2use core::iter;
3use core::marker::PhantomData;
4use core::ops::Deref;
5
6use p3_field::{ExtensionField, Field, PackedValue};
7
8use crate::Matrix;
9use crate::bitrev::BitReversibleMatrix;
10
11#[derive(Debug)]
17pub struct FlatMatrixView<F, EF, Inner>(Inner, PhantomData<(F, EF)>);
18
19impl<F, EF, Inner> FlatMatrixView<F, EF, Inner> {
20 pub const fn new(inner: Inner) -> Self {
21 Self(inner, PhantomData)
22 }
23}
24
25impl<F, EF, Inner> Deref for FlatMatrixView<F, EF, Inner> {
26 type Target = Inner;
27
28 fn deref(&self) -> &Self::Target {
29 &self.0
30 }
31}
32
33impl<F, EF, Inner> Matrix<F> for FlatMatrixView<F, EF, Inner>
34where
35 F: Field,
36 EF: ExtensionField<F>,
37 Inner: Matrix<EF>,
38{
39 fn width(&self) -> usize {
40 self.0.width() * EF::DIMENSION
41 }
42
43 fn height(&self) -> usize {
44 self.0.height()
45 }
46
47 unsafe fn get_unchecked(&self, r: usize, c: usize) -> F {
48 let c_inner = c / EF::DIMENSION;
51 let inner = unsafe {
52 self.0.get_unchecked(r, c_inner)
55 };
56 inner.as_basis_coefficients_slice()[c % EF::DIMENSION]
57 }
58
59 unsafe fn row_unchecked(
60 &self,
61 r: usize,
62 ) -> impl IntoIterator<Item = F, IntoIter = impl Iterator<Item = F> + Send + Sync> {
63 unsafe {
64 FlatIter {
66 inner: self.0.row_unchecked(r).into_iter().peekable(),
67 idx: 0,
68 _phantom: PhantomData,
69 }
70 }
71 }
72
73 unsafe fn row_subseq_unchecked(
74 &self,
75 r: usize,
76 start: usize,
77 end: usize,
78 ) -> impl IntoIterator<Item = F, IntoIter = impl Iterator<Item = F> + Send + Sync> {
79 let len = end - start;
81 let inner_start = start / EF::DIMENSION;
82 unsafe {
83 FlatIter {
85 inner: self
86 .0
87 .row_subseq_unchecked(r, inner_start, self.0.width())
90 .into_iter()
91 .peekable(),
92 idx: start % EF::DIMENSION,
93 _phantom: PhantomData,
94 }
95 .take(len)
96 }
97 }
98
99 unsafe fn row_slice_unchecked(&self, r: usize) -> impl Deref<Target = [F]> {
100 unsafe {
101 self.0
103 .row_slice_unchecked(r)
104 .iter()
105 .flat_map(|val| val.as_basis_coefficients_slice())
106 .copied()
107 .collect::<Vec<_>>()
108 }
109 }
110
111 fn vertically_packed_row<P>(&self, r: usize) -> impl Iterator<Item = P>
112 where
113 F: Copy,
114 P: PackedValue<Value = F>,
115 {
116 let rows = self.0.wrapping_row_slices(r, P::WIDTH);
117 (0..self.width()).map(move |c| {
118 P::from_fn(|lane| {
119 rows[lane][c / EF::DIMENSION].as_basis_coefficients_slice()[c % EF::DIMENSION]
120 })
121 })
122 }
123}
124
125pub struct FlatIter<F, I: Iterator> {
126 inner: iter::Peekable<I>,
127 idx: usize,
128 _phantom: PhantomData<F>,
129}
130
131impl<F, EF, I> Iterator for FlatIter<F, I>
132where
133 F: Field,
134 EF: ExtensionField<F>,
135 I: Iterator<Item = EF>,
136{
137 type Item = F;
138 fn next(&mut self) -> Option<Self::Item> {
139 if self.idx == EF::DIMENSION {
140 self.idx = 0;
141 self.inner.next();
142 }
143 let value = self.inner.peek()?.as_basis_coefficients_slice()[self.idx];
144 self.idx += 1;
145 Some(value)
146 }
147}
148
149impl<F, EF, Inner> BitReversibleMatrix<F> for FlatMatrixView<F, EF, Inner>
150where
151 F: Field,
152 EF: ExtensionField<F>,
153 Inner: BitReversibleMatrix<EF>,
154{
155 type BitRev = FlatMatrixView<F, EF, Inner::BitRev>;
156
157 fn bit_reverse_rows(self) -> Self::BitRev {
158 FlatMatrixView::new(self.0.bit_reverse_rows())
159 }
160}
161
162#[cfg(test)]
163mod tests {
164 use alloc::vec;
165 use alloc::vec::Vec;
166
167 use itertools::Itertools;
168 use p3_baby_bear::BabyBear;
169 use p3_field::extension::{BinomialExtensionField, Complex, CubicTrinomialExtensionField};
170 use p3_field::{BasedVectorSpace, PackedValue, PrimeCharacteristicRing};
171 use p3_goldilocks::Goldilocks;
172 use p3_mersenne_31::Mersenne31;
173
174 use super::*;
175 use crate::dense::RowMajorMatrix;
176 type F = Mersenne31;
177 type EF = Complex<Mersenne31>;
178
179 fn assert_vertical_packing<F, EF, Inner, P>(flat: &FlatMatrixView<F, EF, Inner>)
180 where
181 F: Field,
182 EF: ExtensionField<F>,
183 Inner: Matrix<EF>,
184 P: PackedValue<Value = F>,
185 {
186 for r in [0, flat.height() - 1, flat.height() + 1] {
187 let mut packed_width = 0;
188 for (c, packed) in flat.vertically_packed_row::<P>(r).enumerate() {
189 packed_width += 1;
190 for (lane, value) in packed.as_slice().iter().enumerate() {
191 assert_eq!(*value, flat.get((r + lane) % flat.height(), c).unwrap());
192 }
193 }
194 assert_eq!(packed_width, flat.width());
195 }
196 }
197
198 fn extension_matrix<F, EF>(height: usize, width: usize) -> RowMajorMatrix<EF>
199 where
200 F: Field,
201 EF: ExtensionField<F>,
202 {
203 let values = (0..height * width)
204 .map(|i| {
205 EF::from_basis_coefficients_fn(|coeff| F::from_usize(i * EF::DIMENSION + coeff + 1))
206 })
207 .collect();
208 RowMajorMatrix::new(values, width)
209 }
210
211 fn assert_dense_vertical_packing<F, EF>()
212 where
213 F: Field,
214 EF: ExtensionField<F>,
215 {
216 for height in [1, 3, 16] {
217 for width in [1, 3, 17] {
218 let flat = FlatMatrixView::<F, EF, _>::new(extension_matrix(height, width));
219 assert_vertical_packing::<F, EF, _, F>(&flat);
220 assert_vertical_packing::<F, EF, _, F::Packing>(&flat);
221 }
222 }
223 }
224
225 fn assert_bit_reversed_vertical_packing<F, EF>()
226 where
227 F: Field,
228 EF: ExtensionField<F>,
229 {
230 for height in [1, 16] {
232 for width in [1, 3, 17] {
233 let inner = extension_matrix::<F, EF>(height, width).bit_reverse_rows();
234 let flat = FlatMatrixView::<F, EF, _>::new(inner);
235 assert_vertical_packing::<F, EF, _, F>(&flat);
236 assert_vertical_packing::<F, EF, _, F::Packing>(&flat);
237 }
238 }
239 }
240
241 #[test]
242 fn test_vertically_packed_row_extension_degrees() {
243 type EF2 = BinomialExtensionField<Goldilocks, 2>;
244 type EF3 = CubicTrinomialExtensionField<Goldilocks>;
245 type EF4 = BinomialExtensionField<BabyBear, 4>;
246 type EF5 = BinomialExtensionField<BabyBear, 5>;
247
248 assert_dense_vertical_packing::<Goldilocks, EF2>();
249 assert_dense_vertical_packing::<Goldilocks, EF3>();
250 assert_dense_vertical_packing::<BabyBear, EF4>();
251 assert_dense_vertical_packing::<BabyBear, EF5>();
252
253 assert_bit_reversed_vertical_packing::<Goldilocks, EF2>();
254 assert_bit_reversed_vertical_packing::<Goldilocks, EF3>();
255 assert_bit_reversed_vertical_packing::<BabyBear, EF4>();
256 assert_bit_reversed_vertical_packing::<BabyBear, EF5>();
257 }
258
259 #[test]
260 #[should_panic]
261 fn test_vertically_packed_row_empty_height_panics() {
262 let flat = FlatMatrixView::<BabyBear, BinomialExtensionField<BabyBear, 4>, _>::new(
263 RowMajorMatrix::new(vec![], 1),
264 );
265 let _ = flat.vertically_packed_row::<BabyBear>(0).next();
266 }
267
268 #[test]
269 fn flat_matrix() {
270 let values = vec![
271 EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + 10)),
272 EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + 20)),
273 EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + 30)),
274 EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + 40)),
275 EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + 50)),
276 EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + 60)),
277 ];
278 let ext = RowMajorMatrix::<EF>::new(values, 2);
279 let flat = FlatMatrixView::<F, EF, _>::new(ext);
280
281 assert_eq!(flat.width(), 4);
282 assert_eq!(flat.height(), 3);
283
284 assert_eq!(flat.get(0, 2), Some(F::from_u8(20)));
285 assert_eq!(flat.get(1, 3), Some(F::from_u8(41)));
286 assert_eq!(flat.get(2, 0), Some(F::from_u8(50)));
287
288 unsafe {
289 assert_eq!(flat.get_unchecked(0, 1), F::from_u8(11));
290 assert_eq!(flat.get_unchecked(1, 0), F::from_u8(30));
291 assert_eq!(flat.get_unchecked(2, 2), F::from_u8(60));
292 }
293
294 assert_eq!(
295 &*flat.row_slice(0).unwrap(),
296 &[10, 11, 20, 21].map(F::from_u8)
297 );
298 unsafe {
299 assert_eq!(
300 &*flat.row_slice_unchecked(1),
301 &[30, 31, 40, 41].map(F::from_u8)
302 );
303 assert_eq!(
304 &*flat.row_subslice_unchecked(2, 0, 3),
305 &[50, 51, 60].map(F::from_u8)
306 );
307 }
308
309 assert_eq!(
310 flat.row(2).unwrap().into_iter().collect_vec(),
311 [50, 51, 60, 61].map(F::from_u8)
312 );
313 unsafe {
314 assert_eq!(
315 flat.row_unchecked(1).into_iter().collect_vec(),
316 [30, 31, 40, 41].map(F::from_u8)
317 );
318 assert_eq!(
319 flat.row_subseq_unchecked(0, 1, 4).into_iter().collect_vec(),
320 [11, 20, 21].map(F::from_u8)
321 );
322 }
323
324 assert!(flat.get(0, 4).is_none()); assert!(flat.get(3, 0).is_none()); assert!(flat.row(3).is_none()); assert!(flat.row_slice(3).is_none()); }
329
330 #[test]
331 fn test_flat_matrix_width() {
332 let matrix = RowMajorMatrix::<EF>::new(vec![EF::default(); 4], 2);
336 let flat = FlatMatrixView::<F, EF, _>::new(matrix);
337 assert_eq!(flat.width(), 2 * <EF as BasedVectorSpace<F>>::DIMENSION);
338 }
339
340 #[test]
341 fn test_flat_matrix_height() {
342 let matrix = RowMajorMatrix::<EF>::new(vec![EF::default(); 6], 3);
345 let flat = FlatMatrixView::<F, EF, _>::new(matrix);
346 assert_eq!(flat.height(), 2);
347 }
348
349 #[test]
350 fn test_flat_matrix_row_iterator() {
351 let values = vec![
354 EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + 1)),
355 EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + 10)),
356 ];
357 let matrix = RowMajorMatrix::new(values, 2);
358 let flat = FlatMatrixView::<F, EF, _>::new(matrix);
359
360 let row: Vec<_> = flat.first_row().unwrap().into_iter().collect();
362 let expected = [1, 2, 10, 11].map(F::from_u8).to_vec();
363
364 assert_eq!(row, expected);
365 }
366
367 #[test]
368 fn test_flat_matrix_row_slice_correctness() {
369 let ef = |offset| EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + offset));
372 let matrix = RowMajorMatrix::new(vec![ef(1), ef(10)], 2);
373 let flat = FlatMatrixView::<F, EF, _>::new(matrix);
374
375 assert_eq!(
376 &*flat.row_slice(0).unwrap(),
377 &[1, 2, 10, 11].map(F::from_u8)
378 );
379 }
380
381 #[test]
382 fn test_flat_matrix_empty() {
383 let matrix = RowMajorMatrix::<EF>::new(vec![], 0);
386 let flat = FlatMatrixView::<F, EF, _>::new(matrix);
387
388 assert_eq!(flat.height(), 0);
389 assert_eq!(flat.width(), 0);
390 }
391
392 #[test]
393 fn test_flat_iter_length_and_values() {
394 let ef = |offset| EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + offset));
397 let values = vec![ef(0), ef(10), ef(20)];
398 let matrix = RowMajorMatrix::new(values, 3); let flat = FlatMatrixView::<F, EF, _>::new(matrix);
400
401 let row: Vec<_> = flat.first_row().unwrap().into_iter().collect();
402 let expected = [0, 1, 10, 11, 20, 21].map(F::from_u8).to_vec();
403 assert_eq!(row, expected);
404 }
405
406 #[test]
407 fn test_flat_matrix_multiple_rows() {
408 let ef = |base| EF::from_basis_coefficients_fn(|i| F::from_u8(base + i as u8));
412 let matrix = RowMajorMatrix::new(vec![ef(0), ef(10), ef(20), ef(30)], 2);
413 let flat = FlatMatrixView::<F, EF, _>::new(matrix);
414
415 let row0: Vec<_> = flat.first_row().unwrap().into_iter().collect();
416 let row1: Vec<_> = flat.row(1).unwrap().into_iter().collect();
417
418 assert_eq!(row0, [0, 1, 10, 11].map(F::from_u8).to_vec());
419 assert_eq!(row1, [20, 21, 30, 31].map(F::from_u8).to_vec());
420 }
421
422 #[test]
423 fn test_flat_iter_yields_across_multiple_efs() {
424 let ef = |offset| EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + offset));
432 let matrix = RowMajorMatrix::new(vec![ef(0), ef(10), ef(20)], 3); let flat = FlatMatrixView::<F, EF, _>::new(matrix);
434
435 let mut row_iter = flat.row(0).unwrap().into_iter();
436
437 let expected = [0, 1, 10, 11, 20, 21].map(F::from_u8);
439
440 for expected_val in expected {
441 assert_eq!(row_iter.next(), Some(expected_val));
442 }
443
444 assert_eq!(row_iter.next(), None);
446 }
447
448 #[test]
449 fn test_row_subseq_start_ge_dimension() {
450 let ef = |offset| EF::from_basis_coefficients_fn(|i| F::from_u8(i as u8 + offset));
451 let values = vec![ef(10), ef(20), ef(30)];
452 let matrix = RowMajorMatrix::new(values, 3);
453 let flat = FlatMatrixView::<F, EF, _>::new(matrix);
454
455 unsafe {
456 let result: Vec<_> = flat.row_subseq_unchecked(0, 2, 5).into_iter().collect();
457 assert_eq!(result, [20, 21, 30].map(F::from_u8).to_vec());
458
459 let result: Vec<_> = flat.row_subseq_unchecked(0, 3, 6).into_iter().collect();
460 assert_eq!(result, [21, 30, 31].map(F::from_u8).to_vec());
461 }
462 }
463}