1use 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
32pub(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 if dims.iter().any(|&d| d == 0) {
48 return Ok(());
49 }
50 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
82pub 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
95pub 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
108pub 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 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 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 #[inline]
204 pub fn dims(&self) -> &[usize] {
205 &self.dims
206 }
207
208 #[inline]
210 pub fn strides(&self) -> &[isize] {
211 &self.strides
212 }
213
214 #[inline]
216 pub fn offset(&self) -> isize {
217 self.offset
218 }
219
220 #[inline]
222 pub fn ndim(&self) -> usize {
223 self.dims.len()
224 }
225
226 #[inline]
228 pub fn len(&self) -> usize {
229 self.dims.iter().product()
230 }
231
232 #[inline]
234 pub fn is_empty(&self) -> bool {
235 self.dims.iter().any(|&d| d == 0)
236 }
237
238 #[inline]
240 pub fn data(&self) -> &'a [T] {
241 self.data
242 }
243
244 #[inline]
246 pub fn ptr(&self) -> *const T {
247 self.ptr
248 }
249
250 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 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 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
364impl<'a, T: Copy + ElementOpApply, Op: ComposableElementOp<T>> StridedView<'a, T, Op> {
366 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 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 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
415impl<'a, T: Copy, Op: ElementOp<T>> StridedView<'a, T, Op> {
417 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 #[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
447pub 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 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 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 #[inline]
524 pub fn dims(&self) -> &[usize] {
525 &self.dims
526 }
527
528 #[inline]
530 pub fn strides(&self) -> &[isize] {
531 &self.strides
532 }
533
534 #[inline]
536 pub fn offset(&self) -> isize {
537 self.offset
538 }
539
540 #[inline]
542 pub fn ndim(&self) -> usize {
543 self.dims.len()
544 }
545
546 #[inline]
548 pub fn len(&self) -> usize {
549 self.dims.iter().product()
550 }
551
552 #[inline]
554 pub fn is_empty(&self) -> bool {
555 self.dims.iter().any(|&d| d == 0)
556 }
557
558 #[inline]
560 pub fn ptr(&self) -> *const T {
561 self.ptr as *const T
562 }
563
564 #[inline]
566 pub fn as_mut_ptr(&self) -> *mut T {
567 self.ptr
568 }
569
570 #[inline]
572 pub fn data(&self) -> &[T] {
573 &self.data
574 }
575
576 #[inline]
578 pub fn data_mut(&mut self) -> &mut [T] {
579 &mut self.data
580 }
581
582 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 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 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 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
651pub 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 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 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 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 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 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 #[inline]
787 pub fn dims(&self) -> &[usize] {
788 &self.dims
789 }
790
791 #[inline]
793 pub fn strides(&self) -> &[isize] {
794 &self.strides
795 }
796
797 #[inline]
799 pub fn ndim(&self) -> usize {
800 self.dims.len()
801 }
802
803 #[inline]
805 pub fn len(&self) -> usize {
806 self.dims.iter().product()
807 }
808
809 #[inline]
811 pub fn is_empty(&self) -> bool {
812 self.dims.iter().any(|&d| d == 0)
813 }
814
815 #[inline]
817 pub fn data(&self) -> &[T] {
818 &self.data
819 }
820
821 #[inline]
823 pub fn data_mut(&mut self) -> &mut [T] {
824 &mut self.data
825 }
826
827 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 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 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 pub fn into_data(self) -> Vec<T> {
891 self.data
892 }
893
894 pub fn iter(&self) -> std::slice::Iter<'_, T> {
896 self.data.iter()
897 }
898
899 pub fn iter_mut(&mut self) -> std::slice::IterMut<'_, T> {
901 self.data.iter_mut()
902 }
903}
904
905impl<T: Default> StridedArray<T> {
906 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 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 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 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 pub fn get(&self, indices: &[usize]) -> T {
972 self.view().get(indices)
973 }
974
975 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#[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 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); assert_eq!(view.get(&[1, 0]), 2.0); assert_eq!(view.get(&[0, 1]), 3.0); assert_eq!(view.get(&[1, 1]), 4.0); assert_eq!(view.get(&[0, 2]), 5.0); assert_eq!(view.get(&[1, 2]), 6.0); }
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 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 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]); 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]); 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]); 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 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); assert_eq!(diag.get(&[1]), 4.0); assert_eq!(diag.get(&[2]), 8.0); }
1280
1281 #[test]
1282 fn test_diagonal_view_3d_adjacent() {
1283 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 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 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 #[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 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 assert_eq!(arr.get(&[0, 0]), MyScalar(1.0));
1341 assert_eq!(arr.get(&[1, 0]), MyScalar(3.0));
1342
1343 let view: StridedView<MyScalar> = arr.view();
1345 assert_eq!(view.get(&[0, 1]), MyScalar(2.0));
1346
1347 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)); }
1352}