1use crate::{Error, Result, UnaryOp};
2use ruda_core::tensor::DType;
3use ruda_kernel::{dsl::Runtime, tensor::RudaTensor};
4
5pub type Mode = i32;
7
8#[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 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 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#[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}