use crate::fuse::compute_costs;
use crate::{BLOCK_MEMORY_SIZE, CACHE_LINE_SIZE};
use strided_view::auxiliary::index_order;
pub(crate) fn compute_block_sizes(
dims: &[usize],
order: &[usize],
strides_list: &[&[isize]],
elem_size: usize,
) -> Vec<usize> {
if order.is_empty() {
return Vec::new();
}
let ordered_dims: Vec<usize> = order.iter().map(|&i| dims[i]).collect();
let byte_strides: Vec<Vec<isize>> = strides_list
.iter()
.map(|strides| {
order
.iter()
.map(|&i| strides[i] * elem_size as isize)
.collect()
})
.collect();
let stride_orders: Vec<Vec<usize>> = byte_strides.iter().map(|bs| index_order(bs)).collect();
let reordered_strides: Vec<Vec<isize>> = strides_list
.iter()
.map(|strides| order.iter().map(|&i| strides[i]).collect())
.collect();
let reordered_refs: Vec<&[isize]> = reordered_strides.iter().map(|s| s.as_slice()).collect();
let costs = compute_costs(&reordered_refs);
let byte_stride_refs: Vec<&[isize]> = byte_strides.iter().map(|s| s.as_slice()).collect();
let stride_order_refs: Vec<&[usize]> = stride_orders.iter().map(|s| s.as_slice()).collect();
compute_blocks(
&ordered_dims,
&costs,
&byte_stride_refs,
&stride_order_refs,
BLOCK_MEMORY_SIZE,
)
}
fn compute_blocks(
dims: &[usize],
costs: &[isize],
byte_strides: &[&[isize]],
stride_orders: &[&[usize]],
block_size: usize,
) -> Vec<usize> {
let n = dims.len();
if n == 0 {
return vec![];
}
if total_memory_region(dims, byte_strides) <= block_size {
return dims.to_vec();
}
let min_order = stride_orders
.iter()
.filter_map(|orders| orders.iter().min().copied())
.min()
.unwrap_or(1);
if stride_orders
.iter()
.all(|orders| !orders.is_empty() && orders[0] == min_order)
{
let tail_dims: Vec<usize> = dims[1..].to_vec();
let tail_costs: Vec<isize> = costs[1..].to_vec();
let tail_byte_strides: Vec<&[isize]> = byte_strides.iter().map(|s| &s[1..]).collect();
let tail_stride_orders: Vec<&[usize]> = stride_orders.iter().map(|s| &s[1..]).collect();
let tail_blocks = compute_blocks(
&tail_dims,
&tail_costs,
&tail_byte_strides,
&tail_stride_orders,
block_size,
);
let mut result = vec![dims[0]];
result.extend(tail_blocks);
return result;
}
let min_stride = byte_strides
.iter()
.filter_map(|s| s.iter().map(|x| x.unsigned_abs()).min())
.min()
.unwrap_or(0);
if min_stride > block_size {
return vec![1; n];
}
let mut blocks = dims.to_vec();
while total_memory_region(&blocks, byte_strides) >= 2 * block_size {
let i = last_argmax_weighted(&blocks, costs);
if i.is_none() || blocks[i.unwrap()] <= 1 {
break;
}
let i = i.unwrap();
blocks[i] = blocks[i].div_ceil(2);
}
while total_memory_region(&blocks, byte_strides) > block_size {
let i = last_argmax_weighted(&blocks, costs);
if i.is_none() || blocks[i.unwrap()] <= 1 {
break;
}
let i = i.unwrap();
blocks[i] -= 1;
}
blocks
}
fn total_memory_region(dims: &[usize], byte_strides: &[&[isize]]) -> usize {
let cache_line = CACHE_LINE_SIZE;
let mut memory_region = 0usize;
for strides in byte_strides {
let mut num_contiguous_cache_lines = 0isize;
let mut num_cache_line_blocks = 1usize;
for (&d, &s) in dims.iter().zip(strides.iter()) {
let s_abs = s.unsigned_abs();
if s_abs < cache_line {
num_contiguous_cache_lines += (d.saturating_sub(1) as isize) * (s_abs as isize);
} else {
num_cache_line_blocks *= d;
}
}
let contiguous_lines = (num_contiguous_cache_lines as usize / cache_line) + 1;
memory_region += cache_line * contiguous_lines * num_cache_line_blocks;
}
memory_region
}
fn last_argmax_weighted(blocks: &[usize], costs: &[isize]) -> Option<usize> {
if blocks.is_empty() {
return None;
}
let mut max_score = 0isize;
let mut max_idx = None;
for (i, (&b, &c)) in blocks.iter().zip(costs.iter()).enumerate() {
if b <= 1 {
continue;
}
let score = (b as isize - 1) * c;
if score >= max_score {
max_score = score;
max_idx = Some(i);
}
}
max_idx
}
#[cfg(test)]
mod tests {
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);
}
}
}