use crate::{
CubeTuneId,
kernel::{
autotune_bounds,
interpolate::{execute_interpolate, map_options},
},
ops::permute_nchw_to_nhwc_shape,
tensor::CubeTensor,
};
use burn_backend::cubecl::dtype_to_elem_type;
use burn_backend::ops::InterpolateOptions;
use cubecl::{
std::throughput::roofline_bounds,
tune::{LocalTuner, Tunable, TunableSet, local_tuner},
};
use cubek::interpolate::{
InterpolateStrategy,
definition::{InterpolateCost, InterpolateForwardProblem, InterpolateProblem},
tune_key::InterpolateAutotuneKey,
};
type Inputs = (CubeTensor, [usize; 2], InterpolateOptions);
const STRATEGIES: [(&str, InterpolateStrategy); 2] = [
(
"maximize_throughput",
InterpolateStrategy::MaximizeThroughput,
),
("minimize_latency", InterpolateStrategy::MinimizeLatency),
];
pub fn interpolate_autotune(
input: CubeTensor,
output_size: [usize; 2],
options: InterpolateOptions,
) -> CubeTensor {
let client = input.client.clone();
static TUNER: LocalTuner<InterpolateAutotuneKey, CubeTuneId> = local_tuner!();
let tune_id = CubeTuneId::new(&client, &input.device);
let tunables = TUNER.init(&tune_id, move || {
let mut set = with_bounds(TunableSet::new(create_key, input_gen));
for (name, strategy) in STRATEGIES {
set = set.with(Tunable::new(name, move |(input, output_size, options)| {
execute_interpolate(input, output_size, options, strategy)
}));
}
set
});
TUNER.execute(&tune_id, &client, tunables, (input, output_size, options))
}
fn with_bounds<Out: 'static>(
set: TunableSet<InterpolateAutotuneKey, Inputs, Out>,
) -> TunableSet<InterpolateAutotuneKey, Inputs, Out> {
autotune_bounds::with_bounds(
set,
|_key, (input, output_size, options): &Inputs, thresholds| {
let problem = forward_problem(input, output_size, options);
let cost = InterpolateCost::new(
InterpolateProblem::Forward(problem),
dtype_to_elem_type(input.dtype),
);
roofline_bounds(&input.client, cost.compute_key(), cost.work(), thresholds)
},
)
}
fn forward_problem(
input: &CubeTensor,
output_size: &[usize; 2],
options: &InterpolateOptions,
) -> InterpolateForwardProblem {
let shape = permute_nchw_to_nhwc_shape(input.meta.shape().clone());
InterpolateForwardProblem::from_input_output_shapes(
&shape,
output_size,
map_options(options.clone()),
)
}
fn create_key((input, output_size, options): &Inputs) -> InterpolateAutotuneKey {
let elem = dtype_to_elem_type(input.dtype);
let problem = forward_problem(input, output_size, options);
InterpolateAutotuneKey::generate(
elem,
elem,
problem.options.mode,
problem.options.align_corners,
problem.input_height,
problem.input_width,
problem.channels,
problem.output_height,
problem.output_width,
)
}
fn input_gen(_key: &InterpolateAutotuneKey, (input, output_size, options): &Inputs) -> Inputs {
(input.clone(), *output_size, options.clone())
}