1use alloc::vec::Vec;
2use core::marker::PhantomData;
3use core::ops::Range;
4
5use p3_field::PackedValue;
6
7use crate::Matrix;
8use crate::bitrev::BitReversibleMatrix;
9
10#[derive(Clone)]
16pub struct HorizontallyTruncated<T, Inner> {
17 inner: Inner,
19 column_range: Range<usize>,
21 _phantom: PhantomData<T>,
23}
24
25impl<T, Inner: Matrix<T>> HorizontallyTruncated<T, Inner>
26where
27 T: Send + Sync + Clone,
28{
29 pub fn new(inner: Inner, truncated_width: usize) -> Option<Self> {
39 Self::new_with_range(inner, 0..truncated_width)
40 }
41
42 pub fn new_with_range(inner: Inner, column_range: Range<usize>) -> Option<Self> {
50 let valid = column_range.start <= column_range.end && column_range.end <= inner.width();
54 valid.then(|| Self {
55 inner,
56 column_range,
57 _phantom: PhantomData,
58 })
59 }
60}
61
62impl<T, Inner> Matrix<T> for HorizontallyTruncated<T, Inner>
63where
64 T: Send + Sync + Clone,
65 Inner: Matrix<T>,
66{
67 #[inline(always)]
69 fn width(&self) -> usize {
70 self.column_range.len()
71 }
72
73 #[inline(always)]
75 fn height(&self) -> usize {
76 self.inner.height()
77 }
78
79 #[inline(always)]
80 unsafe fn get_unchecked(&self, r: usize, c: usize) -> T {
81 unsafe {
82 self.inner.get_unchecked(r, self.column_range.start + c)
86 }
87 }
88
89 unsafe fn row_unchecked(
90 &self,
91 r: usize,
92 ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
93 unsafe {
94 self.inner
96 .row_subseq_unchecked(r, self.column_range.start, self.column_range.end)
97 }
98 }
99
100 unsafe fn row_subseq_unchecked(
101 &self,
102 r: usize,
103 start: usize,
104 end: usize,
105 ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
106 unsafe {
107 self.inner.row_subseq_unchecked(
111 r,
112 self.column_range.start + start,
113 self.column_range.start + end,
114 )
115 }
116 }
117
118 unsafe fn row_subslice_unchecked(
119 &self,
120 r: usize,
121 start: usize,
122 end: usize,
123 ) -> impl core::ops::Deref<Target = [T]> {
124 unsafe {
125 self.inner.row_subslice_unchecked(
129 r,
130 self.column_range.start + start,
131 self.column_range.start + end,
132 )
133 }
134 }
135
136 #[inline]
137 fn vertically_packed_row<P>(&self, r: usize) -> impl Iterator<Item = P>
138 where
139 T: Copy,
140 P: PackedValue<Value = T>,
141 {
142 self.inner
143 .vertically_packed_row::<P>(r)
144 .skip(self.column_range.start)
145 .take(self.width())
146 }
147
148 #[inline]
149 fn vertically_packed_row_pair<P>(&self, r: usize, step: usize) -> Vec<P>
150 where
151 T: Copy,
152 P: PackedValue<Value = T>,
153 {
154 self.vertically_packed_row::<P>(r)
155 .chain(self.vertically_packed_row::<P>(r + step))
156 .collect()
157 }
158}
159
160impl<T: Clone + Send + Sync, Inner: BitReversibleMatrix<T>> BitReversibleMatrix<T>
161 for HorizontallyTruncated<T, Inner>
162{
163 type BitRev = HorizontallyTruncated<T, Inner::BitRev>;
164
165 fn bit_reverse_rows(self) -> Self::BitRev {
166 HorizontallyTruncated {
167 inner: self.inner.bit_reverse_rows(),
168 column_range: self.column_range,
169 _phantom: PhantomData,
170 }
171 }
172}
173
174#[cfg(test)]
175mod tests {
176 use alloc::vec;
177 use alloc::vec::Vec;
178
179 use super::*;
180 use crate::dense::RowMajorMatrix;
181
182 #[test]
183 fn test_truncate_width_by_one() {
184 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12], 4);
189
190 let truncated = HorizontallyTruncated::new(inner, 3).unwrap();
192
193 assert_eq!(truncated.width(), 3);
195
196 assert_eq!(truncated.height(), 3);
198
199 assert_eq!(truncated.get(0, 0), Some(1)); assert_eq!(truncated.get(1, 1), Some(6)); unsafe {
203 assert_eq!(truncated.get_unchecked(0, 1), 2); assert_eq!(truncated.get_unchecked(2, 2), 11); }
206
207 let row0: Vec<_> = truncated.row(0).unwrap().into_iter().collect();
209 assert_eq!(row0, vec![1, 2, 3]);
210 unsafe {
211 let row1: Vec<_> = truncated.row_unchecked(1).into_iter().collect();
213 assert_eq!(row1, vec![5, 6, 7]);
214
215 let row3_subset: Vec<_> = truncated
217 .row_subseq_unchecked(2, 1, 2)
218 .into_iter()
219 .collect();
220 assert_eq!(row3_subset, vec![10]);
221 }
222
223 unsafe {
224 let row1 = truncated.row_slice(1).unwrap();
225 assert_eq!(&*row1, &[5, 6, 7]);
226
227 let row2 = truncated.row_slice_unchecked(2);
228 assert_eq!(&*row2, &[9, 10, 11]);
229
230 let row0_subslice = truncated.row_subslice_unchecked(0, 0, 2);
231 assert_eq!(&*row0_subslice, &[1, 2]);
232 }
233
234 assert!(truncated.get(0, 3).is_none()); assert!(truncated.get(3, 0).is_none()); assert!(truncated.row(3).is_none()); assert!(truncated.row_slice(3).is_none()); let as_matrix = truncated.to_row_major_matrix();
241
242 let expected = RowMajorMatrix::new(vec![1, 2, 3, 5, 6, 7, 9, 10, 11], 3);
247
248 assert_eq!(as_matrix, expected);
249 }
250
251 #[test]
252 fn test_no_truncation() {
253 let inner = RowMajorMatrix::new(vec![7, 8, 9, 10], 2);
257
258 let truncated = HorizontallyTruncated::new(inner, 2).unwrap();
260
261 assert_eq!(truncated.width(), 2);
262 assert_eq!(truncated.height(), 2);
263 assert_eq!(truncated.get(0, 1).unwrap(), 8);
264 assert_eq!(truncated.get(1, 0).unwrap(), 9);
265
266 unsafe {
267 assert_eq!(truncated.get_unchecked(0, 0), 7);
268 assert_eq!(truncated.get_unchecked(1, 1), 10);
269 }
270
271 let row0: Vec<_> = truncated.row(0).unwrap().into_iter().collect();
272 assert_eq!(row0, vec![7, 8]);
273
274 let row1: Vec<_> = unsafe { truncated.row_unchecked(1).into_iter().collect() };
275 assert_eq!(row1, vec![9, 10]);
276
277 assert!(truncated.get(0, 2).is_none()); assert!(truncated.get(2, 0).is_none()); assert!(truncated.row(2).is_none()); assert!(truncated.row_slice(2).is_none()); }
282
283 #[test]
284 fn test_truncate_to_zero_width() {
285 let inner = RowMajorMatrix::new(vec![11, 12, 13], 3);
287
288 let truncated = HorizontallyTruncated::new(inner, 0).unwrap();
290
291 assert_eq!(truncated.width(), 0);
292 assert_eq!(truncated.height(), 1);
293
294 assert!(truncated.row(0).unwrap().into_iter().next().is_none());
296
297 assert!(truncated.get(0, 0).is_none()); assert!(truncated.get(1, 0).is_none()); assert!(truncated.row(1).is_none()); assert!(truncated.row_slice(1).is_none()); }
302
303 #[test]
304 fn test_invalid_truncation_width() {
305 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4], 2);
309
310 assert!(HorizontallyTruncated::new(inner, 5).is_none());
312 }
313
314 #[test]
315 fn test_column_range_middle() {
316 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15], 5);
321
322 let view = HorizontallyTruncated::new_with_range(inner, 1..4).unwrap();
324
325 assert_eq!(view.width(), 3);
327
328 assert_eq!(view.height(), 3);
330
331 assert_eq!(view.get(0, 0), Some(2)); assert_eq!(view.get(0, 1), Some(3)); assert_eq!(view.get(0, 2), Some(4)); assert_eq!(view.get(1, 0), Some(7)); assert_eq!(view.get(2, 2), Some(14)); unsafe {
339 assert_eq!(view.get_unchecked(1, 1), 8); assert_eq!(view.get_unchecked(2, 0), 12); }
342
343 let row0: Vec<_> = view.row(0).unwrap().into_iter().collect();
345 assert_eq!(row0, vec![2, 3, 4]);
346
347 let row1: Vec<_> = view.row(1).unwrap().into_iter().collect();
349 assert_eq!(row1, vec![7, 8, 9]);
350
351 unsafe {
352 let row2: Vec<_> = view.row_unchecked(2).into_iter().collect();
354 assert_eq!(row2, vec![12, 13, 14]);
355
356 let row1_subseq: Vec<_> = view.row_subseq_unchecked(1, 1, 3).into_iter().collect();
358 assert_eq!(row1_subseq, vec![8, 9]);
359 }
360
361 assert!(view.get(0, 3).is_none()); assert!(view.get(3, 0).is_none()); let as_matrix = view.to_row_major_matrix();
367
368 let expected = RowMajorMatrix::new(vec![2, 3, 4, 7, 8, 9, 12, 13, 14], 3);
373
374 assert_eq!(as_matrix, expected);
375 }
376
377 #[test]
378 fn test_column_range_end() {
379 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8], 4);
383
384 let view = HorizontallyTruncated::new_with_range(inner, 2..4).unwrap();
386
387 assert_eq!(view.width(), 2);
388 assert_eq!(view.height(), 2);
389
390 let row0: Vec<_> = view.row(0).unwrap().into_iter().collect();
392 assert_eq!(row0, vec![3, 4]);
393
394 let row1: Vec<_> = view.row(1).unwrap().into_iter().collect();
396 assert_eq!(row1, vec![7, 8]);
397
398 assert_eq!(view.get(0, 0), Some(3));
399 assert_eq!(view.get(1, 1), Some(8));
400 }
401
402 #[test]
403 fn test_column_range_single_column() {
404 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12], 4);
409
410 let view = HorizontallyTruncated::new_with_range(inner, 2..3).unwrap();
412
413 assert_eq!(view.width(), 1);
414 assert_eq!(view.height(), 3);
415
416 assert_eq!(view.get(0, 0), Some(3));
417 assert_eq!(view.get(1, 0), Some(7));
418 assert_eq!(view.get(2, 0), Some(11));
419
420 let row0: Vec<_> = view.row(0).unwrap().into_iter().collect();
422 assert_eq!(row0, vec![3]);
423 }
424
425 #[test]
426 fn test_column_range_empty() {
427 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
431
432 let view = HorizontallyTruncated::new_with_range(inner, 2..2).unwrap();
434
435 assert_eq!(view.width(), 0);
436 assert_eq!(view.height(), 2);
437
438 assert!(view.row(0).unwrap().into_iter().next().is_none());
440 }
441
442 #[test]
443 fn test_invalid_column_range() {
444 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
448
449 assert!(HorizontallyTruncated::new_with_range(inner, 1..5).is_none());
451 }
452
453 #[test]
454 fn test_inverted_column_range_is_rejected() {
455 let inner = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
459
460 let inverted = Range { start: 2, end: 1 };
467 assert!(HorizontallyTruncated::new_with_range(inner, inverted).is_none());
468 }
469
470 #[test]
471 fn test_vertically_packed_row_scalar_width_1() {
472 use p3_baby_bear::BabyBear;
473
474 type Packed = BabyBear;
475
476 let inner = RowMajorMatrix::new((1..17).map(BabyBear::new).collect::<Vec<_>>(), 4);
482
483 let truncated = HorizontallyTruncated::new(inner, 2).unwrap();
485
486 let packed = truncated
487 .vertically_packed_row::<Packed>(2)
488 .collect::<Vec<_>>();
489
490 assert_eq!(packed, vec![BabyBear::new(9), BabyBear::new(10)]);
491 }
492
493 #[test]
494 fn test_vertically_packed_row_pair_middle_column_range() {
495 use p3_baby_bear::BabyBear;
496 use p3_field::FieldArray;
497
498 type Packed = FieldArray<BabyBear, 2>;
499
500 let inner = RowMajorMatrix::new((1..17).map(BabyBear::new).collect::<Vec<_>>(), 4);
501
502 let view = HorizontallyTruncated::new_with_range(inner, 1..3).unwrap();
505
506 let packed = view.vertically_packed_row_pair::<Packed>(0, 2);
507
508 assert_eq!(
509 packed,
510 vec![
511 [BabyBear::new(2), BabyBear::new(6)].into(),
512 [BabyBear::new(3), BabyBear::new(7)].into(),
513 [BabyBear::new(10), BabyBear::new(14)].into(),
514 [BabyBear::new(11), BabyBear::new(15)].into(),
515 ]
516 );
517 }
518}