ruccl/tensor_device/
storage.rs1use super::{Primitive, TensorBuffer, TensorDevice, TensorDeviceError, TensorElement};
2use ruda_tensor::{Backend, DType, FloatDType, IntDType, Slice, TensorData, read_sync};
3use std::ops::Range;
4
5pub(super) fn checked_length<T: TensorElement>(length: usize) -> Result<(), TensorDeviceError> {
6 let bytes = length.checked_mul(std::mem::size_of::<T>())
7 .ok_or(TensorDeviceError::InvalidBuffer("collective buffer byte length overflow"))?;
8 if bytes > isize::MAX as usize {
9 return Err(TensorDeviceError::InvalidBuffer("collective buffer exceeds addressable size"));
10 }
11 Ok(())
12}
13
14pub(super) fn checked_range(
15 total: usize,
16 offset: usize,
17 length: usize,
18) -> Result<Range<usize>, TensorDeviceError> {
19 let end = offset.checked_add(length)
20 .ok_or(TensorDeviceError::InvalidBuffer("collective element range overflow"))?;
21 if end > total {
22 return Err(TensorDeviceError::InvalidBuffer("collective element range out of bounds"));
23 }
24 Ok(offset..end)
25}
26
27impl<B: Backend> Primitive<B> {
28 pub(super) fn slice(self, range: Range<usize>) -> Self {
29 let slice = Slice::from(range);
30 match self {
31 Self::Float(value) => Self::Float(B::float_slice(value, &[slice])),
32 Self::Int(value) => Self::Int(B::int_slice(value, &[slice])),
33 }
34 }
35
36 pub(super) fn assign(self, range: Range<usize>, value: Self) -> Self {
37 let slice = Slice::from(range);
38 match (self, value) {
39 (Self::Float(tensor), Self::Float(value)) => {
40 Self::Float(B::float_slice_assign(tensor, &[slice], value))
41 }
42 (Self::Int(tensor), Self::Int(value)) => {
43 Self::Int(B::int_slice_assign(tensor, &[slice], value))
44 }
45 _ => unreachable!("typed collective buffer storage kind"),
46 }
47 }
48}
49
50impl<B: Backend> TensorDevice<B> {
51 pub fn allocate<T: TensorElement>(
52 &self,
53 length: usize,
54 ) -> Result<TensorBuffer<B, T>, TensorDeviceError> {
55 self.validate_type::<T>()?;
56 checked_length::<T>(length)?;
57 let shape = [length].into();
58 let value = match T::dtype() {
59 DType::F32 => Primitive::Float(B::float_empty(shape, &self.device, FloatDType::F32)),
60 DType::F16 => Primitive::Float(B::float_empty(shape, &self.device, FloatDType::F16)),
61 DType::BF16 => Primitive::Float(B::float_empty(shape, &self.device, FloatDType::BF16)),
62 DType::I32 => Primitive::Int(B::int_empty(shape, &self.device, IntDType::I32)),
63 _ => unreachable!("sealed collective element type"),
64 };
65 Ok(self.wrap(value, length))
66 }
67
68 pub(super) fn from_values<T: TensorElement>(&self, values: &[T]) -> Primitive<B> {
69 let data = TensorData::new(values.to_vec(), [values.len()]);
70 if T::dtype() == DType::I32 {
71 Primitive::Int(B::int_from_data(data, &self.device))
72 } else {
73 Primitive::Float(B::float_from_data(data, &self.device))
74 }
75 }
76
77 pub fn write<T: TensorElement>(
78 &self,
79 buffer: &TensorBuffer<B, T>,
80 values: &[T],
81 ) -> Result<(), TensorDeviceError> {
82 if values.len() != buffer.length {
83 return Err(TensorDeviceError::InvalidBuffer("collective full write length mismatch"));
84 }
85 self.write_at(buffer, 0, values)
86 }
87
88 pub fn write_at<T: TensorElement>(
89 &self,
90 buffer: &TensorBuffer<B, T>,
91 offset: usize,
92 values: &[T],
93 ) -> Result<(), TensorDeviceError> {
94 self.validate_buffer(buffer)?;
95 let range = checked_range(buffer.length, offset, values.len())?;
96 if range.is_empty() {
97 return Ok(());
98 }
99 let incoming = self.from_values(values);
100 buffer.update(|current| {
101 if range.start == 0 && range.end == buffer.length {
102 incoming
103 } else {
104 current.assign(range, incoming)
105 }
106 })
107 }
108
109 pub fn download<T: TensorElement>(
110 &self,
111 buffer: &TensorBuffer<B, T>,
112 ) -> Result<Vec<T>, TensorDeviceError> {
113 self.download_at(buffer, 0, buffer.length)
114 }
115
116 pub fn download_at<T: TensorElement>(
117 &self,
118 buffer: &TensorBuffer<B, T>,
119 offset: usize,
120 length: usize,
121 ) -> Result<Vec<T>, TensorDeviceError> {
122 self.validate_buffer(buffer)?;
123 let range = checked_range(buffer.length, offset, length)?;
124 if range.is_empty() {
125 return Ok(Vec::new());
126 }
127 let mut value = buffer.snapshot()?;
128 if range.start != 0 || range.end != buffer.length {
129 value = value.slice(range);
130 }
131 let data = match value {
132 Primitive::Float(value) => read_sync(B::float_into_data(value))?,
133 Primitive::Int(value) => read_sync(B::int_into_data(value))?,
134 };
135 data.into_vec::<T>().map_err(|error| TensorDeviceError::Data(format!("{error:?}")))
136 }
137}