Skip to main content

p3_matrix/
lib.rs

1#![doc = include_str!("../README.md")]
2#![no_std]
3
4extern crate alloc;
5
6use alloc::vec::Vec;
7use core::fmt::{Debug, Display, Formatter};
8use core::ops::Deref;
9
10use itertools::Itertools;
11use p3_field::{
12    BasedVectorSpace, ExtensionField, Field, FieldArray, PackedField, PackedFieldExtension,
13    PackedValue, PrimeCharacteristicRing,
14};
15use p3_maybe_rayon::prelude::*;
16use strided::{VerticallyStridedMatrixView, VerticallyStridedRowIndexMap};
17use tracing::instrument;
18
19use crate::dense::RowMajorMatrix;
20
21pub mod bitrev;
22pub mod dense;
23pub mod extension;
24pub mod horizontally_truncated;
25pub mod interpolation;
26pub mod row_index_mapped;
27pub mod stack;
28pub mod strided;
29pub mod util;
30
31/// A simple struct representing the shape of a matrix.
32///
33/// The `Dimensions` type stores the number of columns (`width`) and rows (`height`)
34/// of a matrix. It is commonly used for querying and displaying matrix shapes.
35#[derive(Copy, Clone, PartialEq, Eq)]
36pub struct Dimensions {
37    /// Number of columns in the matrix.
38    pub width: usize,
39    /// Number of rows in the matrix.
40    pub height: usize,
41}
42
43impl Debug for Dimensions {
44    fn fmt(&self, f: &mut Formatter<'_>) -> core::fmt::Result {
45        write!(f, "{}x{}", self.width, self.height)
46    }
47}
48
49impl Display for Dimensions {
50    fn fmt(&self, f: &mut Formatter<'_>) -> core::fmt::Result {
51        write!(f, "{}x{}", self.width, self.height)
52    }
53}
54
55/// A generic trait for two-dimensional matrix-like data structures.
56///
57/// The `Matrix` trait provides a uniform interface for accessing rows, elements,
58/// and computing with matrices in both sequential and parallel contexts. It supports
59/// packing strategies for SIMD optimizations and interaction with extension fields.
60pub trait Matrix<T: Send + Sync + Clone>: Send + Sync {
61    /// Returns the number of columns in the matrix.
62    fn width(&self) -> usize;
63
64    /// Returns the number of rows in the matrix.
65    fn height(&self) -> usize;
66
67    /// Returns the dimensions (width, height) of the matrix.
68    fn dimensions(&self) -> Dimensions {
69        Dimensions {
70            width: self.width(),
71            height: self.height(),
72        }
73    }
74
75    // The methods:
76    // get, get_unchecked, row, row_unchecked, row_subseq_unchecked, row_slice, row_slice_unchecked, row_subslice_unchecked
77    // are all defined in a circular manner so you only need to implement a subset of them.
78    // In particular is is enough to implement just one of: row_unchecked, row_subseq_unchecked
79    //
80    // That being said, most implementations will want to implement several methods for performance reasons.
81
82    /// Returns the element at the given row and column.
83    ///
84    /// Returns `None` if either `r >= height()` or `c >= width()`.
85    #[inline]
86    fn get(&self, r: usize, c: usize) -> Option<T> {
87        (r < self.height() && c < self.width()).then(|| unsafe {
88            // Safety: Clearly `r < self.height()` and `c < self.width()`.
89            self.get_unchecked(r, c)
90        })
91    }
92
93    /// Returns the element at the given row and column.
94    ///
95    /// For a safe alternative, see [`Self::get`].
96    ///
97    /// # Safety
98    /// The caller must ensure that `r < self.height()` and `c < self.width()`.
99    /// Breaking any of these assumptions is considered undefined behaviour.
100    #[inline]
101    unsafe fn get_unchecked(&self, r: usize, c: usize) -> T {
102        unsafe { self.row_slice_unchecked(r)[c].clone() }
103    }
104
105    /// Returns an iterator over the elements of the `r`-th row.
106    ///
107    /// The iterator will have `self.width()` elements.
108    ///
109    /// Returns `None` if `r >= height()`.
110    #[inline]
111    fn row(
112        &self,
113        r: usize,
114    ) -> Option<impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync>> {
115        (r < self.height()).then(|| unsafe {
116            // Safety: Clearly `r < self.height()`.
117            self.row_unchecked(r)
118        })
119    }
120
121    /// Returns an iterator over the elements of the `r`-th row.
122    ///
123    /// The iterator will have `self.width()` elements.
124    ///
125    /// For a safe alternative, see [`Self::row`].
126    ///
127    /// # Safety
128    /// The caller must ensure that `r < self.height()`.
129    /// Breaking this assumption is considered undefined behaviour.
130    #[inline]
131    unsafe fn row_unchecked(
132        &self,
133        r: usize,
134    ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
135        unsafe { self.row_subseq_unchecked(r, 0, self.width()) }
136    }
137
138    /// Returns an iterator over the elements of the `r`-th row from position `start` to `end`.
139    ///
140    /// When `start = 0` and `end = width()`, this is equivalent to [`Self::row_unchecked`].
141    ///
142    /// For a safe alternative, use [`Self::row`], along with the `skip` and `take` iterator methods.
143    ///
144    /// # Safety
145    /// The caller must ensure that `r < self.height()` and `start <= end <= self.width()`.
146    /// Breaking any of these assumptions is considered undefined behaviour.
147    #[inline]
148    unsafe fn row_subseq_unchecked(
149        &self,
150        r: usize,
151        start: usize,
152        end: usize,
153    ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
154        unsafe {
155            self.row_unchecked(r)
156                .into_iter()
157                .skip(start)
158                .take(end - start)
159        }
160    }
161
162    /// Returns the elements of the `r`-th row as something which can be coerced to a slice.
163    ///
164    /// Returns `None` if `r >= height()`.
165    #[inline]
166    fn row_slice(&self, r: usize) -> Option<impl Deref<Target = [T]>> {
167        (r < self.height()).then(|| unsafe {
168            // Safety: Clearly `r < self.height()`.
169            self.row_slice_unchecked(r)
170        })
171    }
172
173    /// Returns the elements of the `r`-th row as something which can be coerced to a slice.
174    ///
175    /// For a safe alternative, see [`Self::row_slice`].
176    ///
177    /// # Safety
178    /// The caller must ensure that `r < self.height()`.
179    /// Breaking this assumption is considered undefined behaviour.
180    #[inline]
181    unsafe fn row_slice_unchecked(&self, r: usize) -> impl Deref<Target = [T]> {
182        unsafe { self.row_subslice_unchecked(r, 0, self.width()) }
183    }
184
185    /// Returns a subset of elements of the `r`-th row as something which can be coerced to a slice.
186    ///
187    /// When `start = 0` and `end = width()`, this is equivalent to [`Self::row_slice_unchecked`].
188    ///
189    /// For a safe alternative, see [`Self::row_slice`].
190    ///
191    /// # Safety
192    /// The caller must ensure that `r < self.height()` and `start <= end <= self.width()`.
193    /// Breaking any of these assumptions is considered undefined behaviour.
194    #[inline]
195    unsafe fn row_subslice_unchecked(
196        &self,
197        r: usize,
198        start: usize,
199        end: usize,
200    ) -> impl Deref<Target = [T]> {
201        unsafe {
202            self.row_subseq_unchecked(r, start, end)
203                .into_iter()
204                .collect_vec()
205        }
206    }
207
208    /// Returns an iterator over all rows in the matrix.
209    #[inline]
210    fn rows(&self) -> impl Iterator<Item = impl Iterator<Item = T>> + Send + Sync {
211        unsafe {
212            // Safety: `r` always satisfies `r < self.height()`.
213            (0..self.height()).map(move |r| self.row_unchecked(r).into_iter())
214        }
215    }
216
217    /// Returns a parallel iterator over all rows in the matrix.
218    #[inline]
219    fn par_rows(
220        &self,
221    ) -> impl IndexedParallelIterator<Item = impl Iterator<Item = T>> + Send + Sync {
222        unsafe {
223            // Safety: `r` always satisfies `r < self.height()`.
224            (0..self.height())
225                .into_par_iter()
226                .map(move |r| self.row_unchecked(r).into_iter())
227        }
228    }
229
230    /// Collect the elements of the rows `r` through `r + c`. If anything is larger than `self.height()`
231    /// simply wrap around to the beginning of the matrix.
232    fn wrapping_row_slices(&self, r: usize, c: usize) -> Vec<impl Deref<Target = [T]>> {
233        unsafe {
234            // Safety: Thank to the `%`, the rows index is always less than `self.height()`.
235            (0..c)
236                .map(|i| self.row_slice_unchecked((r + i) % self.height()))
237                .collect_vec()
238        }
239    }
240
241    /// Returns an iterator over the first row of the matrix.
242    ///
243    /// Returns None if `height() == 0`.
244    #[inline]
245    fn first_row(
246        &self,
247    ) -> Option<impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync>> {
248        self.row(0)
249    }
250
251    /// Returns an iterator over the last row of the matrix.
252    ///
253    /// Returns None if `height() == 0`.
254    #[inline]
255    fn last_row(
256        &self,
257    ) -> Option<impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync>> {
258        if self.height() == 0 {
259            None
260        } else {
261            // Safety: Clearly `self.height() - 1 < self.height()`.
262            unsafe { Some(self.row_unchecked(self.height() - 1)) }
263        }
264    }
265
266    /// Converts the matrix into a `RowMajorMatrix` by collecting all rows into a single vector.
267    fn to_row_major_matrix(self) -> RowMajorMatrix<T>
268    where
269        Self: Sized,
270        T: Clone,
271    {
272        RowMajorMatrix::new(self.rows().flatten().collect(), self.width())
273    }
274
275    /// Get a packed iterator over the `r`-th row.
276    ///
277    /// If the row length is not divisible by the packing width, the final elements
278    /// are returned as a base iterator with length `<= P::WIDTH - 1`.
279    ///
280    /// # Panics
281    /// Panics if `r >= height()`.
282    fn horizontally_packed_row<'a, P>(
283        &'a self,
284        r: usize,
285    ) -> (
286        impl Iterator<Item = P> + Send + Sync,
287        impl Iterator<Item = T> + Send + Sync,
288    )
289    where
290        P: PackedValue<Value = T>,
291        T: Clone + 'a,
292    {
293        assert!(r < self.height(), "Row index out of bounds.");
294        let num_packed = self.width() / P::WIDTH;
295        unsafe {
296            // Safety: We have already checked that `r < height()`.
297            let mut iter = self
298                .row_subseq_unchecked(r, 0, num_packed * P::WIDTH)
299                .into_iter();
300
301            // array::from_fn is guaranteed to always call in order.
302            let packed =
303                (0..num_packed).map(move |_| P::from_fn(|_| iter.next().unwrap_unchecked()));
304
305            let sfx = self
306                .row_subseq_unchecked(r, num_packed * P::WIDTH, self.width())
307                .into_iter();
308            (packed, sfx)
309        }
310    }
311
312    /// Get a packed iterator over the `r`-th row.
313    ///
314    /// If the row length is not divisible by the packing width, the final entry will be zero-padded.
315    ///
316    /// # Panics
317    /// Panics if `r >= height()`.
318    fn padded_horizontally_packed_row<'a, P>(
319        &'a self,
320        r: usize,
321    ) -> impl Iterator<Item = P> + Send + Sync
322    where
323        P: PackedValue<Value = T>,
324        T: Clone + Default + 'a,
325    {
326        let mut row_iter = self.row(r).expect("Row index out of bounds.").into_iter();
327        let num_elems = self.width().div_ceil(P::WIDTH);
328        // array::from_fn is guaranteed to always call in order.
329        (0..num_elems).map(move |_| P::from_fn(|_| row_iter.next().unwrap_or_default()))
330    }
331
332    /// Get a parallel iterator over all packed rows of the matrix.
333    ///
334    /// If the matrix width is not divisible by the packing width, the final elements
335    /// of each row are returned as a base iterator with length `<= P::WIDTH - 1`.
336    fn par_horizontally_packed_rows<'a, P>(
337        &'a self,
338    ) -> impl IndexedParallelIterator<
339        Item = (
340            impl Iterator<Item = P> + Send + Sync,
341            impl Iterator<Item = T> + Send + Sync,
342        ),
343    >
344    where
345        P: PackedValue<Value = T>,
346        T: Clone + 'a,
347    {
348        (0..self.height())
349            .into_par_iter()
350            .map(|r| self.horizontally_packed_row(r))
351    }
352
353    /// Get a parallel iterator over all packed rows of the matrix.
354    ///
355    /// If the matrix width is not divisible by the packing width, the final entry of each row will be zero-padded.
356    fn par_padded_horizontally_packed_rows<'a, P>(
357        &'a self,
358    ) -> impl IndexedParallelIterator<Item = impl Iterator<Item = P> + Send + Sync>
359    where
360        P: PackedValue<Value = T>,
361        T: Clone + Default + 'a,
362    {
363        (0..self.height())
364            .into_par_iter()
365            .map(|r| self.padded_horizontally_packed_row(r))
366    }
367
368    /// Pack together a collection of adjacent rows from the matrix.
369    ///
370    /// Returns an iterator whose i'th element is packing of the i'th element of the
371    /// rows r through r + P::WIDTH - 1. If we exceed the height of the matrix,
372    /// wrap around and include initial rows.
373    #[inline]
374    fn vertically_packed_row<P>(&self, r: usize) -> impl Iterator<Item = P>
375    where
376        T: Copy,
377        P: PackedValue<Value = T>,
378    {
379        // Precompute row slices once to minimize redundant calls and improve performance.
380        let rows = self.wrapping_row_slices(r, P::WIDTH);
381
382        // Using precomputed rows avoids repeatedly calling `row_slice`, which is costly.
383        (0..self.width()).map(move |c| P::from_fn(|i| rows[i][c]))
384    }
385
386    /// Pack together a collection of rows and "next" rows from the matrix.
387    ///
388    /// Returns a vector corresponding to 2 packed rows. The i'th element of the first
389    /// row contains the packing of the i'th element of the rows r through r + P::WIDTH - 1.
390    /// The i'th element of the second row contains the packing of the i'th element of the
391    /// rows r + step through r + step + P::WIDTH - 1. If at some point we exceed the
392    /// height of the matrix, wrap around and include initial rows.
393    #[inline]
394    fn vertically_packed_row_pair<P>(&self, r: usize, step: usize) -> Vec<P>
395    where
396        T: Copy,
397        P: PackedValue<Value = T>,
398    {
399        // Whilst it would appear that this can be replaced by two calls to vertically_packed_row
400        // tests seem to indicate that combining them in the same function is slightly faster.
401        // It's probably allowing the compiler to make some optimizations on the fly.
402
403        let rows = self.wrapping_row_slices(r, P::WIDTH);
404        let next_rows = self.wrapping_row_slices(r + step, P::WIDTH);
405
406        (0..self.width())
407            .map(|c| P::from_fn(|i| rows[i][c]))
408            .chain((0..self.width()).map(|c| P::from_fn(|i| next_rows[i][c])))
409            .collect_vec()
410    }
411
412    /// Returns a view over a vertically strided submatrix.
413    ///
414    /// The view selects rows using `r = offset + i * stride` for each `i`.
415    fn vertically_strided(self, stride: usize, offset: usize) -> VerticallyStridedMatrixView<Self>
416    where
417        Self: Sized,
418    {
419        VerticallyStridedRowIndexMap::new_view(self, stride, offset)
420    }
421
422    /// Compute Mᵀv, aka premultiply this matrix by the given vector,
423    /// aka scale each row by the corresponding entry in `v` and take the sum across rows.
424    /// `v` can be a vector of extension elements.
425    #[instrument(level = "debug", skip_all, fields(dims = %self.dimensions()))]
426    fn columnwise_dot_product<EF>(&self, v: &[EF]) -> Vec<EF>
427    where
428        T: Field,
429        EF: ExtensionField<T>,
430    {
431        // Below this many total elements, the rayon fork-join and SIMD-packing machinery
432        // costs more than the dot product itself; fall back to a plain scalar accumulation.
433        // Gating on total elements (rather than height alone) also covers wide-but-short
434        // matrices, where a per-row cost proportional to width still adds up.
435        const SMALL_ELEMS: usize = 256;
436        if self.height().saturating_mul(self.width()) <= SMALL_ELEMS {
437            let mut acc = EF::zero_vec(self.width());
438            for (row, &scale) in self.rows().zip(v) {
439                for (l, r) in acc.iter_mut().zip(row) {
440                    *l += scale * r;
441                }
442            }
443            return acc;
444        }
445
446        let packed_width = self.width().div_ceil(T::Packing::WIDTH);
447
448        let packed_result = self
449            .par_padded_horizontally_packed_rows::<T::Packing>()
450            .zip(v)
451            .par_fold_reduce(
452                || EF::ExtensionPacking::zero_vec(packed_width),
453                |mut acc, (row, &scale)| {
454                    let scale: EF::ExtensionPacking = scale.into();
455                    acc.iter_mut().zip(row).for_each(|(l, r)| *l += scale * r);
456                    acc
457                },
458                |mut acc_l, acc_r| {
459                    acc_l.iter_mut().zip(&acc_r).for_each(|(l, r)| *l += *r);
460                    acc_l
461                },
462            );
463
464        EF::ExtensionPacking::to_ext_iter(packed_result)
465            .take(self.width())
466            .collect()
467    }
468
469    /// Compute Mᵀ · [v₀, v₁, ..., vₙ₋₁] for N weight vectors simultaneously.
470    ///
471    /// Computes `result[col][j] = Σᵣ M[r, col] · vⱼ[r]` for all columns and all j ∈ [0, N).
472    ///
473    /// Batching N weight vectors reduces memory bandwidth: each matrix row is loaded once
474    /// instead of N times. Uses SIMD packing (width W) to process W columns in parallel.
475    #[instrument(level = "debug", skip_all, fields(dims = %self.dimensions()))]
476    fn columnwise_dot_product_batched<EF, const N: usize>(
477        &self,
478        vs: &[FieldArray<EF, N>],
479    ) -> Vec<FieldArray<EF, N>>
480    where
481        T: Field,
482        EF: ExtensionField<T>,
483    {
484        let packed_width = self.width().div_ceil(T::Packing::WIDTH);
485        let height = self.height().min(vs.len());
486
487        // Split the rows into a bounded number of contiguous chunks; each task runs the
488        // field's columnwise kernel serially over its chunk (letting it defer modular
489        // reductions across rows) and the per-task accumulators are summed at the end.
490        let num_chunks = (4 * current_num_threads()).clamp(1, height.max(1));
491        let chunk_rows = height.div_ceil(num_chunks);
492
493        let packed_results: Vec<EF::ExtensionPacking> =
494            (0..num_chunks).into_par_iter().par_fold_reduce(
495                || EF::ExtensionPacking::zero_vec(packed_width * N),
496                |mut acc, chunk| {
497                    let rows = chunk * chunk_rows..((chunk + 1) * chunk_rows).min(height);
498                    T::batched_columnwise_dot_product::<EF, _, _, N>(
499                        &mut acc,
500                        rows.map(|r| {
501                            (
502                                self.padded_horizontally_packed_row::<T::Packing>(r),
503                                vs[r].0,
504                            )
505                        }),
506                    );
507                    acc
508                },
509                |mut acc_l, acc_r| {
510                    acc_l.iter_mut().zip(&acc_r).for_each(|(lj, rj)| *lj += *rj);
511                    acc_l
512                },
513            );
514
515        // Unpack: chunk[j].lane(i) → result[c·W + i][j] for column batch c
516        packed_results
517            .chunks(N)
518            .flat_map(|chunk| {
519                (0..T::Packing::WIDTH)
520                    .map(move |lane| FieldArray::from_fn(|j| chunk[j].extract(lane)))
521            })
522            .take(self.width())
523            .collect()
524    }
525
526    /// Compute the matrix vector product `M . vec`, aka take the dot product of each
527    /// row of `M` by `vec`. If the length of `vec` is longer than the width of `M`,
528    /// `vec` is truncated to the first `width()` elements.
529    ///
530    /// We make use of `PackedFieldExtension` to speed up computations. Thus `vec` is passed in as
531    /// a slice of `PackedFieldExtension` elements.
532    ///
533    /// # Panics
534    /// This function panics if the length of `vec` is less than `self.width().div_ceil(T::Packing::WIDTH)`.
535    fn rowwise_packed_dot_product<EF>(
536        &self,
537        vec: &[EF::ExtensionPacking],
538    ) -> impl IndexedParallelIterator<Item = EF>
539    where
540        T: Field,
541        EF: ExtensionField<T>,
542    {
543        // The length of a `padded_horizontally_packed_row` is `self.width().div_ceil(T::Packing::WIDTH)`.
544        assert!(vec.len() >= self.width().div_ceil(T::Packing::WIDTH));
545
546        // Instead of creating N intermediate ExtPacking products and summing them,
547        // we track D separate BasePacking accumulators (one per extension coefficient).
548        self.par_padded_horizontally_packed_rows::<T::Packing>()
549            .map(move |row_packed| {
550                // Get the extension dimension from the first vec element's coefficients
551                let d = <EF::ExtensionPacking as BasedVectorSpace<T::Packing>>::DIMENSION;
552
553                // Accumulate coefficient-wise: for each (v, r) pair, acc[i] += v.coefficient(i) * r
554                let coeff_accs = T::Packing::coeffwise_dot_product(
555                    d,
556                    vec.iter()
557                        .zip(row_packed)
558                        .map(|(v, r)| (v.as_basis_coefficients_slice(), r)),
559                );
560
561                // Construct the result ExtPacking from the accumulators and sum the coefficients.
562                let packed_result =
563                    EF::ExtensionPacking::from_basis_coefficients_fn(|i| coeff_accs[i]);
564                EF::ExtensionPacking::to_ext_iter([packed_result]).sum()
565            })
566    }
567}
568
569#[cfg(test)]
570mod tests {
571    use alloc::vec::Vec;
572    use alloc::{format, vec};
573
574    use itertools::izip;
575    use p3_baby_bear::BabyBear;
576    use p3_field::PrimeCharacteristicRing;
577    use p3_field::extension::BinomialExtensionField;
578    use rand::SeedableRng;
579    use rand::rngs::SmallRng;
580
581    use super::*;
582
583    #[test]
584    fn test_columnwise_dot_product() {
585        type F = BabyBear;
586        type EF = BinomialExtensionField<BabyBear, 4>;
587
588        let mut rng = SmallRng::seed_from_u64(1);
589        let m = RowMajorMatrix::<F>::rand(&mut rng, 1 << 8, 1 << 4);
590        let v = RowMajorMatrix::<EF>::rand(&mut rng, 1 << 8, 1).values;
591
592        let mut expected = EF::zero_vec(m.width());
593        for (row, &scale) in izip!(m.rows(), &v) {
594            for (l, r) in izip!(&mut expected, row) {
595                *l += scale * r;
596            }
597        }
598
599        assert_eq!(m.columnwise_dot_product(&v), expected);
600    }
601
602    #[test]
603    fn test_columnwise_dot_product_small_height() {
604        type F = BabyBear;
605        type EF = BinomialExtensionField<BabyBear, 4>;
606
607        let mut rng = SmallRng::seed_from_u64(2);
608
609        // Cover heights below, at, and just above the small-height serial threshold.
610        for height in [0, 1, 3, 16, 17] {
611            let m = RowMajorMatrix::<F>::rand(&mut rng, height, 1 << 4);
612            let v = RowMajorMatrix::<EF>::rand(&mut rng, height, 1).values;
613
614            let mut expected = EF::zero_vec(m.width());
615            for (row, &scale) in izip!(m.rows(), &v) {
616                for (l, r) in izip!(&mut expected, row) {
617                    *l += scale * r;
618                }
619            }
620
621            assert_eq!(m.columnwise_dot_product(&v), expected, "height = {height}");
622        }
623    }
624
625    #[test]
626    fn test_columnwise_dot_product_batched() {
627        type F = BabyBear;
628        type EF = BinomialExtensionField<BabyBear, 4>;
629
630        let mut rng = SmallRng::seed_from_u64(1);
631        let m = RowMajorMatrix::<F>::rand(&mut rng, 1 << 8, 1 << 4);
632        let v1 = RowMajorMatrix::<EF>::rand(&mut rng, 1 << 8, 1).values;
633        let v2 = RowMajorMatrix::<EF>::rand(&mut rng, 1 << 8, 1).values;
634
635        // Compute expected via two separate calls
636        let expected1 = m.columnwise_dot_product(&v1);
637        let expected2 = m.columnwise_dot_product(&v2);
638
639        // Compute via batched call - returns Vec<[EF; 2]> where result[col] = [dot1, dot2]
640        let vs: Vec<FieldArray<EF, 2>> = v1
641            .into_iter()
642            .zip(v2)
643            .map(|(a, b)| FieldArray([a, b]))
644            .collect();
645        let results = m.columnwise_dot_product_batched::<EF, 2>(&vs);
646
647        // Extract each point's results
648        let result1: Vec<EF> = results.iter().map(|r| r[0]).collect();
649        let result2: Vec<EF> = results.iter().map(|r| r[1]).collect();
650
651        assert_eq!(result1, expected1);
652        assert_eq!(result2, expected2);
653    }
654
655    // Mock implementation for testing purposes
656    struct MockMatrix {
657        data: Vec<Vec<u32>>,
658        width: usize,
659        height: usize,
660    }
661
662    impl Matrix<u32> for MockMatrix {
663        fn width(&self) -> usize {
664            self.width
665        }
666
667        fn height(&self) -> usize {
668            self.height
669        }
670
671        unsafe fn row_unchecked(
672            &self,
673            r: usize,
674        ) -> impl IntoIterator<Item = u32, IntoIter = impl Iterator<Item = u32> + Send + Sync>
675        {
676            // Just a mock implementation so we just do the easy safe thing.
677            self.data[r].clone()
678        }
679    }
680
681    #[test]
682    fn test_dimensions() {
683        let dims = Dimensions {
684            width: 3,
685            height: 5,
686        };
687        assert_eq!(dims.width, 3);
688        assert_eq!(dims.height, 5);
689        assert_eq!(format!("{dims:?}"), "3x5");
690        assert_eq!(format!("{dims}"), "3x5");
691    }
692
693    #[test]
694    fn test_mock_matrix_dimensions() {
695        let matrix = MockMatrix {
696            data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
697            width: 3,
698            height: 3,
699        };
700        assert_eq!(matrix.width(), 3);
701        assert_eq!(matrix.height(), 3);
702        assert_eq!(
703            matrix.dimensions(),
704            Dimensions {
705                width: 3,
706                height: 3
707            }
708        );
709    }
710
711    #[test]
712    fn test_first_row() {
713        let matrix = MockMatrix {
714            data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
715            width: 3,
716            height: 3,
717        };
718        let mut first_row = matrix.first_row().unwrap().into_iter();
719        assert_eq!(first_row.next(), Some(1));
720        assert_eq!(first_row.next(), Some(2));
721        assert_eq!(first_row.next(), Some(3));
722    }
723
724    #[test]
725    fn test_last_row() {
726        let matrix = MockMatrix {
727            data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
728            width: 3,
729            height: 3,
730        };
731        let mut last_row = matrix.last_row().unwrap().into_iter();
732        assert_eq!(last_row.next(), Some(7));
733        assert_eq!(last_row.next(), Some(8));
734        assert_eq!(last_row.next(), Some(9));
735    }
736
737    #[test]
738    fn test_first_last_row_empty_matrix() {
739        let matrix = MockMatrix {
740            data: vec![],
741            width: 3,
742            height: 0,
743        };
744        let first_row = matrix.first_row();
745        let last_row = matrix.last_row();
746        assert!(first_row.is_none());
747        assert!(last_row.is_none());
748    }
749
750    #[test]
751    fn test_to_row_major_matrix() {
752        let matrix = MockMatrix {
753            data: vec![vec![1, 2], vec![3, 4]],
754            width: 2,
755            height: 2,
756        };
757        let row_major = matrix.to_row_major_matrix();
758        assert_eq!(row_major.values, vec![1, 2, 3, 4]);
759        assert_eq!(row_major.width, 2);
760    }
761
762    #[test]
763    fn test_matrix_get_methods() {
764        let matrix = MockMatrix {
765            data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
766            width: 3,
767            height: 3,
768        };
769        assert_eq!(matrix.get(0, 0), Some(1));
770        assert_eq!(matrix.get(1, 2), Some(6));
771        assert_eq!(matrix.get(2, 1), Some(8));
772
773        unsafe {
774            assert_eq!(matrix.get_unchecked(0, 1), 2);
775            assert_eq!(matrix.get_unchecked(1, 0), 4);
776            assert_eq!(matrix.get_unchecked(2, 2), 9);
777        }
778
779        assert_eq!(matrix.get(3, 0), None); // Height out of bounds
780        assert_eq!(matrix.get(0, 3), None); // Width out of bounds
781    }
782
783    #[test]
784    fn test_matrix_row_methods_iteration() {
785        let matrix = MockMatrix {
786            data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
787            width: 3,
788            height: 3,
789        };
790
791        let mut row_iter = matrix.row(1).unwrap().into_iter();
792        assert_eq!(row_iter.next(), Some(4));
793        assert_eq!(row_iter.next(), Some(5));
794        assert_eq!(row_iter.next(), Some(6));
795        assert_eq!(row_iter.next(), None);
796
797        unsafe {
798            let mut row_iter_unchecked = matrix.row_unchecked(2).into_iter();
799            assert_eq!(row_iter_unchecked.next(), Some(7));
800            assert_eq!(row_iter_unchecked.next(), Some(8));
801            assert_eq!(row_iter_unchecked.next(), Some(9));
802            assert_eq!(row_iter_unchecked.next(), None);
803
804            let mut row_iter_subset = matrix.row_subseq_unchecked(0, 1, 3).into_iter();
805            assert_eq!(row_iter_subset.next(), Some(2));
806            assert_eq!(row_iter_subset.next(), Some(3));
807            assert_eq!(row_iter_subset.next(), None);
808        }
809
810        assert!(matrix.row(3).is_none()); // Height out of bounds
811    }
812
813    #[test]
814    fn test_row_slice_methods() {
815        let matrix = MockMatrix {
816            data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
817            width: 3,
818            height: 3,
819        };
820        let row_slice = matrix.row_slice(1).unwrap();
821        assert_eq!(*row_slice, [4, 5, 6]);
822        unsafe {
823            let row_slice_unchecked = matrix.row_slice_unchecked(2);
824            assert_eq!(*row_slice_unchecked, [7, 8, 9]);
825
826            let row_subslice = matrix.row_subslice_unchecked(0, 1, 2);
827            assert_eq!(*row_subslice, [2]);
828        }
829
830        assert!(matrix.row_slice(3).is_none()); // Height out of bounds
831    }
832
833    #[test]
834    fn test_matrix_rows() {
835        let matrix = MockMatrix {
836            data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
837            width: 3,
838            height: 3,
839        };
840
841        let all_rows: Vec<Vec<u32>> = matrix.rows().map(|row| row.collect()).collect();
842        assert_eq!(all_rows, vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]]);
843    }
844
845    #[test]
846    fn test_rowwise_packed_dot_product() {
847        use p3_field::PackedFieldExtension;
848
849        type F = BabyBear;
850        type EF = BinomialExtensionField<BabyBear, 4>;
851        type PF = <F as p3_field::Field>::Packing;
852        type EFPacked = <EF as p3_field::ExtensionField<F>>::ExtensionPacking;
853
854        let mut rng = SmallRng::seed_from_u64(42);
855
856        // Test with various matrix dimensions to cover edge cases.
857        for (height, width) in [(32, 16), (64, 128), (128, 17), (256, 255)] {
858            let m = RowMajorMatrix::<F>::rand(&mut rng, height, width);
859            let v = RowMajorMatrix::<EF>::rand(&mut rng, width, 1).values;
860
861            // Compute expected result naively: for each row, compute dot product with v.
862            let expected: Vec<EF> = m
863                .rows()
864                .map(|row| {
865                    row.into_iter()
866                        .zip(v.iter())
867                        .map(|(r, &ve)| ve * r)
868                        .sum::<EF>()
869                })
870                .collect();
871
872            // Pack the vector for the optimized function.
873            let packed_v: Vec<EFPacked> = v
874                .chunks(<PF as PackedValue>::WIDTH)
875                .map(|chunk| {
876                    let mut padded = EF::zero_vec(<PF as PackedValue>::WIDTH);
877                    padded[..chunk.len()].copy_from_slice(chunk);
878                    EFPacked::from_ext_slice(&padded)
879                })
880                .collect();
881
882            // Compute using the optimized function.
883            let result: Vec<EF> = m.rowwise_packed_dot_product::<EF>(&packed_v).collect();
884
885            assert_eq!(result, expected, "Mismatch for matrix {}x{}", height, width);
886        }
887    }
888}