Skip to main content

slop_tensor/
inner.rs

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    /// Constructs a tensor from initialized storage and dimensions.
27    ///
28    /// Returns an error when the physical storage length does not exactly match the declared
29    /// dimensions.
30    #[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    /// Returns whether the physical storage length matches the declared dimensions.
44    #[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    /// # Safety
119    ///
120    /// The caller must ensure that the new dimensions are compatible with the existing dimensions.
121    #[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    /// Returns the dimensions of the tensor.
163    #[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    /// # Safety
179    ///
180    /// This function is unsafe because it enables bypassing the lifetime of the tensor.
181    #[inline]
182    pub unsafe fn owned_unchecked(&self) -> ManuallyDrop<Self> {
183        self.owned_unchecked_in(self.storage.allocator().clone())
184    }
185
186    /// # Safety
187    ///
188    /// This function is unsafe because it enables bypassing the lifetime of the tensor.
189    #[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    /// # Safety
258    ///
259    /// See [std::mem::MaybeUninit::assume_init].
260    #[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 to ensure that the view is not used after the original tensor is freed.
399    _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    /// # Safety
420    ///
421    /// The caller must ensure that the pointer is valid for the given dimensions and backend.
422    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 to ensure that we get an exlusive reference, and that the view is not used after the
535    /// original tensor is freed.
536    _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// A macro to create a 1D or 2D tensor from a list of elements.
651#[macro_export]
652macro_rules! tensor {
653    // ----- 2D pattern: e.g. tensor![[1,2,3], [4,5,6]] -----
654    //
655    // Matches a top-level array of sub-arrays: [ [a,b,c], [d,e,f], ... ].
656    // Each sub-array is 1D. We gather them all in a Vec<Vec<_>>,
657    // check that all rows have the same length, flatten them,
658    // and reshape into a 2D Tensor.
659
660    ($([$($elem:expr),* $(,)?]),+ $(,)?) => {{
661        // Gather each sub-array into a temporary Vec<Vec<T>>.
662        let rows = vec![
663            $(
664                vec![$($elem,)*]
665            ),*
666        ];
667
668        // Check that all rows have the same length.
669        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        // Flatten everything into a single Vec<T>.
676        let flattened = rows.into_iter().flatten().collect::<Vec<_>>();
677
678        // Build the Tensor and reshape it to [rows_count, row_len].
679        // (We assume .reshape([..]) returns Self in your code.)
680        $crate::Tensor::from(flattened).reshape([rows_count, row_len])
681    }};
682
683    // ----- 1D pattern with outer brackets: e.g. tensor!([1, 2, 3]) -----
684    //
685    // If you do want “bare” bracket usage to produce a 1D Tensor (shape = [3]).
686
687    ([$($elem:expr),* $(,)?]) => {{
688        let v = vec![$($elem,)*];
689        $crate::Tensor::from(v)
690    }};
691
692    // ----- 1D “bare” comma‐separated: e.g. tensor![1, 2, 3] -----
693    //
694    // Matches a simple comma list at top-level.
695
696    ($($elem:expr),+ $(,)?) => {{
697        let v = vec![$($elem,)*];
698        $crate::Tensor::from(v)
699    }};
700}
701
702// Make a serialize and deserialize for Tensor<T> using the fact that we can serialize the buffer
703// and the dimensions.
704
705impl<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}