#[cfg(feature = "autotune")]
use super::{autotune_reduce, autotune_reduce_with_indices, autotune_sum};
use crate::ops::permute;
use crate::{
ops::numeric::{empty_device_contiguous_dtype, fill_device_dtype, zeros_client},
tensor::CubeTensor,
};
use burn_backend::cubecl::{dtype_to_elem_type, dtype_to_storage_type, elem_type_to_dtype};
use burn_backend::{DType, TensorMetadata};
use burn_std::{BoolDType, Metadata};
use burn_std::{Shape, Strides};
use cubecl::{AutotuneKey, client::Client, features::AtomicUsage, ir::Type, prelude::InputScalar};
use cubek::reduce::{
ReduceDtypes, ReduceError, ReduceStrategy, ReduceWithIndicesDtypes,
components::instructions::ReduceOperationConfig,
launch::{RoutineStrategy, VectorizationStrategy},
routines::{BlueprintStrategy, unit::UnitStrategy},
shared_sum,
};
use serde::{Deserialize, Serialize};
#[derive(Hash, Eq, PartialEq, Debug, Clone, Serialize, Deserialize, AutotuneKey)]
pub struct SumAutotuneKey {
dtype: burn_backend::DType,
#[autotune(anchor)]
length: usize,
}
fn empty_reduce_identity(config: ReduceOperationConfig, dtype: DType) -> Option<f64> {
match config {
ReduceOperationConfig::Sum | ReduceOperationConfig::Any => Some(0.0),
ReduceOperationConfig::Prod | ReduceOperationConfig::All => Some(1.0),
ReduceOperationConfig::Mean => dtype.is_float().then_some(f64::NAN),
ReduceOperationConfig::Max
| ReduceOperationConfig::Min
| ReduceOperationConfig::MaxAbs
| ReduceOperationConfig::TopK(_)
| ReduceOperationConfig::ArgMax
| ReduceOperationConfig::ArgMin
| ReduceOperationConfig::ArgTopK(_) => None,
}
}
fn reduce_empty_axis(
output: CubeTensor,
axis_length: usize,
config: ReduceOperationConfig,
) -> Result<CubeTensor, ReduceError> {
let identity =
empty_reduce_identity(config, output.dtype).ok_or(ReduceError::ReduceAxisTooSmall {
axis_length,
k: accumulator_len(config),
})?;
if output.meta.num_elements() == 0 {
return Ok(output);
}
let identity = InputScalar::new(identity, dtype_to_storage_type(output.dtype));
Ok(fill_device_dtype(output, identity))
}
fn supports_atomic_add(client: &Client, dtype: DType) -> bool {
client
.properties()
.atomic_type_usage(Type::atomic(dtype_to_elem_type(dtype)))
.contains(AtomicUsage::Add)
}
pub fn sum_fallback(
tensor: CubeTensor,
mut strategy: SumStrategy,
) -> Result<CubeTensor, ReduceError> {
if matches!(strategy, SumStrategy::OneShot(_))
&& !supports_atomic_add(&tensor.client, tensor.dtype)
{
strategy = SumStrategy::Chained(Default::default());
}
sum(tensor, strategy)
}
pub fn sum(tensor: CubeTensor, strategy: SumStrategy) -> Result<CubeTensor, ReduceError> {
let client = tensor.client.clone();
let device = tensor.device.clone();
if tensor.meta.num_elements() == 0 {
return Ok(zeros_client(client, device, [1].into(), tensor.dtype));
}
match strategy {
SumStrategy::OneShot(cube_count) => {
let output = zeros_client(client.clone(), device, [1].into(), tensor.dtype);
let dtype = tensor.dtype;
shared_sum(
&client,
tensor.binding(),
output.clone().binding(),
cube_count,
dtype_to_elem_type(dtype),
)?;
Ok(output)
}
SumStrategy::Chained(strategy) => {
reduce(tensor, None, strategy, ReduceOperationConfig::Sum)
}
#[cfg(feature = "autotune")]
SumStrategy::Autotune => Ok(autotune_sum(&client, tensor)),
}
}
pub enum SumStrategy {
OneShot(u32),
Chained(KernelReduceStrategy),
#[cfg(feature = "autotune")]
Autotune,
}
impl Default for SumStrategy {
fn default() -> Self {
#[cfg(feature = "autotune")]
return Self::Autotune;
#[cfg(not(feature = "autotune"))]
return Self::OneShot(4);
}
}
pub fn reduce(
mut tensor: CubeTensor,
output_dtype: Option<DType>,
strategy: KernelReduceStrategy,
config: ReduceOperationConfig,
) -> Result<CubeTensor, cubek::reduce::ReduceError> {
let sorted_axis = argsort(tensor.meta.shape());
for axis in sorted_axis {
tensor = reduce_dim(tensor, output_dtype, axis, strategy.clone(), config)?;
}
*tensor.meta = Metadata::new([1], [1]);
Ok(tensor)
}
pub fn reduce_dims(
input: CubeTensor,
output_dtype: Option<DType>,
dims: &[usize],
strategy: KernelReduceStrategy,
config: ReduceOperationConfig,
) -> Result<CubeTensor, ReduceError> {
let rank = input.meta.num_dims();
let mut shape = input.meta.shape().clone();
let reduced: Vec<usize> = (0..rank).filter(|dim| dims.contains(dim)).collect();
let empty: Vec<usize> = reduced
.iter()
.copied()
.filter(|dim| shape[*dim] == 0)
.collect();
let mut left: Vec<usize> = reduced
.iter()
.copied()
.filter(|dim| shape[*dim] > 1)
.collect();
if empty.is_empty() {
match (reduced.first(), left.len()) {
(None, _) => return Ok(input),
(Some(&dim), 0) => return reduce_dim(input, output_dtype, dim, strategy, config),
(_, 1) => return reduce_dim(input, output_dtype, left[0], strategy, config),
_ => {}
}
}
let mut tensor = input;
for dim in empty {
tensor = reduce_dim(tensor, output_dtype, dim, strategy.clone(), config)?;
shape[dim] = 1;
}
while !left.is_empty() {
let run =
largest_run_memory_holds_together(tensor.meta.shape(), tensor.meta.strides(), &left);
let rest = (0..rank).filter(|dim| !run.contains(dim));
let presented_dims: Vec<usize> = run.iter().copied().chain(rest).collect();
let mut presented = permute(tensor, &presented_dims);
fold_leading_dims(&mut presented, run.len());
tensor = reduce_dim(presented, output_dtype, 0, strategy.clone(), config)?;
for dim in &run {
shape[*dim] = 1;
}
*tensor.meta = Metadata::new(shape.clone(), burn_std::tensor::contiguous_strides(&shape));
left.retain(|dim| !run.contains(dim));
}
Ok(tensor)
}
fn largest_run_memory_holds_together(
shape: &Shape,
strides: &Strides,
left: &[usize],
) -> Vec<usize> {
let mut memory_order: Vec<usize> = (0..shape.num_dims()).collect();
memory_order.sort_by(|a, b| strides[*b].cmp(&strides[*a]).then(a.cmp(b)));
let elements = |run: &[usize]| run.iter().map(|dim| shape[*dim]).product::<usize>();
let mut largest: Vec<usize> = Vec::new();
let mut run: Vec<usize> = Vec::new();
for dim in memory_order {
if shape[dim] == 1 {
continue;
}
if !left.contains(&dim) {
run.clear();
continue;
}
if let Some(&outside) = run.last()
&& strides[outside] != strides[dim] * shape[dim]
{
run.clear();
}
run.push(dim);
if elements(&run) > elements(&largest) {
largest.clone_from(&run);
}
}
largest
}
fn fold_leading_dims(tensor: &mut CubeTensor, count: usize) {
let shape = tensor.meta.shape();
let strides = tensor.meta.strides();
let mut folded_shape: Vec<usize> = vec![shape[..count].iter().product()];
folded_shape.extend_from_slice(&shape[count..]);
let mut folded_strides: Vec<usize> = vec![strides[count - 1]];
folded_strides.extend_from_slice(&strides[count..]);
*tensor.meta = Metadata::new(Shape::from(folded_shape), Strides::new(&folded_strides));
}
pub fn reduce_logical(
tensor: CubeTensor,
dim: Option<usize>,
config: ReduceOperationConfig,
out_dtype: BoolDType,
) -> CubeTensor {
debug_assert!(
matches!(
config,
ReduceOperationConfig::Any | ReduceOperationConfig::All
),
"reduce_logical only supports Any / All, got {config:?}"
);
let out_bool = DType::Bool(out_dtype);
let backing = elem_type_to_dtype(dtype_to_elem_type(out_bool));
let mut out = match dim {
Some(d) => reduce_dim(tensor, Some(backing), d, Default::default(), config),
None => reduce(tensor, Some(backing), Default::default(), config),
}
.expect("Any/All reduce on a valid axis cannot fail");
out.dtype = out_bool; out
}
pub(crate) fn accumulator_len(config: ReduceOperationConfig) -> usize {
match config {
ReduceOperationConfig::TopK(k) | ReduceOperationConfig::ArgTopK(k) => k,
_ => 1,
}
}
fn argsort(shape: &[usize]) -> Vec<usize> {
let mut indices = (0..shape.len()).collect::<Vec<_>>();
indices.sort_by_key(|&i| &shape[i]);
indices
}
pub fn reduce_dim(
input: CubeTensor,
output_dtype: Option<DType>,
dim: usize,
strategy: KernelReduceStrategy,
config: ReduceOperationConfig,
) -> Result<CubeTensor, cubek::reduce::ReduceError> {
let input = crate::kernel::untile(input);
debug_assert!(
!matches!(
config,
ReduceOperationConfig::ArgMax
| ReduceOperationConfig::ArgMin
| ReduceOperationConfig::ArgTopK(_)
| ReduceOperationConfig::Any
| ReduceOperationConfig::All
) || output_dtype.is_some(),
"The `output_dtype` has to be `Some` when the `config` is `ArgMax`, `ArgMin`, `ArgTopK`, `Any` or `All`.
"
);
let accumulator_len = accumulator_len(config);
let dtypes = config.precision(
dtype_to_elem_type(input.dtype),
output_dtype.map(dtype_to_elem_type),
);
let client = input.client.clone();
let output = init_reduce_output(&input, dim, &dtypes, accumulator_len).ok_or(
cubek::reduce::ReduceError::InvalidAxis {
axis: dim,
rank: input.meta.num_dims(),
},
)?;
let axis_length = input.meta.shape[dim];
if axis_length == 0 {
return reduce_empty_axis(output, axis_length, config);
}
let result = match strategy {
KernelReduceStrategy::Unspecified => cubek::reduce::reduce(
&client,
input.binding(),
output.clone().binding(),
dim,
ReduceStrategy {
routine: RoutineStrategy::Unit(BlueprintStrategy::Inferred(UnitStrategy)),
vectorization: VectorizationStrategy {
parallel_output_vectorization: false,
},
autotune_level: Default::default(),
},
config,
dtypes,
),
KernelReduceStrategy::Specific(strategy) => cubek::reduce::reduce(
&client,
input.binding(),
output.clone().binding(),
dim,
strategy,
config,
dtypes,
),
#[cfg(feature = "autotune")]
KernelReduceStrategy::Autotune => {
autotune_reduce(&client, input, output.clone(), dim, config, dtypes);
Ok(())
}
};
result.map(|_| output)
}
pub fn reduce_dim_with_indices(
input: CubeTensor,
indices_dtype: DType,
dim: usize,
strategy: KernelReduceStrategy,
config: ReduceOperationConfig,
) -> Result<(CubeTensor, CubeTensor), ReduceError> {
let unsupported = |operation| ReduceError::IndicesUnsupported { operation };
let config = match config {
ReduceOperationConfig::ArgMax => ReduceOperationConfig::Max,
ReduceOperationConfig::ArgMin => ReduceOperationConfig::Min,
ReduceOperationConfig::ArgTopK(k) => ReduceOperationConfig::TopK(k),
ReduceOperationConfig::Max
| ReduceOperationConfig::Min
| ReduceOperationConfig::TopK(_) => config,
ReduceOperationConfig::Sum => return Err(unsupported("Sum")),
ReduceOperationConfig::Prod => return Err(unsupported("Prod")),
ReduceOperationConfig::Mean => return Err(unsupported("Mean")),
ReduceOperationConfig::MaxAbs => return Err(unsupported("MaxAbs")),
ReduceOperationConfig::Any => return Err(unsupported("Any")),
ReduceOperationConfig::All => return Err(unsupported("All")),
};
let out_len = accumulator_len(config);
let value_dtypes = config.precision(dtype_to_elem_type(input.dtype), None);
let dtypes = ReduceWithIndicesDtypes {
input: value_dtypes.input,
values: value_dtypes.output,
indices: dtype_to_elem_type(indices_dtype),
accumulation: value_dtypes.accumulation,
};
let invalid_axis = || ReduceError::InvalidAxis {
axis: dim,
rank: input.meta.num_dims(),
};
let values = init_reduce_output_dtype(&input, dim, elem_type_to_dtype(dtypes.values), out_len)
.ok_or_else(invalid_axis)?;
let indices =
init_reduce_output_dtype(&input, dim, indices_dtype, out_len).ok_or_else(invalid_axis)?;
if input.meta.shape[dim] == 0 {
return Err(ReduceError::ReduceAxisTooSmall {
axis_length: 0,
k: out_len,
});
}
let client = input.client.clone();
let result = match strategy {
KernelReduceStrategy::Unspecified => cubek::reduce::reduce_with_indices(
&client,
input.binding(),
values.clone().binding(),
indices.clone().binding(),
dim,
ReduceStrategy {
routine: RoutineStrategy::Unit(BlueprintStrategy::Inferred(UnitStrategy)),
vectorization: VectorizationStrategy {
parallel_output_vectorization: false,
},
autotune_level: Default::default(),
},
config,
dtypes,
),
KernelReduceStrategy::Specific(strategy) => cubek::reduce::reduce_with_indices(
&client,
input.binding(),
values.clone().binding(),
indices.clone().binding(),
dim,
strategy,
config,
dtypes,
),
#[cfg(feature = "autotune")]
KernelReduceStrategy::Autotune => {
autotune_reduce_with_indices(
&client,
input,
values.clone(),
indices.clone(),
dim,
config,
dtypes,
);
Ok(())
}
};
result.map(|_| (values, indices))
}
pub fn init_reduce_output(
input: &CubeTensor,
dim: usize,
dtypes: &ReduceDtypes,
accumulator_len: usize,
) -> Option<CubeTensor> {
init_reduce_output_dtype(
input,
dim,
elem_type_to_dtype(dtypes.output),
accumulator_len,
)
}
pub fn init_reduce_output_dtype(
input: &CubeTensor,
dim: usize,
dtype: DType,
accumulator_len: usize,
) -> Option<CubeTensor> {
(dim < input.meta.num_dims()).then(|| {
let mut shape_out = input.shape();
shape_out[dim] = accumulator_len;
empty_device_contiguous_dtype(input.client.clone(), input.device.clone(), shape_out, dtype)
})
}
#[derive(Clone, Debug)]
pub enum KernelReduceStrategy {
Unspecified,
Specific(cubek::reduce::launch::ReduceStrategy),
#[cfg(feature = "autotune")]
Autotune,
}
impl Default for KernelReduceStrategy {
fn default() -> Self {
#[cfg(feature = "autotune")]
return Self::Autotune;
#[cfg(not(feature = "autotune"))]
return Self::Unspecified;
}
}