Skip to main content

tract_cuda/kernels/
iff.rs

1use cudarc::driver::{CudaStream, LaunchConfig, PushKernelArg};
2use std::fmt;
3use tract_core::internal::*;
4use tract_gpu::tensor::DeviceTensor;
5
6use crate::context::{TractCudaStream, cuda_context};
7use crate::kernels::launch_args::TractLaunchArgs;
8use crate::kernels::utils::compute_broadcast_strides;
9use crate::kernels::{LibraryName, MAX_THREADS, get_cuda_view};
10
11static TERNARY_MAX_RANK: usize = 5;
12
13#[derive(Debug, PartialEq)]
14pub struct Iff;
15
16impl fmt::Display for Iff {
17    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
18        write!(f, "Iff")
19    }
20}
21
22impl Iff {
23    pub fn name(&self) -> Cow<'_, str> {
24        format!("{self}").into()
25    }
26
27    pub fn output_shape<D: DimLike>(&self, a: &[D], b: &[D], c: &[D]) -> TractResult<TVec<D>> {
28        tract_core::broadcast::multi_broadcast(&[a, b, c])
29            .with_context(|| format!("Error while broadcasting {a:?} {b:?} {c:?}"))
30    }
31
32    pub fn kernel_name(&self, dt: DatumType, variant: &str) -> TractResult<String> {
33        Ok(format!("iff_{variant}_{}", tract_gpu::utils::BroadcastKind::copy_tname(dt)))
34    }
35
36    pub fn eval(
37        &self,
38        stream: &TractCudaStream,
39        cond: &DeviceTensor,
40        then_value: &DeviceTensor,
41        else_value: &DeviceTensor,
42    ) -> TractResult<DeviceTensor> {
43        let out_shape = self.output_shape(cond.shape(), then_value.shape(), else_value.shape())?;
44        ensure!(then_value.datum_type() == else_value.datum_type());
45        let out_dt = then_value.datum_type();
46        let output = unsafe { DeviceTensor::uninitialized_dt(out_dt, &out_shape)? };
47
48        self.dispatch_eval(stream, cond, then_value, else_value, &output)?;
49
50        stream.synchronize()?;
51        Ok(output)
52    }
53
54    pub fn dispatch_eval(
55        &self,
56        stream: &TractCudaStream,
57        cond: &DeviceTensor,
58        then_value: &DeviceTensor,
59        else_value: &DeviceTensor,
60        output: &DeviceTensor,
61    ) -> TractResult<()> {
62        let inputs = [cond, then_value, else_value];
63        let rank = *[cond.rank(), then_value.rank(), else_value.rank()].iter().max().unwrap();
64        ensure!(rank <= TERNARY_MAX_RANK);
65
66        let rank_pad = TERNARY_MAX_RANK - rank;
67        let mut strides = [[0isize; TERNARY_MAX_RANK]; 3];
68        let mut out_shape = [1usize; TERNARY_MAX_RANK];
69        let mut out_strides = [0isize; TERNARY_MAX_RANK];
70
71        for axis in 0..rank {
72            out_shape[rank_pad + axis] = output.shape()[axis];
73            out_strides[rank_pad + axis] = output.strides()[axis];
74            for input in 0..3 {
75                strides[input][rank_pad + axis] =
76                    if inputs[input].shape()[axis] < output.shape()[axis] {
77                        0
78                    } else {
79                        inputs[input].strides()[axis]
80                    };
81            }
82        }
83
84        let total_elems: usize = out_shape.iter().product();
85        let block_dim = (128_u32, 1, 1);
86        let (grid_dim, variant) =
87        //     if out_shape[TERNARY_MAX_RANK - 1] >= 256 && total_elems >= 4096 {
88        //     (
89        //         (
90        //             out_shape[TERNARY_MAX_RANK - 2] as u32,
91        //             out_shape[TERNARY_MAX_RANK - 3] as u32,
92        //             out_shape[..TERNARY_MAX_RANK - 3].iter().product::<usize>() as u32,
93        //         ),
94        //         "large",
95        //     )
96        // } else {
97            ((total_elems.div_ceil(block_dim.0 as usize) as u32, 1, 1), "generic")
98        // };
99        ;
100
101        let kernel_name = self.kernel_name(output.datum_type(), variant)?;
102        let func = cuda_context().load_pipeline(LibraryName::Binary, kernel_name)?;
103
104        let cfg = LaunchConfig { grid_dim, block_dim, shared_mem_bytes: 0 };
105
106        let cond_view = get_cuda_view(cond);
107        let then_view = get_cuda_view(then_value);
108        let else_view = get_cuda_view(else_value);
109        let o_view = get_cuda_view(output);
110
111        let mut launch_args = TractLaunchArgs::new(stream, &func);
112        launch_args.push_view(&cond_view);
113        launch_args.push_view(&then_view);
114        launch_args.push_view(&else_view);
115        launch_args.push_view(&o_view);
116        launch_args.push_slice_i32(&out_shape);
117        for stride in &strides {
118            launch_args.push_slice_i32(stride);
119        }
120        launch_args.push_slice_i32(&out_strides);
121
122        launch_args.launch(cfg)?;
123
124        Ok(())
125    }
126}
127
128#[allow(clippy::too_many_arguments)]
129pub fn cuda_iff_dispatch(
130    cond: &DeviceTensor,
131    then_value: &DeviceTensor,
132    else_value: &DeviceTensor,
133    cond_strides: &[isize],
134    then_strides: &[isize],
135    else_strides: &[isize],
136    output: &DeviceTensor,
137    output_shape: &[usize],
138    output_strides: &[isize],
139) -> TractResult<()> {
140    crate::with_cuda_stream(|stream| {
141        let total_elems: usize = output_shape.iter().product();
142        let block_dim = (128_u32, 1, 1);
143        let grid_dim = (total_elems.div_ceil(block_dim.0 as usize) as u32, 1, 1);
144
145        let kernel_name = format!(
146            "iff_generic_{}",
147            tract_gpu::utils::BroadcastKind::copy_tname(output.datum_type())
148        );
149        let func = cuda_context().load_pipeline(LibraryName::Binary, kernel_name)?;
150        let cfg = LaunchConfig { grid_dim, block_dim, shared_mem_bytes: 0 };
151
152        let cond_view = get_cuda_view(cond);
153        let then_view = get_cuda_view(then_value);
154        let else_view = get_cuda_view(else_value);
155        let o_view = get_cuda_view(output);
156
157        let mut launch_args = TractLaunchArgs::new(stream, &func);
158        launch_args.push_view(&cond_view);
159        launch_args.push_view(&then_view);
160        launch_args.push_view(&else_view);
161        launch_args.push_view(&o_view);
162        launch_args.push_slice_i32(output_shape);
163        launch_args.push_slice_i32(cond_strides);
164        launch_args.push_slice_i32(then_strides);
165        launch_args.push_slice_i32(else_strides);
166        launch_args.push_slice_i32(output_strides);
167
168        launch_args.launch(cfg)
169    })
170}
171
172crate::register_cuda_op!(tract_core::ops::logic::Iff, |_source, _node, _op| {
173    Ok(Some(Box::new(tract_gpu::ops::iff::GpuIff::new("Cuda", cuda_iff_dispatch))))
174});