use crate::context::cuda_context;
use crate::kernels::launch_args::TractLaunchArgs;
use crate::kernels::{LibraryName, MAX_THREADS, get_cuda_view, launch_args, utils};
use cudarc::driver::{CudaStream, LaunchConfig, PushKernelArg};
use tract_core::internal::*;
use tract_gpu::tensor::DeviceTensor;
pub use tract_gpu::ops::reduce::Reducer;
fn cuda_reduce_is_supported_dt(reducer: &Reducer, dt: DatumType) -> bool {
reducer.is_supported_dt(dt)
|| (matches!(reducer, Reducer::Sum | Reducer::Prod) && dt.is::<i64>())
}
pub fn kernel_name(reducer: &Reducer, dt: DatumType, n_cols: usize) -> TractResult<String> {
ensure!(
cuda_reduce_is_supported_dt(reducer, dt),
"Unsupported dt {dt:?} for cuda reduceop {:?}",
reducer
);
let tname = DeviceTensor::tname(dt)?;
if n_cols < 1024 {
Ok(format!("reduce_{}_small_{tname}", reducer))
} else {
Ok(format!("reduce_{}_{tname}", reducer))
}
}
pub fn cuda_reduce_launch(
reducer: &Reducer,
input: &DeviceTensor,
axis: usize,
output: &DeviceTensor,
) -> TractResult<()> {
crate::with_cuda_stream(|stream| {
ensure!(output.datum_type() == input.datum_type());
ensure!(output.shape()[axis] == 1);
let input_shape_nd3 = utils::reshape_to_rank_3(input.shape(), axis);
let input_strides_nd3 = Tensor::natural_strides(&input_shape_nd3);
let output_shape_nd3 = utils::reshape_to_rank_3(output.shape(), axis);
let output_strides_nd3 = Tensor::natural_strides(&output_shape_nd3);
let total = (input_shape_nd3[0] as u64) * (input_shape_nd3[2] as u64);
let i_view = get_cuda_view(input);
let o_view = get_cuda_view(output);
let func = cuda_context().load_pipeline(
LibraryName::NN,
kernel_name(reducer, input.datum_type(), input_shape_nd3[1])?,
)?;
let mut launch_args = TractLaunchArgs::new(stream, &func);
launch_args.push_view(&i_view);
launch_args.push_view(&o_view);
launch_args.push_slice_i32(&input_shape_nd3);
launch_args.push_slice_i32(&input_strides_nd3);
launch_args.push_slice_i32(&output_strides_nd3);
let cfg = LaunchConfig {
grid_dim: (total as u32, 1, 1),
block_dim: if input_shape_nd3[1] < MAX_THREADS {
(32, 1, 1)
} else {
(MAX_THREADS as _, 1, 1)
},
shared_mem_bytes: 0,
};
launch_args.launch(cfg)
})
}
crate::register_cuda_op!(tract_core::ops::nn::Reduce, |source, node, op| {
let dt = source.node_input_facts(node.id)?[0].datum_type;
if let Ok(gpu_op) =
tract_gpu::ops::reduce::GpuReduce::from_tract_core(op, "Cuda", cuda_reduce_launch)
&& cuda_reduce_is_supported_dt(&gpu_op.reducer, dt)
{
return Ok(Some(Box::new(gpu_op)));
}
Ok(None)
});
#[cfg(test)]
mod tests {
use super::*;
use derive_new::new;
use num_traits::AsPrimitive;
use num_traits::Float;
use proptest::collection::vec;
use proptest::prelude::*;
use tract_core::internal::Tensor;
use tract_core::ops::nn::Reducer as TractReducer;
use tract_core::tract_data::itertools::Itertools;
use tract_gpu::tensor::IntoDevice;
fn test_case<F>(
reducer: Reducer,
tract_reducer: TractReducer,
shape: &[usize],
axis: usize,
scale: f32,
) -> TractResult<()>
where
F: Float + Datum,
usize: AsPrimitive<f32>,
f32: AsPrimitive<F>,
{
crate::with_cuda_stream(|stream| {
let len = shape.iter().product::<usize>();
let a = Tensor::from_shape(
shape,
&(0..len)
.map(|f| -> F {
let v: f32 = f.as_();
(v * scale).as_()
})
.collect::<Vec<_>>(),
)?
.into_device()?;
let cpu_output = tract_reducer.reduce(&[axis], &a.to_host()?.into_tensor())?;
let mut o_shape = a.shape().to_vec();
o_shape[axis] = 1;
let cuda_output_dt =
unsafe { DeviceTensor::uninitialized_dt(a.datum_type(), &o_shape)? };
cuda_reduce_launch(&reducer, &a, axis, &cuda_output_dt)?;
stream.synchronize()?;
let cuda_output = cuda_output_dt;
cpu_output
.close_enough(&cuda_output.to_host()?.into_tensor(), Approximation::Approximate)
.with_context(|| {
format!(
"A: {:?}, scale: {:?} Cpu: {:?}, Cuda: {:?}",
a.to_host().and_then(|it| it.dump(true)),
scale,
cpu_output.dump(true),
cuda_output.to_host().and_then(|it| it.dump(true))
)
})?;
Ok(())
})
}
#[test]
fn test_reduce_mean_of_squares() -> TractResult<()> {
test_case::<f32>(Reducer::MeanOfSquares, TractReducer::MeanOfSquares, &[4, 4], 1, 1.0)?;
test_case::<f16>(
Reducer::MeanOfSquares,
TractReducer::MeanOfSquares,
&[4, 4],
1,
1.0 / 100.0,
)?;
test_case::<f16>(
Reducer::MeanOfSquares,
TractReducer::MeanOfSquares,
&[1, 10],
0,
1.0 / 100.0,
)?;
test_case::<f32>(
Reducer::MeanOfSquares,
TractReducer::MeanOfSquares,
&[1, 10],
0,
1.0 / 100.0,
)?;
test_case::<f16>(
Reducer::MeanOfSquares,
TractReducer::MeanOfSquares,
&[2, 1],
1,
1.0 / 100.0,
)?;
test_case::<f32>(
Reducer::MeanOfSquares,
TractReducer::MeanOfSquares,
&[2, 1],
1,
1.0 / 100.0,
)?;
test_case::<f16>(
Reducer::MeanOfSquares,
TractReducer::MeanOfSquares,
&[2, 2, 82, 38],
1,
1.0 / 100.0,
)?;
test_case::<f16>(
Reducer::MeanOfSquares,
TractReducer::MeanOfSquares,
&[2, 2, 82, 38],
2,
1.0 / 100.0,
)?;
test_case::<f32>(
Reducer::MeanOfSquares,
TractReducer::MeanOfSquares,
&[2, 2, 82, 38],
1,
1.0 / 100.0,
)?;
test_case::<f32>(
Reducer::MeanOfSquares,
TractReducer::MeanOfSquares,
&[2, 2, 82, 38],
2,
1.0 / 100.0,
)?;
Ok(())
}
#[test]
fn test_reduce_sum() -> TractResult<()> {
test_case::<f32>(Reducer::Sum, TractReducer::Sum, &[4, 4], 1, 1.0)?;
test_case::<f16>(Reducer::Sum, TractReducer::Sum, &[4, 4], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Sum, TractReducer::Sum, &[1, 10], 0, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Sum, TractReducer::Sum, &[1, 10], 0, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Sum, TractReducer::Sum, &[2, 1], 1, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Sum, TractReducer::Sum, &[2, 1], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Sum, TractReducer::Sum, &[2, 2, 82, 38], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Sum, TractReducer::Sum, &[2, 2, 82, 38], 2, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Sum, TractReducer::Sum, &[2, 2, 82, 38], 1, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Sum, TractReducer::Sum, &[2, 2, 82, 38], 2, 1.0 / 100.0)?;
Ok(())
}
#[test]
fn test_reduce_prod() -> TractResult<()> {
test_case::<f32>(Reducer::Prod, TractReducer::Prod, &[4, 4], 1, 1.0)?;
test_case::<f16>(Reducer::Prod, TractReducer::Prod, &[4, 4], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Prod, TractReducer::Prod, &[1, 10], 0, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Prod, TractReducer::Prod, &[1, 10], 0, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Prod, TractReducer::Prod, &[2, 1], 1, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Prod, TractReducer::Prod, &[2, 1], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Prod, TractReducer::Prod, &[2, 2, 82, 38], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Prod, TractReducer::Prod, &[2, 2, 82, 38], 2, 1.0 / 100000.0)?;
test_case::<f32>(Reducer::Prod, TractReducer::Prod, &[2, 2, 82, 38], 1, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Prod, TractReducer::Prod, &[2, 2, 82, 38], 2, 1.0 / 1000.0)?;
Ok(())
}
#[test]
fn test_reduce_max() -> TractResult<()> {
test_case::<f32>(Reducer::Max, TractReducer::Max, &[2, 2], 1, 1.0)?;
test_case::<f16>(Reducer::Max, TractReducer::Max, &[4, 4], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Max, TractReducer::Max, &[1, 10], 0, -1.0 / 100.0)?;
test_case::<f32>(Reducer::Max, TractReducer::Max, &[1, 10], 0, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Max, TractReducer::Max, &[2, 1], 1, -1.0 / 100.0)?;
test_case::<f32>(Reducer::Max, TractReducer::Max, &[2, 1], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Max, TractReducer::Max, &[2, 2, 82, 38], 1, -1.0 / 100.0)?;
test_case::<f16>(Reducer::Max, TractReducer::Max, &[2, 2, 82, 38], 2, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Max, TractReducer::Max, &[2, 2, 82, 38], 1, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Max, TractReducer::Max, &[2, 2, 82, 38], 2, -1.0 / 100.0)?;
Ok(())
}
#[test]
fn test_reduce_min() -> TractResult<()> {
test_case::<f32>(Reducer::Min, TractReducer::Min, &[4, 4], 1, 1.0)?;
test_case::<f16>(Reducer::Min, TractReducer::Min, &[4, 4], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Min, TractReducer::Min, &[1, 10], 0, -1.0 / 100.0)?;
test_case::<f32>(Reducer::Min, TractReducer::Min, &[1, 10], 0, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Min, TractReducer::Min, &[2, 1], 1, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Min, TractReducer::Min, &[2, 1], 1, 1.0 / 100.0)?;
test_case::<f16>(Reducer::Min, TractReducer::Min, &[2, 2, 82, 38], 1, -1.0 / 100.0)?;
test_case::<f16>(Reducer::Min, TractReducer::Min, &[2, 2, 82, 38], 2, 1.0 / 100.0)?;
test_case::<f32>(Reducer::Min, TractReducer::Min, &[2, 2, 82, 38], 1, -1.0 / 100.0)?;
test_case::<f32>(Reducer::Min, TractReducer::Min, &[2, 2, 82, 38], 2, 1.0 / 100.0)?;
Ok(())
}
proptest::proptest! {
#[test]
fn reduce_prop_f32(pb in any::<ReduceProblem<f32>>()) {
fn run(pb: ReduceProblem<f32>) -> TractResult<()> {
let out = pb.run()?;
let reference = pb.reference()?;
out.close_enough(&reference, Approximation::Approximate)
.with_context(|| format!("Cpu: {:?}, Cuda: {:?}", reference.dump(true), out.dump(true)))
}
run(pb).map_err(|e| TestCaseError::Fail(format!("{:?}", e).into()))?;
}
#[test]
fn reduce_prop_f16(pb in any::<ReduceProblem<f16>>()) {
fn run(pb: ReduceProblem<f16>) -> TractResult<()> {
let out = pb.run()?;
let reference = pb.reference()?;
out.close_enough(&reference, Approximation::Approximate)
.with_context(|| format!("Cpu: {:?}, Cuda: {:?}", reference.dump(true), out.dump(true)))
}
run(pb).map_err(|e| TestCaseError::Fail(format!("{:?}", e).into()))?;
}
}
#[derive(Debug, new)]
pub struct ReduceProblem<F: Datum + Float>
where
F: Datum + Float,
usize: AsPrimitive<F>,
{
pub op: Reducer,
pub shape: Vec<usize>,
pub axis: usize,
pub input: Vec<F>,
}
impl<F> Arbitrary for ReduceProblem<F>
where
F: Datum + Float,
usize: AsPrimitive<F>,
{
type Parameters = ();
type Strategy = BoxedStrategy<Self>;
fn arbitrary_with(_: ()) -> Self::Strategy {
let reducers = Reducer::ALL.into_iter().filter(|r| !r.is_logic()).collect_vec();
(0..reducers.len(), 0usize..3, 0usize..3)
.prop_flat_map(move |(op_ix, left, right)| {
let axis = left;
let shape_len = usize::min(left + right + 1, 4);
let shape = 1usize..10;
(Just(reducers[op_ix]), vec(shape, shape_len..=shape_len), Just(axis))
})
.prop_map(|(op, shape, axis)| {
let input = (0..shape.iter().product::<usize>())
.map(|f| f.as_() / 1000.as_())
.collect::<Vec<_>>();
Self { op, shape, axis, input }
})
.boxed()
}
}
impl<F> ReduceProblem<F>
where
F: Datum + Float + std::ops::AddAssign,
usize: AsPrimitive<F>,
{
pub fn reference(&self) -> TractResult<Tensor> {
let a = Tensor::from_shape(self.shape.as_slice(), &self.input)?;
let cpu_output = match self.op {
Reducer::Sum => TractReducer::Sum.reduce(&[self.axis], &a)?,
Reducer::Prod => TractReducer::Prod.reduce(&[self.axis], &a)?,
Reducer::MeanOfSquares => TractReducer::MeanOfSquares.reduce(&[self.axis], &a)?,
Reducer::Min => TractReducer::Min.reduce(&[self.axis], &a)?,
Reducer::Max => TractReducer::Max.reduce(&[self.axis], &a)?,
Reducer::Any => TractReducer::Any.reduce(&[self.axis], &a)?,
Reducer::All => TractReducer::All.reduce(&[self.axis], &a)?,
};
Ok(cpu_output)
}
pub fn run(&self) -> TractResult<Tensor> {
crate::with_cuda_stream(|stream| {
let a = Tensor::from_shape(self.shape.as_slice(), &self.input)?.into_device()?;
let mut o_shape = a.shape().to_vec();
o_shape[self.axis] = 1;
let output = unsafe { DeviceTensor::uninitialized_dt(a.datum_type(), &o_shape)? };
cuda_reduce_launch(&self.op, &a, self.axis, &output)?;
stream.synchronize()?;
Ok(output.to_host()?.into_tensor())
})
}
}
}