#![allow(
clippy::type_complexity,
reason = "Too sensitive, triggers on tuple of vector."
)]
pub mod components;
pub mod launch;
pub mod routines;
mod error;
#[cfg(any(feature = "cpu-reference", feature = "benchmarks"))]
pub mod eval;
pub use crate::launch::{ReduceStrategy, ReduceWithIndicesDtypes};
use crate::{
components::instructions::ReduceOperationConfig,
launch::{launch_reduce, launch_reduce_with_indices},
};
pub use components::{
args::init_tensors,
config::*,
instructions::{ReduceFamily, ReduceInstruction},
precision::ReducePrecision,
};
use cubecl::prelude::*;
pub use error::*;
pub use launch::{ReduceDtypes, reduce_kernel};
pub use routines::shared_sum::shared_sum;
pub fn reduce<R: Runtime>(
client: &ComputeClient<R>,
input: TensorBinding<R>,
output: TensorBinding<R>,
axis: usize,
strategy: ReduceStrategy,
operation: ReduceOperationConfig,
dtypes: ReduceDtypes,
) -> Result<(), ReduceError> {
validate_axis(input.shape.len(), axis)?;
validate_shapes(
&input.shape,
&output.shape,
axis,
match operation {
ReduceOperationConfig::ArgTopK(k) => Some(k),
ReduceOperationConfig::TopK(k) => Some(k),
_ => None,
},
)?;
launch_reduce::<R>(client, input, output, axis, strategy, dtypes, operation)
}
#[allow(clippy::too_many_arguments)]
pub fn reduce_with_indices<R: Runtime>(
client: &ComputeClient<R>,
input: TensorBinding<R>,
values: TensorBinding<R>,
indices: TensorBinding<R>,
axis: usize,
strategy: ReduceStrategy,
operation: ReduceOperationConfig,
dtypes: ReduceWithIndicesDtypes,
) -> Result<(), ReduceError> {
let k = match operation {
ReduceOperationConfig::TopK(k) | ReduceOperationConfig::ArgTopK(k) => k,
ReduceOperationConfig::Max
| ReduceOperationConfig::ArgMax
| ReduceOperationConfig::Min
| ReduceOperationConfig::ArgMin => 1,
other => {
return Err(ReduceError::IndicesUnsupported {
operation: operation_name(&other),
});
}
};
validate_axis(input.shape.len(), axis)?;
validate_shapes(&input.shape, &values.shape, axis, Some(k))?;
if indices.shape.as_slice() != values.shape.as_slice() {
return Err(ReduceError::MismatchIndicesShape {
values_shape: values.shape.to_vec(),
indices_shape: indices.shape.to_vec(),
});
}
if indices.strides != values.strides {
return Err(ReduceError::MismatchIndicesStrides {
values_strides: values.strides.to_vec(),
indices_strides: indices.strides.to_vec(),
});
}
launch_reduce_with_indices::<R>(
client, input, values, indices, axis, strategy, dtypes, operation,
)
}
fn operation_name(operation: &ReduceOperationConfig) -> &'static str {
match operation {
ReduceOperationConfig::Sum => "Sum",
ReduceOperationConfig::Prod => "Prod",
ReduceOperationConfig::Mean => "Mean",
ReduceOperationConfig::MaxAbs => "MaxAbs",
ReduceOperationConfig::ArgMax => "ArgMax",
ReduceOperationConfig::ArgMin => "ArgMin",
ReduceOperationConfig::Max => "Max",
ReduceOperationConfig::Min => "Min",
ReduceOperationConfig::ArgTopK(_) => "ArgTopK",
ReduceOperationConfig::TopK(_) => "TopK",
ReduceOperationConfig::Any => "Any",
ReduceOperationConfig::All => "All",
}
}
fn validate_axis(rank: usize, axis: usize) -> Result<(), ReduceError> {
if axis > rank {
return Err(ReduceError::InvalidAxis { axis, rank });
}
Ok(())
}
fn validate_shapes(
input_shape: &[usize],
output_shape: &[usize],
axis: usize,
k: Option<usize>,
) -> Result<(), ReduceError> {
let mut expected_shape = input_shape.to_vec();
let k = k.unwrap_or(1);
if expected_shape[axis] < k {
return Err(ReduceError::ReduceAxisTooSmall {
axis_length: expected_shape[axis],
k,
});
}
expected_shape[axis] = k;
if output_shape != expected_shape {
return Err(ReduceError::MismatchOutputShape {
expected_shape,
output_shape: output_shape.to_vec(),
});
}
Ok(())
}