Skip to main content

tract_cuda/kernels/array/
copy.rs

1use cudarc::driver::{CudaStream, LaunchConfig, PushKernelArg};
2use derive_new::new;
3use std::fmt;
4use tract_core::internal::*;
5use tract_gpu::tensor::DeviceTensor;
6
7use crate::context::{TractCudaStream, cuda_context};
8use crate::kernels::{LibraryName, get_cuda_view, get_cuda_view_mut, get_sliced_cuda_view};
9
10#[derive(Debug, Clone, new, PartialEq, Eq, Hash)]
11pub struct Memcpy;
12
13impl fmt::Display for Memcpy {
14    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
15        write!(f, "{self:?}")
16    }
17}
18
19impl Memcpy {
20    pub fn is_supported_dt(dt: DatumType) -> bool {
21        matches!(
22            dt,
23            DatumType::F32
24                | DatumType::F16
25                | DatumType::U8
26                | DatumType::U16
27                | DatumType::U32
28                | DatumType::U64
29                | DatumType::I8
30                | DatumType::I16
31                | DatumType::I32
32                | DatumType::I64
33                | DatumType::Bool
34        )
35    }
36
37    pub fn dispatch_eval(
38        &self,
39        stream: &TractCudaStream,
40        input: &DeviceTensor,
41        input_offset: usize,
42        output: &DeviceTensor,
43    ) -> TractResult<()> {
44        ensure!(input_offset % input.datum_type().size_of() == 0);
45        ensure!(output.len() <= input.len() - (input_offset / input.datum_type().size_of()));
46        ensure!(
47            Self::is_supported_dt(input.datum_type()),
48            "Unsupported dt {:?} for cuda memcpy",
49            input.datum_type()
50        );
51
52        let i_view = get_sliced_cuda_view(
53            input,
54            input_offset,
55            input.len() * input.datum_type().size_of() - input_offset,
56        )?;
57        let mut o_view = get_cuda_view_mut(output);
58        let len = output.len();
59        stream.memcpy_dtod(&i_view, &mut o_view);
60
61        Ok(())
62    }
63
64    pub fn eval(
65        &self,
66        stream: &TractCudaStream,
67        input: &DeviceTensor,
68        input_offset: usize,
69        output_shape: &[usize],
70    ) -> TractResult<DeviceTensor> {
71        let output = unsafe { DeviceTensor::uninitialized_dt(input.datum_type(), output_shape)? };
72        self.dispatch_eval(stream, input, input_offset, &output)?;
73        stream.synchronize()?;
74        Ok(output)
75    }
76}
77
78pub fn cuda_memcpy_dispatch(
79    input: &DeviceTensor,
80    input_offset: usize,
81    output: &DeviceTensor,
82) -> TractResult<()> {
83    crate::with_cuda_stream(|stream| Memcpy.dispatch_eval(stream, input, input_offset, output))
84}
85
86#[cfg(test)]
87mod tests {
88
89    use super::*;
90    use tract_gpu::tensor::IntoDevice;
91    use tract_itertools::Itertools;
92
93    use num_traits::Zero;
94
95    use tract_core::internal::Tensor;
96
97    fn run_test_case(shape: &[usize], offset: usize) -> TractResult<()> {
98        crate::with_cuda_stream(|stream| {
99            let len = shape.iter().product::<usize>();
100            let data = (0..len).map(|f| f as f32).collect::<Vec<_>>();
101            let input = Tensor::from_shape(shape, &data)?;
102
103            let output = Memcpy {}.eval(
104                stream,
105                &input.clone().into_device()?,
106                offset,
107                &[len - (offset / size_of::<f32>())],
108            )?;
109
110            assert_eq!(
111                output.to_host()?.into_tensor(),
112                input.into_shape(&[len])?.slice(0, offset / size_of::<f32>(), len)?
113            );
114            Ok(())
115        })
116    }
117
118    #[test]
119    fn test_cpy() -> TractResult<()> {
120        run_test_case(&[3, 4], 0)?;
121        run_test_case(&[2, 5], 2 * size_of::<f32>())?;
122        Ok(())
123    }
124}