#[cfg(feature = "autotune")]
use super::{autotune_reduce, autotune_reduce_with_indices, autotune_sum};
use crate::{
CubeRuntime,
ops::numeric::{empty_device_contiguous_dtype, zeros_client},
tensor::CubeTensor,
};
use burn_backend::cubecl::{dtype_to_elem_type, elem_type_to_dtype};
use burn_backend::{DType, TensorMetadata};
use burn_std::{BoolDType, Metadata};
use cubecl::{AutotuneKey, client::ComputeClient, features::AtomicUsage, ir::Type};
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 supports_atomic_add<R: CubeRuntime>(client: &ComputeClient<R>, dtype: DType) -> bool {
client
.properties()
.atomic_type_usage(Type::atomic(dtype_to_elem_type(dtype)))
.contains(AtomicUsage::Add)
}
pub fn sum_fallback<R: CubeRuntime>(
tensor: CubeTensor<R>,
mut strategy: SumStrategy,
) -> Result<CubeTensor<R>, ReduceError> {
if matches!(strategy, SumStrategy::OneShot(_))
&& !supports_atomic_add(&tensor.client, tensor.dtype)
{
strategy = SumStrategy::Chained(Default::default());
}
sum(tensor, strategy)
}
pub fn sum<Run: CubeRuntime>(
tensor: CubeTensor<Run>,
strategy: SumStrategy,
) -> Result<CubeTensor<Run>, ReduceError> {
let client = tensor.client.clone();
let device = tensor.device.clone();
match strategy {
SumStrategy::OneShot(cube_count) => {
let output = zeros_client(client.clone(), device, [1].into(), tensor.dtype);
let dtype = tensor.dtype;
shared_sum::<Run>(
&client,
tensor.binding(),
output.clone().binding(),
cube_count,
dtype_to_elem_type(dtype),
)?;
Ok(output)
}
SumStrategy::Chained(strategy) => {
reduce::<Run>(tensor, None, strategy, ReduceOperationConfig::Sum)
}
#[cfg(feature = "autotune")]
SumStrategy::Autotune => Ok(autotune_sum::<Run>(&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<Run: CubeRuntime>(
mut tensor: CubeTensor<Run>,
output_dtype: Option<DType>,
strategy: KernelReduceStrategy,
config: ReduceOperationConfig,
) -> Result<CubeTensor<Run>, cubek::reduce::ReduceError> {
let sorted_axis = argsort(tensor.meta.shape());
for axis in sorted_axis {
tensor = reduce_dim::<Run>(tensor, output_dtype, axis, strategy.clone(), config)?;
}
*tensor.meta = Metadata::new([1], [1]);
Ok(tensor)
}
pub fn reduce_logical<Run: CubeRuntime>(
tensor: CubeTensor<Run>,
dim: Option<usize>,
config: ReduceOperationConfig,
out_dtype: BoolDType,
) -> CubeTensor<Run> {
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::<Run>(tensor, Some(backing), d, Default::default(), config),
None => reduce::<Run>(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<Run: CubeRuntime>(
input: CubeTensor<Run>,
output_dtype: Option<DType>,
dim: usize,
strategy: KernelReduceStrategy,
config: ReduceOperationConfig,
) -> Result<CubeTensor<Run>, cubek::reduce::ReduceError> {
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::<Run>(&input, dim, &dtypes, accumulator_len).ok_or(
cubek::reduce::ReduceError::InvalidAxis {
axis: dim,
rank: input.meta.num_dims(),
},
)?;
let result = match strategy {
KernelReduceStrategy::Unspecified => cubek::reduce::reduce::<Run>(
&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::<Run>(
&client,
input.binding(),
output.clone().binding(),
dim,
strategy,
config,
dtypes,
),
#[cfg(feature = "autotune")]
KernelReduceStrategy::Autotune => {
autotune_reduce::<Run>(&client, input, output.clone(), dim, config, dtypes);
Ok(())
}
};
result.map(|_| output)
}
pub fn reduce_dim_with_indices<Run: CubeRuntime>(
input: CubeTensor<Run>,
indices_dtype: DType,
dim: usize,
strategy: KernelReduceStrategy,
config: ReduceOperationConfig,
) -> Result<(CubeTensor<Run>, CubeTensor<Run>), 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).into(),
accumulation: value_dtypes.accumulation,
};
let invalid_axis = || ReduceError::InvalidAxis {
axis: dim,
rank: input.meta.num_dims(),
};
let values = init_reduce_output_dtype::<Run>(
&input,
dim,
elem_type_to_dtype(dtypes.values.elem_type()),
out_len,
)
.ok_or_else(invalid_axis)?;
let indices = init_reduce_output_dtype::<Run>(&input, dim, indices_dtype, out_len)
.ok_or_else(invalid_axis)?;
let client = input.client.clone();
let result = match strategy {
KernelReduceStrategy::Unspecified => cubek::reduce::reduce_with_indices::<Run>(
&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::<Run>(
&client,
input.binding(),
values.clone().binding(),
indices.clone().binding(),
dim,
strategy,
config,
dtypes,
),
#[cfg(feature = "autotune")]
KernelReduceStrategy::Autotune => {
autotune_reduce_with_indices::<Run>(
&client,
input,
values.clone(),
indices.clone(),
dim,
config,
dtypes,
);
Ok(())
}
};
result.map(|_| (values, indices))
}
pub fn init_reduce_output<Run: CubeRuntime>(
input: &CubeTensor<Run>,
dim: usize,
dtypes: &ReduceDtypes,
accumulator_len: usize,
) -> Option<CubeTensor<Run>> {
init_reduce_output_dtype::<Run>(
input,
dim,
elem_type_to_dtype(dtypes.output.elem_type()),
accumulator_len,
)
}
pub fn init_reduce_output_dtype<Run: CubeRuntime>(
input: &CubeTensor<Run>,
dim: usize,
dtype: DType,
accumulator_len: usize,
) -> Option<CubeTensor<Run>> {
(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;
}
}