use crate::kernels::{CubeclKernelError, Result};
pub(crate) fn validate_axis(rank: usize, axis: usize) -> Result<()> {
if axis >= rank {
return Err(CubeclKernelError::InvalidAxis { axis, rank });
}
Ok(())
}
pub(crate) fn keepdims_output_shape(input_shape: &[usize], axis: usize) -> Result<Vec<usize>> {
validate_axis(input_shape.len(), axis)?;
let mut output = input_shape.to_vec();
output[axis] = 1;
Ok(output)
}
pub(crate) fn validate_keepdims_output_shape(
input_shape: &[usize],
output_shape: &[usize],
axis: usize,
) -> Result<()> {
let expected = keepdims_output_shape(input_shape, axis)?;
if output_shape != expected {
return Err(CubeclKernelError::MismatchOutputShape {
expected,
actual: output_shape.to_vec(),
});
}
Ok(())
}