1use alloc::vec::Vec;
2use core::ops::Deref;
3
4use p3_field::PackedValue;
5
6use crate::Matrix;
7use crate::dense::RowMajorMatrix;
8
9const MAX_RESOLVED_LANES: usize = 8;
14
15pub trait RowIndexMap: Send + Sync {
20 fn height(&self) -> usize;
22
23 fn map_row_index(&self, r: usize) -> usize;
30
31 fn to_row_major_matrix<T: Clone + Send + Sync, Inner: Matrix<T>>(
36 &self,
37 inner: Inner,
38 ) -> RowMajorMatrix<T> {
39 RowMajorMatrix::new(
40 unsafe {
41 (0..self.height())
43 .flat_map(|r| inner.row_unchecked(self.map_row_index(r)))
44 .collect()
45 },
46 inner.width(),
47 )
48 }
49}
50
51#[derive(Copy, Clone, Debug)]
56pub struct RowIndexMappedView<IndexMap, Inner> {
57 pub index_map: IndexMap,
59 pub inner: Inner,
61}
62
63impl<T: Send + Sync + Clone, IndexMap: RowIndexMap, Inner: Matrix<T>> Matrix<T>
64 for RowIndexMappedView<IndexMap, Inner>
65{
66 fn width(&self) -> usize {
67 self.inner.width()
68 }
69
70 fn height(&self) -> usize {
71 self.index_map.height()
72 }
73
74 unsafe fn get_unchecked(&self, r: usize, c: usize) -> T {
75 unsafe {
76 self.inner.get_unchecked(self.index_map.map_row_index(r), c)
78 }
79 }
80
81 unsafe fn row_unchecked(
82 &self,
83 r: usize,
84 ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
85 unsafe {
86 self.inner.row_unchecked(self.index_map.map_row_index(r))
88 }
89 }
90
91 unsafe fn row_subseq_unchecked(
92 &self,
93 r: usize,
94 start: usize,
95 end: usize,
96 ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
97 unsafe {
98 self.inner
100 .row_subseq_unchecked(self.index_map.map_row_index(r), start, end)
101 }
102 }
103
104 unsafe fn row_slice_unchecked(&self, r: usize) -> impl Deref<Target = [T]> {
105 unsafe {
106 self.inner
108 .row_slice_unchecked(self.index_map.map_row_index(r))
109 }
110 }
111
112 unsafe fn row_subslice_unchecked(
113 &self,
114 r: usize,
115 start: usize,
116 end: usize,
117 ) -> impl Deref<Target = [T]> {
118 unsafe {
119 self.inner
121 .row_subslice_unchecked(self.index_map.map_row_index(r), start, end)
122 }
123 }
124
125 fn to_row_major_matrix(self) -> RowMajorMatrix<T>
126 where
127 Self: Sized,
128 T: Clone,
129 {
130 self.index_map.to_row_major_matrix(self.inner)
132 }
133
134 fn horizontally_packed_row<'a, P>(
135 &'a self,
136 r: usize,
137 ) -> (
138 impl Iterator<Item = P> + Send + Sync,
139 impl Iterator<Item = T> + Send + Sync,
140 )
141 where
142 P: PackedValue<Value = T>,
143 T: Clone + 'a,
144 {
145 self.inner
146 .horizontally_packed_row(self.index_map.map_row_index(r))
147 }
148
149 fn padded_horizontally_packed_row<'a, P>(
150 &'a self,
151 r: usize,
152 ) -> impl Iterator<Item = P> + Send + Sync
153 where
154 P: PackedValue<Value = T>,
155 T: Clone + Default + 'a,
156 {
157 self.inner
158 .padded_horizontally_packed_row(self.index_map.map_row_index(r))
159 }
160
161 #[inline]
162 fn vertically_packed_row<P>(&self, r: usize) -> impl Iterator<Item = P>
163 where
164 T: Copy,
165 P: PackedValue<Value = T>,
166 {
167 let height = self.height();
168 let width = self.width();
169 let no_wrap = P::WIDTH != 1 && r + P::WIDTH <= height;
176 let lane_rows: [Option<_>; MAX_RESOLVED_LANES] = core::array::from_fn(|i| {
177 (P::WIDTH <= MAX_RESOLVED_LANES && i < P::WIDTH && width != 0).then(|| {
178 let lane_row = if no_wrap { r + i } else { (r + i) % height };
179 let row = unsafe { self.row_slice_unchecked(lane_row) };
181 assert_eq!(
182 row.len(),
183 width,
184 "inner row length differs from the matrix width"
185 );
186 row
187 })
188 });
189 (0..width).map(move |c| {
190 if P::WIDTH <= MAX_RESOLVED_LANES {
191 P::from_fn(|i| unsafe {
196 *lane_rows
197 .get_unchecked(i)
198 .as_deref()
199 .unwrap_unchecked()
200 .get_unchecked(c)
201 })
202 } else if no_wrap {
203 P::from_fn(|i| unsafe { self.get_unchecked(r + i, c) })
205 } else {
206 P::from_fn(|i| unsafe { self.get_unchecked((r + i) % height, c) })
208 }
209 })
210 }
211
212 #[inline]
213 fn vertically_packed_row_pair<P>(&self, r: usize, step: usize) -> Vec<P>
214 where
215 T: Copy,
216 P: PackedValue<Value = T>,
217 {
218 self.vertically_packed_row::<P>(r)
219 .chain(self.vertically_packed_row::<P>(r + step))
220 .collect()
221 }
222}
223
224#[cfg(test)]
225mod tests {
226 use alloc::vec;
227 use alloc::vec::Vec;
228 use core::fmt::Debug;
229
230 use itertools::Itertools;
231 use p3_baby_bear::BabyBear;
232 use p3_field::{Field, FieldArray, Vectorized};
233 use p3_goldilocks::Goldilocks;
234 use rand::SeedableRng;
235 use rand::distr::{Distribution, StandardUniform};
236 use rand::rngs::SmallRng;
237
238 use super::*;
239 use crate::bitrev::BitReversibleMatrix;
240 use crate::dense::RowMajorMatrix;
241 use crate::stack::HorizontalPair;
242
243 struct IdentityMap(usize);
245
246 impl RowIndexMap for IdentityMap {
247 fn height(&self) -> usize {
248 self.0
249 }
250
251 fn map_row_index(&self, r: usize) -> usize {
252 r
253 }
254 }
255
256 struct ReverseMap(usize);
258
259 impl RowIndexMap for ReverseMap {
260 fn height(&self) -> usize {
261 self.0
262 }
263
264 fn map_row_index(&self, r: usize) -> usize {
265 self.0 - 1 - r
266 }
267 }
268
269 struct ConstantMap;
271
272 impl RowIndexMap for ConstantMap {
273 fn height(&self) -> usize {
274 1
275 }
276
277 fn map_row_index(&self, _r: usize) -> usize {
278 0
279 }
280 }
281
282 #[test]
283 fn test_identity_row_index_map() {
284 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
289
290 let mapped_view = RowIndexMappedView {
292 index_map: IdentityMap(inner.height()),
293 inner,
294 };
295
296 assert_eq!(mapped_view.height(), 2);
298 assert_eq!(mapped_view.width(), 3);
299
300 assert_eq!(mapped_view.get(0, 0).unwrap(), 1);
302 assert_eq!(mapped_view.get(1, 2).unwrap(), 6);
303
304 unsafe {
305 assert_eq!(mapped_view.get_unchecked(0, 1), 2);
306 assert_eq!(mapped_view.get_unchecked(1, 0), 4);
307 }
308
309 let rows: Vec<Vec<_>> = mapped_view.rows().map(|row| row.collect()).collect();
311 assert_eq!(rows, vec![vec![1, 2, 3], vec![4, 5, 6]]);
312
313 let dense = mapped_view.to_row_major_matrix();
315 assert_eq!(dense.values, vec![1, 2, 3, 4, 5, 6]);
316 }
317
318 #[test]
319 fn test_reverse_row_index_map() {
320 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
325
326 let mapped_view = RowIndexMappedView {
328 index_map: ReverseMap(inner.height()),
329 inner,
330 };
331
332 assert_eq!(mapped_view.height(), 2);
334 assert_eq!(mapped_view.width(), 3);
335
336 assert_eq!(mapped_view.get(0, 0).unwrap(), 4);
338 assert_eq!(mapped_view.get(1, 2).unwrap(), 3);
340
341 unsafe {
342 assert_eq!(mapped_view.get_unchecked(0, 1), 5);
343 assert_eq!(mapped_view.get_unchecked(1, 0), 1);
344 }
345
346 let rows: Vec<Vec<_>> = mapped_view.rows().map(|row| row.collect()).collect();
348 assert_eq!(rows, vec![vec![4, 5, 6], vec![1, 2, 3]]);
349
350 let dense = mapped_view.to_row_major_matrix();
352 assert_eq!(dense.values, vec![4, 5, 6, 1, 2, 3]);
353 }
354
355 #[test]
356 fn test_horizontally_packed_row() {
357 type Packed = FieldArray<BabyBear, 2>;
359
360 let inner = RowMajorMatrix::new(
365 vec![
366 BabyBear::new(1),
367 BabyBear::new(2),
368 BabyBear::new(3),
369 BabyBear::new(4),
370 ],
371 2,
372 );
373
374 let mapped_view = RowIndexMappedView {
376 index_map: ReverseMap(inner.height()),
377 inner,
378 };
379
380 let (packed_iter, mut suffix_iter) = mapped_view.horizontally_packed_row::<Packed>(0);
382
383 let packed: Vec<_> = packed_iter.collect();
385
386 assert_eq!(
388 packed,
389 &[Packed::from([BabyBear::new(3), BabyBear::new(4)])]
390 );
391
392 assert!(suffix_iter.next().is_none());
394 }
395
396 #[test]
397 fn test_padded_horizontally_packed_row() {
398 type Packed = FieldArray<BabyBear, 3>;
400
401 let inner = RowMajorMatrix::new(
405 vec![
406 BabyBear::new(1),
407 BabyBear::new(2),
408 BabyBear::new(3),
409 BabyBear::new(4),
410 ],
411 2,
412 );
413
414 let mapped_view = RowIndexMappedView {
416 index_map: IdentityMap(inner.height()),
417 inner,
418 };
419
420 let packed: Vec<_> = mapped_view
422 .padded_horizontally_packed_row::<Packed>(1)
423 .collect();
424
425 assert_eq!(
427 packed,
428 vec![Packed::from([
429 BabyBear::new(3),
430 BabyBear::new(4),
431 BabyBear::new(0),
432 ])]
433 );
434 }
435
436 #[test]
437 fn test_vertically_packed_row() {
438 type Packed = FieldArray<BabyBear, 2>;
440
441 let inner = RowMajorMatrix::new((1..=8).map(BabyBear::new).collect::<Vec<_>>(), 2);
447
448 let mapped_view = RowIndexMappedView {
451 index_map: ReverseMap(inner.height()),
452 inner,
453 };
454
455 let packed: Vec<_> = mapped_view.vertically_packed_row::<Packed>(0).collect();
457 assert_eq!(
458 packed,
459 vec![
460 Packed::from([BabyBear::new(7), BabyBear::new(5)]),
461 Packed::from([BabyBear::new(8), BabyBear::new(6)]),
462 ]
463 );
464
465 let packed: Vec<_> = mapped_view.vertically_packed_row::<Packed>(3).collect();
467 assert_eq!(
468 packed,
469 vec![
470 Packed::from([BabyBear::new(1), BabyBear::new(7)]),
471 Packed::from([BabyBear::new(2), BabyBear::new(8)]),
472 ]
473 );
474 }
475
476 #[test]
477 fn test_vertically_packed_row_pair() {
478 type Packed = FieldArray<BabyBear, 2>;
480
481 let inner = RowMajorMatrix::new((1..=8).map(BabyBear::new).collect::<Vec<_>>(), 2);
487
488 let mapped_view = RowIndexMappedView {
491 index_map: ReverseMap(inner.height()),
492 inner,
493 };
494
495 let packed = mapped_view.vertically_packed_row_pair::<Packed>(0, 2);
497 assert_eq!(
498 packed,
499 vec![
500 Packed::from([BabyBear::new(7), BabyBear::new(5)]),
501 Packed::from([BabyBear::new(8), BabyBear::new(6)]),
502 Packed::from([BabyBear::new(3), BabyBear::new(1)]),
503 Packed::from([BabyBear::new(4), BabyBear::new(2)]),
504 ]
505 );
506
507 let packed = mapped_view.vertically_packed_row_pair::<Packed>(1, 2);
509 assert_eq!(
510 packed,
511 vec![
512 Packed::from([BabyBear::new(5), BabyBear::new(3)]),
513 Packed::from([BabyBear::new(6), BabyBear::new(4)]),
514 Packed::from([BabyBear::new(1), BabyBear::new(7)]),
515 Packed::from([BabyBear::new(2), BabyBear::new(8)]),
516 ]
517 );
518
519 let packed = mapped_view.vertically_packed_row_pair::<Packed>(3, 2);
521 assert_eq!(
522 packed,
523 vec![
524 Packed::from([BabyBear::new(1), BabyBear::new(7)]),
525 Packed::from([BabyBear::new(2), BabyBear::new(8)]),
526 Packed::from([BabyBear::new(5), BabyBear::new(3)]),
527 Packed::from([BabyBear::new(6), BabyBear::new(4)]),
528 ]
529 );
530 }
531
532 fn scalar_packed_row<T, P, M>(m: &M, r: usize) -> Vec<Vec<T>>
534 where
535 T: Copy + Send + Sync,
536 P: PackedValue<Value = T>,
537 M: Matrix<T>,
538 {
539 let height = m.height();
540 (0..m.width())
541 .map(|c| {
542 (0..P::WIDTH)
543 .map(|i| m.get((r + i) % height, c).unwrap())
544 .collect()
545 })
546 .collect()
547 }
548
549 fn lanes<P: PackedValue>(packed: &[P]) -> Vec<Vec<P::Value>> {
550 packed.iter().map(|p| p.as_slice().to_vec()).collect()
551 }
552
553 fn assert_packing_matches_scalar_gather<T, P, M>(m: &M)
556 where
557 T: Copy + Send + Sync + PartialEq + Debug,
558 P: PackedValue<Value = T>,
559 M: Matrix<T>,
560 {
561 let height = m.height();
562 for r in 0..2 * height + P::WIDTH {
563 let expected = scalar_packed_row::<T, P, M>(m, r);
564 let packed: Vec<P> = m.vertically_packed_row::<P>(r).collect();
565 assert_eq!(lanes(&packed), expected, "r = {r}, height = {height}");
566
567 for step in [0, 1, 2, height] {
568 let mut expected_pair = expected.clone();
569 expected_pair.extend(scalar_packed_row::<T, P, M>(m, r + step));
570 let pair = m.vertically_packed_row_pair::<P>(r, step);
571 assert_eq!(lanes(&pair), expected_pair, "r = {r}, step = {step}");
572 }
573 }
574 }
575
576 fn assert_packings_match_scalar_gather<F, M>(m: &M)
579 where
580 F: Field,
581 M: Matrix<F>,
582 {
583 assert_packing_matches_scalar_gather::<F, F::Packing, _>(m);
584 assert_packing_matches_scalar_gather::<F, F, _>(m);
585 assert_packing_matches_scalar_gather::<F, Vectorized<F, 2>, _>(m);
586 assert_packing_matches_scalar_gather::<F, FieldArray<F, 3>, _>(m);
587 assert_packing_matches_scalar_gather::<F, FieldArray<F, 16>, _>(m);
588 }
589
590 fn check_packings_for<F>()
595 where
596 F: Field,
597 StandardUniform: Distribution<F>,
598 {
599 let mut rng = SmallRng::seed_from_u64(1);
600 for width in [1, 2, 3, 5, 8, 13] {
601 for height in [1, 2, 3, 4, 5, 7, 8, 9, 15, 16, 17] {
602 let view = RowIndexMappedView {
603 index_map: ReverseMap(height),
604 inner: RowMajorMatrix::<F>::rand(&mut rng, height, width),
605 };
606 assert_packings_match_scalar_gather(&view);
607 }
608 for log_height in 0..6 {
609 let view =
610 RowMajorMatrix::<F>::rand(&mut rng, 1 << log_height, width).bit_reverse_rows();
611 assert_packings_match_scalar_gather(&view);
612 }
613 for (inner_height, stride, offset) in [(5, 1, 2), (9, 2, 1), (17, 3, 0), (17, 4, 3)] {
614 let view = RowMajorMatrix::<F>::rand(&mut rng, inner_height, width)
615 .vertically_strided(stride, offset);
616 assert_packings_match_scalar_gather(&view);
617 }
618 for height in [1, 5, 9] {
619 let view = RowIndexMappedView {
620 index_map: ReverseMap(height),
621 inner: HorizontalPair::new(
622 RowMajorMatrix::<F>::rand(&mut rng, height, width),
623 RowMajorMatrix::<F>::rand(&mut rng, height, 2),
624 ),
625 };
626 assert_packings_match_scalar_gather(&view);
627 }
628 }
629 }
630
631 #[test]
632 fn test_vertically_packed_row_matches_scalar_gather_babybear() {
633 check_packings_for::<BabyBear>();
634 }
635
636 #[test]
637 fn test_vertically_packed_row_matches_scalar_gather_goldilocks() {
638 check_packings_for::<Goldilocks>();
639 }
640
641 struct ShortRows(RowMajorMatrix<BabyBear>);
643
644 impl Matrix<BabyBear> for ShortRows {
645 fn width(&self) -> usize {
646 self.0.width()
647 }
648
649 fn height(&self) -> usize {
650 self.0.height()
651 }
652
653 unsafe fn row_subseq_unchecked(
654 &self,
655 r: usize,
656 start: usize,
657 end: usize,
658 ) -> impl IntoIterator<Item = BabyBear, IntoIter = impl Iterator<Item = BabyBear> + Send + Sync>
659 {
660 unsafe {
663 self.0
664 .row_subseq_unchecked(r, start, start.max(end.saturating_sub(1)))
665 }
666 }
667 }
668
669 #[test]
670 #[should_panic(expected = "inner row length differs from the matrix width")]
671 fn test_vertically_packed_row_rejects_short_inner_rows() {
672 let height = 4;
673 let view = RowIndexMappedView {
674 index_map: ReverseMap(height),
675 inner: ShortRows(RowMajorMatrix::new(
676 (0..height as u32 * 3).map(BabyBear::new).collect(),
677 3,
678 )),
679 };
680 let _ = view
681 .vertically_packed_row::<FieldArray<BabyBear, 4>>(0)
682 .collect::<Vec<_>>();
683 }
684
685 #[test]
686 fn test_row_and_row_slice_methods() {
687 let inner = RowMajorMatrix::new(vec![10, 20, 30, 40, 50, 60], 3);
691
692 let mapped_view = RowIndexMappedView {
694 index_map: ReverseMap(inner.height()),
695 inner,
696 };
697
698 assert_eq!(mapped_view.row_slice(0).unwrap().deref(), &[40, 50, 60]); assert_eq!(
701 mapped_view.row(1).unwrap().into_iter().collect_vec(),
702 vec![10, 20, 30]
703 ); unsafe {
706 assert_eq!(
708 mapped_view.row_unchecked(0).into_iter().collect_vec(),
709 vec![40, 50, 60]
710 ); assert_eq!(mapped_view.row_slice_unchecked(1).deref(), &[10, 20, 30]); assert_eq!(
714 mapped_view.row_subslice_unchecked(0, 1, 3).deref(),
715 &[50, 60]
716 ); assert_eq!(
718 mapped_view
719 .row_subseq_unchecked(1, 0, 2)
720 .into_iter()
721 .collect_vec(),
722 vec![10, 20]
723 ); }
725
726 assert!(mapped_view.row(2).is_none()); assert!(mapped_view.row_slice(2).is_none()); }
729
730 #[test]
731 fn test_out_of_bounds_access() {
732 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4], 2);
736
737 let mapped_view = RowIndexMappedView {
739 index_map: IdentityMap(inner.height()),
740 inner,
741 };
742
743 assert_eq!(mapped_view.get(2, 1), None);
745 assert!(mapped_view.row(5).is_none());
746 assert!(mapped_view.row_slice(11).is_none());
747 assert_eq!(mapped_view.get(0, 20), None);
748 }
749
750 #[test]
751 fn test_out_of_bounds_access_with_bad_map() {
752 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4], 4);
756
757 let mapped_view = RowIndexMappedView {
759 index_map: ConstantMap,
760 inner,
761 };
762
763 assert_eq!(mapped_view.get(0, 2), Some(3));
764
765 assert_eq!(mapped_view.get(1, 0), None);
767 assert!(mapped_view.row(1).is_none());
768 assert!(mapped_view.row_slice(1).is_none());
769 }
770}