use crate::{
ReduceError, ReducePrecision, VectorizationMode,
components::{
args::{NumericVector, ReduceArgs, TensorArgs, init_tensors},
global::{
cube::GlobalFullCubeReduce, plane::GlobalFullPlaneReduce, unit::GlobalFullUnitReduce,
},
instructions::*,
},
launch::{ReduceStrategy, RoutineStrategy, generate_vector_size},
output_vectorization_axis,
routines::{
GlobalReduceBlueprint, ReduceBlueprint, ReduceLaunchSettings, ReduceProblem,
ReduceVectorSettings, Routine, cube::CubeRoutine, plane::PlaneRoutine, unit::UnitRoutine,
},
};
use cubecl::{prelude::*, std::tensor::r#virtual::VirtualTensor};
#[derive(Clone, Copy, Debug)]
pub struct ReduceDtypes {
pub input: StorageType,
pub output: StorageType,
pub accumulation: StorageType,
}
#[derive(Clone, Copy, Debug)]
pub struct ReduceWithIndicesDtypes {
pub input: StorageType,
pub values: StorageType,
pub indices: StorageType,
pub accumulation: StorageType,
}
impl ReduceWithIndicesDtypes {
fn values_dtypes(&self) -> ReduceDtypes {
ReduceDtypes {
input: self.input,
output: self.values,
accumulation: self.accumulation,
}
}
}
#[allow(clippy::too_many_arguments)]
fn prepare_reduce_launch<Run: Runtime>(
client: &ComputeClient<Run>,
input: &TensorBinding<Run>,
output: &TensorBinding<Run>,
reduce_axis: usize,
strategy: ReduceStrategy,
dtypes: ReduceDtypes,
inst: ReduceOperationConfig,
address_type: AddressType,
second_output: Option<StorageType>,
) -> Result<(ReduceBlueprint, ReduceLaunchSettings, usize), ReduceError> {
let reduce_len = input.shape[reduce_axis];
let input_elems: usize = input.shape.iter().copied().product();
let reduce_count = input_elems / reduce_len;
let problem = ReduceProblem {
reduce_len,
reduce_count,
axis: reduce_axis,
dtypes,
instruction: inst,
address_type,
};
let vectorization_mode = match input.strides[reduce_axis] {
1 => VectorizationMode::Parallel,
_ => VectorizationMode::Perpendicular,
};
let out_vec_axis = output_vectorization_axis(&input.strides, reduce_axis, vectorization_mode);
let (vector_size_input, vector_size_output) = generate_vector_size::<Run>(
client,
input,
output,
reduce_axis,
problem.dtypes.input,
vectorization_mode,
&strategy.vectorization,
);
let vector_size_output = match second_output {
None => vector_size_output,
Some(index_dtype) => client
.io_optimized_vector_sizes(index_dtype.size())
.filter(|&width| width <= vector_size_output)
.max()
.unwrap_or(1),
};
let settings = ReduceVectorSettings {
vectorization_mode,
vector_size_input,
vector_size_output,
unchecked_fast_paths: matches!(
strategy.autotune_level,
cubecl::config::autotune::AutotuneLevel::Full
),
fuse_on_read: false,
};
let (blueprint, settings) = match strategy.routine {
RoutineStrategy::Unit(strategy) => {
let routine = UnitRoutine;
routine.prepare(client, problem, settings, strategy)?
}
RoutineStrategy::Plane(strategy) => {
let routine = PlaneRoutine;
routine.prepare(client, problem, settings, strategy)?
}
RoutineStrategy::Cube(strategy) => {
let routine = CubeRoutine;
routine.prepare(client, problem, settings, strategy)?
}
};
Ok((blueprint, settings, out_vec_axis))
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn launch_reduce<Run: Runtime>(
client: &ComputeClient<Run>,
input: TensorBinding<Run>,
output: TensorBinding<Run>,
reduce_axis: usize,
strategy: ReduceStrategy,
dtypes: ReduceDtypes,
inst: ReduceOperationConfig,
) -> Result<(), ReduceError> {
let address_type = input
.required_address_type(dtypes.input.size())
.max(output.required_address_type(dtypes.output.size()));
let (blueprint, settings, out_vec_axis) = prepare_reduce_launch::<Run>(
client,
&input,
&output,
reduce_axis,
strategy,
dtypes,
inst,
address_type,
None,
)?;
unsafe {
reduce_kernel::launch_unchecked::<TensorArgs, Run>(
client,
settings.cube_count,
settings.cube_dim,
settings.address_type,
settings.vector.vector_size_input,
settings.vector.vector_size_output,
input.into_tensor_arg(),
output.into_tensor_arg(),
reduce_axis,
out_vec_axis,
blueprint,
inst,
dtypes.input,
dtypes.output,
dtypes.accumulation,
)
};
Ok(())
}
#[cube(launch_unchecked, address_type = "dynamic")]
pub fn reduce_kernel<
In: Numeric,
InSize: Size,
Out: Numeric,
OutSize: Size,
Acc: Numeric,
RA: ReduceArgs,
>(
input: &RA::Input<In, InSize>,
output: &mut RA::Output<Out, OutSize>,
reduce_axis: usize,
out_vec_axis: usize,
#[comptime] blueprint: ReduceBlueprint,
#[comptime] config: ReduceOperationConfig,
#[define(In)] _input_dtype: StorageType,
#[define(Out)] _output_dtype: StorageType,
#[define(Acc)] _acc_dtype: StorageType,
) {
let (input, mut output) = init_tensors::<RA, In, InSize, Out, OutSize>(input, output);
reduce_kernel_virtual::<In, InSize, Out, OutSize, Acc>(
&input,
&mut output,
reduce_axis,
out_vec_axis,
blueprint,
config,
);
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn launch_reduce_with_indices<Run: Runtime>(
client: &ComputeClient<Run>,
input: TensorBinding<Run>,
values: TensorBinding<Run>,
indices: TensorBinding<Run>,
reduce_axis: usize,
strategy: ReduceStrategy,
dtypes: ReduceWithIndicesDtypes,
operation: ReduceOperationConfig,
) -> Result<(), ReduceError> {
match operation {
ReduceOperationConfig::TopK(k) | ReduceOperationConfig::ArgTopK(k) => {
launch_fused::<Run, TopK>(
client,
input,
values,
indices,
reduce_axis,
strategy,
dtypes,
TopKConfig {
k,
output: ReduceOutputMode::Indices,
},
ReduceOperationConfig::ArgTopK(k),
)
}
ReduceOperationConfig::Max | ReduceOperationConfig::ArgMax => launch_fused::<Run, Max>(
client,
input,
values,
indices,
reduce_axis,
strategy,
dtypes,
ReduceOutputMode::Indices,
ReduceOperationConfig::ArgMax,
),
ReduceOperationConfig::Min | ReduceOperationConfig::ArgMin => launch_fused::<Run, Min>(
client,
input,
values,
indices,
reduce_axis,
strategy,
dtypes,
ReduceOutputMode::Indices,
ReduceOperationConfig::ArgMin,
),
_ => unreachable!("reduce_with_indices rejects operations without indices"),
}
}
#[allow(clippy::too_many_arguments)]
fn launch_fused<Run: Runtime, R: ReduceWithIndicesFamily>(
client: &ComputeClient<Run>,
input: TensorBinding<Run>,
values: TensorBinding<Run>,
indices: TensorBinding<Run>,
reduce_axis: usize,
strategy: ReduceStrategy,
dtypes: ReduceWithIndicesDtypes,
config: R::Config,
blueprint_operation: ReduceOperationConfig,
) -> Result<(), ReduceError> {
let address_type = input
.required_address_type(dtypes.input.size())
.max(values.required_address_type(dtypes.values.size()))
.max(indices.required_address_type(dtypes.indices.size()));
let (blueprint, settings, out_vec_axis) = prepare_reduce_launch::<Run>(
client,
&input,
&values,
reduce_axis,
strategy,
dtypes.values_dtypes(),
blueprint_operation,
address_type,
Some(dtypes.indices),
)?;
unsafe {
reduce_with_indices_kernel::launch_unchecked::<TensorArgs, R, Run>(
client,
settings.cube_count,
settings.cube_dim,
settings.address_type,
settings.vector.vector_size_input,
settings.vector.vector_size_output,
settings.vector.vector_size_output,
input.into_tensor_arg(),
values.into_tensor_arg(),
indices.into_tensor_arg(),
reduce_axis,
out_vec_axis,
blueprint,
config,
dtypes.input,
dtypes.values,
dtypes.indices,
dtypes.accumulation,
)
};
Ok(())
}
#[cube(launch_unchecked, address_type = "dynamic")]
pub fn reduce_with_indices_kernel<
In: Numeric,
InSize: Size,
Out: Numeric,
OutSize: Size,
Idx: Numeric,
IdxSize: Size,
Acc: Numeric,
RA: ReduceArgs,
R: ReduceWithIndicesFamily,
>(
input: &RA::Input<In, InSize>,
output: &mut RA::Output<Out, OutSize>,
indices: &mut RA::Output<Idx, IdxSize>,
reduce_axis: usize,
out_vec_axis: usize,
#[comptime] blueprint: ReduceBlueprint,
#[comptime] config: R::Config,
#[define(In)] _input_dtype: StorageType,
#[define(Out)] _output_dtype: StorageType,
#[define(Idx)] _indices_dtype: StorageType,
#[define(Acc)] _acc_dtype: StorageType,
) {
let (input_values, mut output) = init_tensors::<RA, In, InSize, Out, OutSize>(input, output);
let (_input_indices, mut indices) =
init_tensors::<RA, In, InSize, Idx, IdxSize>(input, indices);
reduce_with_indices_kernel_inner::<(In, InSize, Acc), (Out, OutSize), (Idx, IdxSize), R>(
&input_values,
&mut output,
&mut indices,
reduce_axis,
out_vec_axis,
blueprint,
config,
);
}
#[cube]
fn reduce_with_indices_kernel_inner<
P: ReducePrecision,
Out: NumericVector,
Idx: NumericVector,
R: ReduceWithIndicesFamily,
>(
input: &VirtualTensor<P::EI, P::SI>,
output: &mut VirtualTensor<Out::T, Out::N, ReadWrite>,
indices: &mut VirtualTensor<Idx::T, Idx::N, ReadWrite>,
reduce_axis: usize,
out_vec_axis: usize,
#[comptime] blueprint: ReduceBlueprint,
#[comptime] config: R::Config,
) {
let inst = R::Instruction::<P>::from_config(config);
match blueprint.global {
GlobalReduceBlueprint::Cube(cube) => {
GlobalFullCubeReduce::execute_with_indices::<P, Out, Idx, R::Instruction<P>>(
input,
output,
indices,
reduce_axis,
out_vec_axis,
&inst,
blueprint.vectorization_mode,
cube,
)
}
GlobalReduceBlueprint::Plane(plane) => {
GlobalFullPlaneReduce::execute_with_indices::<P, Out, Idx, R::Instruction<P>>(
input,
output,
indices,
reduce_axis,
out_vec_axis,
&inst,
blueprint.vectorization_mode,
plane,
)
}
GlobalReduceBlueprint::Unit(unit) => {
GlobalFullUnitReduce::execute_with_indices::<P, Out, Idx, R::Instruction<P>>(
input,
output,
indices,
reduce_axis,
out_vec_axis,
&inst,
blueprint.vectorization_mode,
unit,
)
}
};
}
#[cube]
pub fn reduce_kernel_virtual<
In: Numeric,
InSize: Size,
Out: Numeric,
OutSize: Size,
Acc: Numeric,
>(
input: &VirtualTensor<In, InSize>,
output: &mut VirtualTensor<Out, OutSize, ReadWrite>,
reduce_axis: usize,
out_vec_axis: usize,
#[comptime] blueprint: ReduceBlueprint,
#[comptime] config: ReduceOperationConfig,
) {
reduce_kernel_inner::<(In, InSize, Acc), (Out, OutSize), ReduceOperation>(
input,
output,
reduce_axis,
out_vec_axis,
blueprint,
config,
)
}
#[cube]
fn reduce_kernel_inner<P: ReducePrecision, Out: NumericVector, R: ReduceFamily>(
input: &VirtualTensor<P::EI, P::SI>,
output: &mut VirtualTensor<Out::T, Out::N, ReadWrite>,
reduce_axis: usize,
out_vec_axis: usize,
#[comptime] blueprint: ReduceBlueprint,
#[comptime] config: R::Config,
) {
let inst = R::Instruction::<P>::from_config(config);
match blueprint.global {
GlobalReduceBlueprint::Cube(cube) => {
GlobalFullCubeReduce::execute::<P, Out, R::Instruction<P>>(
input,
output,
reduce_axis,
out_vec_axis,
&inst,
blueprint.vectorization_mode,
cube,
)
}
GlobalReduceBlueprint::Plane(plane) => {
GlobalFullPlaneReduce::execute::<P, Out, R::Instruction<P>>(
input,
output,
reduce_axis,
out_vec_axis,
&inst,
blueprint.vectorization_mode,
plane,
)
}
GlobalReduceBlueprint::Unit(unit) => {
GlobalFullUnitReduce::execute::<P, Out, R::Instruction<P>>(
input,
output,
reduce_axis,
out_vec_axis,
&inst,
blueprint.vectorization_mode,
unit,
)
}
};
}