tract_cuda/kernels/array/
copy.rs1use 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}