Skip to main content

ruccl/tensor_device/
storage.rs

1use 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}