use super::*;
fn check_sort_last_dim(rows: usize, cols: usize) {
let n = rows * cols;
let src: Vec<f32> = (0..n)
.map(|i| ((i * 1664525 + 1013904223) % 1000) as f32)
.collect();
let mut data = src.clone();
let shape = Shape::new([rows, cols]);
sort_along_dim(&mut data, &shape, 1, false, f32::total_cmp);
for r in 0..rows {
let row = &data[r * cols..(r + 1) * cols];
for w in row.windows(2) {
assert!(w[0] <= w[1], "row {r} not sorted: {:?}", row);
}
let mut expected: Vec<f32> = src[r * cols..(r + 1) * cols].to_vec();
expected.sort_unstable_by(f32::total_cmp);
assert_eq!(row, expected.as_slice());
}
}
#[test]
fn sort_along_last_dim_small_serial() {
check_sort_last_dim(64, 64);
}
#[cfg(feature = "rayon")]
#[test]
fn sort_along_last_dim_large_parallel() {
let cols = 1024;
let rows = (PARALLEL_THRESHOLD / cols) + 1;
check_sort_last_dim(rows, cols);
}
#[test]
fn sort_along_last_dim_descending() {
let mut data: Vec<f32> = (0..4096).map(|i| (i % 17) as f32).collect();
let shape = Shape::new([128, 32]);
sort_along_dim(&mut data, &shape, 1, true, f32::total_cmp);
for r in 0..128 {
let row = &data[r * 32..(r + 1) * 32];
for w in row.windows(2) {
assert!(w[0] >= w[1]);
}
}
}
fn check_sort_with_indices_last_dim(rows: usize, cols: usize, descending: bool) {
let src: Vec<f32> = (0..rows * cols).map(|i| (i as f32 * 0.37).sin()).collect();
let mut values = src.clone();
let mut indices = vec![0isize; rows * cols];
let shape = Shape::new([rows, cols]);
sort_along_dim_with_indices(
&mut values,
&mut indices,
&shape,
1,
descending,
f32::total_cmp,
);
for r in 0..rows {
let vs = &values[r * cols..(r + 1) * cols];
let idx_row = &indices[r * cols..(r + 1) * cols];
let orig = &src[r * cols..(r + 1) * cols];
let want_order = if descending {
core::cmp::Ordering::Less
} else {
core::cmp::Ordering::Greater
};
for w in vs.windows(2) {
assert_ne!(f32::total_cmp(&w[0], &w[1]), want_order);
}
let mut seen = vec![false; cols];
for (i, &orig_idx) in idx_row.iter().enumerate() {
let j = orig_idx as usize;
assert_eq!(vs[i], orig[j]);
assert!(!seen[j], "row {r}: index {j} repeated");
seen[j] = true;
}
}
}
#[test]
fn sort_with_indices_last_dim_ascending() {
check_sort_with_indices_last_dim(512, 512, false);
}
#[test]
fn sort_with_indices_last_dim_descending() {
check_sort_with_indices_last_dim(512, 512, true);
}
fn check_argsort_last_dim(rows: usize, cols: usize, descending: bool) {
let src: Vec<f32> = (0..rows * cols)
.map(|i| ((i * 7919) % 997) as f32)
.collect();
let mut indices = vec![0isize; rows * cols];
let shape = Shape::new([rows, cols]);
argsort_along_dim(&src, &mut indices, &shape, 1, descending, f32::total_cmp);
for r in 0..rows {
let idx_row = &indices[r * cols..(r + 1) * cols];
let orig = &src[r * cols..(r + 1) * cols];
let sorted: Vec<f32> = idx_row.iter().map(|&i| orig[i as usize]).collect();
for w in sorted.windows(2) {
if descending {
assert!(w[0] >= w[1]);
} else {
assert!(w[0] <= w[1]);
}
}
let mut seen = vec![false; cols];
for &i in idx_row {
let j = i as usize;
assert!(!seen[j], "row {r}: index {j} repeated");
seen[j] = true;
}
}
}
#[test]
fn argsort_last_dim_ascending() {
check_argsort_last_dim(200, 1500, false);
}
#[test]
fn argsort_last_dim_descending() {
check_argsort_last_dim(200, 1500, true);
}