use super::*;
#[test]
fn test_total_memory_region_contiguous() {
let dims = [100usize];
let strides = [8isize]; let byte_strides: Vec<&[isize]> = vec![&strides];
let region = total_memory_region(&dims, &byte_strides);
assert_eq!(region, 832);
}
#[test]
fn test_total_memory_region_strided() {
let dims = [10usize];
let strides = [128isize]; let byte_strides: Vec<&[isize]> = vec![&strides];
let region = total_memory_region(&dims, &byte_strides);
assert_eq!(region, 640);
}
#[test]
fn test_compute_blocks_small() {
let dims = [10usize, 10];
let costs = [2isize, 2];
let strides = [8isize, 80];
let orders = [1usize, 2];
let byte_strides: Vec<&[isize]> = vec![&strides];
let stride_orders: Vec<&[usize]> = vec![&orders];
let blocks = compute_blocks(
&dims,
&costs,
&byte_strides,
&stride_orders,
BLOCK_MEMORY_SIZE,
);
assert_eq!(blocks, vec![10, 10]);
}
#[test]
fn test_compute_blocks_large() {
let dims = [1000usize, 1000];
let costs = [2isize, 2];
let strides = [8isize, 8000]; let orders = [1usize, 2];
let byte_strides: Vec<&[isize]> = vec![&strides];
let stride_orders: Vec<&[usize]> = vec![&orders];
let blocks = compute_blocks(
&dims,
&costs,
&byte_strides,
&stride_orders,
BLOCK_MEMORY_SIZE,
);
assert!(blocks[0] <= dims[0]);
assert!(blocks[1] <= dims[1]);
assert!(blocks[0] >= 1);
assert!(blocks[1] >= 1);
}
#[test]
fn test_last_argmax_weighted() {
let blocks = [10usize, 20, 5];
let costs = [1isize, 1, 2];
let idx = last_argmax_weighted(&blocks, &costs);
assert_eq!(idx, Some(1));
}
#[test]
fn test_last_argmax_weighted_tie() {
let blocks = [10usize, 10];
let costs = [1isize, 1];
let idx = last_argmax_weighted(&blocks, &costs);
assert_eq!(idx, Some(1));
}
#[test]
fn test_compute_block_sizes_full_pipeline() {
let dims = [100usize, 100];
let order = [0usize, 1];
let strides = [1isize, 100];
let strides_list: Vec<&[isize]> = vec![&strides];
let blocks = compute_block_sizes(&dims, &order, &strides_list, 8);
assert_eq!(blocks.len(), 2);
assert!(blocks[0] >= 1 && blocks[0] <= 100);
assert!(blocks[1] >= 1 && blocks[1] <= 100);
}
#[test]
fn test_total_memory_region_julia_match_2d() {
let dims = [10usize, 10];
let strides = [8isize, 80];
let byte_strides: Vec<&[isize]> = vec![&strides];
let region = total_memory_region(&dims, &byte_strides);
assert_eq!(region, 1280);
}
#[test]
fn test_total_memory_region_julia_match_multiple_arrays() {
let dims = [10usize, 10];
let strides1 = [8isize, 80]; let strides2 = [80isize, 8]; let byte_strides: Vec<&[isize]> = vec![&strides1, &strides2];
let region = total_memory_region(&dims, &byte_strides);
assert_eq!(region, 2560);
}
#[test]
fn test_total_memory_region_all_contiguous() {
let dims = [10usize, 5];
let strides = [8isize, 40]; let byte_strides: Vec<&[isize]> = vec![&strides];
let region = total_memory_region(&dims, &byte_strides);
assert_eq!(region, 256);
}
#[test]
fn test_total_memory_region_all_large_strides() {
let dims = [5usize, 4];
let strides = [64isize, 320]; let byte_strides: Vec<&[isize]> = vec![&strides];
let region = total_memory_region(&dims, &byte_strides);
assert_eq!(region, 1280);
}
#[test]
fn test_compute_blocks_first_dim_smallest_stride() {
let dims = [100usize, 10, 10];
let costs = [2isize, 2, 2];
let strides1 = [8isize, 800, 8000];
let orders1 = [1usize, 2, 3]; let byte_strides: Vec<&[isize]> = vec![&strides1];
let stride_orders: Vec<&[usize]> = vec![&orders1];
let blocks = compute_blocks(
&dims,
&costs,
&byte_strides,
&stride_orders,
BLOCK_MEMORY_SIZE,
);
assert_eq!(blocks[0], 100);
}
#[test]
fn test_compute_blocks_min_stride_larger_than_blocksize() {
let dims = [100usize, 100];
let costs = [2isize, 2];
let strides = [40000000isize, 40000]; let orders = [2usize, 1]; let byte_strides: Vec<&[isize]> = vec![&strides];
let stride_orders: Vec<&[usize]> = vec![&orders];
let initial_mem = total_memory_region(&dims, &byte_strides);
assert!(
initial_mem > BLOCK_MEMORY_SIZE,
"Initial memory {} should exceed {}",
initial_mem,
BLOCK_MEMORY_SIZE
);
let min_stride = strides.iter().map(|s| s.unsigned_abs()).min().unwrap();
assert!(
min_stride > BLOCK_MEMORY_SIZE,
"Min stride {} should exceed {}",
min_stride,
BLOCK_MEMORY_SIZE
);
let blocks = compute_blocks(
&dims,
&costs,
&byte_strides,
&stride_orders,
BLOCK_MEMORY_SIZE,
);
assert_eq!(blocks, vec![1, 1]);
}
#[test]
fn test_compute_blocks_4d_array_first_dim_smallest() {
let dims = [10usize, 10, 10, 10];
let costs = [2isize, 2, 2, 2];
let strides = [8isize, 80, 800, 8000];
let orders = [1usize, 2, 3, 4];
let byte_strides: Vec<&[isize]> = vec![&strides];
let stride_orders: Vec<&[usize]> = vec![&orders];
let blocks = compute_blocks(
&dims,
&costs,
&byte_strides,
&stride_orders,
BLOCK_MEMORY_SIZE,
);
assert_eq!(blocks.len(), 4);
for (i, &b) in blocks.iter().enumerate() {
assert!(b >= 1 && b <= dims[i], "Block {} out of range", i);
}
}
#[test]
fn test_compute_blocks_4d_needs_reduction() {
let dims = [100usize, 100, 100, 100];
let costs = [2isize, 2, 2, 2];
let strides = [80isize, 8, 8000, 800];
let orders = [2usize, 1, 4, 3]; let byte_strides: Vec<&[isize]> = vec![&strides];
let stride_orders: Vec<&[usize]> = vec![&orders];
let initial_mem = total_memory_region(&dims, &byte_strides);
assert!(initial_mem > BLOCK_MEMORY_SIZE);
let blocks = compute_blocks(
&dims,
&costs,
&byte_strides,
&stride_orders,
BLOCK_MEMORY_SIZE,
);
assert_eq!(blocks.len(), 4);
let total_elements: usize = blocks.iter().product();
let original_elements: usize = dims.iter().product();
assert!(
total_elements < original_elements,
"Blocks should be smaller than original"
);
}
#[test]
fn test_compute_blocks_4d_permuted() {
let dims = [10usize, 10, 10, 10];
let costs = [2isize, 4, 8, 16]; let strides = [8000isize, 800, 80, 8];
let orders = [4usize, 3, 2, 1];
let byte_strides: Vec<&[isize]> = vec![&strides];
let stride_orders: Vec<&[usize]> = vec![&orders];
let blocks = compute_blocks(
&dims,
&costs,
&byte_strides,
&stride_orders,
BLOCK_MEMORY_SIZE,
);
assert_eq!(blocks.len(), 4);
for (i, &b) in blocks.iter().enumerate() {
assert!(b >= 1 && b <= dims[i]);
}
}
#[test]
fn test_compute_blocks_mixed_strides_two_arrays() {
let dims = [100usize, 100];
let costs = [2isize, 2];
let strides1 = [8isize, 800]; let strides2 = [800isize, 8]; let orders1 = [1usize, 2];
let orders2 = [2usize, 1];
let byte_strides: Vec<&[isize]> = vec![&strides1, &strides2];
let stride_orders: Vec<&[usize]> = vec![&orders1, &orders2];
let blocks = compute_blocks(
&dims,
&costs,
&byte_strides,
&stride_orders,
BLOCK_MEMORY_SIZE,
);
assert_eq!(blocks.len(), 2);
assert!(blocks[0] >= 1 && blocks[0] <= 100);
assert!(blocks[1] >= 1 && blocks[1] <= 100);
}
#[test]
fn test_last_argmax_weighted_all_ones() {
let blocks = [1usize, 1, 1];
let costs = [1isize, 2, 3];
let idx = last_argmax_weighted(&blocks, &costs);
assert_eq!(idx, None);
}
#[test]
fn test_last_argmax_weighted_mixed() {
let blocks = [1usize, 5, 3];
let costs = [100isize, 1, 1];
let idx = last_argmax_weighted(&blocks, &costs);
assert_eq!(idx, Some(1));
}
#[test]
fn test_last_argmax_weighted_cost_matters() {
let blocks = [3usize, 10];
let costs = [10isize, 1];
let idx = last_argmax_weighted(&blocks, &costs);
assert_eq!(idx, Some(0));
}
#[test]
fn test_compute_blocks_negative_strides() {
let dims = [10usize, 10];
let costs = [2isize, 2];
let strides = [-8isize, -80]; let orders = [1usize, 2];
let byte_strides: Vec<&[isize]> = vec![&strides];
let stride_orders: Vec<&[usize]> = vec![&orders];
let blocks = compute_blocks(
&dims,
&costs,
&byte_strides,
&stride_orders,
BLOCK_MEMORY_SIZE,
);
assert_eq!(blocks.len(), 2);
assert!(blocks[0] >= 1 && blocks[0] <= 10);
assert!(blocks[1] >= 1 && blocks[1] <= 10);
}
#[test]
fn test_compute_block_sizes_4d_column_major() {
let dims = [32usize, 32, 32, 32];
let order = [0usize, 1, 2, 3]; let strides = [1isize, 32, 1024, 32768]; let strides_list: Vec<&[isize]> = vec![&strides];
let blocks = compute_block_sizes(&dims, &order, &strides_list, 8);
assert_eq!(blocks.len(), 4);
for (i, &b) in blocks.iter().enumerate() {
assert!(b >= 1 && b <= dims[i], "Block {} = {} out of range", i, b);
}
}
#[test]
fn test_compute_block_sizes_4d_permuted_strides() {
let dims = [32usize, 32, 32, 32];
let order = [3usize, 2, 1, 0]; let strides = [32768isize, 1024, 32, 1];
let strides_list: Vec<&[isize]> = vec![&strides];
let blocks = compute_block_sizes(&dims, &order, &strides_list, 8);
assert_eq!(blocks.len(), 4);
for (i, &b) in blocks.iter().enumerate() {
assert!(b >= 1, "Block {} must be >= 1", i);
}
}