Skip to main content

rutensor/
descriptor.rs

1use crate::{Error, Result, UnaryOp};
2use ruda_core::tensor::DType;
3use ruda_kernel::{dsl::Runtime, tensor::RudaTensor};
4
5/// A named tensor index. Names identify axes independently of their position.
6pub type Mode = i32;
7
8/// Extents and element strides, relative to the start of a tensor's buffer view.
9#[derive(Clone, Debug, PartialEq, Eq)]
10pub struct TensorDescriptor {
11    pub(crate) extents: Vec<usize>,
12    pub(crate) strides: Vec<usize>,
13    pub(crate) dtype: DType,
14    elements: usize,
15    bytes: usize,
16}
17
18pub(crate) fn product(values: &[usize]) -> Result<usize> {
19    if values.contains(&0) { return Ok(0); }
20    values.iter().try_fold(1usize, |n, &d| n.checked_mul(d).ok_or(Error::Overflow))
21}
22
23pub(crate) fn check_dtype(dtype: DType) -> Result<()> {
24    match dtype {
25        DType::F16 | DType::BF16 | DType::F32 | DType::F64 => Ok(()),
26        _ => Err(Error::UnsupportedDType(format!("{dtype:?}; use F16, BF16, F32 or F64"))),
27    }
28}
29
30impl TensorDescriptor {
31    pub fn new(extents: &[usize], strides: &[usize], dtype: DType) -> Result<Self> {
32        check_dtype(dtype)?;
33        if extents.len() != strides.len() {
34            return Err(Error::InvalidDescriptor("extent and stride ranks differ".into()));
35        }
36        let elements = product(extents)?;
37        let span = if elements == 0 { 0 } else {
38            extents.iter().zip(strides).try_fold(1usize, |span, (&d, &s)| {
39                span.checked_add((d - 1).checked_mul(s).ok_or(Error::Overflow)?)
40                    .ok_or(Error::Overflow)
41            })?
42        };
43        let bytes = span.checked_mul(dtype.size()).ok_or(Error::Overflow)?;
44        Ok(Self { extents: extents.to_vec(), strides: strides.to_vec(), dtype, elements, bytes })
45    }
46
47    /// Packed row-major layout; an empty extent list denotes a scalar.
48    pub fn contiguous(extents: &[usize], dtype: DType) -> Result<Self> {
49        if extents.contains(&0) {
50            return Self::new(extents, &vec![0; extents.len()], dtype);
51        }
52        let mut strides = vec![1; extents.len()];
53        let mut stride = 1usize;
54        for axis in (0..extents.len()).rev() {
55            strides[axis] = stride;
56            stride = stride.checked_mul(extents[axis].max(1)).ok_or(Error::Overflow)?;
57        }
58        Self::new(extents, &strides, dtype)
59    }
60
61    pub fn from_tensor<R: Runtime>(tensor: &RudaTensor<R>) -> Result<Self> {
62        if tensor.qparams.is_some() {
63            return Err(Error::UnsupportedDType("quantized tensor".into()));
64        }
65        let descriptor = Self::new(tensor.meta.shape(), tensor.meta.strides(), tensor.dtype)?;
66        descriptor.check_buffer(tensor)?;
67        Ok(descriptor)
68    }
69
70    pub fn extents(&self) -> &[usize] { &self.extents }
71    pub fn strides(&self) -> &[usize] { &self.strides }
72    pub fn dtype(&self) -> DType { self.dtype }
73    pub fn rank(&self) -> usize { self.extents.len() }
74    pub fn num_elements(&self) -> usize { self.elements }
75    pub fn storage_bytes(&self) -> usize { self.bytes }
76
77    /// Whether the axes form a writable, non-overlapping strided layout.
78    pub fn is_nonoverlapping(&self) -> bool {
79        if self.elements == 0 { return true; }
80        let mut axes: Vec<_> = self.extents.iter().zip(&self.strides)
81            .filter(|(d, _)| **d > 1).collect();
82        axes.sort_by_key(|(_, s)| **s);
83        let mut span = 1usize;
84        for (&d, &s) in axes {
85            if s < span { return false; }
86            span += (d - 1) * s;
87        }
88        true
89    }
90
91    pub(crate) fn matches<R: Runtime>(&self, tensor: &RudaTensor<R>) -> bool {
92        self.dtype == tensor.dtype && tensor.qparams.is_none()
93            && self.extents.as_slice() == &tensor.meta.shape()[..]
94            && self.strides.as_slice() == &tensor.meta.strides()[..]
95    }
96
97    pub(crate) fn check_buffer<R: Runtime>(&self, tensor: &RudaTensor<R>) -> Result<()> {
98        if (tensor.handle.size() as u128) < self.bytes as u128 {
99            return Err(Error::BufferTooSmall);
100        }
101        Ok(())
102    }
103}
104
105/// Associates index names and a unary transform with a tensor descriptor.
106#[derive(Clone, Debug, PartialEq, Eq)]
107pub struct OperandDescriptor {
108    pub(crate) tensor: TensorDescriptor,
109    pub(crate) modes: Vec<Mode>,
110    pub(crate) unary: UnaryOp,
111}
112
113impl OperandDescriptor {
114    pub fn new(tensor: TensorDescriptor, modes: &[Mode]) -> Result<Self> {
115        if tensor.rank() != modes.len() {
116            return Err(Error::InvalidDescriptor("one mode is required for each axis".into()));
117        }
118        for (axis, mode) in modes.iter().enumerate() {
119            for previous in 0..axis {
120                if modes[previous] == *mode && tensor.extents[previous] != tensor.extents[axis] {
121                    return Err(Error::InvalidDescriptor("repeated modes require equal diagonal extents".into()));
122                }
123            }
124        }
125        Ok(Self { tensor, modes: modes.to_vec(), unary: UnaryOp::Identity })
126    }
127
128    pub fn from_tensor<R: Runtime>(tensor: &RudaTensor<R>, modes: &[Mode]) -> Result<Self> {
129        Self::new(TensorDescriptor::from_tensor(tensor)?, modes)
130    }
131
132    pub fn with_unary(mut self, unary: UnaryOp) -> Self { self.unary = unary; self }
133    pub fn tensor(&self) -> &TensorDescriptor { &self.tensor }
134    pub fn modes(&self) -> &[Mode] { &self.modes }
135    pub fn unary(&self) -> UnaryOp { self.unary }
136}