use cubecl::zspace::Strides;
#[derive(Debug, PartialEq, Eq, Clone, Copy, Hash)]
pub enum VectorizationMode {
Parallel,
Perpendicular,
}
pub fn output_vectorization_axis(
input_strides: &Strides,
reduce_axis: usize,
_vectorization_mode: VectorizationMode,
) -> usize {
if input_strides.len() < 2 {
return 0;
}
let mut min1 = (usize::MAX, 0); let mut min2 = (usize::MAX, 0);
for (i, &s) in input_strides.iter().enumerate() {
if s < min1.0 {
min2 = min1;
min1 = (s, i);
} else if s < min2.0 {
min2 = (s, i);
}
}
if min1.1 == reduce_axis {
min2.1
} else {
min1.1
}
}
#[derive(Debug, PartialEq, Eq, Clone, Copy, Hash)]
pub enum BoundChecks {
None,
Mask,
Branch,
}
impl BoundChecks {
pub fn idle(self) -> Self {
Self::Mask
}
}
#[derive(Debug, PartialEq, Eq, Clone, Copy, Hash)]
pub enum IdleMode {
None,
Mask,
Terminate,
}
impl IdleMode {
pub fn is_enabled(&self) -> bool {
!matches!(self, Self::None)
}
}