use ruda_kernel::dsl as kernel_dsl;
use crate::reduce::{VectorizationMode, launch::VectorizationStrategy};
use ruda_kernel::dsl::ir::HardwareProperties;
use ruda_kernel::dsl::prelude::*;
use ruda_kernel::library::tensor::is_contiguous;
use ruda_kernel::dsl::tensor_vector_size_parallel;
use ruda_kernel::dsl::tensor_vector_size_perpendicular;
pub fn calculate_plane_count_per_ruda(
working_units: usize,
plane_dim: u32,
properties: &HardwareProperties,
) -> u32 {
let plane_count = match properties.num_cpu_cores {
Some(num_cores) => core::cmp::min(num_cores, working_units as u32),
None => {
let plane_count_max = core::cmp::max(1, working_units / plane_dim as usize);
const NUM_PLANE_MAX: u32 = 8u32;
const NUM_PLANE_MAX_LOG2: u32 = NUM_PLANE_MAX.ilog2();
let plane_count_max_log2 =
core::cmp::min(NUM_PLANE_MAX_LOG2, usize::ilog2(plane_count_max));
2u32.pow(plane_count_max_log2)
}
};
let max_plane_per_ruda = properties.max_units_per_ruda / plane_dim;
plane_count.min(max_plane_per_ruda)
}
pub fn generate_vector_size<R: Runtime>(
client: &ComputeClient<R>,
input: &TensorBinding<R>,
output: &TensorBinding<R>,
axis: usize,
dtype: StorageType,
vectorization_mode: VectorizationMode,
strategy: &VectorizationStrategy,
) -> (usize, usize) {
generate_vector_size_with_dtypes(client, input, output, axis, [dtype, dtype], vectorization_mode, strategy)
}
pub(crate) fn generate_vector_size_with_dtypes<R: Runtime>(
client: &ComputeClient<R>,
input: &TensorBinding<R>,
output: &TensorBinding<R>,
axis: usize,
dtypes: [StorageType; 2],
vectorization_mode: VectorizationMode,
strategy: &VectorizationStrategy,
) -> (usize, usize) {
let [dtype, output_dtype] = dtypes;
let vector_size_input = match vectorization_mode {
VectorizationMode::Parallel => tensor_vector_size_parallel(
client.io_optimized_vector_sizes(dtype.size()),
&input.shape,
&input.strides,
axis,
),
VectorizationMode::Perpendicular => {
let mut input_axis_and_strides = input.strides.iter().enumerate().collect::<Vec<_>>();
input_axis_and_strides.sort_by_key(|(_, stride)| *stride);
let input_sorted_axis = input_axis_and_strides
.into_iter()
.map(|(a, _)| a)
.take_while(|a| *a != axis);
let mut output_axis_and_strides = output.strides.iter().enumerate().collect::<Vec<_>>();
output_axis_and_strides.sort_by_key(|(_, stride)| *stride);
let output_sorted_axis = output_axis_and_strides
.into_iter()
.filter_map(|(a, _)| (a != axis).then_some(a));
let max_vector_size = input_sorted_axis
.zip(output_sorted_axis)
.take_while(|(i, o)| i == o)
.map(|(i, _)| output.shape[i])
.product();
match client.properties().hardware.num_cpu_cores.is_some() {
true => {
let supported_vector_sizes =
client.io_optimized_vector_sizes(1).filter(|size| {
*size <= max_vector_size && max_vector_size.is_multiple_of(*size)
});
tensor_vector_size_perpendicular(
supported_vector_sizes,
&input.shape,
&input.strides,
axis,
)
}
false => {
let supported_vector_sizes = client
.io_optimized_vector_sizes(dtype.size())
.filter(|&size| {
size <= max_vector_size && max_vector_size.is_multiple_of(size)
});
tensor_vector_size_perpendicular(
supported_vector_sizes,
&input.shape,
&input.strides,
axis,
)
}
}
}
};
let mut vector_size_output = 1;
let max_run = output_contiguous_run(&output.shape, &output.strides, axis);
if vector_size_input > 1 && vectorization_mode == VectorizationMode::Perpendicular {
let rank = output.strides.len();
let is_contiguous = is_contiguous(&output.shape[axis..rank], &output.strides[axis..rank])
&& output.strides[rank - 1] == 1;
let shape = output.shape.get(axis + 1).copied().unwrap_or(1);
if is_contiguous {
vector_size_output = client
.io_optimized_vector_sizes(output_dtype.size())
.filter(|&vector_size| {
vector_size_input.is_multiple_of(vector_size)
&& shape.is_multiple_of(vector_size)
&& output_vector_is_aligned(&output.shape, &output.strides, axis, max_run, vector_size)
})
.max()
.unwrap_or(1);
}
}
if strategy.parallel_output_vectorization
&& vectorization_mode == VectorizationMode::Parallel
&& vector_size_input > 1
&& is_contiguous(&input.shape, &input.strides)
&& axis == input.shape.len() - 1
{
let supported_vector_sizes = client.io_optimized_vector_sizes(output_dtype.size());
let num_reduce = output.shape.iter().copied().product::<usize>();
vector_size_output = supported_vector_sizes
.filter(|&vector_size| {
num_reduce % vector_size == 0
&& output_vector_is_aligned(&output.shape, &output.strides, axis, max_run, vector_size)
})
.max()
.unwrap_or(1);
}
(vector_size_input, vector_size_output)
}
fn output_vector_is_aligned(
shape: &[usize],
strides: &[usize],
reduce_axis: usize,
run: usize,
vector_size: usize,
) -> bool {
run.is_multiple_of(vector_size)
&& strides.iter().enumerate().all(|(axis, &stride)| {
shape[axis] <= 1
|| (axis != reduce_axis && stride < run)
|| stride.is_multiple_of(vector_size)
})
}
fn output_contiguous_run(shape: &[usize], strides: &[usize], reduce_axis: usize) -> usize {
let mut dims: Vec<(usize, usize)> = (0..strides.len())
.filter(|&d| d != reduce_axis)
.map(|d| (strides[d], shape[d]))
.collect();
dims.sort_by_key(|&(stride, size)| (stride, size));
let mut run = 1;
for (stride, size) in dims {
if stride != run {
break;
}
run *= size;
}
run
}