Skip to main content

p3_matrix/
extension.rs

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/// A view that flattens a matrix of extension field elements into a matrix of base field elements.
12///
13/// Each element of the original matrix is an extension field element `EF`, composed of several
14/// base field elements `F`. This view expands each `EF` element into its base field components,
15/// effectively increasing the number of columns (width) while keeping the number of rows unchanged.
16#[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        // The c'th base field element in a row of extension field elements is
49        // at index c % EF::DIMENSION in the c / EF::DIMENSION'th extension element.
50        let c_inner = c / EF::DIMENSION;
51        let inner = unsafe {
52            // Safety: The caller must ensure that r < self.height() and c < self.width().
53            // Assuming this, c / EF::DIMENSION < self.0.width().
54            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            // Safety: The caller must ensure that r < self.height().
65            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        // We can skip the first start / EF::DIMENSION elements in the row.
80        let len = end - start;
81        let inner_start = start / EF::DIMENSION;
82        unsafe {
83            // Safety: The caller must ensure that r < self.height() and start <= end <= self.width().
84            FlatIter {
85                inner: self
86                    .0
87                    // We set end to be the width of the inner matrix and use take to ensure we get the right
88                    // number of elements.
89                    .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            // Safety: The caller must ensure that r < self.height().
102            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        // Bit-reversal views require power-of-two heights.
231        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()); // Width out of bounds
325        assert!(flat.get(3, 0).is_none()); // Height out of bounds
326        assert!(flat.row(3).is_none()); // Height out of bounds
327        assert!(flat.row_slice(3).is_none()); // Height out of bounds
328    }
329
330    #[test]
331    fn test_flat_matrix_width() {
332        // Create a 2-column, 2-row matrix of EF elements.
333        // Each EF element expands to EF::DIMENSION base field elements when flattened.
334        // Therefore, the flattened width should be 2 * EF::DIMENSION.
335        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        // Construct a 3-column matrix with 6 EF elements (2 rows).
343        // The flattened view should preserve the original number of rows.
344        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        // Create a single row of two EF elements:
352        // First EF = [1, 2], second EF = [10, 11] (in base field representation).
353        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        // Flattened row should concatenate basis coefficients of both EF elements.
361        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        // Construct a row with two EF values: [1, 2] and [10, 11].
370        // Verify that row_slice() correctly returns a flat &[F] of base field values.
371        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        // Edge case: test behavior on empty matrix.
384        // Expect zero width and height in the flattened view.
385        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        // Create a row with three EF values, each with offset base coefficients:
395        // [0,1], [10,11], [20,21] -> flattened row should be [0,1,10,11,20,21].
396        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); // 1 row
399        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        // Construct a 2-column, 2-row matrix of EF values, with varying offsets per row.
409        // Row 0: [0,1], [10,11]; Row 1: [20,21], [30,31].
410        // Verify that the flattening preserves row structure and ordering.
411        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        // Build 1 row with 3 EF elements:
425        // - ef(0)   = [0, 1]
426        // - ef(10)  = [10, 11]
427        // - ef(20)  = [20, 21]
428        //
429        // The flattened row should yield:
430        // [0, 1, 10, 11, 20, 21] as base field elements (F)
431        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); // 1 row, 3 EF elements
433        let flat = FlatMatrixView::<F, EF, _>::new(matrix);
434
435        let mut row_iter = flat.row(0).unwrap().into_iter();
436
437        // Expected flattened result
438        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        // Iterator should now be exhausted
445        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}