tract_cuda/kernels/
iff.rs1use 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 ((total_elems.div_ceil(block_dim.0 as usize) as u32, 1, 1), "generic")
98 ;
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});