Skip to main content

tract_cuda/kernels/
binary.rs

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}