1use std::{
2 marker::PhantomData,
3 mem::ManuallyDrop,
4 ops::{Index, IndexMut},
5};
6
7use derive_where::derive_where;
8use rand::{distributions::Standard, prelude::Distribution, Rng};
9use serde::{ser::SerializeStruct, Deserialize, Deserializer, Serialize, Serializer};
10use slop_algebra::{ExtensionField, Field};
11use slop_alloc::{
12 Backend, Buffer, CpuBackend, HasBackend, Init, TryReserveError, GLOBAL_CPU_BACKEND,
13};
14use slop_matrix::Matrix;
15
16use crate::{Dimensions, DimensionsError};
17
18#[derive(Debug, Clone)]
19#[derive_where(PartialEq, Eq; Buffer<T, A>)]
20pub struct Tensor<T, A: Backend = CpuBackend> {
21 pub storage: Buffer<T, A>,
22 pub dimensions: Dimensions,
23}
24
25impl<T, A: Backend> Tensor<T, A> {
26 #[inline]
31 pub fn try_from_parts(
32 storage: Buffer<T, A>,
33 dimensions: Dimensions,
34 ) -> Result<Self, DimensionsError> {
35 let expected = dimensions.total_len();
36 let actual = storage.len();
37 if actual != expected {
38 return Err(DimensionsError::NumElementsMismatch(expected, actual));
39 }
40 Ok(Self { storage, dimensions })
41 }
42
43 #[inline]
45 pub fn has_valid_shape(&self) -> bool {
46 self.storage.len() == self.dimensions.total_len()
47 }
48
49 #[inline]
50 pub fn with_sizes_in(sizes: impl AsRef<[usize]>, allocator: A) -> Self {
51 Self::try_with_sizes_in(sizes, allocator).unwrap()
52 }
53
54 #[inline]
55 pub fn zeros_in(sizes: impl AsRef<[usize]>, allocator: A) -> Self {
56 let mut tensor = Self::with_sizes_in(sizes, allocator);
57 tensor.storage.write_bytes(0, tensor.total_len() * std::mem::size_of::<T>()).unwrap();
58 tensor
59 }
60
61 #[inline]
62 pub fn zeros_in_with_total_capacity(sizes: impl AsRef<[usize]>, allocator: A) -> Self {
63 let mut tensor = Self::with_sizes_in(sizes, allocator);
64 tensor.storage.write_bytes(0, tensor.total_len() * std::mem::size_of::<T>()).unwrap();
65 tensor
66 }
67
68 #[inline]
69 pub fn try_with_sizes_in(
70 sizes: impl AsRef<[usize]>,
71 allocator: A,
72 ) -> Result<Self, TryReserveError> {
73 let dimensions = Dimensions::try_from(sizes.as_ref()).unwrap();
74 Ok(Self {
75 storage: Buffer::try_with_capacity_in(dimensions.total_len(), allocator)?,
76 dimensions,
77 })
78 }
79
80 #[track_caller]
81 pub fn reshape_in_place(&mut self, sizes: impl AsRef<[usize]>) {
82 #[cold]
83 #[track_caller]
84 #[inline(never)]
85 fn dimension_fail(new_dimensions: &Dimensions, old_dimensions: &Dimensions) -> ! {
86 panic!(
87 "TensorView::reshape: dimension mismatch: {new_dimensions:?} vs {old_dimensions:?}"
88 );
89 }
90
91 let dimensions: Dimensions = sizes.as_ref().try_into().unwrap();
92 if self.dimensions.compatible(&dimensions).is_err() {
93 dimension_fail(&dimensions, &self.dimensions);
94 }
95 self.dimensions = dimensions;
96 }
97
98 #[inline]
99 #[track_caller]
100 pub fn reshape(mut self, sizes: impl AsRef<[usize]>) -> Self {
101 #[cold]
102 #[track_caller]
103 #[inline(never)]
104 fn dimension_fail(new_dimensions: &Dimensions, old_dimensions: &Dimensions) -> ! {
105 panic!(
106 "TensorView::reshape: dimension mismatch: {new_dimensions:?} vs {old_dimensions:?}"
107 );
108 }
109
110 let dimensions: Dimensions = sizes.as_ref().try_into().unwrap();
111 if self.dimensions.compatible(&dimensions).is_err() {
112 dimension_fail(&dimensions, &self.dimensions);
113 }
114 self.dimensions = dimensions;
115 self
116 }
117
118 #[inline]
122 pub unsafe fn reshape_unchecked(mut self, dimensions: Dimensions) {
123 self.dimensions = dimensions;
124 }
125
126 #[inline]
127 pub fn flatten_in_place(&mut self) {
128 self.reshape_in_place([self.dimensions.total_len()]);
129 }
130
131 #[inline]
132 pub fn flatten(mut self) -> Self {
133 self.flatten_in_place();
134 self
135 }
136
137 #[inline]
138 pub fn into_buffer(self) -> Buffer<T, A> {
139 self.storage
140 }
141
142 #[inline]
143 pub fn as_buffer(&self) -> &Buffer<T, A> {
144 &self.storage
145 }
146
147 #[inline]
148 pub fn as_mut_buffer(&mut self) -> &mut Buffer<T, A> {
149 &mut self.storage
150 }
151
152 #[inline]
153 pub fn backend(&self) -> &A {
154 self.storage.allocator()
155 }
156
157 #[inline]
158 pub fn shape(&self) -> &Dimensions {
159 &self.dimensions
160 }
161
162 #[inline]
164 pub fn sizes(&self) -> &[usize] {
165 self.dimensions.sizes()
166 }
167
168 #[inline]
169 pub fn strides(&self) -> &[usize] {
170 self.dimensions.strides()
171 }
172
173 #[inline]
174 pub fn as_ptr(&self) -> *const T {
175 self.storage.as_ptr()
176 }
177
178 #[inline]
182 pub unsafe fn owned_unchecked(&self) -> ManuallyDrop<Self> {
183 self.owned_unchecked_in(self.storage.allocator().clone())
184 }
185
186 #[inline]
190 pub unsafe fn owned_unchecked_in(&self, storage_allocator: A) -> ManuallyDrop<Self> {
191 let dimensions = self.dimensions.clone();
192 let storage = self.storage.owned_unchecked_in(storage_allocator);
193 let storage = ManuallyDrop::into_inner(storage);
194 ManuallyDrop::new(Self { storage, dimensions })
195 }
196
197 #[inline]
198 pub fn total_len(&self) -> usize {
199 self.dimensions.total_len()
200 }
201
202 pub fn as_mut_ptr(&mut self) -> *mut T {
203 self.storage.as_mut_ptr()
204 }
205
206 #[inline]
207 pub fn as_view(&'_ self) -> TensorView<'_, T, A> {
208 assert!(
209 self.has_valid_shape(),
210 "Tensor::as_view: storage length {} does not match declared length {}",
211 self.storage.len(),
212 self.dimensions.total_len()
213 );
214 TensorView {
215 ptr: self.as_ptr(),
216 dimensions: self.dimensions.clone(),
217 backend: self.backend().clone(),
218 _marker: PhantomData,
219 }
220 }
221
222 #[inline]
223 pub fn as_view_mut(&'_ mut self) -> TensorViewMut<'_, T, A> {
224 assert!(
225 self.has_valid_shape(),
226 "Tensor::as_view_mut: storage length {} does not match declared length {}",
227 self.storage.len(),
228 self.dimensions.total_len()
229 );
230 TensorViewMut {
231 ptr: self.as_mut_ptr(),
232 dimensions: self.dimensions.clone(),
233 _marker: PhantomData,
234 }
235 }
236
237 #[inline]
238 pub fn get(&'_ self, index: usize) -> Option<TensorView<'_, T, A>> {
239 self.as_view().get(index)
240 }
241
242 #[inline]
243 pub fn get_mut(&'_ mut self, index: usize) -> Option<TensorViewMut<'_, T, A>> {
244 self.as_view_mut().get(index)
245 }
246
247 #[inline]
248 pub fn split(&'_ self) -> impl Iterator<Item = TensorView<'_, T, A>> {
249 self.as_view().split()
250 }
251
252 #[inline]
253 pub fn split_mut(&'_ mut self) -> impl Iterator<Item = TensorViewMut<'_, T, A>> {
254 self.as_view_mut().split_mut()
255 }
256
257 #[inline]
261 pub unsafe fn assume_init(&mut self) {
262 self.storage.set_len(self.storage.capacity());
263 }
264
265 pub fn flatten_to_base<F: Field>(self) -> Tensor<F, A>
266 where
267 T: ExtensionField<F>,
268 {
269 let [height, width]: [usize; 2] = self.sizes().try_into().unwrap();
270 let dimensions = Dimensions::try_from([height, T::D * width]).unwrap();
271 let data_storage = self.into_buffer().flatten_to_base();
272 Tensor { storage: data_storage, dimensions }
273 }
274
275 pub fn into_extension<ET: ExtensionField<T>>(self) -> Tensor<ET, A>
276 where
277 T: Field,
278 {
279 let [height, width]: [usize; 2] = self.sizes().try_into().unwrap();
280 let dimensions = Dimensions::try_from([height, width / ET::D]).unwrap();
281 let extension_storage = self.into_buffer().into_extension();
282 Tensor { storage: extension_storage, dimensions }
283 }
284}
285
286impl<T, A: Backend, I: AsRef<[usize]>> Index<I> for Tensor<T, A> {
287 type Output = Init<T, A>;
288
289 #[track_caller]
290 fn index(&self, index: I) -> &Self::Output {
291 #[cold]
292 #[track_caller]
293 #[inline(never)]
294 fn dimension_fail(index_len: usize, sizes_len: usize) -> ! {
295 panic!(
296 "Index length ({index_len}) does not match tensor dimensions length ({sizes_len})"
297 );
298 }
299
300 if index.as_ref().len() != self.dimensions.sizes().len() {
301 dimension_fail(index.as_ref().len(), self.dimensions.sizes().len());
302 }
303 let index = self.dimensions.index_map(index);
304 &self.storage[index]
305 }
306}
307
308impl<T, A: Backend, I: AsRef<[usize]>> IndexMut<I> for Tensor<T, A> {
309 fn index_mut(&mut self, index: I) -> &mut Self::Output {
310 let index = self.dimensions.index_map(index);
311 &mut self.storage[index]
312 }
313}
314
315impl<T, A: Backend> From<Buffer<T, A>> for Tensor<T, A> {
316 #[inline]
317 fn from(buffer: Buffer<T, A>) -> Self {
318 let dims = [buffer.len()].into_iter().collect();
319 Self { storage: buffer, dimensions: dims }
320 }
321}
322
323impl<T, A: Backend> HasBackend for Tensor<T, A> {
324 type Backend = A;
325
326 fn backend(&self) -> &Self::Backend {
327 self.backend()
328 }
329}
330
331impl<T> From<Vec<T>> for Tensor<T, CpuBackend> {
332 #[inline]
333 fn from(vec: Vec<T>) -> Self {
334 Self::from(Buffer::from(vec))
335 }
336}
337
338impl<T> FromIterator<T> for Tensor<T, CpuBackend> {
339 #[inline]
340 fn from_iter<I: IntoIterator<Item = T>>(iter: I) -> Self {
341 Self::from(iter.into_iter().collect::<Vec<_>>())
342 }
343}
344
345impl<T: Clone + Send + Sync> From<slop_matrix::dense::RowMajorMatrix<T>> for Tensor<T, CpuBackend> {
346 fn from(value: slop_matrix::dense::RowMajorMatrix<T>) -> Self {
347 let dimensions: Dimensions = [value.height(), value.width()].try_into().unwrap();
348 let storage = Buffer::from(value.values);
349 Self { storage, dimensions }
350 }
351}
352
353impl<T: Clone + Send + Sync> TryFrom<Tensor<T, CpuBackend>>
354 for slop_matrix::dense::RowMajorMatrix<T>
355{
356 type Error = DimensionsError;
357 fn try_from(value: Tensor<T, CpuBackend>) -> Result<Self, Self::Error> {
358 if value.sizes().len() != 2 {
359 return Err(DimensionsError::TooManyDimensions(value.sizes().len()));
360 }
361 let width = value.sizes()[1];
362 let values = value.storage.into_vec();
363 Ok(Self::new(values, width))
364 }
365}
366
367impl<T> Tensor<T, CpuBackend> {
368 pub fn rand<R: Rng>(rng: &mut R, sizes: impl AsRef<[usize]>) -> Self
369 where
370 Standard: Distribution<T>,
371 {
372 let dimensions: Dimensions = sizes.as_ref().try_into().unwrap();
373 let values = rng.sample_iter(Standard).take(dimensions.total_len()).collect::<Vec<_>>();
374 Self { storage: Buffer::from(values), dimensions }
375 }
376
377 #[inline]
378 pub fn with_sizes(sizes: impl AsRef<[usize]>) -> Self {
379 Tensor::with_sizes_in(sizes, GLOBAL_CPU_BACKEND)
380 }
381
382 #[inline]
383 pub fn as_slice(&self) -> &[T] {
384 &self.storage[..]
385 }
386
387 #[inline]
388 pub fn as_mut_slice(&mut self) -> &mut [T] {
389 &mut self.storage[..]
390 }
391}
392
393#[derive(Debug)]
394pub struct TensorView<'a, T, A: Backend = CpuBackend> {
395 ptr: *const T,
396 dimensions: Dimensions,
397 backend: A,
398 _marker: PhantomData<&'a Tensor<T, A>>,
400}
401
402impl<'a, T, A: Backend> TensorView<'a, T, A> {
403 #[inline]
404 pub fn as_ptr(&self) -> *const T {
405 self.ptr
406 }
407
408 #[inline]
409 pub fn sizes(&self) -> &[usize] {
410 self.dimensions.sizes()
411 }
412
413 #[inline]
414 pub fn backend(&self) -> &A {
415 &self.backend
416 }
417
418 #[inline]
419 pub unsafe fn from_raw_parts(ptr: *const T, dimensions: Dimensions, backend: A) -> Self {
423 Self { ptr, dimensions, backend, _marker: PhantomData }
424 }
425
426 #[inline]
427 pub fn strides(&self) -> &[usize] {
428 self.dimensions.strides()
429 }
430
431 #[inline]
432 pub fn total_len(&self) -> usize {
433 self.dimensions.total_len()
434 }
435
436 #[inline]
437 pub fn shape(&self) -> &Dimensions {
438 &self.dimensions
439 }
440
441 #[inline]
442 pub fn flatten(self) -> TensorView<'a, T, A> {
443 let total_len = self.total_len();
444 self.reshape([total_len])
445 }
446
447 #[inline]
448 #[track_caller]
449 pub fn reshape(self, sizes: impl AsRef<[usize]>) -> TensorView<'a, T, A> {
450 #[cold]
451 #[track_caller]
452 #[inline(never)]
453 fn dimension_fail(new_dimensions: &Dimensions, old_dimensions: &Dimensions) -> ! {
454 panic!(
455 "TensorView::reshape: dimension mismatch: {new_dimensions:?} vs {old_dimensions:?}"
456 );
457 }
458
459 let dimensions: Dimensions = sizes.as_ref().try_into().unwrap();
460 if self.dimensions.compatible(&dimensions).is_err() {
461 dimension_fail(&dimensions, &self.dimensions);
462 }
463 TensorView {
464 ptr: self.ptr,
465 dimensions,
466 backend: self.backend.clone(),
467 _marker: PhantomData,
468 }
469 }
470
471 #[inline]
472 pub fn get(mut self, index: usize) -> Option<Self> {
473 let size = self.dimensions.sizes_mut().remove(0);
474 if index >= size {
475 return None;
476 }
477 let stride = self.dimensions.strides_mut().remove(0);
478 let offset = index * stride;
479
480 let ptr = unsafe { self.ptr.add(offset) };
481 Some(Self {
482 ptr,
483 dimensions: self.dimensions,
484 backend: self.backend.clone(),
485 _marker: PhantomData,
486 })
487 }
488
489 pub fn split(self) -> impl Iterator<Item = Self> {
490 (0..self.dimensions.sizes()[0]).map(move |i| self.clone().get(i).unwrap())
491 }
492}
493
494impl<'a, T, A: Backend> Clone for TensorView<'a, T, A> {
495 fn clone(&self) -> Self {
496 Self {
497 ptr: self.ptr,
498 dimensions: self.dimensions.clone(),
499 backend: self.backend.clone(),
500 _marker: PhantomData,
501 }
502 }
503}
504
505impl<'a, T, A: Backend> From<&'a Tensor<T, A>> for TensorView<'a, T, A> {
506 fn from(tensor: &'a Tensor<T, A>) -> Self {
507 tensor.as_view()
508 }
509}
510
511impl<'a, T, A: Backend, I: AsRef<[usize]>> Index<I> for TensorView<'a, T, A> {
512 type Output = Init<T, A>;
513
514 #[inline]
515 fn index(&self, index: I) -> &Self::Output {
516 let index = self.dimensions.index_map(index);
517 unsafe {
518 let ptr = self.ptr.add(index) as *const Init<T, A>;
519 ptr.as_ref().unwrap()
520 }
521 }
522}
523
524impl<T> Default for Tensor<T, CpuBackend> {
525 fn default() -> Self {
526 Self::from(Buffer::default())
527 }
528}
529
530#[derive(Debug)]
531pub struct TensorViewMut<'a, T, A: Backend = CpuBackend> {
532 ptr: *mut T,
533 dimensions: Dimensions,
534 _marker: PhantomData<&'a mut Tensor<T, A>>,
537}
538
539impl<'a, T, A: Backend> TensorViewMut<'a, T, A> {
540 #[inline]
541 pub fn as_mut_ptr(&mut self) -> *mut T {
542 self.ptr
543 }
544
545 #[inline]
546 pub fn sizes(&self) -> &[usize] {
547 self.dimensions.sizes()
548 }
549
550 #[inline]
551 pub fn shape(&self) -> &Dimensions {
552 &self.dimensions
553 }
554
555 #[inline]
556 pub fn strides(&self) -> &[usize] {
557 self.dimensions.strides()
558 }
559
560 #[inline]
561 pub fn flatten(self) -> TensorViewMut<'a, T, A> {
562 let total_len = self.total_len();
563 self.reshape([total_len])
564 }
565
566 #[inline]
567 pub fn reshape(self, sizes: impl AsRef<[usize]>) -> TensorViewMut<'a, T, A> {
568 let dimensions: Dimensions = sizes.as_ref().try_into().unwrap();
569 self.dimensions.compatible(&dimensions).unwrap();
570 TensorViewMut { ptr: self.ptr, dimensions, _marker: PhantomData }
571 }
572
573 #[inline]
574 pub fn get(mut self, index: usize) -> Option<Self> {
575 let size = self.dimensions.sizes_mut().remove(0);
576 if index >= size {
577 return None;
578 }
579 let stride = self.dimensions.strides_mut().remove(0);
580 let offset = index * stride;
581
582 let ptr = unsafe { self.ptr.add(offset) };
583 Some(Self { ptr, dimensions: self.dimensions, _marker: PhantomData })
584 }
585
586 #[inline]
587 pub fn split_mut(self) -> impl Iterator<Item = Self> {
588 (0..self.dimensions.sizes()[0]).map(move |i| {
589 let self_copy =
590 Self { ptr: self.ptr, dimensions: self.dimensions.clone(), _marker: PhantomData };
591 self_copy.get(i).unwrap()
592 })
593 }
594
595 #[inline]
596 pub fn total_len(&self) -> usize {
597 self.dimensions.total_len()
598 }
599}
600
601impl<'a, T> TensorView<'a, T, CpuBackend> {
602 #[inline]
603 pub fn as_slice(self) -> &'a [T] {
604 unsafe { std::slice::from_raw_parts(self.ptr, self.dimensions.total_len()) }
605 }
606}
607
608impl<'a, T> TensorViewMut<'a, T, CpuBackend> {
609 #[inline]
610 pub fn as_slice(self) -> &'a [T] {
611 unsafe { std::slice::from_raw_parts(self.ptr, self.dimensions.total_len()) }
612 }
613
614 #[inline]
615 pub fn as_mut_slice(self) -> &'a mut [T] {
616 unsafe { std::slice::from_raw_parts_mut(self.ptr, self.dimensions.total_len()) }
617 }
618}
619
620impl<'a, T, A: Backend> From<&'a mut Tensor<T, A>> for TensorViewMut<'a, T, A> {
621 fn from(tensor: &'a mut Tensor<T, A>) -> Self {
622 tensor.as_view_mut()
623 }
624}
625
626impl<'a, T, A: Backend, I: AsRef<[usize]>> Index<I> for TensorViewMut<'a, T, A> {
627 type Output = Init<T, A>;
628
629 #[inline]
630 fn index(&self, index: I) -> &Self::Output {
631 let index = self.dimensions.index_map(index);
632 unsafe {
633 let ptr = self.ptr.add(index) as *const T as *const Init<T, A>;
634 ptr.as_ref().unwrap()
635 }
636 }
637}
638
639impl<'a, T, A: Backend, I: AsRef<[usize]>> IndexMut<I> for TensorViewMut<'a, T, A> {
640 #[inline]
641 fn index_mut(&mut self, index: I) -> &mut Self::Output {
642 let index = self.dimensions.index_map(index);
643 unsafe {
644 let ptr = self.ptr.add(index) as *mut Init<T, A>;
645 ptr.as_mut().unwrap()
646 }
647 }
648}
649
650#[macro_export]
652macro_rules! tensor {
653 ($([$($elem:expr),* $(,)?]),+ $(,)?) => {{
661 let rows = vec![
663 $(
664 vec![$($elem,)*]
665 ),*
666 ];
667
668 let row_len = rows[0].len();
670 let rows_count = rows.len();
671 if !rows.iter().all(|r| r.len() == row_len) {
672 panic!("All sub-lists must have the same length to form a 2D tensor.");
673 }
674
675 let flattened = rows.into_iter().flatten().collect::<Vec<_>>();
677
678 $crate::Tensor::from(flattened).reshape([rows_count, row_len])
681 }};
682
683 ([$($elem:expr),* $(,)?]) => {{
688 let v = vec![$($elem,)*];
689 $crate::Tensor::from(v)
690 }};
691
692 ($($elem:expr),+ $(,)?) => {{
697 let v = vec![$($elem,)*];
698 $crate::Tensor::from(v)
699 }};
700}
701
702impl<T: Serialize> Serialize for Tensor<T> {
706 fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
707 let mut state = serializer.serialize_struct("Tensor", 2)?;
708 state.serialize_field("storage", &self.storage)?;
709 state.serialize_field("dimensions", &self.dimensions)?;
710 state.end()
711 }
712}
713
714impl<'de, T: Deserialize<'de>> Deserialize<'de> for Tensor<T> {
715 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
716 #[derive(Deserialize)]
717 #[serde(field_identifier, rename_all = "lowercase")]
718 enum Field {
719 Storage,
720 Dimensions,
721 }
722
723 struct TensorVisitor<T>(PhantomData<T>);
724
725 impl<'de, T: Deserialize<'de>> serde::de::Visitor<'de> for TensorVisitor<T> {
726 type Value = Tensor<T>;
727
728 fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
729 formatter.write_str("struct Tensor")
730 }
731
732 fn visit_seq<V>(self, mut seq: V) -> Result<Self::Value, V::Error>
733 where
734 V: serde::de::SeqAccess<'de>,
735 {
736 let storage: Buffer<T> = seq
737 .next_element()?
738 .ok_or_else(|| serde::de::Error::invalid_length(0, &self))?;
739 let dimensions: Dimensions = seq
740 .next_element()?
741 .ok_or_else(|| serde::de::Error::invalid_length(1, &self))?;
742 Tensor::try_from_parts(storage, dimensions).map_err(serde::de::Error::custom)
743 }
744
745 fn visit_map<V>(self, mut map: V) -> Result<Self::Value, V::Error>
746 where
747 V: serde::de::MapAccess<'de>,
748 {
749 let mut storage = None;
750 let mut dimensions = None;
751
752 while let Some(key) = map.next_key()? {
753 match key {
754 Field::Storage => {
755 if storage.is_some() {
756 return Err(serde::de::Error::duplicate_field("storage"));
757 }
758 storage = Some(map.next_value()?);
759 }
760 Field::Dimensions => {
761 if dimensions.is_some() {
762 return Err(serde::de::Error::duplicate_field("dimensions"));
763 }
764 dimensions = Some(map.next_value()?);
765 }
766 }
767 }
768
769 let storage = storage.ok_or_else(|| serde::de::Error::missing_field("storage"))?;
770 let dimensions =
771 dimensions.ok_or_else(|| serde::de::Error::missing_field("dimensions"))?;
772 Tensor::try_from_parts(storage, dimensions).map_err(serde::de::Error::custom)
773 }
774 }
775
776 deserializer.deserialize_struct(
777 "Tensor",
778 &["storage", "dimensions"],
779 TensorVisitor(PhantomData),
780 )
781 }
782}
783
784#[cfg(test)]
785mod tests {
786
787 use slop_alloc::buffer;
788
789 use super::*;
790
791 #[test]
792 fn test_tensor_element_index() {
793 let tensor = Tensor::<u32>::from(buffer![1, 2, 3, 4, 5, 6, 7, 8, 9, 10]).reshape([2, 5]);
794 assert_eq!(*tensor[[0, 0]], 1);
795 assert_eq!(*tensor[[0, 1]], 2);
796 assert_eq!(*tensor[[0, 2]], 3);
797 assert_eq!(*tensor[[0, 3]], 4);
798 assert_eq!(*tensor[[0, 4]], 5);
799 assert_eq!(*tensor[[1, 0]], 6);
800 assert_eq!(*tensor[[1, 1]], 7);
801 assert_eq!(*tensor[[1, 2]], 8);
802 assert_eq!(*tensor[[1, 3]], 9);
803 assert_eq!(*tensor[[1, 4]], 10);
804 }
805
806 #[test]
807 fn test_tensor_slice_index() {
808 let tensor = Tensor::<u32>::from(buffer![1, 2, 3, 4, 5, 6, 7, 8, 9, 10]).reshape([2, 5]);
809
810 let first_row = tensor.get(0).unwrap();
811 assert_eq!(first_row.sizes(), [5]);
812 assert_eq!(first_row.strides(), [1]);
813 assert_eq!(*first_row[[0]], 1);
814 assert_eq!(*first_row[[1]], 2);
815 assert_eq!(*first_row[[2]], 3);
816 assert_eq!(*first_row[[3]], 4);
817 assert_eq!(*first_row[[4]], 5);
818
819 let second_row = tensor.get(1).unwrap();
820 assert_eq!(*second_row[[0]], 6);
821 assert_eq!(*second_row[[1]], 7);
822 assert_eq!(*second_row[[2]], 8);
823 assert_eq!(*second_row[[3]], 9);
824 assert_eq!(*second_row[[4]], 10);
825
826 let tensor = Tensor::<u32>::from((0..24).collect::<Vec<_>>()).reshape([2, 3, 4]);
827 assert_eq!(*tensor[[0, 0, 0]], 0);
828 assert_eq!(*tensor[[0, 0, 1]], 1);
829 assert_eq!(*tensor[[0, 0, 2]], 2);
830 assert_eq!(*tensor[[0, 0, 3]], 3);
831 assert_eq!(*tensor[[0, 1, 0]], 4);
832 assert_eq!(*tensor[[0, 1, 1]], 5);
833 assert_eq!(*tensor[[0, 1, 2]], 6);
834 assert_eq!(*tensor[[0, 1, 3]], 7);
835 assert_eq!(*tensor[[0, 2, 0]], 8);
836 assert_eq!(*tensor[[0, 2, 1]], 9);
837 assert_eq!(*tensor[[0, 2, 2]], 10);
838 assert_eq!(*tensor[[0, 2, 3]], 11);
839 assert_eq!(*tensor[[1, 0, 0]], 12);
840 assert_eq!(*tensor[[1, 0, 1]], 13);
841 assert_eq!(*tensor[[1, 0, 2]], 14);
842 assert_eq!(*tensor[[1, 0, 3]], 15);
843 assert_eq!(*tensor[[1, 1, 0]], 16);
844 assert_eq!(*tensor[[1, 1, 1]], 17);
845 assert_eq!(*tensor[[1, 1, 2]], 18);
846 assert_eq!(*tensor[[1, 1, 3]], 19);
847 assert_eq!(*tensor[[1, 2, 0]], 20);
848 assert_eq!(*tensor[[1, 2, 1]], 21);
849 assert_eq!(*tensor[[1, 2, 2]], 22);
850 assert_eq!(*tensor[[1, 2, 3]], 23);
851 }
852
853 #[test]
854 fn test_p3_matrix_to_tensor() {
855 let mut rng = rand::thread_rng();
856 let matrix = slop_matrix::dense::RowMajorMatrix::<u32>::rand(&mut rng, 100, 400);
857 let tensor = Tensor::from(matrix.clone());
858
859 assert_eq!(tensor.sizes(), [100, 400]);
860
861 let matrix_back = slop_matrix::dense::RowMajorMatrix::<u32>::try_from(tensor).unwrap();
862 assert_eq!(matrix_back.values, matrix.values);
863 }
864
865 #[test]
866 fn test_tensor_macro() {
867 let tensor = tensor![1, 2, 3, 4, 5, 6];
868 assert_eq!(tensor.sizes(), [6]);
869 assert_eq!(tensor.as_slice(), [1, 2, 3, 4, 5, 6]);
870
871 let tensor = tensor![[1, 2, 3], [4, 5, 6]];
872 assert_eq!(tensor.sizes(), [2, 3]);
873 assert_eq!(tensor.as_slice(), [1, 2, 3, 4, 5, 6]);
874
875 let tensor = tensor![[1, 2, 3, 4, 5]];
876 assert_eq!(tensor.sizes(), [1, 5]);
877 assert_eq!(tensor.as_slice(), [1, 2, 3, 4, 5]);
878
879 let tensor = tensor![[1], [2], [3], [4], [5]];
880 assert_eq!(tensor.sizes(), [5, 1]);
881 assert_eq!(tensor.as_slice(), [1, 2, 3, 4, 5]);
882 }
883
884 #[test]
885 fn test_tensor_serialize_deserialize() {
886 let tensor = Tensor::<u32>::from(buffer![1, 2, 3, 4, 5, 6, 7, 8, 9, 10]).reshape([2, 5]);
887 let serialized = serde_json::to_string(&tensor).unwrap();
888 let deserialized: Tensor<u32> = serde_json::from_str(&serialized).unwrap();
889 assert_eq!(deserialized, tensor);
890 }
891
892 #[test]
893 fn test_tensor_deserialize_rejects_appended_storage() {
894 let result =
895 serde_json::from_str::<Tensor<u32>>(r#"{"storage":[1,2,3,4,5],"dimensions":[2,2]}"#);
896 assert!(result.is_err());
897 }
898
899 #[test]
900 fn test_tensor_deserialize_rejects_truncated_storage() {
901 let result =
902 serde_json::from_str::<Tensor<u32>>(r#"{"storage":[1,2,3],"dimensions":[2,2]}"#);
903 assert!(result.is_err());
904 }
905
906 #[test]
907 fn test_tensor_bincode_round_trip() {
908 let tensor = Tensor::<u32>::from(buffer![1, 2, 3, 4]).reshape([2, 2]);
909 let serialized = bincode::serialize(&tensor).unwrap();
910 let deserialized: Tensor<u32> = bincode::deserialize(&serialized).unwrap();
911 assert_eq!(deserialized, tensor);
912 }
913
914 #[test]
915 fn test_tensor_bincode_rejects_appended_storage() {
916 let tensor = Tensor {
917 storage: buffer![1, 2, 3, 4, 5],
918 dimensions: Dimensions::try_from([2, 2]).unwrap(),
919 };
920 let serialized = bincode::serialize(&tensor).unwrap();
921 assert!(bincode::deserialize::<Tensor<u32>>(&serialized).is_err());
922 }
923
924 #[test]
925 fn test_tensor_bincode_rejects_truncated_storage() {
926 let tensor =
927 Tensor { storage: buffer![1, 2, 3], dimensions: Dimensions::try_from([2, 2]).unwrap() };
928 let serialized = bincode::serialize(&tensor).unwrap();
929 assert!(bincode::deserialize::<Tensor<u32>>(&serialized).is_err());
930 }
931
932 #[test]
933 #[should_panic(expected = "storage length 1 does not match declared length 2")]
934 fn test_tensor_as_view_rejects_invalid_shape() {
935 let tensor = Tensor { storage: buffer![1], dimensions: Dimensions::try_from([2]).unwrap() };
936 let _ = tensor.as_view();
937 }
938
939 #[test]
940 #[should_panic(expected = "storage length 2 does not match declared length 1")]
941 fn test_tensor_as_view_mut_rejects_invalid_shape() {
942 let mut tensor =
943 Tensor { storage: buffer![1, 2], dimensions: Dimensions::try_from([1]).unwrap() };
944 let _ = tensor.as_view_mut();
945 }
946}