use core::borrow::BorrowMut;
use p3_maybe_rayon::PARALLEL_ENABLED;
use p3_maybe_rayon::prelude::*;
use p3_util::{log2_strict_usize, reverse_bits_len, reverse_slice_index_bits};
use tracing::instrument;
use crate::Matrix;
use crate::dense::{DenseMatrix, DenseStorage, RowMajorMatrix};
#[instrument(level = "debug", skip_all)]
pub fn reverse_matrix_index_bits<'a, F, S>(mat: &mut DenseMatrix<F, S>)
where
F: Clone + Send + Sync + 'a,
S: DenseStorage<F> + BorrowMut<[F]>,
{
let w = mat.width();
let h = mat.height();
let log_h = log2_strict_usize(h);
let values = mat.values.borrow_mut();
let total_bytes = core::mem::size_of_val(values);
let use_slice = w == 1 && core::mem::size_of::<F>() <= 8;
#[cfg(target_arch = "aarch64")]
let use_slice = use_slice && (total_bytes <= 64 * 1024 || current_num_threads() == 1);
if use_slice {
reverse_slice_index_bits(values);
return;
}
let values = values.as_mut_ptr() as usize;
let swap = |i, use_outline| {
let values = values as *mut F;
let j = reverse_bits_len(i, log_h);
if i < j {
unsafe { swap_rows_raw(values, w, i, j, use_outline) };
}
};
if total_bytes <= 32 * 1024 {
(0..h).for_each(|i| swap(i, true));
} else {
const CHUNK_BYTES: usize = 64 * 1024;
let rows_per_chunk = CHUNK_BYTES.div_ceil(w * core::mem::size_of::<F>());
(0..h.div_ceil(rows_per_chunk))
.into_par_iter()
.for_each(|chunk| {
let start = chunk * rows_per_chunk;
let end = start + rows_per_chunk.min(h - start);
for i in start..end {
swap(i, !PARALLEL_ENABLED);
}
});
}
}
pub fn swap_rows<F: Clone + Send + Sync>(mat: &mut RowMajorMatrix<F>, i: usize, j: usize) {
let w = mat.width();
let (upper, lower) = mat.values.split_at_mut(j * w);
let row_i = &mut upper[i * w..(i + 1) * w];
let row_j = &mut lower[..w];
row_i.swap_with_slice(row_j);
}
unsafe fn swap_rows_raw<F>(mat: *mut F, w: usize, i: usize, j: usize, use_outline: bool) {
unsafe {
let row_i = core::slice::from_raw_parts_mut(mat.add(i * w), w);
let row_j = core::slice::from_raw_parts_mut(mat.add(j * w), w);
if use_outline && core::mem::size_of_val(row_i) >= 64 {
swap_row_slices(row_i, row_j);
} else {
row_i.swap_with_slice(row_j);
}
}
}
#[inline(never)]
fn swap_row_slices<F>(row_i: &mut [F], row_j: &mut [F]) {
row_i.swap_with_slice(row_j);
}
#[cfg(test)]
mod tests {
use alloc::vec;
use super::*;
use crate::dense::RowMajorMatrix;
#[test]
fn test_swap_rows_basic() {
let mut matrix = RowMajorMatrix::new(
vec![
1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, ],
3,
);
swap_rows(&mut matrix, 0, 2);
assert_eq!(
matrix.values,
vec![
7, 8, 9, 4, 5, 6, 1, 2, 3, 10, 11, 12, ]
);
}
#[test]
fn test_swap_rows_raw_basic() {
let mut matrix = RowMajorMatrix::new(
vec![
1, 2, 3, 4, 5, 6, 7, 8, 9, ],
3,
);
let ptr = matrix.values.as_mut_ptr();
unsafe {
swap_rows_raw(ptr, matrix.width(), 0, 2, false);
}
assert_eq!(
matrix.values,
vec![
7, 8, 9, 4, 5, 6, 1, 2, 3, ]
);
}
#[test]
fn test_reverse_matrix_index_bits_pow2_height() {
let mut matrix = RowMajorMatrix::new(
vec![
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, ],
2,
);
reverse_matrix_index_bits(&mut matrix);
assert_eq!(
matrix.values,
vec![
0, 1, 8, 9, 4, 5, 12, 13, 2, 3, 10, 11, 6, 7, 14, 15, ]
);
}
#[test]
fn test_reverse_matrix_index_bits_strings() {
use alloc::string::ToString;
use alloc::vec::Vec;
for width in [1, 3, 17] {
for log_h in 0..=14 {
let height = 1 << log_h;
let original: Vec<_> = (0..height * width).map(|i| i.to_string()).collect();
let expected: Vec<_> = (0..height)
.flat_map(|row| {
let start = reverse_bits_len(row, log_h) * width;
original[start..start + width].iter().cloned()
})
.collect();
let mut matrix = RowMajorMatrix::new(original.clone(), width);
reverse_matrix_index_bits(&mut matrix);
assert_eq!(matrix.values, expected, "width={width}, log_h={log_h}");
reverse_matrix_index_bits(&mut matrix);
assert_eq!(
matrix.values, original,
"involution: width={width}, log_h={log_h}"
);
}
}
}
#[test]
fn test_reverse_matrix_index_bits_preserves_owners_without_cloning() {
use alloc::vec::Vec;
use core::sync::atomic::{AtomicUsize, Ordering};
struct Owner<'a>(&'a AtomicUsize);
impl Clone for Owner<'_> {
fn clone(&self) -> Self {
panic!("bit reversal must not clone elements");
}
}
impl Drop for Owner<'_> {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
for width in [1, 3, 17] {
for log_h in 0..=15 {
let height = 1 << log_h;
let drops: Vec<_> = (0..height * width).map(|_| AtomicUsize::new(0)).collect();
let mut matrix = RowMajorMatrix::new(drops.iter().map(Owner).collect(), width);
reverse_matrix_index_bits(&mut matrix);
for (i, value) in matrix.values.iter().enumerate() {
let source = reverse_bits_len(i / width, log_h) * width + i % width;
assert!(core::ptr::eq(value.0, &drops[source]));
}
reverse_matrix_index_bits(&mut matrix);
for (value, count) in matrix.values.iter().zip(&drops) {
assert!(core::ptr::eq(value.0, count));
assert_eq!(count.load(Ordering::Relaxed), 0);
}
drop(matrix);
assert!(drops.iter().all(|count| count.load(Ordering::Relaxed) == 1));
}
}
}
#[test]
fn test_reverse_matrix_index_bits_work_unit_boundaries() {
use alloc::vec::Vec;
for (width, height) in [(3, 8192), (16_383, 8), (16_384, 8), (16_385, 8)] {
let original: Vec<u32> = (0..width * height).map(|i| i as u32).collect();
let mut matrix = RowMajorMatrix::new(original.clone(), width);
reverse_matrix_index_bits(&mut matrix.as_view_mut());
for (row, values) in matrix.values.chunks_exact(width).enumerate() {
let source = row.reverse_bits() >> (usize::BITS - height.ilog2());
assert_eq!(
values,
&original[source * width..(source + 1) * width],
"width={width}, height={height}, row={row}"
);
}
reverse_matrix_index_bits(&mut matrix);
assert_eq!(matrix.values, original, "width={width}, height={height}");
}
}
#[test]
fn test_reverse_matrix_index_bits_zero_sized_elements() {
let mut matrix = RowMajorMatrix::new(vec![(); 3 * 8192], 3);
reverse_matrix_index_bits(&mut matrix);
assert_eq!(matrix.width(), 3);
assert_eq!(matrix.height(), 8192);
assert_eq!(matrix.values.len(), 3 * 8192);
}
#[test]
fn test_reverse_matrix_index_bits_height_1() {
let mut matrix = RowMajorMatrix::new(
vec![
42, 43, ],
2,
);
reverse_matrix_index_bits(&mut matrix);
assert_eq!(
matrix.values,
vec![
42, 43, ]
);
}
#[test]
#[should_panic]
fn test_reverse_matrix_index_bits_non_power_of_two_should_panic() {
let mut matrix = RowMajorMatrix::new(
vec![
1, 2, 3, 4, 5, 6, ],
2,
);
reverse_matrix_index_bits(&mut matrix);
}
}