1use super::BitSet;
2#[cfg(target_arch = "x86_64")]
3use super::{SimdBackend, simd_backend};
4use std::ops::{BitXorAssign, Index, IndexMut, Mul};
5
6#[derive(Clone, Debug, PartialEq, Eq)]
8pub struct BitMatrix {
9 pub shape: (usize, usize),
10 pub data: Vec<BitSet>,
11}
12
13#[derive(Clone, Debug, PartialEq, Eq)]
14pub struct BitMatrixSolution {
15 pub particular: BitSet,
16 pub basis: Vec<BitSet>,
17}
18
19impl BitMatrix {
20 pub fn zeros(shape: (usize, usize)) -> Self {
21 Self {
22 shape,
23 data: vec![BitSet::new(shape.1); shape.0],
24 }
25 }
26
27 pub fn from_vec(data: Vec<BitSet>) -> Self {
28 let shape = (data.len(), data.first().map_or(0, BitSet::len));
29 assert!(data.iter().all(|row| row.len() == shape.1));
30 Self { shape, data }
31 }
32
33 pub fn new_with(shape: (usize, usize), mut f: impl FnMut(usize, usize) -> bool) -> Self {
34 let mut a = Self::zeros(shape);
35 for (i, row) in a.data.iter_mut().enumerate() {
36 for (w, word) in row.words_mut().iter_mut().enumerate() {
37 for j in w * 64..shape.1.min((w + 1) * 64) {
38 *word |= u64::from(f(i, j)) << (j & 63);
39 }
40 }
41 }
42 a
43 }
44
45 pub fn eye(shape: (usize, usize)) -> Self {
46 let mut a = Self::zeros(shape);
47 for i in 0..shape.0.min(shape.1) {
48 a[i].set(i, true);
49 }
50 a
51 }
52
53 pub fn transpose(&self) -> Self {
54 let mut a = Self::zeros((self.shape.1, self.shape.0));
55 for (i, row) in self.data.iter().enumerate() {
56 for j in row.iter_ones() {
57 a[j].set(i, true);
58 }
59 }
60 a
61 }
62
63 pub fn row_reduction(&mut self) -> Vec<usize> {
65 self.eliminate(self.shape.1, true, false)
66 }
67
68 pub fn rank(&mut self) -> usize {
70 self.eliminate(self.shape.1, false, false).len()
71 }
72
73 pub fn determinant(&mut self) -> bool {
75 assert_eq!(self.shape.0, self.shape.1);
76 self.eliminate(self.shape.1, false, true).len() == self.shape.0
77 }
78
79 pub fn inverse(&self) -> Option<Self> {
81 let (n, m) = self.shape;
82 assert_eq!(n, m);
83 let mut a = Self::zeros((n, 2 * n));
84 for i in 0..n {
85 a[i].words_mut()[..self[i].words().len()].copy_from_slice(self[i].words());
86 a[i].set(n + i, true);
87 }
88 if a.eliminate(n, true, true).len() != n {
89 return None;
90 }
91 for row in &mut a.data {
92 *row >>= n;
93 }
94 let mut inverse = Self::zeros((n, n));
95 for (row, source) in inverse.data.iter_mut().zip(&a.data) {
96 let len = row.words().len();
97 row.words_mut().copy_from_slice(&source.words()[..len]);
98 }
99 Some(inverse)
100 }
101
102 pub fn solve_system_of_linear_equations(&self, b: &BitSet) -> Option<BitMatrixSolution> {
105 let (n, m) = self.shape;
106 assert_eq!(b.len(), n);
107 let mut a = Self::zeros((n, m + 1));
108 for i in 0..n {
109 a[i].words_mut()[..self[i].words().len()].copy_from_slice(self[i].words());
110 a[i].set(m, b.get(i));
111 }
112 let pivots = a.eliminate(m, true, false);
113 if a.data[pivots.len()..].iter().any(|row| row.get(m)) {
114 return None;
115 }
116 let mut particular = BitSet::new(m);
117 let mut free = BitSet::ones(m);
118 for (i, &c) in pivots.iter().enumerate() {
119 particular.set(c, a[i].get(m));
120 free.set(c, false);
121 }
122 let columns: Vec<_> = free.iter_ones().collect();
123 let mut basis: Vec<_> = columns
124 .iter()
125 .map(|&c| {
126 let mut row = BitSet::new(m);
127 row.set(c, true);
128 row
129 })
130 .collect();
131 for (i, &p) in pivots.iter().enumerate() {
132 for (row, &c) in basis.iter_mut().zip(&columns) {
133 row.words_mut()[p / 64] |= u64::from(a[i].get(c)) << (p & 63);
134 }
135 }
136 Some(BitMatrixSolution { particular, basis })
137 }
138
139 pub fn pow(self, mut n: usize) -> Self {
140 assert_eq!(self.shape.0, self.shape.1);
141 let mut result = Self::eye(self.shape);
142 let mut a = self;
143 while n != 0 {
144 if n & 1 != 0 {
145 result = &result * &a;
146 }
147 n >>= 1;
148 if n != 0 {
149 a = &a * &a;
150 }
151 }
152 result
153 }
154
155 fn eliminate(&mut self, cols: usize, full: bool, require_full_rank: bool) -> Vec<usize> {
156 #[cfg(target_arch = "x86_64")]
157 match simd_backend() {
158 SimdBackend::Avx512 => {
160 return unsafe { self.eliminate_avx512(cols, full, require_full_rank) };
161 }
162 SimdBackend::Avx2 => {
164 return unsafe { self.eliminate_avx2(cols, full, require_full_rank) };
165 }
166 SimdBackend::Scalar => {}
167 }
168 self.eliminate_impl(cols, full, require_full_rank)
169 }
170
171 #[inline(always)]
173 fn eliminate_impl(&mut self, cols: usize, full: bool, require_full_rank: bool) -> Vec<usize> {
174 let n = self.shape.0;
175 let mut pivots = Vec::with_capacity(n.min(cols));
176 if n < 32 {
177 let mut c = 0;
178 while c < cols {
179 let r = pivots.len();
180 if r == n {
181 break;
182 }
183 let Some(p) = (r..n).find(|&i| self[i].get(c)) else {
184 if require_full_rank {
185 return pivots;
186 }
187 c = self.next_column(r, c + 1, cols);
188 continue;
189 };
190 self.data.swap(r, p);
191 let (upper, lower) = self.data.split_at_mut(r);
192 let (pivot, lower) = lower.split_first_mut().unwrap();
193 for row in lower
194 .iter_mut()
195 .chain(upper.iter_mut().take(if full { r } else { 0 }))
196 {
197 if row.get(c) {
198 xor(&mut row.words_mut()[c / 64..], &pivot.words()[c / 64..]);
199 }
200 }
201 pivots.push(c);
202 c += 1;
203 }
204 return pivots;
205 }
206 if self
208 .data
209 .iter()
210 .all(|row| row.iter_ones().take_while(|&c| c < cols).take(3).count() <= 2)
211 {
212 return self.eliminate_sparse(cols, full, require_full_rank);
213 }
214 let block: usize = if n < 512 {
215 4
216 } else if n < 1536 {
217 8
218 } else {
219 32
220 };
221 let mut table = vec![BitSet::new(self.shape.1); (1 << block.min(8)) * block.div_ceil(8)];
222 let mut reduced = vec![0; n];
223 let mut start = 0;
224 while start < cols {
225 let first = pivots.len();
226 let word = start / 64;
227 reduced.fill(first);
228 let end = cols.min(start + block);
229 let mut c = start;
230 while c < end {
231 let r = pivots.len();
232 let mut pivot = None;
233 for (i, reduced) in reduced.iter_mut().enumerate().skip(r) {
234 let (upper, lower) = self.data.split_at_mut(i);
235 let row = &mut lower[0];
236 for p in *reduced..r {
237 if row.get(pivots[p]) {
238 xor(&mut row.words_mut()[word..], &upper[p].words()[word..]);
239 }
240 }
241 *reduced = r;
242 if row.get(c) {
243 pivot = Some(i);
244 break;
245 }
246 }
247 if let Some(p) = pivot {
248 self.data.swap(r, p);
249 reduced.swap(r, p);
250 pivots.push(c);
251 c += 1;
252 } else if require_full_rank {
253 return pivots;
254 } else {
255 c = self.next_column(r, c + 1, end);
256 }
257 }
258 let rank = pivots.len();
259 if first == rank {
260 let next = self.next_column(rank, start + block, cols);
261 if next == cols {
262 break;
263 }
264 start = next / block * block;
265 continue;
266 }
267 for r in (first..rank).rev() {
269 let (upper, lower) = self.data.split_at_mut(r);
270 for row in &mut upper[first..] {
271 if row.get(pivots[r]) {
272 xor(&mut row.words_mut()[word..], &lower[0].words()[word..]);
273 }
274 }
275 }
276 let mut indices = [[0usize; 256]; 4];
277 let mut masks = [0usize; 4];
278 for group in 0..block.div_ceil(8) {
279 let mut keys = [0usize; 256];
280 let mut count = 0;
281 for (row, &c) in pivots[first..].iter().enumerate() {
282 if (c - start) / 8 != group {
283 continue;
284 }
285 let half = 1 << count;
286 count += 1;
287 for index in 0..half {
288 keys[index + half] = keys[index] | (1 << ((c - start) % 8));
289 let (lower, upper) = table.split_at_mut(group * 256 + index + half);
290 let target = &mut upper[0].words_mut()[word..];
291 let source = &lower[group * 256 + index].words()[word..];
292 let pivot = &self[first + row].words()[word..];
293 for ((x, y), z) in target.iter_mut().zip(source).zip(pivot) {
294 *x = y ^ z;
295 }
296 }
297 }
298 masks[group] = keys[(1 << count) - 1];
299 for (i, &key) in keys[..1 << count].iter().enumerate() {
300 indices[group][key] = i;
301 }
302 }
303 for i in (rank..n).chain(0..if full { first } else { 0 }) {
304 let key = (self[i].words()[word] >> (start & 63)) as usize;
305 let x = indices[0][key & masks[0]];
306 if block <= 8 {
307 if x != 0 {
308 xor(&mut self[i].words_mut()[word..], &table[x].words()[word..]);
309 }
310 continue;
311 }
312 let y = indices[1][(key >> 8) & masks[1]];
313 let z = indices[2][(key >> 16) & masks[2]];
314 let w = indices[3][(key >> 24) & masks[3]];
315 if x | y | z | w != 0 {
316 let p = &table[x].words()[word..];
317 let q = &table[256 + y].words()[word..];
318 let r = &table[512 + z].words()[word..];
319 let s = &table[768 + w].words()[word..];
320 for ((((x, y), z), r), s) in self[i].words_mut()[word..]
321 .iter_mut()
322 .zip(p)
323 .zip(q)
324 .zip(r)
325 .zip(s)
326 {
327 *x ^= y ^ z ^ r ^ s;
328 }
329 }
330 }
331 if rank == n {
332 break;
333 }
334 start += block;
335 }
336 pivots
337 }
338
339 fn next_column(&self, row: usize, mut start: usize, cols: usize) -> usize {
340 while start < cols {
341 let word = start / 64;
342 let bits = self.data[row..]
343 .iter()
344 .fold(0, |x, row| x | row.words()[word])
345 & (u64::MAX << (start & 63));
346 if bits != 0 {
347 return (word * 64 + bits.trailing_zeros() as usize).min(cols);
348 }
349 start = (word + 1) * 64;
350 }
351 cols
352 }
353
354 #[inline(always)]
355 fn eliminate_sparse(&mut self, cols: usize, full: bool, require_full_rank: bool) -> Vec<usize> {
356 let n = self.shape.0;
357 let mut basis = vec![n; cols];
358 let mut pivots = Vec::new();
359 for i in 0..n {
360 loop {
361 let Some(c) = self[i].iter_ones().next().filter(|&c| c < cols) else {
362 if require_full_rank {
363 return pivots;
364 }
365 break;
366 };
367 if basis[c] == n {
368 basis[c] = i;
369 pivots.push(c);
370 break;
371 }
372 let (upper, lower) = self.data.split_at_mut(i);
373 xor(
374 &mut lower[0].words_mut()[c / 64..],
375 &upper[basis[c]].words()[c / 64..],
376 );
377 }
378 }
379 pivots.sort_unstable();
380 self.data
381 .sort_by_cached_key(|row| row.iter_ones().next().filter(|&c| c < cols).unwrap_or(cols));
382 if full {
383 for (i, &c) in pivots.iter().enumerate() {
384 basis[c] = i;
385 }
386 for i in (0..pivots.len()).rev() {
387 let next = self[i].iter_ones().take_while(|&c| c < cols).nth(1);
388 if let Some(c) = next
389 && basis[c] != n
390 {
391 let (upper, lower) = self.data.split_at_mut(basis[c]);
392 xor(
393 &mut upper[i].words_mut()[c / 64..],
394 &lower[0].words()[c / 64..],
395 );
396 }
397 }
398 }
399 pivots
400 }
401
402 #[inline(always)]
403 fn mul_impl(&self, rhs: &Self) -> Self {
404 let mut result = Self::zeros((self.shape.0, rhs.shape.1));
405 let ones = self.data.iter().map(BitSet::count_ones).sum::<u64>();
406 let size = self.shape.0 as u64 * self.shape.1 as u64;
407 if self.shape.0 < 256 || self.shape.1 < 32 || ones <= size / 8 {
408 for (a, c) in self.data.iter().zip(&mut result.data) {
409 for j in a.iter_ones() {
410 xor(c.words_mut(), rhs[j].words());
411 }
412 }
413 return result;
414 }
415 if size - ones <= size / 8 {
416 let mut sum = BitSet::new(rhs.shape.1);
417 for row in &rhs.data {
418 sum ^= row;
419 }
420 for (a, c) in self.data.iter().zip(&mut result.data) {
421 c.words_mut().copy_from_slice(sum.words());
422 for j in (!a.clone()).iter_ones() {
423 xor(c.words_mut(), rhs[j].words());
424 }
425 }
426 return result;
427 }
428 let width = rhs.shape.1.div_ceil(64);
429 if width == 0 {
430 return result;
431 }
432
433 let group = 256 * width + 8;
435 let mut storage = BitSet::new(8 * group * 64);
436 let table = storage.words_mut();
437 for start in (0..self.shape.1).step_by(64) {
438 for (t, table) in table.chunks_exact_mut(group).enumerate() {
439 let col = start + t * 8;
440 for bit in 0..self.shape.1.saturating_sub(col).min(8) {
441 let row = rhs[col + bit].words();
442 let half = (1 << bit) * width;
443 let (lower, upper) = table.split_at_mut(half);
444 for (source, dest) in
445 lower.chunks_exact(width).zip(upper.chunks_exact_mut(width))
446 {
447 for ((x, y), z) in dest.iter_mut().zip(source).zip(row) {
448 *x = y ^ z;
449 }
450 }
451 }
452 }
453 for (a, c) in self.data.iter().zip(&mut result.data) {
454 let key = a.words()[start / 64];
455 let offset = (key & 255) as usize * width;
456 let p0 = &table[offset..offset + width];
457 let offset = group + (key >> 8 & 255) as usize * width;
458 let p1 = &table[offset..offset + width];
459 let offset = 2 * group + (key >> 16 & 255) as usize * width;
460 let p2 = &table[offset..offset + width];
461 let offset = 3 * group + (key >> 24 & 255) as usize * width;
462 let p3 = &table[offset..offset + width];
463 let offset = 4 * group + (key >> 32 & 255) as usize * width;
464 let p4 = &table[offset..offset + width];
465 let offset = 5 * group + (key >> 40 & 255) as usize * width;
466 let p5 = &table[offset..offset + width];
467 let offset = 6 * group + (key >> 48 & 255) as usize * width;
468 let p6 = &table[offset..offset + width];
469 let offset = 7 * group + (key >> 56 & 255) as usize * width;
470 let p7 = &table[offset..offset + width];
471 for ((((((((x, p0), p1), p2), p3), p4), p5), p6), p7) in c
472 .words_mut()
473 .iter_mut()
474 .zip(p0)
475 .zip(p1)
476 .zip(p2)
477 .zip(p3)
478 .zip(p4)
479 .zip(p5)
480 .zip(p6)
481 .zip(p7)
482 {
483 *x ^= p0 ^ p1 ^ p2 ^ p3 ^ p4 ^ p5 ^ p6 ^ p7;
484 }
485 }
486 }
487 result
488 }
489
490 #[cfg(target_arch = "x86_64")]
491 #[target_feature(enable = "avx2")]
492 unsafe fn eliminate_avx2(
493 &mut self,
494 cols: usize,
495 full: bool,
496 require_full_rank: bool,
497 ) -> Vec<usize> {
498 self.eliminate_impl(cols, full, require_full_rank)
499 }
500 #[cfg(target_arch = "x86_64")]
501 #[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
502 unsafe fn eliminate_avx512(
503 &mut self,
504 cols: usize,
505 full: bool,
506 require_full_rank: bool,
507 ) -> Vec<usize> {
508 self.eliminate_impl(cols, full, require_full_rank)
509 }
510 #[cfg(target_arch = "x86_64")]
511 #[target_feature(enable = "avx2")]
512 unsafe fn mul_avx2(&self, rhs: &Self) -> Self {
513 self.mul_impl(rhs)
514 }
515 #[cfg(target_arch = "x86_64")]
516 #[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
517 unsafe fn mul_avx512(&self, rhs: &Self) -> Self {
518 self.mul_impl(rhs)
519 }
520}
521
522#[inline(always)]
523fn xor(row: &mut [u64], pivot: &[u64]) {
524 for (x, y) in row.iter_mut().zip(pivot) {
525 *x ^= y;
526 }
527}
528
529impl Index<usize> for BitMatrix {
530 type Output = BitSet;
531 fn index(&self, i: usize) -> &Self::Output {
532 &self.data[i]
533 }
534}
535impl IndexMut<usize> for BitMatrix {
536 fn index_mut(&mut self, i: usize) -> &mut Self::Output {
537 &mut self.data[i]
538 }
539}
540impl BitXorAssign<&Self> for BitMatrix {
541 fn bitxor_assign(&mut self, rhs: &Self) {
542 assert_eq!(self.shape, rhs.shape);
543 for (a, b) in self.data.iter_mut().zip(&rhs.data) {
544 *a ^= b;
545 }
546 }
547}
548impl Mul<&BitMatrix> for &BitMatrix {
549 type Output = BitMatrix;
550 fn mul(self, rhs: &BitMatrix) -> BitMatrix {
551 assert_eq!(self.shape.1, rhs.shape.0);
552 #[cfg(target_arch = "x86_64")]
553 match simd_backend() {
554 SimdBackend::Avx512 => return unsafe { self.mul_avx512(rhs) },
556 SimdBackend::Avx2 => return unsafe { self.mul_avx2(rhs) },
558 SimdBackend::Scalar => {}
559 }
560 self.mul_impl(rhs)
561 }
562}
563impl Mul for BitMatrix {
564 type Output = Self;
565 fn mul(self, rhs: Self) -> Self {
566 &self * &rhs
567 }
568}
569
570#[cfg(test)]
571mod tests {
572 use super::*;
573 use crate::tools::Xorshift;
574 use std::collections::BTreeSet;
575
576 #[test]
577 fn test_random_small_linear_algebra() {
578 let mut rng = Xorshift::new_with_seed(918217);
579 for _ in 0..256 {
580 let n = rng.random(0..=8);
581 let m = rng.random(0..=8);
582 let density = rng.rand(1010);
583 let rows: Vec<Vec<bool>> = (0..n)
584 .map(|_| (0..m).map(|_| rng.rand(1009) < density).collect())
585 .collect();
586 let a = BitMatrix::new_with((n, m), |i, j| rows[i][j]);
587 let images: Vec<usize> = (0..1usize << m)
589 .map(|x| {
590 rows.iter().enumerate().fold(0, |image, (i, row)| {
591 let bit = row
592 .iter()
593 .enumerate()
594 .fold(false, |v, (j, &b)| v ^ (b && x >> j & 1 != 0));
595 image | (usize::from(bit) << i)
596 })
597 })
598 .collect();
599 let image: BTreeSet<_> = images.iter().copied().collect();
600 let rank = image.len().ilog2() as usize;
601 assert_eq!(a.clone().rank(), rank);
602 if n == m {
603 assert_eq!(a.clone().determinant(), rank == n);
604 let inverse = a.inverse();
605 assert_eq!(inverse.is_some(), rank == n);
606 if let Some(inverse) = inverse {
607 assert_eq!(&a * &inverse, BitMatrix::eye((n, n)));
608 assert_eq!(&inverse * &a, BitMatrix::eye((n, n)));
609 }
610 }
611 for b in 0..1usize << n {
612 let rhs = (0..n).map(|i| b >> i & 1 != 0).collect();
613 let solution = a.solve_system_of_linear_equations(&rhs);
614 assert_eq!(solution.is_some(), image.contains(&b));
615 if let Some(solution) = solution {
616 assert_eq!(solution.basis.len(), m - rank);
617 let solutions: BTreeSet<_> = (0..1 << solution.basis.len())
618 .map(|mask| {
619 let mut x = solution.particular.clone();
620 for (i, row) in solution.basis.iter().enumerate() {
621 if mask >> i & 1 != 0 {
622 x ^= row;
623 }
624 }
625 x.iter_ones().map(|i| 1 << i).sum::<usize>()
626 })
627 .collect();
628 let expected: BTreeSet<_> = images
629 .iter()
630 .enumerate()
631 .filter_map(|(x, &y)| (y == b).then_some(x))
632 .collect();
633 assert_eq!(solutions, expected);
634 }
635 }
636 }
637 }
638
639 #[test]
640 fn test_random_rectangular_matrices() {
641 let mut rng = Xorshift::new_with_seed(1234567);
642 for case in 0..184 {
643 let (n, m, k) = match case {
645 0..128 => (rng.random(0..=32), rng.random(0..=32), rng.random(0..=32)),
646 128..160 => (rng.random(0..256), rng.random(0..256), rng.random(0..256)),
647 _ => match case % 3 {
648 0 => (
649 rng.random(256..800),
650 rng.random(32..160),
651 rng.random(128..1600),
652 ),
653 1 => (
654 rng.random(32..160),
655 rng.random(256..1800),
656 rng.random(0..48),
657 ),
658 _ => (
659 rng.random(512..800),
660 rng.random(256..800),
661 rng.random(0..32),
662 ),
663 },
664 };
665 let density = rng.rand(1010);
666 let mut rows: Vec<Vec<bool>> = (0..n)
667 .map(|_| (0..m).map(|_| rng.rand(1009) < density).collect())
668 .collect();
669 match rng.rand(5) {
670 0 => {
671 for row in &mut rows {
672 row.fill(false);
673 if m != 0 {
674 for _ in 0..rng.rand(3) {
675 row[rng.random(0..m)] = true;
676 }
677 }
678 }
679 }
680 1 => {
681 let rank_bound = rng.random(0..=n.min(16));
682 for i in rank_bound..n {
683 rows[i] = if rank_bound == 0 {
684 vec![false; m]
685 } else {
686 let p = rng.random(0..rank_bound);
687 let q = rng.random(0..rank_bound);
688 rows[p].iter().zip(&rows[q]).map(|(x, y)| x ^ y).collect()
689 };
690 }
691 }
692 2 => {
693 let start = rng.random(0..=m);
694 let active: Vec<_> = (0..m).map(|j| j >= start && rng.rand(3) != 0).collect();
695 for row in &mut rows {
696 for (x, active) in row.iter_mut().zip(&active) {
697 *x &= active;
698 }
699 }
700 }
701 _ => {}
702 }
703 let a = BitMatrix::new_with((n, m), |i, j| rows[i][j]);
704 let density = rng.rand(1010);
705 let right: Vec<Vec<bool>> = (0..m)
706 .map(|_| (0..k).map(|_| rng.rand(1009) < density).collect())
707 .collect();
708 let b = BitMatrix::new_with((m, k), |i, j| right[i][j]);
709 let mut product = vec![vec![false; k]; n];
710 for (row, result) in rows.iter().zip(&mut product) {
711 for (&bit, source) in row.iter().zip(&right) {
712 if bit {
713 for (x, y) in result.iter_mut().zip(source) {
714 *x ^= y;
715 }
716 }
717 }
718 }
719 let expected = BitMatrix::new_with((n, k), |i, j| product[i][j]);
720 assert_eq!(&a * &b, expected, "case {case}, shape {:?}", (n, m, k));
721 let complement = BitMatrix::new_with((n, m), |i, j| !rows[i][j]);
722 let complement_expected = BitMatrix::new_with((n, k), |i, j| {
723 (0..m).fold(false, |v, t| v ^ (!rows[i][t] && right[t][j]))
724 });
725 assert_eq!(&complement * &b, complement_expected, "case {case}");
726 assert_eq!(
727 a.transpose(),
728 BitMatrix::new_with((m, n), |i, j| rows[j][i])
729 );
730
731 let x: Vec<bool> = (0..m).map(|_| rng.rand(1009) < 504).collect();
734 let rhs: Vec<[bool; 2]> = rows
735 .iter()
736 .map(|row| {
737 [
738 row.iter().zip(&x).fold(false, |v, (a, b)| v ^ (a & b)),
739 rng.rand(1009) < 504,
740 ]
741 })
742 .collect();
743 let mut reduced_rows = rows.clone();
744 for (row, rhs) in reduced_rows.iter_mut().zip(&rhs) {
745 row.extend(rhs);
746 }
747 let mut pivots = Vec::new();
748 for c in 0..m {
749 let r = pivots.len();
750 if let Some(p) = (r..n).find(|&i| reduced_rows[i][c]) {
751 reduced_rows.swap(r, p);
752 let pivot = reduced_rows[r].clone();
753 for (i, row) in reduced_rows.iter_mut().enumerate() {
754 if i != r && row[c] {
755 for (x, y) in row.iter_mut().zip(&pivot) {
756 *x ^= y;
757 }
758 }
759 }
760 pivots.push(c);
761 }
762 }
763 let rref = BitMatrix::new_with((n, m), |i, j| reduced_rows[i][j]);
764 let mut reduced = a.clone();
765 assert_eq!(reduced.row_reduction(), pivots, "case {case}");
766 assert_eq!(reduced, rref, "case {case}");
767 assert_eq!(a.clone().rank(), pivots.len(), "case {case}");
768 let free: Vec<_> = (0..m).filter(|j| !pivots.contains(j)).collect();
769 for t in 0..2 {
770 let rhs_bits = rhs.iter().map(|b| b[t]).collect();
771 let solution = a.solve_system_of_linear_equations(&rhs_bits);
772 let consistent = reduced_rows[pivots.len()..].iter().all(|row| !row[m + t]);
773 assert_eq!(solution.is_some(), consistent, "case {case}, rhs {t}");
774 if let Some(solution) = solution {
775 assert_eq!(solution.particular.len(), m);
776 assert_eq!(solution.basis.len(), free.len());
777 for (row, rhs) in rows.iter().zip(&rhs) {
778 assert_eq!(
779 row.iter().enumerate().fold(false, |v, (j, &b)| {
780 v ^ (b && solution.particular.get(j))
781 }),
782 rhs[t]
783 );
784 }
785 for (i, vector) in solution.basis.iter().enumerate() {
786 assert_eq!(vector.len(), m);
787 for (j, &c) in free.iter().enumerate() {
788 assert_eq!(vector.get(c), i == j);
789 }
790 for (r, &c) in pivots.iter().enumerate() {
791 assert_eq!(vector.get(c), reduced_rows[r][free[i]]);
792 }
793 }
794 }
795 }
796 #[cfg(target_arch = "x86_64")]
797 {
798 if is_x86_feature_detected!("avx2") {
799 unsafe {
801 assert_eq!(a.mul_avx2(&b), expected);
802 assert_eq!(complement.mul_avx2(&b), complement_expected);
803 let mut reduced = a.clone();
804 assert_eq!(reduced.eliminate_avx2(m, true, false), pivots);
805 assert_eq!(reduced, rref);
806 }
807 }
808 if crate::tools::avx512_supported() {
809 unsafe {
811 assert_eq!(a.mul_avx512(&b), expected);
812 assert_eq!(complement.mul_avx512(&b), complement_expected);
813 let mut reduced = a.clone();
814 assert_eq!(reduced.eliminate_avx512(m, true, false), pivots);
815 assert_eq!(reduced, rref);
816 }
817 }
818 assert_eq!(a.mul_impl(&b), expected);
819 assert_eq!(complement.mul_impl(&b), complement_expected);
820 let mut reduced = a.clone();
821 assert_eq!(reduced.eliminate_impl(m, true, false), pivots);
822 assert_eq!(reduced, rref);
823 }
824 }
825 }
826
827 #[test]
828 fn test_random_square_matrices() {
829 let mut rng = Xorshift::new_with_seed(7712389);
830 for case in 0..80 {
831 let n = if case < 64 {
832 rng.random(0..160)
833 } else {
834 rng.random(256..1200)
835 };
836 let mut a = BitMatrix::eye((n, n));
837 match rng.rand(3) {
838 0 => {
839 for i in 0..n {
840 for j in i + 1..n {
841 a[i].set(j, rng.rand(1009) < 504);
842 }
843 }
844 }
845 1 => {
846 for i in 0..n.saturating_sub(1) {
847 a[i].set(rng.random(i + 1..n), true);
848 }
849 }
850 _ => {}
851 }
852 rng.shuffle(&mut a.data);
853 let inv = a.inverse().unwrap();
854 assert!(a.clone().determinant());
855 let identity = BitMatrix::eye((n, n));
856 assert_eq!(&a * &inv, identity);
857 assert_eq!(&inv * &a, identity);
858 let b: BitSet = (0..n).map(|_| rng.rand(1009) < 504).collect();
859 let sol = a.solve_system_of_linear_equations(&b).unwrap();
860 assert!(sol.basis.is_empty());
861 for i in 0..n {
862 assert_eq!(
863 (0..n).fold(false, |v, j| v ^ (a[i].get(j) && sol.particular.get(j))),
864 b.get(i)
865 );
866 }
867 let exponent = rng.random(0..10);
868 let mut power = identity;
869 for _ in 0..exponent {
870 power = &power * &a;
871 }
872 assert_eq!(a.clone().pow(exponent), power);
873 let other = BitMatrix::new_with((n, n), |_, _| rng.rand(1009) < 504);
874 let mut sum = a.clone();
875 sum ^= &other;
876 assert_eq!(
877 sum,
878 BitMatrix::new_with((n, n), |i, j| a[i].get(j) ^ other[i].get(j))
879 );
880 if n != 0 {
881 let i = rng.random(0..n);
882 if n > 1 && rng.rand(2) != 0 {
883 let j = (i + rng.random(1..n)) % n;
884 a.data[i] = a[j].clone();
885 } else {
886 a[i].reset();
887 }
888 assert!(!a.clone().determinant());
889 assert!(a.inverse().is_none());
890 assert_eq!(a.rank(), n - 1);
891 }
892 }
893 }
894}