ruccl/tensor_device/
device.rs1use super::{TensorBuffer, TensorDevice, TensorDeviceError, TensorElement, TensorReductionLaunch, element};
2use crate::in_process::device::InProcessDevice;
3use crate::rank::{ReductionOperation, device::RankDevice};
4use ruda_tensor::Backend;
5
6impl<B: Backend, T: TensorElement> RankDevice<T> for TensorDevice<B> {
7 type Buffer = TensorBuffer<B, T>;
8 type Kernel = ReductionOperation;
9 type ReductionLaunch = TensorReductionLaunch;
10 type Error = TensorDeviceError;
11
12 const ELEMENT_SIZE: usize = std::mem::size_of::<T>();
13
14 fn buffer_len(&self, buffer: &Self::Buffer) -> usize { buffer.len() }
15 fn buffer_bytes(&self, buffer: &Self::Buffer) -> usize { buffer.len() * std::mem::size_of::<T>() }
16 fn encode(values: &[T]) -> Vec<u8> { element::encode(values) }
17 fn decode(bytes: &[u8]) -> Result<Vec<T>, Self::Error> { element::decode(bytes) }
18
19 fn alloc(&self, length: usize) -> Result<Self::Buffer, Self::Error> { self.allocate(length) }
20
21 fn copy_to_device(&self, buffer: &Self::Buffer, values: &[T]) -> Result<(), Self::Error> {
22 self.write(buffer, values)
23 }
24
25 fn copy_from_device(&self, buffer: &Self::Buffer) -> Result<Vec<T>, Self::Error> {
26 self.download(buffer)
27 }
28
29 fn copy_from_device_at(&self, buffer: &Self::Buffer, element_offset: usize, length: usize) -> Result<Vec<T>, Self::Error> {
30 self.download_at(buffer, element_offset, length)
31 }
32
33 fn copy_bytes_to_device_at(&self, buffer: &Self::Buffer, element_offset: usize, values: &[u8]) -> Result<(), Self::Error> {
34 self.write_at(buffer, element_offset, &element::decode::<T>(values)?)
35 }
36
37 fn copy_bytes_from_device_at(&self, buffer: &Self::Buffer, element_offset: usize, length: usize) -> Result<Vec<u8>, Self::Error> {
38 Ok(element::encode(&self.download_at(buffer, element_offset, length)?))
39 }
40
41 fn prepare_reduction(&self, length: usize, destination_offset: usize) -> Result<Self::ReductionLaunch, Self::Error> {
42 self.prepare::<T>(length, destination_offset)
43 }
44
45 fn launch_reduction(&self, kernel: &Self::Kernel, launch: &Self::ReductionLaunch, source: &Self::Buffer, destination: &Self::Buffer) -> Result<(), Self::Error> {
46 self.reduce(*kernel, launch, source, destination)
47 }
48}
49
50impl<B: Backend, T: TensorElement> InProcessDevice<T> for TensorDevice<B> {
51 type Context = Self;
52 type Buffer = TensorBuffer<B, T>;
53 type Kernel = ReductionOperation;
54 type ReductionLaunch = TensorReductionLaunch;
55 type Error = TensorDeviceError;
56
57 const ELEMENT_SIZE: usize = std::mem::size_of::<T>();
58
59 fn buffer_len(buffer: &Self::Buffer) -> usize { buffer.len() }
60 fn buffer_is_empty(buffer: &Self::Buffer) -> bool { buffer.is_empty() }
61
62 fn alloc(context: &Self::Context, length: usize) -> Result<Self::Buffer, Self::Error> {
63 context.allocate(length)
64 }
65
66 fn copy_to_device(context: &Self::Context, buffer: &Self::Buffer, values: &[T]) -> Result<(), Self::Error> {
67 context.write(buffer, values)
68 }
69
70 fn copy_to_device_at(context: &Self::Context, buffer: &Self::Buffer, offset: usize, values: &[T]) -> Result<(), Self::Error> {
71 context.write_at(buffer, offset, values)
72 }
73
74 fn copy_from_device(context: &Self::Context, buffer: &Self::Buffer) -> Result<Vec<T>, Self::Error> {
75 context.download(buffer)
76 }
77
78 fn copy_from_device_at(context: &Self::Context, buffer: &Self::Buffer, offset: usize, length: usize) -> Result<Vec<T>, Self::Error> {
79 context.download_at(buffer, offset, length)
80 }
81
82 fn prepare_reduction(length: u32, destination_offset: u32) -> Self::ReductionLaunch {
83 TensorReductionLaunch { length: length as usize, destination_offset: destination_offset as usize }
84 }
85
86 fn launch_reduction(context: &Self::Context, kernel: &Self::Kernel, launch: &Self::ReductionLaunch, source: &Self::Buffer, destination: &Self::Buffer) -> Result<(), Self::Error> {
87 context.reduce(*kernel, launch, source, destination)
88 }
89}