1use alloc::vec::Vec;
2use core::ops::Deref;
3
4use p3_field::PackedValue;
5
6use crate::Matrix;
7use crate::dense::RowMajorMatrix;
8
9pub trait RowIndexMap: Send + Sync {
14 fn height(&self) -> usize;
16
17 fn map_row_index(&self, r: usize) -> usize;
24
25 fn to_row_major_matrix<T: Clone + Send + Sync, Inner: Matrix<T>>(
30 &self,
31 inner: Inner,
32 ) -> RowMajorMatrix<T> {
33 RowMajorMatrix::new(
34 unsafe {
35 (0..self.height())
37 .flat_map(|r| inner.row_unchecked(self.map_row_index(r)))
38 .collect()
39 },
40 inner.width(),
41 )
42 }
43}
44
45#[derive(Copy, Clone, Debug)]
50pub struct RowIndexMappedView<IndexMap, Inner> {
51 pub index_map: IndexMap,
53 pub inner: Inner,
55}
56
57impl<T: Send + Sync + Clone, IndexMap: RowIndexMap, Inner: Matrix<T>> Matrix<T>
58 for RowIndexMappedView<IndexMap, Inner>
59{
60 fn width(&self) -> usize {
61 self.inner.width()
62 }
63
64 fn height(&self) -> usize {
65 self.index_map.height()
66 }
67
68 unsafe fn get_unchecked(&self, r: usize, c: usize) -> T {
69 unsafe {
70 self.inner.get_unchecked(self.index_map.map_row_index(r), c)
72 }
73 }
74
75 unsafe fn row_unchecked(
76 &self,
77 r: usize,
78 ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
79 unsafe {
80 self.inner.row_unchecked(self.index_map.map_row_index(r))
82 }
83 }
84
85 unsafe fn row_subseq_unchecked(
86 &self,
87 r: usize,
88 start: usize,
89 end: usize,
90 ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
91 unsafe {
92 self.inner
94 .row_subseq_unchecked(self.index_map.map_row_index(r), start, end)
95 }
96 }
97
98 unsafe fn row_slice_unchecked(&self, r: usize) -> impl Deref<Target = [T]> {
99 unsafe {
100 self.inner
102 .row_slice_unchecked(self.index_map.map_row_index(r))
103 }
104 }
105
106 unsafe fn row_subslice_unchecked(
107 &self,
108 r: usize,
109 start: usize,
110 end: usize,
111 ) -> impl Deref<Target = [T]> {
112 unsafe {
113 self.inner
115 .row_subslice_unchecked(self.index_map.map_row_index(r), start, end)
116 }
117 }
118
119 fn to_row_major_matrix(self) -> RowMajorMatrix<T>
120 where
121 Self: Sized,
122 T: Clone,
123 {
124 self.index_map.to_row_major_matrix(self.inner)
126 }
127
128 fn horizontally_packed_row<'a, P>(
129 &'a self,
130 r: usize,
131 ) -> (
132 impl Iterator<Item = P> + Send + Sync,
133 impl Iterator<Item = T> + Send + Sync,
134 )
135 where
136 P: PackedValue<Value = T>,
137 T: Clone + 'a,
138 {
139 self.inner
140 .horizontally_packed_row(self.index_map.map_row_index(r))
141 }
142
143 fn padded_horizontally_packed_row<'a, P>(
144 &'a self,
145 r: usize,
146 ) -> impl Iterator<Item = P> + Send + Sync
147 where
148 P: PackedValue<Value = T>,
149 T: Clone + Default + 'a,
150 {
151 self.inner
152 .padded_horizontally_packed_row(self.index_map.map_row_index(r))
153 }
154
155 #[inline]
156 fn vertically_packed_row<P>(&self, r: usize) -> impl Iterator<Item = P>
157 where
158 T: Copy,
159 P: PackedValue<Value = T>,
160 {
161 let height = self.height();
162 let width = self.width();
163 let no_wrap = P::WIDTH != 1 && r + P::WIDTH <= height;
168 (0..width).map(move |c| {
169 if no_wrap {
170 P::from_fn(|i| unsafe { self.get_unchecked(r + i, c) })
172 } else {
173 P::from_fn(|i| unsafe { self.get_unchecked((r + i) % height, c) })
175 }
176 })
177 }
178
179 #[inline]
180 fn vertically_packed_row_pair<P>(&self, r: usize, step: usize) -> Vec<P>
181 where
182 T: Copy,
183 P: PackedValue<Value = T>,
184 {
185 let height = self.height();
186 let width = self.width();
187 let no_wrap = P::WIDTH != 1 && r + P::WIDTH <= height;
188 let next_no_wrap = P::WIDTH != 1 && r + step + P::WIDTH <= height;
189
190 (0..width)
191 .map(move |c| {
192 if no_wrap {
193 P::from_fn(|i| unsafe { self.get_unchecked(r + i, c) })
195 } else {
196 P::from_fn(|i| unsafe { self.get_unchecked((r + i) % height, c) })
198 }
199 })
200 .chain((0..width).map(move |c| {
201 if next_no_wrap {
202 P::from_fn(|i| unsafe { self.get_unchecked(r + step + i, c) })
204 } else {
205 P::from_fn(|i| unsafe { self.get_unchecked((r + step + i) % height, c) })
207 }
208 }))
209 .collect()
210 }
211}
212
213#[cfg(test)]
214mod tests {
215 use alloc::vec;
216 use alloc::vec::Vec;
217
218 use itertools::Itertools;
219 use p3_baby_bear::BabyBear;
220 use p3_field::FieldArray;
221
222 use super::*;
223 use crate::dense::RowMajorMatrix;
224
225 struct IdentityMap(usize);
227
228 impl RowIndexMap for IdentityMap {
229 fn height(&self) -> usize {
230 self.0
231 }
232
233 fn map_row_index(&self, r: usize) -> usize {
234 r
235 }
236 }
237
238 struct ReverseMap(usize);
240
241 impl RowIndexMap for ReverseMap {
242 fn height(&self) -> usize {
243 self.0
244 }
245
246 fn map_row_index(&self, r: usize) -> usize {
247 self.0 - 1 - r
248 }
249 }
250
251 struct ConstantMap;
253
254 impl RowIndexMap for ConstantMap {
255 fn height(&self) -> usize {
256 1
257 }
258
259 fn map_row_index(&self, _r: usize) -> usize {
260 0
261 }
262 }
263
264 #[test]
265 fn test_identity_row_index_map() {
266 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
271
272 let mapped_view = RowIndexMappedView {
274 index_map: IdentityMap(inner.height()),
275 inner,
276 };
277
278 assert_eq!(mapped_view.height(), 2);
280 assert_eq!(mapped_view.width(), 3);
281
282 assert_eq!(mapped_view.get(0, 0).unwrap(), 1);
284 assert_eq!(mapped_view.get(1, 2).unwrap(), 6);
285
286 unsafe {
287 assert_eq!(mapped_view.get_unchecked(0, 1), 2);
288 assert_eq!(mapped_view.get_unchecked(1, 0), 4);
289 }
290
291 let rows: Vec<Vec<_>> = mapped_view.rows().map(|row| row.collect()).collect();
293 assert_eq!(rows, vec![vec![1, 2, 3], vec![4, 5, 6]]);
294
295 let dense = mapped_view.to_row_major_matrix();
297 assert_eq!(dense.values, vec![1, 2, 3, 4, 5, 6]);
298 }
299
300 #[test]
301 fn test_reverse_row_index_map() {
302 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
307
308 let mapped_view = RowIndexMappedView {
310 index_map: ReverseMap(inner.height()),
311 inner,
312 };
313
314 assert_eq!(mapped_view.height(), 2);
316 assert_eq!(mapped_view.width(), 3);
317
318 assert_eq!(mapped_view.get(0, 0).unwrap(), 4);
320 assert_eq!(mapped_view.get(1, 2).unwrap(), 3);
322
323 unsafe {
324 assert_eq!(mapped_view.get_unchecked(0, 1), 5);
325 assert_eq!(mapped_view.get_unchecked(1, 0), 1);
326 }
327
328 let rows: Vec<Vec<_>> = mapped_view.rows().map(|row| row.collect()).collect();
330 assert_eq!(rows, vec![vec![4, 5, 6], vec![1, 2, 3]]);
331
332 let dense = mapped_view.to_row_major_matrix();
334 assert_eq!(dense.values, vec![4, 5, 6, 1, 2, 3]);
335 }
336
337 #[test]
338 fn test_horizontally_packed_row() {
339 type Packed = FieldArray<BabyBear, 2>;
341
342 let inner = RowMajorMatrix::new(
347 vec![
348 BabyBear::new(1),
349 BabyBear::new(2),
350 BabyBear::new(3),
351 BabyBear::new(4),
352 ],
353 2,
354 );
355
356 let mapped_view = RowIndexMappedView {
358 index_map: ReverseMap(inner.height()),
359 inner,
360 };
361
362 let (packed_iter, mut suffix_iter) = mapped_view.horizontally_packed_row::<Packed>(0);
364
365 let packed: Vec<_> = packed_iter.collect();
367
368 assert_eq!(
370 packed,
371 &[Packed::from([BabyBear::new(3), BabyBear::new(4)])]
372 );
373
374 assert!(suffix_iter.next().is_none());
376 }
377
378 #[test]
379 fn test_padded_horizontally_packed_row() {
380 type Packed = FieldArray<BabyBear, 3>;
382
383 let inner = RowMajorMatrix::new(
387 vec![
388 BabyBear::new(1),
389 BabyBear::new(2),
390 BabyBear::new(3),
391 BabyBear::new(4),
392 ],
393 2,
394 );
395
396 let mapped_view = RowIndexMappedView {
398 index_map: IdentityMap(inner.height()),
399 inner,
400 };
401
402 let packed: Vec<_> = mapped_view
404 .padded_horizontally_packed_row::<Packed>(1)
405 .collect();
406
407 assert_eq!(
409 packed,
410 vec![Packed::from([
411 BabyBear::new(3),
412 BabyBear::new(4),
413 BabyBear::new(0),
414 ])]
415 );
416 }
417
418 #[test]
419 fn test_vertically_packed_row() {
420 type Packed = FieldArray<BabyBear, 2>;
422
423 let inner = RowMajorMatrix::new((1..=8).map(BabyBear::new).collect::<Vec<_>>(), 2);
429
430 let mapped_view = RowIndexMappedView {
433 index_map: ReverseMap(inner.height()),
434 inner,
435 };
436
437 let packed: Vec<_> = mapped_view.vertically_packed_row::<Packed>(0).collect();
439 assert_eq!(
440 packed,
441 vec![
442 Packed::from([BabyBear::new(7), BabyBear::new(5)]),
443 Packed::from([BabyBear::new(8), BabyBear::new(6)]),
444 ]
445 );
446
447 let packed: Vec<_> = mapped_view.vertically_packed_row::<Packed>(3).collect();
449 assert_eq!(
450 packed,
451 vec![
452 Packed::from([BabyBear::new(1), BabyBear::new(7)]),
453 Packed::from([BabyBear::new(2), BabyBear::new(8)]),
454 ]
455 );
456 }
457
458 #[test]
459 fn test_vertically_packed_row_pair() {
460 type Packed = FieldArray<BabyBear, 2>;
462
463 let inner = RowMajorMatrix::new((1..=8).map(BabyBear::new).collect::<Vec<_>>(), 2);
469
470 let mapped_view = RowIndexMappedView {
473 index_map: ReverseMap(inner.height()),
474 inner,
475 };
476
477 let packed = mapped_view.vertically_packed_row_pair::<Packed>(0, 2);
479 assert_eq!(
480 packed,
481 vec![
482 Packed::from([BabyBear::new(7), BabyBear::new(5)]),
483 Packed::from([BabyBear::new(8), BabyBear::new(6)]),
484 Packed::from([BabyBear::new(3), BabyBear::new(1)]),
485 Packed::from([BabyBear::new(4), BabyBear::new(2)]),
486 ]
487 );
488
489 let packed = mapped_view.vertically_packed_row_pair::<Packed>(1, 2);
491 assert_eq!(
492 packed,
493 vec![
494 Packed::from([BabyBear::new(5), BabyBear::new(3)]),
495 Packed::from([BabyBear::new(6), BabyBear::new(4)]),
496 Packed::from([BabyBear::new(1), BabyBear::new(7)]),
497 Packed::from([BabyBear::new(2), BabyBear::new(8)]),
498 ]
499 );
500
501 let packed = mapped_view.vertically_packed_row_pair::<Packed>(3, 2);
503 assert_eq!(
504 packed,
505 vec![
506 Packed::from([BabyBear::new(1), BabyBear::new(7)]),
507 Packed::from([BabyBear::new(2), BabyBear::new(8)]),
508 Packed::from([BabyBear::new(5), BabyBear::new(3)]),
509 Packed::from([BabyBear::new(6), BabyBear::new(4)]),
510 ]
511 );
512 }
513
514 #[test]
515 fn test_row_and_row_slice_methods() {
516 let inner = RowMajorMatrix::new(vec![10, 20, 30, 40, 50, 60], 3);
520
521 let mapped_view = RowIndexMappedView {
523 index_map: ReverseMap(inner.height()),
524 inner,
525 };
526
527 assert_eq!(mapped_view.row_slice(0).unwrap().deref(), &[40, 50, 60]); assert_eq!(
530 mapped_view.row(1).unwrap().into_iter().collect_vec(),
531 vec![10, 20, 30]
532 ); unsafe {
535 assert_eq!(
537 mapped_view.row_unchecked(0).into_iter().collect_vec(),
538 vec![40, 50, 60]
539 ); assert_eq!(mapped_view.row_slice_unchecked(1).deref(), &[10, 20, 30]); assert_eq!(
543 mapped_view.row_subslice_unchecked(0, 1, 3).deref(),
544 &[50, 60]
545 ); assert_eq!(
547 mapped_view
548 .row_subseq_unchecked(1, 0, 2)
549 .into_iter()
550 .collect_vec(),
551 vec![10, 20]
552 ); }
554
555 assert!(mapped_view.row(2).is_none()); assert!(mapped_view.row_slice(2).is_none()); }
558
559 #[test]
560 fn test_out_of_bounds_access() {
561 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4], 2);
565
566 let mapped_view = RowIndexMappedView {
568 index_map: IdentityMap(inner.height()),
569 inner,
570 };
571
572 assert_eq!(mapped_view.get(2, 1), None);
574 assert!(mapped_view.row(5).is_none());
575 assert!(mapped_view.row_slice(11).is_none());
576 assert_eq!(mapped_view.get(0, 20), None);
577 }
578
579 #[test]
580 fn test_out_of_bounds_access_with_bad_map() {
581 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4], 4);
585
586 let mapped_view = RowIndexMappedView {
588 index_map: ConstantMap,
589 inner,
590 };
591
592 assert_eq!(mapped_view.get(0, 2), Some(3));
593
594 assert_eq!(mapped_view.get(1, 0), None);
596 assert!(mapped_view.row(1).is_none());
597 assert!(mapped_view.row_slice(1).is_none());
598 }
599}