Skip to main content

strided_view/
view.rs

1//! Julia-like dynamic-rank strided view types.
2//!
3//! This module provides the canonical view types for strided operations,
4//! matching Julia's StridedViews.jl data model:
5//!
6//! - [`StridedView`]: Immutable dynamic-rank strided view with lazy element operations
7//! - [`StridedViewMut`]: Mutable dynamic-rank strided view (Identity op only)
8//! - [`StridedArray`]: Owned strided multidimensional array
9
10use std::marker::PhantomData;
11use std::ops::{Index, IndexMut};
12use std::sync::Arc;
13
14use crate::element_op::{ComposableElementOp, ElementOp, ElementOpApply, Identity};
15use crate::{Result, StridedError};
16
17#[inline]
18fn empty_layout(dims: &[usize]) -> bool {
19    dims.iter().any(|&dim| dim == 0)
20}
21
22#[inline]
23fn empty_const_ptr<T>() -> *const T {
24    core::ptr::NonNull::<T>::dangling().as_ptr()
25}
26
27#[inline]
28fn empty_mut_ptr<T>() -> *mut T {
29    core::ptr::NonNull::<T>::dangling().as_ptr()
30}
31
32// ============================================================================
33// Validation helpers
34// ============================================================================
35
36/// Validate that all accessed offsets stay within `[0, len)`.
37pub(crate) fn validate_bounds(
38    len: usize,
39    dims: &[usize],
40    strides: &[isize],
41    offset: isize,
42) -> Result<()> {
43    if dims.len() != strides.len() {
44        return Err(StridedError::StrideLengthMismatch);
45    }
46    // Empty array - no access needed
47    if dims.iter().any(|&d| d == 0) {
48        return Ok(());
49    }
50    // Compute min and max offsets
51    let mut min_offset = offset;
52    let mut max_offset = offset;
53    for (&dim, &stride) in dims.iter().zip(strides.iter()) {
54        if dim > 1 {
55            let last = dim
56                .checked_sub(1)
57                .and_then(|value| isize::try_from(value).ok())
58                .ok_or(StridedError::OffsetOverflow)?;
59            let end = stride
60                .checked_mul(last)
61                .ok_or(StridedError::OffsetOverflow)?;
62            if end >= 0 {
63                max_offset = max_offset
64                    .checked_add(end)
65                    .ok_or(StridedError::OffsetOverflow)?;
66            } else {
67                min_offset = min_offset
68                    .checked_add(end)
69                    .ok_or(StridedError::OffsetOverflow)?;
70            }
71        }
72    }
73    if min_offset < 0 || max_offset < 0 {
74        return Err(StridedError::OffsetOverflow);
75    }
76    if max_offset as usize >= len {
77        return Err(StridedError::OffsetOverflow);
78    }
79    Ok(())
80}
81
82/// Compute column-major strides (Julia default: first index varies fastest).
83pub fn col_major_strides(dims: &[usize]) -> Vec<isize> {
84    let rank = dims.len();
85    if rank == 0 {
86        return vec![];
87    }
88    let mut strides = vec![1isize; rank];
89    for i in 1..rank {
90        strides[i] = strides[i - 1] * dims[i - 1] as isize;
91    }
92    strides
93}
94
95/// Compute row-major strides (C default: last index varies fastest).
96pub fn row_major_strides(dims: &[usize]) -> Vec<isize> {
97    let rank = dims.len();
98    if rank == 0 {
99        return vec![];
100    }
101    let mut strides = vec![1isize; rank];
102    for i in (0..rank - 1).rev() {
103        strides[i] = strides[i + 1] * dims[i + 1] as isize;
104    }
105    strides
106}
107
108// ============================================================================
109// StridedView
110// ============================================================================
111
112/// Dynamic-rank immutable strided view with lazy element operations.
113///
114/// This is the Julia-equivalent `StridedView` type with:
115/// - Dynamic rank (dims/strides are heap-allocated)
116/// - Lazy element operations via the `Op` type parameter
117/// - Zero-copy transformations (permute, transpose, adjoint, conj)
118///
119/// # Type Parameters
120/// - `'a`: Lifetime of the underlying data
121/// - `T`: Element type
122/// - `Op`: Element operation applied lazily on access (default: `Identity`)
123pub struct StridedView<'a, T, Op = Identity> {
124    ptr: *const T,
125    data: &'a [T],
126    dims: Arc<[usize]>,
127    strides: Arc<[isize]>,
128    offset: isize,
129    _op: PhantomData<Op>,
130}
131
132unsafe impl<T: Send, Op: Send> Send for StridedView<'_, T, Op> {}
133unsafe impl<T: Sync, Op: Sync> Sync for StridedView<'_, T, Op> {}
134
135impl<T, Op> Clone for StridedView<'_, T, Op> {
136    fn clone(&self) -> Self {
137        Self {
138            ptr: self.ptr,
139            data: self.data,
140            dims: self.dims.clone(),
141            strides: self.strides.clone(),
142            offset: self.offset,
143            _op: PhantomData,
144        }
145    }
146}
147
148impl<T: std::fmt::Debug, Op> std::fmt::Debug for StridedView<'_, T, Op> {
149    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
150        f.debug_struct("StridedView")
151            .field("dims", &self.dims)
152            .field("strides", &self.strides)
153            .field("offset", &self.offset)
154            .finish()
155    }
156}
157
158impl<'a, T, Op> StridedView<'a, T, Op> {
159    /// Create a new immutable strided view from a borrowed slice.
160    pub fn new(data: &'a [T], dims: &[usize], strides: &[isize], offset: isize) -> Result<Self> {
161        validate_bounds(data.len(), dims, strides, offset)?;
162        let ptr = if empty_layout(dims) {
163            empty_const_ptr()
164        } else {
165            unsafe { data.as_ptr().offset(offset) }
166        };
167        Ok(Self {
168            ptr,
169            data,
170            dims: Arc::from(dims),
171            strides: Arc::from(strides),
172            offset,
173            _op: PhantomData,
174        })
175    }
176
177    /// Create a view without bounds checking.
178    ///
179    /// # Safety
180    /// The caller must ensure all index combinations stay within bounds.
181    pub unsafe fn new_unchecked(
182        data: &'a [T],
183        dims: &[usize],
184        strides: &[isize],
185        offset: isize,
186    ) -> Self {
187        let ptr = if empty_layout(dims) {
188            empty_const_ptr()
189        } else {
190            data.as_ptr().offset(offset)
191        };
192        Self {
193            ptr,
194            data,
195            dims: Arc::from(dims),
196            strides: Arc::from(strides),
197            offset,
198            _op: PhantomData,
199        }
200    }
201
202    /// Returns the shape (dimension sizes) of this view.
203    #[inline]
204    pub fn dims(&self) -> &[usize] {
205        &self.dims
206    }
207
208    /// Returns the strides (in units of `T`) for each dimension.
209    #[inline]
210    pub fn strides(&self) -> &[isize] {
211        &self.strides
212    }
213
214    /// Returns the byte offset into the backing data.
215    #[inline]
216    pub fn offset(&self) -> isize {
217        self.offset
218    }
219
220    /// Returns the number of dimensions (rank).
221    #[inline]
222    pub fn ndim(&self) -> usize {
223        self.dims.len()
224    }
225
226    /// Returns the total number of elements.
227    #[inline]
228    pub fn len(&self) -> usize {
229        self.dims.iter().product()
230    }
231
232    /// Returns `true` if any dimension is zero.
233    #[inline]
234    pub fn is_empty(&self) -> bool {
235        self.dims.iter().any(|&d| d == 0)
236    }
237
238    /// Returns a reference to the backing data slice.
239    #[inline]
240    pub fn data(&self) -> &'a [T] {
241        self.data
242    }
243
244    /// Raw const pointer to element at the view's base offset.
245    #[inline]
246    pub fn ptr(&self) -> *const T {
247        self.ptr
248    }
249
250    /// Permute dimensions.
251    pub fn permute(&self, perm: &[usize]) -> Result<StridedView<'a, T, Op>> {
252        let rank = self.dims.len();
253        if perm.len() != rank {
254            return Err(StridedError::RankMismatch(perm.len(), rank));
255        }
256        let mut seen = vec![false; rank];
257        for &p in perm {
258            if p >= rank {
259                return Err(StridedError::InvalidAxis { axis: p, rank });
260            }
261            if seen[p] {
262                return Err(StridedError::InvalidAxis { axis: p, rank });
263            }
264            seen[p] = true;
265        }
266        let new_dims: Vec<usize> = perm.iter().map(|&p| self.dims[p]).collect();
267        let new_strides: Vec<isize> = perm.iter().map(|&p| self.strides[p]).collect();
268        Ok(StridedView {
269            ptr: self.ptr,
270            data: self.data,
271            dims: Arc::from(new_dims),
272            strides: Arc::from(new_strides),
273            offset: self.offset,
274            _op: PhantomData,
275        })
276    }
277
278    /// Create a diagonal view by fusing repeated axis pairs via stride trick (zero-copy).
279    ///
280    /// For each pair `(a, b)`:
281    /// - New stride = `strides[a] + strides[b]`
282    /// - New dim = `min(dims[a], dims[b])`
283    /// - The higher-numbered axis is removed
284    /// - Pairs use **original** axis numbering
285    ///
286    /// # Example
287    /// `A[i,i,j]` shape=`[n,n,m]` strides=`[s0,s1,s2]` -> shape=`[n,m]` strides=`[s0+s1, s2]`
288    pub fn diagonal_view(&self, axis_pairs: &[(usize, usize)]) -> Result<StridedView<'a, T, Op>> {
289        let ndim = self.ndim();
290        let mut dims: Vec<usize> = self.dims().to_vec();
291        let mut strides: Vec<isize> = self.strides().to_vec();
292
293        let mut axes_to_remove = Vec::new();
294        for &(a, b) in axis_pairs {
295            let (lo, hi) = if a < b { (a, b) } else { (b, a) };
296            if lo >= ndim || hi >= ndim {
297                return Err(StridedError::InvalidAxis {
298                    axis: hi,
299                    rank: ndim,
300                });
301            }
302            if lo == hi {
303                return Err(StridedError::InvalidAxis {
304                    axis: lo,
305                    rank: ndim,
306                });
307            }
308            strides[lo] += strides[hi];
309            dims[lo] = dims[lo].min(dims[hi]);
310            axes_to_remove.push(hi);
311        }
312
313        axes_to_remove.sort_unstable();
314        axes_to_remove.dedup();
315        for &ax in axes_to_remove.iter().rev() {
316            dims.remove(ax);
317            strides.remove(ax);
318        }
319
320        unsafe {
321            Ok(StridedView::new_unchecked(
322                self.data(),
323                &dims,
324                &strides,
325                self.offset(),
326            ))
327        }
328    }
329
330    /// Broadcast this view to a target shape.
331    ///
332    /// Size-1 dimensions are expanded (stride set to 0) to match target.
333    pub fn broadcast(&self, target_dims: &[usize]) -> Result<StridedView<'a, T, Op>> {
334        if self.dims.len() != target_dims.len() {
335            return Err(StridedError::RankMismatch(
336                self.dims.len(),
337                target_dims.len(),
338            ));
339        }
340        let mut new_strides = Vec::with_capacity(self.dims.len());
341        for i in 0..self.dims.len() {
342            if self.dims[i] == target_dims[i] {
343                new_strides.push(self.strides[i]);
344            } else if self.dims[i] == 1 {
345                new_strides.push(0);
346            } else {
347                return Err(StridedError::ShapeMismatch(
348                    self.dims.to_vec(),
349                    target_dims.to_vec(),
350                ));
351            }
352        }
353        Ok(StridedView {
354            ptr: self.ptr,
355            data: self.data,
356            dims: Arc::from(target_dims),
357            strides: Arc::from(new_strides),
358            offset: self.offset,
359            _op: PhantomData,
360        })
361    }
362}
363
364/// Composition methods: require `T: ElementOpApply` and `Op: ComposableElementOp<T>`.
365impl<'a, T: Copy + ElementOpApply, Op: ComposableElementOp<T>> StridedView<'a, T, Op> {
366    /// Transpose a 2D view: reverses dimensions and composes the Transpose element op.
367    ///
368    /// Julia equivalent: `Base.transpose(a::AbstractStridedView{<:Any, 2})`
369    pub fn transpose_2d(&self) -> Result<StridedView<'a, T, Op::ComposeTranspose>> {
370        if self.dims.len() != 2 {
371            return Err(StridedError::RankMismatch(self.dims.len(), 2));
372        }
373        Ok(StridedView {
374            ptr: self.ptr,
375            data: self.data,
376            dims: Arc::new([self.dims[1], self.dims[0]]),
377            strides: Arc::new([self.strides[1], self.strides[0]]),
378            offset: self.offset,
379            _op: PhantomData,
380        })
381    }
382
383    /// Adjoint (conjugate transpose) of a 2D view.
384    ///
385    /// Julia equivalent: `Base.adjoint(a::AbstractStridedView{<:Any, 2})`
386    pub fn adjoint_2d(&self) -> Result<StridedView<'a, T, Op::ComposeAdjoint>> {
387        if self.dims.len() != 2 {
388            return Err(StridedError::RankMismatch(self.dims.len(), 2));
389        }
390        Ok(StridedView {
391            ptr: self.ptr,
392            data: self.data,
393            dims: Arc::new([self.dims[1], self.dims[0]]),
394            strides: Arc::new([self.strides[1], self.strides[0]]),
395            offset: self.offset,
396            _op: PhantomData,
397        })
398    }
399
400    /// Complex conjugate (compose Conj without changing dims/strides).
401    ///
402    /// Julia equivalent: `Base.conj(a::AbstractStridedView)`
403    pub fn conj(&self) -> StridedView<'a, T, Op::ComposeConj> {
404        StridedView {
405            ptr: self.ptr,
406            data: self.data,
407            dims: self.dims.clone(),
408            strides: self.strides.clone(),
409            offset: self.offset,
410            _op: PhantomData,
411        }
412    }
413}
414
415/// Element access: requires `Op: ElementOp<T>`.
416impl<'a, T: Copy, Op: ElementOp<T>> StridedView<'a, T, Op> {
417    /// Get an element with the element operation applied.
418    pub fn get(&self, indices: &[usize]) -> T {
419        assert_eq!(indices.len(), self.dims.len(), "wrong number of indices");
420        let mut idx = 0isize;
421        for (i, &index) in indices.iter().enumerate() {
422            assert!(
423                index < self.dims[i],
424                "index {} out of bounds for dim {}",
425                index,
426                self.dims[i]
427            );
428            idx += index as isize * self.strides[i];
429        }
430        Op::apply(unsafe { *self.ptr.offset(idx) })
431    }
432
433    /// Get an element without bounds checking.
434    ///
435    /// # Safety
436    /// Caller must ensure indices are within bounds.
437    #[inline]
438    pub unsafe fn get_unchecked(&self, indices: &[usize]) -> T {
439        let mut idx = 0isize;
440        for (i, &index) in indices.iter().enumerate() {
441            idx += index as isize * self.strides[i];
442        }
443        Op::apply(*self.ptr.offset(idx))
444    }
445}
446
447// ============================================================================
448// StridedViewMut
449// ============================================================================
450
451/// Dynamic-rank mutable strided view.
452///
453/// Always uses `Identity` element operation for write simplicity.
454/// Julia typically applies ops on the read side.
455pub struct StridedViewMut<'a, T> {
456    ptr: *mut T,
457    data: &'a mut [T],
458    dims: Arc<[usize]>,
459    strides: Arc<[isize]>,
460    offset: isize,
461}
462
463unsafe impl<T: Send> Send for StridedViewMut<'_, T> {}
464
465impl<T: std::fmt::Debug> std::fmt::Debug for StridedViewMut<'_, T> {
466    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
467        f.debug_struct("StridedViewMut")
468            .field("dims", &self.dims)
469            .field("strides", &self.strides)
470            .field("offset", &self.offset)
471            .finish()
472    }
473}
474
475impl<'a, T> StridedViewMut<'a, T> {
476    /// Create a new mutable strided view.
477    pub fn new(
478        data: &'a mut [T],
479        dims: &[usize],
480        strides: &[isize],
481        offset: isize,
482    ) -> Result<Self> {
483        validate_bounds(data.len(), dims, strides, offset)?;
484        let ptr = if empty_layout(dims) {
485            empty_mut_ptr()
486        } else {
487            unsafe { data.as_mut_ptr().offset(offset) }
488        };
489        Ok(Self {
490            ptr,
491            data,
492            dims: Arc::from(dims),
493            strides: Arc::from(strides),
494            offset,
495        })
496    }
497
498    /// Create without bounds checking.
499    ///
500    /// # Safety
501    /// Caller must ensure all index combinations stay within bounds.
502    pub unsafe fn new_unchecked(
503        data: &'a mut [T],
504        dims: &[usize],
505        strides: &[isize],
506        offset: isize,
507    ) -> Self {
508        let ptr = if empty_layout(dims) {
509            empty_mut_ptr()
510        } else {
511            data.as_mut_ptr().offset(offset)
512        };
513        Self {
514            ptr,
515            data,
516            dims: Arc::from(dims),
517            strides: Arc::from(strides),
518            offset,
519        }
520    }
521
522    /// Returns the shape (dimension sizes) of this view.
523    #[inline]
524    pub fn dims(&self) -> &[usize] {
525        &self.dims
526    }
527
528    /// Returns the strides (in units of `T`) for each dimension.
529    #[inline]
530    pub fn strides(&self) -> &[isize] {
531        &self.strides
532    }
533
534    /// Returns the byte offset into the backing data.
535    #[inline]
536    pub fn offset(&self) -> isize {
537        self.offset
538    }
539
540    /// Returns the number of dimensions (rank).
541    #[inline]
542    pub fn ndim(&self) -> usize {
543        self.dims.len()
544    }
545
546    /// Returns the total number of elements.
547    #[inline]
548    pub fn len(&self) -> usize {
549        self.dims.iter().product()
550    }
551
552    /// Returns `true` if any dimension is zero.
553    #[inline]
554    pub fn is_empty(&self) -> bool {
555        self.dims.iter().any(|&d| d == 0)
556    }
557
558    /// Raw const pointer to element at the view's base offset.
559    #[inline]
560    pub fn ptr(&self) -> *const T {
561        self.ptr as *const T
562    }
563
564    /// Raw mutable pointer to element at the view's base offset.
565    #[inline]
566    pub fn as_mut_ptr(&self) -> *mut T {
567        self.ptr
568    }
569
570    /// Returns a reference to the backing data slice.
571    #[inline]
572    pub fn data(&self) -> &[T] {
573        &self.data
574    }
575
576    /// Returns a mutable reference to the backing data slice.
577    #[inline]
578    pub fn data_mut(&mut self) -> &mut [T] {
579        &mut self.data
580    }
581
582    /// Permute dimensions, consuming the mutable view.
583    ///
584    /// Returns a new mutable view with reordered dimensions and strides.
585    /// Takes `self` by value to prevent aliasing of mutable views.
586    pub fn permute(self, perm: &[usize]) -> Result<StridedViewMut<'a, T>> {
587        let rank = self.dims.len();
588        if perm.len() != rank {
589            return Err(StridedError::RankMismatch(perm.len(), rank));
590        }
591        let mut seen = vec![false; rank];
592        for &p in perm {
593            if p >= rank {
594                return Err(StridedError::InvalidAxis { axis: p, rank });
595            }
596            if seen[p] {
597                return Err(StridedError::InvalidAxis { axis: p, rank });
598            }
599            seen[p] = true;
600        }
601        let new_dims: Vec<usize> = perm.iter().map(|&p| self.dims[p]).collect();
602        let new_strides: Vec<isize> = perm.iter().map(|&p| self.strides[p]).collect();
603        Ok(StridedViewMut {
604            ptr: self.ptr,
605            data: self.data,
606            dims: Arc::from(new_dims),
607            strides: Arc::from(new_strides),
608            offset: self.offset,
609        })
610    }
611
612    /// Reborrow as an immutable view.
613    pub fn as_view(&self) -> StridedView<'_, T, Identity> {
614        StridedView {
615            ptr: self.ptr as *const T,
616            data: unsafe { std::slice::from_raw_parts(self.data.as_ptr(), self.data.len()) },
617            dims: self.dims.clone(),
618            strides: self.strides.clone(),
619            offset: self.offset,
620            _op: PhantomData,
621        }
622    }
623}
624
625impl<'a, T: Copy> StridedViewMut<'a, T> {
626    /// Get an element.
627    pub fn get(&self, indices: &[usize]) -> T {
628        assert_eq!(indices.len(), self.dims.len());
629        let mut idx = 0isize;
630        for (i, &index) in indices.iter().enumerate() {
631            assert!(index < self.dims[i]);
632            idx += index as isize * self.strides[i];
633        }
634        unsafe { *self.ptr.offset(idx) }
635    }
636
637    /// Set an element.
638    pub fn set(&mut self, indices: &[usize], value: T) {
639        assert_eq!(indices.len(), self.dims.len());
640        let mut idx = 0isize;
641        for (i, &index) in indices.iter().enumerate() {
642            assert!(index < self.dims[i]);
643            idx += index as isize * self.strides[i];
644        }
645        unsafe {
646            *self.ptr.offset(idx) = value;
647        }
648    }
649}
650
651// ============================================================================
652// StridedArray
653// ============================================================================
654
655/// Owned strided multidimensional array.
656///
657/// Supports both column-major (Julia default) and row-major (C default) layouts.
658pub struct StridedArray<T> {
659    data: Vec<T>,
660    dims: Arc<[usize]>,
661    strides: Arc<[isize]>,
662    offset: isize,
663}
664
665impl<T: std::fmt::Debug> std::fmt::Debug for StridedArray<T> {
666    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
667        f.debug_struct("StridedArray")
668            .field("dims", &self.dims)
669            .field("strides", &self.strides)
670            .field("offset", &self.offset)
671            .finish()
672    }
673}
674
675impl<T: Clone> Clone for StridedArray<T> {
676    fn clone(&self) -> Self {
677        Self {
678            data: self.data.clone(),
679            dims: self.dims.clone(),
680            strides: self.strides.clone(),
681            offset: self.offset,
682        }
683    }
684}
685
686impl<T: Clone + Default> StridedArray<T> {
687    /// Create a column-major (Julia default) tensor filled with Default values.
688    pub fn col_major(dims: &[usize]) -> Self {
689        let total: usize = dims.iter().product();
690        let data = vec![T::default(); total];
691        let strides = col_major_strides(dims);
692        Self {
693            data,
694            dims: Arc::from(dims),
695            strides: Arc::from(strides),
696            offset: 0,
697        }
698    }
699
700    /// Create a row-major (C default) tensor filled with Default values.
701    pub fn row_major(dims: &[usize]) -> Self {
702        let total: usize = dims.iter().product();
703        let data = vec![T::default(); total];
704        let strides = row_major_strides(dims);
705        Self {
706            data,
707            dims: Arc::from(dims),
708            strides: Arc::from(strides),
709            offset: 0,
710        }
711    }
712
713    /// Create a column-major tensor with values produced by a function.
714    ///
715    /// The function is called with indices in column-major iteration order.
716    pub fn from_fn_col_major(dims: &[usize], mut f: impl FnMut(&[usize]) -> T) -> Self {
717        let total: usize = dims.iter().product();
718        let strides = col_major_strides(dims);
719        let rank = dims.len();
720        let mut data = Vec::with_capacity(total);
721        let mut idx = vec![0usize; rank];
722        for _ in 0..total {
723            data.push(f(&idx));
724            for d in 0..rank {
725                idx[d] += 1;
726                if idx[d] < dims[d] {
727                    break;
728                }
729                idx[d] = 0;
730            }
731        }
732        Self {
733            data,
734            dims: Arc::from(dims),
735            strides: Arc::from(strides),
736            offset: 0,
737        }
738    }
739
740    /// Create a row-major tensor with values produced by a function.
741    ///
742    /// The function is called with indices in row-major iteration order.
743    pub fn from_fn_row_major(dims: &[usize], mut f: impl FnMut(&[usize]) -> T) -> Self {
744        let total: usize = dims.iter().product();
745        let strides = row_major_strides(dims);
746        let rank = dims.len();
747        let mut data = Vec::with_capacity(total);
748        let mut idx = vec![0usize; rank];
749        for _ in 0..total {
750            data.push(f(&idx));
751            for d in (0..rank).rev() {
752                idx[d] += 1;
753                if idx[d] < dims[d] {
754                    break;
755                }
756                idx[d] = 0;
757            }
758        }
759        Self {
760            data,
761            dims: Arc::from(dims),
762            strides: Arc::from(strides),
763            offset: 0,
764        }
765    }
766}
767
768impl<T> StridedArray<T> {
769    /// Create from raw parts.
770    pub fn from_parts(
771        data: Vec<T>,
772        dims: &[usize],
773        strides: &[isize],
774        offset: isize,
775    ) -> Result<Self> {
776        validate_bounds(data.len(), dims, strides, offset)?;
777        Ok(Self {
778            data,
779            dims: Arc::from(dims),
780            strides: Arc::from(strides),
781            offset,
782        })
783    }
784
785    /// Returns the shape (dimension sizes) of this array.
786    #[inline]
787    pub fn dims(&self) -> &[usize] {
788        &self.dims
789    }
790
791    /// Returns the strides (in units of `T`) for each dimension.
792    #[inline]
793    pub fn strides(&self) -> &[isize] {
794        &self.strides
795    }
796
797    /// Returns the number of dimensions (rank).
798    #[inline]
799    pub fn ndim(&self) -> usize {
800        self.dims.len()
801    }
802
803    /// Returns the total number of elements.
804    #[inline]
805    pub fn len(&self) -> usize {
806        self.dims.iter().product()
807    }
808
809    /// Returns `true` if any dimension is zero.
810    #[inline]
811    pub fn is_empty(&self) -> bool {
812        self.dims.iter().any(|&d| d == 0)
813    }
814
815    /// Returns a reference to the backing data slice.
816    #[inline]
817    pub fn data(&self) -> &[T] {
818        &self.data
819    }
820
821    /// Returns a mutable reference to the backing data slice.
822    #[inline]
823    pub fn data_mut(&mut self) -> &mut [T] {
824        &mut self.data
825    }
826
827    /// Create an immutable view over this tensor.
828    pub fn view(&self) -> StridedView<'_, T> {
829        let ptr = if self.is_empty() {
830            empty_const_ptr()
831        } else {
832            unsafe { self.data.as_ptr().offset(self.offset) }
833        };
834        StridedView {
835            ptr,
836            data: &self.data,
837            dims: self.dims.clone(),
838            strides: self.strides.clone(),
839            offset: self.offset,
840            _op: PhantomData,
841        }
842    }
843
844    /// Create a mutable view over this tensor.
845    pub fn view_mut(&mut self) -> StridedViewMut<'_, T> {
846        let ptr = if self.is_empty() {
847            empty_mut_ptr()
848        } else {
849            unsafe { self.data.as_mut_ptr().offset(self.offset) }
850        };
851        StridedViewMut {
852            ptr,
853            data: &mut self.data,
854            dims: self.dims.clone(),
855            strides: self.strides.clone(),
856            offset: self.offset,
857        }
858    }
859
860    /// Permute dimensions (metadata-only reorder, no data copy).
861    ///
862    /// Returns a new array with reordered dims and strides.
863    /// The underlying data buffer is not touched.
864    pub fn permuted(self, perm: &[usize]) -> Result<Self> {
865        let rank = self.dims.len();
866        if perm.len() != rank {
867            return Err(StridedError::RankMismatch(perm.len(), rank));
868        }
869        let mut seen = vec![false; rank];
870        for &p in perm {
871            if p >= rank {
872                return Err(StridedError::InvalidAxis { axis: p, rank });
873            }
874            if seen[p] {
875                return Err(StridedError::InvalidAxis { axis: p, rank });
876            }
877            seen[p] = true;
878        }
879        let new_dims: Vec<usize> = perm.iter().map(|&p| self.dims[p]).collect();
880        let new_strides: Vec<isize> = perm.iter().map(|&p| self.strides[p]).collect();
881        Ok(Self {
882            data: self.data,
883            dims: Arc::from(new_dims),
884            strides: Arc::from(new_strides),
885            offset: self.offset,
886        })
887    }
888
889    /// Consume the array and return the backing `Vec<T>`.
890    pub fn into_data(self) -> Vec<T> {
891        self.data
892    }
893
894    /// Iterate over all elements in memory order.
895    pub fn iter(&self) -> std::slice::Iter<'_, T> {
896        self.data.iter()
897    }
898
899    /// Mutable iteration over all elements in memory order.
900    pub fn iter_mut(&mut self) -> std::slice::IterMut<'_, T> {
901        self.data.iter_mut()
902    }
903}
904
905impl<T: Default> StridedArray<T> {
906    /// Create a column-major tensor reusing an existing buffer.
907    ///
908    /// If `buf` has at least `product(dims)` elements, it is truncated and
909    /// zeroed.  Otherwise a fresh buffer is allocated.
910    pub fn col_major_from_buffer(mut buf: Vec<T>, dims: &[usize]) -> Self {
911        let total: usize = dims.iter().product();
912        if buf.len() >= total {
913            buf.truncate(total);
914        } else {
915            buf.resize_with(total, T::default);
916        }
917        // Zero the reused region
918        for v in buf.iter_mut() {
919            *v = T::default();
920        }
921        let strides = col_major_strides(dims);
922        Self {
923            data: buf,
924            dims: Arc::from(dims),
925            strides: Arc::from(strides),
926            offset: 0,
927        }
928    }
929}
930
931impl<T: Copy> StridedArray<T> {
932    /// Create a column-major tensor with **uninitialized** data.
933    ///
934    /// # Safety
935    /// Caller must write every element before reading.
936    pub unsafe fn col_major_uninit(dims: &[usize]) -> Self {
937        let total: usize = dims.iter().product();
938        let mut data = Vec::with_capacity(total);
939        data.set_len(total);
940        let strides = col_major_strides(dims);
941        Self {
942            data,
943            dims: Arc::from(dims),
944            strides: Arc::from(strides),
945            offset: 0,
946        }
947    }
948
949    /// Reuse an existing buffer as a column-major tensor **without zeroing**.
950    ///
951    /// # Safety
952    /// Caller must write every element before reading.
953    pub unsafe fn col_major_from_buffer_uninit(mut buf: Vec<T>, dims: &[usize]) -> Self {
954        let total: usize = dims.iter().product();
955        if buf.capacity() < total {
956            buf.reserve(total - buf.len());
957        }
958        buf.set_len(total);
959        let strides = col_major_strides(dims);
960        Self {
961            data: buf,
962            dims: Arc::from(dims),
963            strides: Arc::from(strides),
964            offset: 0,
965        }
966    }
967}
968
969impl<T: Copy> StridedArray<T> {
970    /// Get an element by multi-dimensional index.
971    pub fn get(&self, indices: &[usize]) -> T {
972        self.view().get(indices)
973    }
974
975    /// Set an element by multi-dimensional index.
976    pub fn set(&mut self, indices: &[usize], value: T) {
977        assert_eq!(indices.len(), self.dims.len());
978        let mut idx = self.offset;
979        for (i, &index) in indices.iter().enumerate() {
980            assert!(index < self.dims[i]);
981            idx += index as isize * self.strides[i];
982        }
983        self.data[idx as usize] = value;
984    }
985}
986
987impl<T: Copy> Index<&[usize]> for StridedArray<T> {
988    type Output = T;
989
990    fn index(&self, indices: &[usize]) -> &T {
991        let mut idx = self.offset;
992        for (i, &index) in indices.iter().enumerate() {
993            assert!(index < self.dims[i]);
994            idx += index as isize * self.strides[i];
995        }
996        &self.data[idx as usize]
997    }
998}
999
1000impl<T: Copy> IndexMut<&[usize]> for StridedArray<T> {
1001    fn index_mut(&mut self, indices: &[usize]) -> &mut T {
1002        let mut idx = self.offset;
1003        for (i, &index) in indices.iter().enumerate() {
1004            assert!(index < self.dims[i]);
1005            idx += index as isize * self.strides[i];
1006        }
1007        &mut self.data[idx as usize]
1008    }
1009}
1010
1011// ============================================================================
1012// Tests
1013// ============================================================================
1014
1015#[cfg(test)]
1016mod tests {
1017    use super::*;
1018
1019    #[test]
1020    fn empty_views_use_dangling_base_pointers_for_extreme_offsets() {
1021        let dims = [0usize];
1022        let strides = [1isize];
1023        let data: [f64; 0] = [];
1024        let view = StridedView::<f64>::new(&data, &dims, &strides, isize::MAX).unwrap();
1025        assert_eq!(view.ptr(), empty_const_ptr());
1026
1027        let mut data: [f64; 0] = [];
1028        let view = StridedViewMut::new(&mut data, &dims, &strides, isize::MAX).unwrap();
1029        assert_eq!(view.ptr(), empty_const_ptr());
1030        assert_eq!(view.as_mut_ptr(), empty_mut_ptr());
1031
1032        let array =
1033            StridedArray::<f64>::from_parts(Vec::new(), &dims, &strides, isize::MAX).unwrap();
1034        assert_eq!(array.view().ptr(), empty_const_ptr());
1035
1036        let mut array =
1037            StridedArray::<f64>::from_parts(Vec::new(), &dims, &strides, isize::MAX).unwrap();
1038        let view = array.view_mut();
1039        assert_eq!(view.ptr(), empty_const_ptr());
1040        assert_eq!(view.as_mut_ptr(), empty_mut_ptr());
1041    }
1042    use num_complex::Complex64;
1043
1044    #[test]
1045    fn test_col_major_strides() {
1046        assert_eq!(col_major_strides(&[3, 4]), vec![1, 3]);
1047        assert_eq!(col_major_strides(&[2, 3, 4]), vec![1, 2, 6]);
1048    }
1049
1050    #[test]
1051    fn test_row_major_strides() {
1052        assert_eq!(row_major_strides(&[3, 4]), vec![4, 1]);
1053        assert_eq!(row_major_strides(&[2, 3, 4]), vec![12, 4, 1]);
1054    }
1055
1056    #[test]
1057    fn test_strided_view_new() {
1058        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
1059        let view = StridedView::<f64>::new(&data, &[2, 3], &[3, 1], 0).unwrap();
1060        assert_eq!(view.ndim(), 2);
1061        assert_eq!(view.dims(), &[2, 3]);
1062        assert_eq!(view.strides(), &[3, 1]);
1063        assert_eq!(view.len(), 6);
1064    }
1065
1066    #[test]
1067    fn test_strided_view_get() {
1068        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
1069        let view = StridedView::<f64>::new(&data, &[2, 3], &[3, 1], 0).unwrap();
1070        assert_eq!(view.get(&[0, 0]), 1.0);
1071        assert_eq!(view.get(&[0, 1]), 2.0);
1072        assert_eq!(view.get(&[0, 2]), 3.0);
1073        assert_eq!(view.get(&[1, 0]), 4.0);
1074        assert_eq!(view.get(&[1, 2]), 6.0);
1075    }
1076
1077    #[test]
1078    fn test_strided_view_col_major() {
1079        // Column-major: strides [1, 2] for 2x3
1080        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
1081        let view = StridedView::<f64>::new(&data, &[2, 3], &[1, 2], 0).unwrap();
1082        assert_eq!(view.get(&[0, 0]), 1.0); // data[0]
1083        assert_eq!(view.get(&[1, 0]), 2.0); // data[1]
1084        assert_eq!(view.get(&[0, 1]), 3.0); // data[2]
1085        assert_eq!(view.get(&[1, 1]), 4.0); // data[3]
1086        assert_eq!(view.get(&[0, 2]), 5.0); // data[4]
1087        assert_eq!(view.get(&[1, 2]), 6.0); // data[5]
1088    }
1089
1090    #[test]
1091    fn test_strided_view_permute() {
1092        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
1093        let view = StridedView::<f64>::new(&data, &[2, 3], &[3, 1], 0).unwrap();
1094        let perm = view.permute(&[1, 0]).unwrap();
1095        assert_eq!(perm.dims(), &[3, 2]);
1096        assert_eq!(perm.strides(), &[1, 3]);
1097        assert_eq!(perm.get(&[0, 0]), 1.0);
1098        assert_eq!(perm.get(&[1, 0]), 2.0);
1099        assert_eq!(perm.get(&[0, 1]), 4.0);
1100    }
1101
1102    #[test]
1103    fn test_strided_view_transpose_2d() {
1104        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
1105        let view = StridedView::<f64>::new(&data, &[2, 3], &[3, 1], 0).unwrap();
1106        let t = view.transpose_2d().unwrap();
1107        assert_eq!(t.dims(), &[3, 2]);
1108        assert_eq!(t.get(&[0, 0]), 1.0);
1109        assert_eq!(t.get(&[1, 0]), 2.0);
1110        assert_eq!(t.get(&[0, 1]), 4.0);
1111    }
1112
1113    #[test]
1114    fn test_strided_view_conj() {
1115        let data = vec![Complex64::new(1.0, 2.0), Complex64::new(3.0, 4.0)];
1116        let view = StridedView::<Complex64>::new(&data, &[2], &[1], 0).unwrap();
1117        let c = view.conj();
1118        assert_eq!(c.get(&[0]), Complex64::new(1.0, -2.0));
1119        assert_eq!(c.get(&[1]), Complex64::new(3.0, -4.0));
1120    }
1121
1122    #[test]
1123    fn test_strided_view_adjoint_2d() {
1124        let data = vec![
1125            Complex64::new(1.0, 2.0),
1126            Complex64::new(3.0, 4.0),
1127            Complex64::new(5.0, 6.0),
1128            Complex64::new(7.0, 8.0),
1129        ];
1130        // 2x2 row-major
1131        let view = StridedView::<Complex64>::new(&data, &[2, 2], &[2, 1], 0).unwrap();
1132        let adj = view.adjoint_2d().unwrap();
1133        assert_eq!(adj.dims(), &[2, 2]);
1134        // Adjoint: conj + transpose
1135        assert_eq!(adj.get(&[0, 0]), Complex64::new(1.0, -2.0));
1136        assert_eq!(adj.get(&[1, 0]), Complex64::new(3.0, -4.0));
1137        assert_eq!(adj.get(&[0, 1]), Complex64::new(5.0, -6.0));
1138    }
1139
1140    #[test]
1141    fn test_strided_view_broadcast() {
1142        let data = vec![1.0, 2.0, 3.0];
1143        let view = StridedView::<f64>::new(&data, &[1, 3], &[3, 1], 0).unwrap();
1144        let broad = view.broadcast(&[4, 3]).unwrap();
1145        assert_eq!(broad.dims(), &[4, 3]);
1146        for i in 0..4 {
1147            assert_eq!(broad.get(&[i, 0]), 1.0);
1148            assert_eq!(broad.get(&[i, 1]), 2.0);
1149            assert_eq!(broad.get(&[i, 2]), 3.0);
1150        }
1151    }
1152
1153    #[test]
1154    fn test_strided_view_mut() {
1155        let mut data = vec![0.0; 6];
1156        {
1157            let mut view = StridedViewMut::<f64>::new(&mut data, &[2, 3], &[3, 1], 0).unwrap();
1158            view.set(&[0, 0], 1.0);
1159            view.set(&[1, 2], 6.0);
1160        }
1161        assert_eq!(data[0], 1.0);
1162        assert_eq!(data[5], 6.0);
1163    }
1164
1165    #[test]
1166    fn test_strided_view_mut_as_view() {
1167        let mut data = vec![1.0, 2.0, 3.0];
1168        let vm = StridedViewMut::<f64>::new(&mut data, &[3], &[1], 0).unwrap();
1169        let v = vm.as_view();
1170        assert_eq!(v.get(&[0]), 1.0);
1171        assert_eq!(v.get(&[2]), 3.0);
1172    }
1173
1174    #[test]
1175    fn test_strided_tensor_col_major() {
1176        let t = StridedArray::<f64>::from_fn_col_major(&[2, 3], |idx| (idx[0] * 3 + idx[1]) as f64);
1177        assert_eq!(t.dims(), &[2, 3]);
1178        assert_eq!(t.strides(), &[1, 2]); // column-major
1179        assert_eq!(t.get(&[0, 0]), 0.0);
1180        assert_eq!(t.get(&[1, 0]), 3.0);
1181        assert_eq!(t.get(&[0, 1]), 1.0);
1182        assert_eq!(t.get(&[1, 2]), 5.0);
1183    }
1184
1185    #[test]
1186    fn test_strided_tensor_row_major() {
1187        let t = StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1]) as f64);
1188        assert_eq!(t.dims(), &[2, 3]);
1189        assert_eq!(t.strides(), &[3, 1]); // row-major
1190        assert_eq!(t.get(&[0, 0]), 0.0);
1191        assert_eq!(t.get(&[0, 1]), 1.0);
1192        assert_eq!(t.get(&[1, 0]), 3.0);
1193        assert_eq!(t.get(&[1, 2]), 5.0);
1194    }
1195
1196    #[test]
1197    fn test_strided_tensor_view() {
1198        let t =
1199            StridedArray::<f64>::from_fn_col_major(&[2, 3], |idx| (idx[0] * 10 + idx[1]) as f64);
1200        let v = t.view();
1201        assert_eq!(v.get(&[0, 0]), 0.0);
1202        assert_eq!(v.get(&[1, 0]), 10.0);
1203        assert_eq!(v.get(&[0, 2]), 2.0);
1204    }
1205
1206    #[test]
1207    fn test_strided_tensor_view_mut() {
1208        let mut t = StridedArray::<f64>::col_major(&[2, 3]);
1209        {
1210            let mut vm = t.view_mut();
1211            vm.set(&[1, 2], 42.0);
1212        }
1213        assert_eq!(t.get(&[1, 2]), 42.0);
1214    }
1215
1216    #[test]
1217    fn test_strided_tensor_index() {
1218        let t =
1219            StridedArray::<f64>::from_fn_row_major(&[3, 4], |idx| (idx[0] * 10 + idx[1]) as f64);
1220        assert_eq!(t[&[0usize, 0] as &[usize]], 0.0);
1221        assert_eq!(t[&[2usize, 3] as &[usize]], 23.0);
1222    }
1223
1224    #[test]
1225    fn test_strided_tensor_index_mut() {
1226        let mut t = StridedArray::<f64>::row_major(&[2, 3]);
1227        t[&[1usize, 2] as &[usize]] = 99.0;
1228        assert_eq!(t.get(&[1, 2]), 99.0);
1229    }
1230
1231    #[test]
1232    fn test_validate_bounds_ok() {
1233        assert!(validate_bounds(6, &[2, 3], &[3, 1], 0).is_ok());
1234        assert!(validate_bounds(6, &[2, 3], &[1, 2], 0).is_ok());
1235    }
1236
1237    #[test]
1238    fn test_validate_bounds_out_of_range() {
1239        assert!(validate_bounds(5, &[2, 3], &[3, 1], 0).is_err());
1240    }
1241
1242    #[test]
1243    fn test_validate_bounds_empty() {
1244        assert!(validate_bounds(0, &[0, 3], &[3, 1], 0).is_ok());
1245    }
1246
1247    #[test]
1248    fn test_validate_bounds_with_offset() {
1249        assert!(validate_bounds(7, &[2, 3], &[3, 1], 1).is_ok());
1250        assert!(validate_bounds(6, &[2, 3], &[3, 1], 1).is_err());
1251    }
1252
1253    #[test]
1254    fn test_strided_tensor_3d() {
1255        let t = StridedArray::<f64>::from_fn_col_major(&[2, 3, 4], |idx| {
1256            (idx[0] * 100 + idx[1] * 10 + idx[2]) as f64
1257        });
1258        assert_eq!(t.ndim(), 3);
1259        assert_eq!(t.strides(), &[1, 2, 6]); // column-major
1260        assert_eq!(t.get(&[0, 0, 0]), 0.0);
1261        assert_eq!(t.get(&[1, 0, 0]), 100.0);
1262        assert_eq!(t.get(&[0, 1, 0]), 10.0);
1263        assert_eq!(t.get(&[0, 0, 1]), 1.0);
1264        assert_eq!(t.get(&[1, 2, 3]), 123.0);
1265    }
1266
1267    #[test]
1268    fn test_diagonal_view_2d() {
1269        // A[i,i] shape=[3,3] row-major strides=[3,1]
1270        // diagonal: shape=[3] strides=[4] (3+1)
1271        let data: Vec<f64> = (0..9).map(|x| x as f64).collect();
1272        let view = StridedView::<f64>::new(&data, &[3, 3], &[3, 1], 0).unwrap();
1273        let diag = view.diagonal_view(&[(0, 1)]).unwrap();
1274        assert_eq!(diag.dims(), &[3]);
1275        assert_eq!(diag.strides(), &[4]);
1276        assert_eq!(diag.get(&[0]), 0.0); // A[0,0]
1277        assert_eq!(diag.get(&[1]), 4.0); // A[1,1]
1278        assert_eq!(diag.get(&[2]), 8.0); // A[2,2]
1279    }
1280
1281    #[test]
1282    fn test_diagonal_view_3d_adjacent() {
1283        // A[i,i,j] shape=[2,2,3] row-major strides=[6,3,1]
1284        // diagonal over (0,1): shape=[2,3] strides=[9,1] (6+3)
1285        let data: Vec<f64> = (0..12).map(|x| x as f64).collect();
1286        let view = StridedView::<f64>::new(&data, &[2, 2, 3], &[6, 3, 1], 0).unwrap();
1287        let diag = view.diagonal_view(&[(0, 1)]).unwrap();
1288        assert_eq!(diag.dims(), &[2, 3]);
1289        assert_eq!(diag.strides(), &[9, 1]);
1290        assert_eq!(diag.get(&[0, 0]), 0.0);
1291        assert_eq!(diag.get(&[0, 2]), 2.0);
1292        assert_eq!(diag.get(&[1, 0]), 9.0);
1293    }
1294
1295    #[test]
1296    fn test_diagonal_view_3d_non_adjacent() {
1297        // A[i,j,i] shape=[2,3,2] row-major strides=[6,2,1]
1298        // diagonal over (0,2): shape=[2,3] strides=[7,2] (6+1)
1299        let data: Vec<f64> = (0..12).map(|x| x as f64).collect();
1300        let view = StridedView::<f64>::new(&data, &[2, 3, 2], &[6, 2, 1], 0).unwrap();
1301        let diag = view.diagonal_view(&[(0, 2)]).unwrap();
1302        assert_eq!(diag.dims(), &[2, 3]);
1303        assert_eq!(diag.strides(), &[7, 2]);
1304        assert_eq!(diag.get(&[0, 0]), 0.0);
1305        assert_eq!(diag.get(&[0, 1]), 2.0);
1306        assert_eq!(diag.get(&[1, 0]), 7.0);
1307        assert_eq!(diag.get(&[1, 2]), 11.0);
1308    }
1309
1310    #[test]
1311    fn test_diagonal_view_two_pairs() {
1312        // A[i,j,i,j] shape=[2,3,2,3] -> A_diag[i,j] shape=[2,3]
1313        let data: Vec<f64> = (0..36).map(|x| x as f64).collect();
1314        let view = StridedView::<f64>::new(&data, &[2, 3, 2, 3], &[18, 6, 3, 1], 0).unwrap();
1315        let diag = view.diagonal_view(&[(0, 2), (1, 3)]).unwrap();
1316        assert_eq!(diag.dims(), &[2, 3]);
1317        assert_eq!(diag.strides(), &[21, 7]);
1318        assert_eq!(diag.get(&[0, 0]), 0.0);
1319        assert_eq!(diag.get(&[1, 1]), 28.0);
1320    }
1321
1322    #[test]
1323    fn test_custom_type_without_element_op_apply() {
1324        // Custom Copy type that does NOT implement ElementOpApply.
1325        // Should work with Identity views (StridedView<T> default).
1326        #[derive(Debug, Clone, Copy, PartialEq)]
1327        struct MyScalar(f64);
1328
1329        impl Default for MyScalar {
1330            fn default() -> Self {
1331                MyScalar(0.0)
1332            }
1333        }
1334
1335        // Create array with custom type
1336        let data = vec![MyScalar(1.0), MyScalar(2.0), MyScalar(3.0), MyScalar(4.0)];
1337        let arr = StridedArray::from_parts(data, &[2, 2], &[2, 1], 0).unwrap();
1338
1339        // Access via Identity view (default)
1340        assert_eq!(arr.get(&[0, 0]), MyScalar(1.0));
1341        assert_eq!(arr.get(&[1, 0]), MyScalar(3.0));
1342
1343        // Create view and access elements
1344        let view: StridedView<MyScalar> = arr.view();
1345        assert_eq!(view.get(&[0, 1]), MyScalar(2.0));
1346
1347        // Permute view (metadata-only, no ElementOpApply needed)
1348        let perm = view.permute(&[1, 0]).unwrap();
1349        assert_eq!(perm.get(&[0, 0]), MyScalar(1.0));
1350        assert_eq!(perm.get(&[0, 1]), MyScalar(3.0)); // transposed
1351    }
1352}