Skip to main content

competitive/math/
bit_matrix.rs

1use super::BitSet;
2#[cfg(target_arch = "x86_64")]
3use super::{SimdBackend, simd_backend};
4use std::ops::{BitXorAssign, Index, IndexMut, Mul};
5
6/// A matrix over GF(2), stored as packed rows.
7#[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    /// Replaces the matrix with reduced row echelon form and returns its pivot columns.
64    pub fn row_reduction(&mut self) -> Vec<usize> {
65        self.eliminate(self.shape.1, true, false)
66    }
67
68    /// Replaces the matrix with row echelon form and returns its rank.
69    pub fn rank(&mut self) -> usize {
70        self.eliminate(self.shape.1, false, false).len()
71    }
72
73    /// Computes the determinant in place. The matrix must be square.
74    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    /// Returns the inverse, or `None` if singular. The matrix must be square.
80    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    /// Returns a particular solution and a basis of the kernel, or `None` if inconsistent.
103    /// `b.len()` must equal the number of rows.
104    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            // SAFETY: the dispatcher checks the required CPU features.
159            SimdBackend::Avx512 => {
160                return unsafe { self.eliminate_avx512(cols, full, require_full_rank) };
161            }
162            // SAFETY: the dispatcher checks AVX2 support.
163            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    // Inlined into each target-feature entry point to vectorize the row operations.
172    #[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        // Eliminating a shared leading coefficient preserves the two-coefficient bound.
207        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            // Make the panel's pivot columns an identity matrix before indexing its table.
268            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        // Separate the table groups by a cache line to avoid mapping them to the same sets.
434        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            // SAFETY: the dispatcher checks the required CPU features.
555            SimdBackend::Avx512 => return unsafe { self.mul_avx512(rhs) },
556            // SAFETY: the dispatcher checks AVX2 support.
557            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            // Enumerate the image and every solution independently of elimination.
588            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            // Sample different scales and aspect ratios while bounding the Boolean oracle's cost.
644            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            // Reduce two right-hand sides with the same independent Boolean operations:
732            // one known to be consistent and one unrestricted random vector.
733            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                    // SAFETY: AVX2 support was checked above.
800                    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                    // SAFETY: all required AVX-512 features were checked above.
810                    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}