1use cudarc::driver::{CudaStream, LaunchConfig, PushKernelArg};
2use tract_core::internal::*;
3use tract_core::ops::binary::BinMiniOp;
4use tract_gpu::tensor::DeviceTensor;
5
6use crate::context::{TractCudaStream, cuda_context};
7use crate::kernels::launch_args::TractLaunchArgs;
8use crate::kernels::{LibraryName, get_cuda_view};
9
10static BINARY_MAX_RANK: usize = 5;
11
12const ALL_OP_NAMES: &[&str] = &[
13 "mul", "add", "div", "sub", "pow", "min", "max", "gt", "gte", "eq", "ne", "lt", "lte", "and",
14 "or", "bitor", "bitand", "bitxor",
15];
16
17pub fn all_functions() -> Vec<String> {
18 ALL_OP_NAMES
19 .iter()
20 .flat_map(|kname| {
21 DeviceTensor::SUPPORTED_DT.into_iter().flat_map(move |dt| {
22 let tname = DeviceTensor::tname(dt).ok()?;
23 Some(
24 ["large", "generic"]
25 .into_iter()
26 .map(move |variant| format!("binary_{kname}_{variant}_{tname}")),
27 )
28 })
29 })
30 .flatten()
31 .collect()
32}
33
34pub fn is_supported(mini_op: &dyn BinMiniOp, dt: DatumType) -> bool {
35 ALL_OP_NAMES.contains(&mini_op.name().to_lowercase().as_str())
36 && (dt.is_number() || dt.is::<bool>())
37}
38
39pub fn dispatch_eval(
40 stream: &TractCudaStream,
41 mini_op: &dyn BinMiniOp,
42 lhs: &DeviceTensor,
43 rhs: &DeviceTensor,
44 output: &DeviceTensor,
45) -> TractResult<()> {
46 let rank = lhs.rank();
47 ensure!(rank == rhs.rank());
48 ensure!(rank <= BINARY_MAX_RANK);
49
50 let rank_offset = BINARY_MAX_RANK - rank;
51 let mut lhs_shape = [1usize; BINARY_MAX_RANK];
52 let mut rhs_shape = [1usize; BINARY_MAX_RANK];
53 let mut out_shape = [1usize; BINARY_MAX_RANK];
54 let mut lhs_strides = [0isize; BINARY_MAX_RANK];
55 let mut rhs_strides = [0isize; BINARY_MAX_RANK];
56 let mut out_strides = [0isize; BINARY_MAX_RANK];
57
58 let base_l_shape = lhs.shape();
59 let base_r_shape = rhs.shape();
60 let base_o_shape = output.shape();
61 let base_l_strides = lhs.strides();
62 let base_r_strides = rhs.strides();
63 let base_o_strides = output.strides();
64 for i in 0..rank {
65 let dst = rank_offset + i;
66 lhs_shape[dst] = base_l_shape[i];
67 rhs_shape[dst] = base_r_shape[i];
68 out_shape[dst] = base_o_shape[i];
69 lhs_strides[dst] =
70 if base_l_shape[i] == 1 && base_r_shape[i] != 1 { 0 } else { base_l_strides[i] };
71 rhs_strides[dst] =
72 if base_r_shape[i] == 1 && base_l_shape[i] != 1 { 0 } else { base_r_strides[i] };
73 out_strides[dst] = base_o_strides[i];
74 }
75
76 let total_elems: usize = out_shape.iter().product();
77 let block_dim = (128_u32, 1, 1);
78 let (grid_dim, variant) = if out_shape[BINARY_MAX_RANK - 1] >= 256 && total_elems >= 4096 {
79 (
80 (
81 out_shape[BINARY_MAX_RANK - 2] as u32,
82 out_shape[BINARY_MAX_RANK - 3] as u32,
83 out_shape[..BINARY_MAX_RANK - 3].iter().product::<usize>() as u32,
84 ),
85 "large",
86 )
87 } else {
88 ((total_elems.div_ceil(block_dim.0 as usize) as u32, 1, 1), "generic")
89 };
90
91 let op_name = mini_op.name().to_lowercase();
92 let tname = DeviceTensor::tname(lhs.datum_type())?;
93 let kname = format!("binary_{op_name}_{variant}_{tname}");
94 let func = cuda_context().load_pipeline(LibraryName::Binary, kname)?;
95
96 let cfg = LaunchConfig { grid_dim, block_dim, shared_mem_bytes: 0 };
97
98 let lhs_view = get_cuda_view(lhs);
99 let rhs_view = get_cuda_view(rhs);
100 let out_view = get_cuda_view(output);
101
102 let mut launch_args = TractLaunchArgs::new(stream, &func);
103 launch_args.push_view(&lhs_view);
104 launch_args.push_view(&rhs_view);
105 launch_args.push_view(&out_view);
106 launch_args.push_slice_i32(&rhs_shape);
107 launch_args.push_slice_i32(&out_shape);
108 launch_args.push_slice_i32(&lhs_strides);
109 launch_args.push_slice_i32(&rhs_strides);
110 launch_args.push_slice_i32(&out_strides);
111
112 launch_args.launch(cfg)?;
113
114 Ok(())
115}
116
117pub fn cuda_bin_op_dispatch(
118 mini_op: &dyn BinMiniOp,
119 lhs: &DeviceTensor,
120 rhs: &DeviceTensor,
121 output: &DeviceTensor,
122) -> TractResult<()> {
123 crate::with_cuda_stream(|stream| dispatch_eval(stream, mini_op, lhs, rhs, output))
124}
125
126pub fn cuda_bin_op(mini_op: Box<dyn BinMiniOp>) -> tract_gpu::ops::binary::GpuBinOp {
127 tract_gpu::ops::binary::GpuBinOp::new(mini_op, "Cuda", cuda_bin_op_dispatch)
128}
129
130crate::register_cuda_op!(tract_core::ops::binary::TypedBinOp, |source, node, op| {
131 rule_if!(is_supported(&*op.0, source.node_input_facts(node.id)?[0].datum_type));
132 Ok(Some(Box::new(cuda_bin_op(op.0.clone()))))
133});
134
135#[cfg(test)]
136mod tests {
137 use tract_gpu::tensor::IntoDevice;
138
139 use super::*;
140 use crate::with_cuda_stream;
141 use derive_new::new;
142 use num_traits::AsPrimitive;
143 use num_traits::Float;
144 use proptest::collection::vec;
145 use proptest::prelude::*;
146 use tract_core::internal::Tensor;
147
148 fn test_case<F>(
149 mini_op: &dyn BinMiniOp,
150 shape: &[usize],
151 offset: f32,
152 scale: f32,
153 ) -> TractResult<()>
154 where
155 F: Float + Datum,
156 usize: AsPrimitive<f32>,
157 f32: AsPrimitive<F>,
158 {
159 with_cuda_stream(|stream| {
160 let len = shape.iter().product::<usize>();
161
162 let a = Tensor::from_shape(
163 shape,
164 &(0..len)
165 .map(|f| -> F {
166 let v: f32 = f.as_();
167 (v * scale + offset).as_()
168 })
169 .collect::<Vec<_>>(),
170 )?
171 .into_device()?;
172
173 let b = Tensor::from_shape(
174 shape,
175 &(0..len)
176 .map(|f| -> F {
177 let v: f32 = f.as_();
178 (v * scale + offset + 1.0).as_()
179 })
180 .collect::<Vec<_>>(),
181 )?
182 .into_device()?;
183
184 let out_dt = mini_op.result_datum_type(a.datum_type(), b.datum_type())?;
185 let output = unsafe { DeviceTensor::uninitialized_dt(out_dt, shape)? };
186 dispatch_eval(stream, mini_op, &a, &b, &output)?;
187 stream.synchronize()?;
188
189 let out = output.to_host()?.into_tensor();
190 assert_eq!(out.shape(), shape);
191 Ok(())
192 })
193 }
194
195 use tract_core::ops::math;
196
197 #[test]
198 fn test_binary_add() -> TractResult<()> {
199 test_case::<f32>(&math::Add, &[4, 4], 0.0, 1.0)?;
200 test_case::<f16>(&math::Add, &[4, 4], 0.0, 1.0 / 100.0)?;
201 Ok(())
202 }
203
204 #[test]
205 fn test_binary_mul() -> TractResult<()> {
206 test_case::<f32>(&math::Mul, &[4, 4], 0.0, 1.0)?;
207 test_case::<f16>(&math::Mul, &[4, 4], 0.0, 1.0 / 100.0)?;
208 Ok(())
209 }
210
211 #[test]
212 fn test_binary_sub() -> TractResult<()> {
213 test_case::<f32>(&math::Sub, &[4, 4], 0.0, 1.0)?;
214 Ok(())
215 }
216
217 #[test]
218 fn test_binary_min() -> TractResult<()> {
219 test_case::<f32>(&math::Min, &[4, 4], 0.0, 1.0)?;
220 Ok(())
221 }
222
223 #[test]
224 fn test_binary_max() -> TractResult<()> {
225 test_case::<f32>(&math::Max, &[4, 4], 0.0, 1.0)?;
226 Ok(())
227 }
228}