1use bytes::Bytes;
11
12use crate::error::{Error, Result};
13use crate::prelude::*;
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
17pub enum DType {
18 F16,
21 F32,
23 F64,
25 I8,
27 I16,
29 I32,
31 I64,
33 U8,
35 U16,
37 U32,
39 U64,
41}
42
43impl DType {
44 pub const fn size(self) -> usize {
46 match self {
47 DType::I8 | DType::U8 => 1,
48 DType::F16 | DType::I16 | DType::U16 => 2,
49 DType::F32 | DType::I32 | DType::U32 => 4,
50 DType::F64 | DType::I64 | DType::U64 => 8,
51 }
52 }
53
54 pub const fn long_name(self) -> &'static str {
56 match self {
57 DType::F16 => "float16",
58 DType::F32 => "float32",
59 DType::F64 => "float64",
60 DType::I8 => "int8",
61 DType::I16 => "int16",
62 DType::I32 => "int32",
63 DType::I64 => "int64",
64 DType::U8 => "uint8",
65 DType::U16 => "uint16",
66 DType::U32 => "uint32",
67 DType::U64 => "uint64",
68 }
69 }
70
71 pub const fn short_name(self) -> &'static str {
73 match self {
74 DType::F16 => "f2",
75 DType::F32 => "f4",
76 DType::F64 => "f8",
77 DType::I8 => "i1",
78 DType::I16 => "i2",
79 DType::I32 => "i4",
80 DType::I64 => "i8",
81 DType::U8 => "u1",
82 DType::U16 => "u2",
83 DType::U32 => "u4",
84 DType::U64 => "u8",
85 }
86 }
87
88 pub fn parse(name: &str) -> Result<DType> {
90 let dt = match name {
91 "float16" | "f2" => DType::F16,
92 "float32" | "f4" => DType::F32,
93 "float64" | "f8" => DType::F64,
94 "int8" | "i1" => DType::I8,
95 "int16" | "i2" => DType::I16,
96 "int32" | "i4" => DType::I32,
97 "int64" | "i8" => DType::I64,
98 "uint8" | "u1" => DType::U8,
99 "uint16" | "u2" => DType::U16,
100 "uint32" | "u4" => DType::U32,
101 "uint64" | "u8" => DType::U64,
102 other => return Err(Error::unsupported(format!("dtype {other}"))),
103 };
104 Ok(dt)
105 }
106
107 pub const fn is_float(self) -> bool {
109 matches!(self, DType::F16 | DType::F32 | DType::F64)
110 }
111}
112
113#[derive(Debug, Clone, PartialEq, Eq)]
118pub struct Tensor {
119 dtype: DType,
120 shape: Vec<usize>,
121 data: Bytes,
122}
123
124impl Tensor {
125 pub fn new(dtype: DType, shape: Vec<usize>, data: impl Into<Bytes>) -> Result<Tensor> {
130 let data = data.into();
131 let expected = shape.iter().product::<usize>() * dtype.size();
132 if data.len() != expected {
133 return Err(Error::format(format!(
134 "tensor data is {} bytes but shape {shape:?} of {} needs {expected}",
135 data.len(),
136 dtype.long_name()
137 )));
138 }
139 Ok(Tensor { dtype, shape, data })
140 }
141
142 pub fn from_f32(values: &[f32]) -> Tensor {
144 Self::from_f32_shaped(values, vec![values.len()])
145 }
146
147 pub fn from_f32_shaped(values: &[f32], shape: Vec<usize>) -> Tensor {
152 assert_eq!(shape.iter().product::<usize>(), values.len(), "shape does not match value count");
153 let mut data = Vec::with_capacity(values.len() * 4);
154 for v in values {
155 data.extend_from_slice(&v.to_ne_bytes());
156 }
157 Tensor { dtype: DType::F32, shape, data: data.into() }
158 }
159
160 pub fn from_u8_shaped(values: Vec<u8>, shape: Vec<usize>) -> Tensor {
165 assert_eq!(shape.iter().product::<usize>(), values.len(), "shape does not match value count");
166 Tensor { dtype: DType::U8, shape, data: values.into() }
167 }
168
169 pub fn dtype(&self) -> DType {
171 self.dtype
172 }
173
174 pub fn shape(&self) -> &[usize] {
176 &self.shape
177 }
178
179 pub fn ndim(&self) -> usize {
181 self.shape.len()
182 }
183
184 pub fn numel(&self) -> usize {
186 self.shape.iter().product()
187 }
188
189 pub fn data(&self) -> &Bytes {
191 &self.data
192 }
193
194 pub fn into_data(self) -> Bytes {
196 self.data
197 }
198
199 pub fn reshape(&self, shape: Vec<usize>) -> Result<Tensor> {
203 Tensor::new(self.dtype, shape, self.data.clone())
204 }
205
206 pub fn to_f64_vec(&self) -> Vec<f64> {
208 let n = self.numel();
209 let mut out = Vec::with_capacity(n);
210 let d = &self.data[..];
211 for i in 0..n {
212 out.push(self.element_at(d, i));
213 }
214 out
215 }
216
217 pub fn to_f32_vec(&self) -> Vec<f32> {
219 self.to_f64_vec().into_iter().map(|v| v as f32).collect()
220 }
221
222 pub fn get(&self, index: usize) -> f64 {
227 assert!(index < self.numel(), "index {index} out of bounds");
228 self.element_at(&self.data[..], index)
229 }
230
231 fn element_at(&self, d: &[u8], i: usize) -> f64 {
232 let s = self.dtype.size();
233 let b = &d[i * s..(i + 1) * s];
234 match self.dtype {
235 DType::U8 => b[0] as f64,
236 DType::I8 => b[0] as i8 as f64,
237 DType::U16 => u16::from_ne_bytes([b[0], b[1]]) as f64,
238 DType::I16 => i16::from_ne_bytes([b[0], b[1]]) as f64,
239 DType::F16 => f16_to_f64(u16::from_ne_bytes([b[0], b[1]])),
240 DType::U32 => u32::from_ne_bytes([b[0], b[1], b[2], b[3]]) as f64,
241 DType::I32 => i32::from_ne_bytes([b[0], b[1], b[2], b[3]]) as f64,
242 DType::F32 => f32::from_ne_bytes([b[0], b[1], b[2], b[3]]) as f64,
243 DType::U64 => u64::from_ne_bytes(b.try_into().expect("8 bytes")) as f64,
244 DType::I64 => i64::from_ne_bytes(b.try_into().expect("8 bytes")) as f64,
245 DType::F64 => f64::from_ne_bytes(b.try_into().expect("8 bytes")),
246 }
247 }
248}
249
250fn f16_to_f64(bits: u16) -> f64 {
252 let sign = if bits & 0x8000 != 0 { -1.0f64 } else { 1.0 };
253 let exponent = ((bits >> 10) & 0x1f) as i32;
254 let mantissa = (bits & 0x3ff) as f64;
255 match exponent {
256 0 => sign * mantissa * exp2(-24),
257 0x1f if mantissa == 0.0 => sign * f64::INFINITY,
258 0x1f => f64::NAN,
259 _ => sign * (1.0 + mantissa / 1024.0) * exp2(exponent - 15),
260 }
261}
262
263fn exp2(n: i32) -> f64 {
268 let mut out = 1.0f64;
269 for _ in 0..n.abs() {
270 if n >= 0 {
271 out *= 2.0;
272 } else {
273 out /= 2.0;
274 }
275 }
276 out
277}
278
279#[cfg(test)]
280mod tests {
281 use super::*;
282
283 #[test]
284 fn round_trips_f32_values() {
285 let t = Tensor::from_f32_shaped(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]);
286 assert_eq!(t.shape(), &[2, 3]);
287 assert_eq!(t.numel(), 6);
288 assert_eq!(t.dtype(), DType::F32);
289 assert_eq!(t.to_f64_vec(), vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
290 assert_eq!(t.get(4), 5.0);
291 }
292
293 #[test]
294 fn rejects_mismatched_shapes() {
295 let err = Tensor::new(DType::F32, vec![3], vec![0u8; 8]);
296 assert!(err.is_err());
297 }
298
299 #[test]
300 fn reshapes_without_copying_semantics() {
301 let t = Tensor::from_f32(&[1.0, 2.0, 3.0, 4.0]);
302 let r = t.reshape(vec![2, 2]).unwrap();
303 assert_eq!(r.shape(), &[2, 2]);
304 assert!(t.reshape(vec![3, 3]).is_err());
305 }
306
307 #[test]
308 fn parses_dtype_names() {
309 assert_eq!(DType::parse("float32").unwrap(), DType::F32);
310 assert_eq!(DType::parse("f4").unwrap(), DType::F32);
311 assert_eq!(DType::parse("u1").unwrap(), DType::U8);
312 assert!(DType::parse("complex64").is_err());
313 }
314}