use crate::{
backward_data::args::{ConcreteArgs, ConcreteInputsFactory, ConcreteOutputFactory},
components::{ConvSetupError, ConvolutionProblem, global::args::RuntimeArgs},
};
use cubecl::{client::Client, prelude::TensorBinding};
use cubek_matmul::{
definition::{AccumulatorOperand, MatmulElems, MatmulVectorSizes},
multi_level::{
BatchMatmulRoutine,
args::{InputArg, OutputArg},
},
routine::BlueprintStrategy,
};
use cubek_std::InputBinding;
#[allow(clippy::result_large_err, clippy::too_many_arguments)]
pub fn launch_kernel_concrete<Args: ConcreteArgs<A>, A: BatchMatmulRoutine<RuntimeArgs>>(
client: &Client,
out_grad: InputBinding,
weights: InputBinding,
in_grad: TensorBinding,
problem: ConvolutionProblem,
vector_sizes: MatmulVectorSizes,
blueprint_strategy: &BlueprintStrategy<RuntimeArgs, A>,
dtypes: &MatmulElems,
) -> Result<(), ConvSetupError> {
let mut view_vector_sizes = vector_sizes;
if let InputBinding::Quantized { scheme, .. } = out_grad {
view_vector_sizes.lhs *= scheme.num_quants();
}
if let InputBinding::Quantized { scheme, .. } = weights {
view_vector_sizes.rhs *= scheme.num_quants();
}
let device_settings = A::device_settings(client, view_vector_sizes);
let expand_info = A::expand_blueprint(
&problem.as_matmul_problem(AccumulatorOperand::Absent),
&device_settings,
blueprint_strategy,
)?;
let problem = Args::adjust_problem(client, problem, &expand_info.blueprint, dtypes);
let launch_info = A::prepare(
&problem.as_matmul_problem(AccumulatorOperand::Absent),
&device_settings,
expand_info,
)?;
let (input, runtime_args) = <InputArg<Args> as ConcreteInputsFactory<A>>::create(
out_grad,
weights,
&launch_info.blueprint,
&problem,
dtypes,
);
let output = <OutputArg<Args> as ConcreteOutputFactory<A>>::create(
in_grad,
&launch_info.blueprint,
&problem,
);
let result = cubek_matmul::multi_level::launch_kernel::<Args, A>(
client,
input,
output,
runtime_args,
launch_info,
);
result.map_err(ConvSetupError::Matmul)
}