1#[cfg(target_has_atomic = "ptr")]
2use alloc::sync::Arc;
3use alloc::vec::Vec;
4use core::fmt;
5#[cfg(not(target_has_atomic = "ptr"))]
6use portable_atomic_util::Arc;
7
8use burn_backend::{DType, Element, TensorData, TensorMetadata};
9use burn_std::{Bytes, Shape, bf16, f16};
10
11use crate::{FlexDevice, layout::Layout};
12
13#[derive(Clone)]
18pub struct FlexTensor {
19 data: Arc<Bytes>,
21 layout: Layout,
23 dtype: DType,
25}
26
27impl fmt::Debug for FlexTensor {
28 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
29 f.debug_struct("FlexTensor")
30 .field("shape", self.layout.shape())
31 .field("dtype", &self.dtype)
32 .field("contiguous", &self.layout.is_contiguous())
33 .field("unique", &self.is_unique())
34 .finish()
35 }
36}
37
38impl FlexTensor {
39 pub fn new(data: Bytes, layout: Layout, dtype: DType) -> Self {
41 Self {
42 data: Arc::new(data),
43 layout,
44 dtype,
45 }
46 }
47
48 pub fn from_data(data: TensorData) -> Self {
50 let shape = data.shape.clone();
51 let layout = Layout::contiguous(shape);
52 let dtype = data.dtype;
53 Self {
54 data: Arc::new(data.bytes),
55 layout,
56 dtype,
57 }
58 }
59
60 pub fn into_data(self) -> TensorData {
64 if self.layout.is_contiguous() && self.layout.start_offset() == 0 {
65 let expected_bytes = self.layout.num_elements() * dtype_size(self.dtype);
66 assert!(
67 expected_bytes <= self.data.len(),
68 "into_data: buffer ({} bytes) too small for {} elements of {:?}",
69 self.data.len(),
70 self.layout.num_elements(),
71 self.dtype
72 );
73 if self.data.len() == expected_bytes {
74 match Arc::try_unwrap(self.data) {
76 Ok(bytes) => TensorData {
77 bytes,
78 shape: self.layout.shape().clone(),
79 dtype: self.dtype,
80 },
81 Err(arc) => {
82 let bytes = Bytes::from_bytes_vec((*arc)[..expected_bytes].to_vec());
83 TensorData {
84 bytes,
85 shape: self.layout.shape().clone(),
86 dtype: self.dtype,
87 }
88 }
89 }
90 } else {
91 let bytes = Bytes::from_bytes_vec(self.data[..expected_bytes].to_vec());
94 TensorData {
95 bytes,
96 shape: self.layout.shape().clone(),
97 dtype: self.dtype,
98 }
99 }
100 } else {
101 self.to_contiguous().into_data()
103 }
104 }
105
106 #[inline]
110 pub fn is_unique(&self) -> bool {
111 Arc::strong_count(&self.data) == 1
112 }
113
114 pub fn layout(&self) -> &Layout {
116 &self.layout
117 }
118
119 pub fn with_layout(self, layout: Layout) -> Self {
123 Self {
124 data: self.data,
125 layout,
126 dtype: self.dtype,
127 }
128 }
129
130 pub fn dtype(&self) -> DType {
132 self.dtype
133 }
134
135 pub fn is_contiguous(&self) -> bool {
137 self.layout.is_contiguous()
138 }
139
140 pub fn bytes(&self) -> &[u8] {
142 &self.data
143 }
144
145 pub fn data_arc(&self) -> Arc<Bytes> {
149 Arc::clone(&self.data)
150 }
151
152 pub fn from_arc(data: Arc<Bytes>, layout: Layout, dtype: DType) -> Self {
156 Self {
157 data,
158 layout,
159 dtype,
160 }
161 }
162
163 pub fn storage<E: Element + bytemuck::Pod>(&self) -> &[E] {
173 assert!(
174 E::dtype() == self.dtype
175 || (matches!(
176 self.dtype,
177 DType::Bool(burn_std::BoolStore::Native | burn_std::BoolStore::U8)
178 ) && E::dtype() == DType::U8),
179 "storage: dtype mismatch (expected {:?}, got {:?})",
180 self.dtype,
181 E::dtype()
182 );
183 bytemuck::cast_slice(&self.data)
184 }
185
186 pub fn storage_mut<E: Element + bytemuck::Pod>(&mut self) -> &mut [E] {
197 assert!(
198 E::dtype() == self.dtype
199 || (matches!(
200 self.dtype,
201 DType::Bool(burn_std::BoolStore::Native | burn_std::BoolStore::U8)
202 ) && E::dtype() == DType::U8),
203 "storage_mut: dtype mismatch (expected {:?}, got {:?})",
204 self.dtype,
205 E::dtype()
206 );
207 let bytes = Arc::make_mut(&mut self.data);
209 bytemuck::cast_slice_mut(bytes)
210 }
211
212 pub fn try_storage_mut<E: Element + bytemuck::Pod>(&mut self) -> Option<&mut [E]> {
219 assert!(
220 E::dtype() == self.dtype
221 || (matches!(
222 self.dtype,
223 DType::Bool(burn_std::BoolStore::Native | burn_std::BoolStore::U8)
224 ) && E::dtype() == DType::U8),
225 "try_storage_mut: dtype mismatch (expected {:?}, got {:?})",
226 self.dtype,
227 E::dtype()
228 );
229 if self.is_unique() {
230 let bytes = Arc::get_mut(&mut self.data)?;
232 Some(bytemuck::cast_slice_mut(bytes))
233 } else {
234 None
235 }
236 }
237
238 pub fn as_slice<E: Element + bytemuck::Pod>(&self) -> Option<&[E]> {
242 if E::dtype() != self.dtype {
243 return None;
244 }
245 let storage: &[E] = self.storage();
246 self.layout
247 .contiguous_offsets()
248 .map(|(start, end)| &storage[start..end])
249 }
250
251 pub fn empty(shape: Shape, dtype: DType) -> Self {
253 let num_elements = shape.num_elements();
254 let elem_size = dtype_size(dtype);
255 let bytes = Bytes::from_bytes_vec(alloc::vec![0u8; num_elements * elem_size]);
256 let layout = Layout::contiguous(shape);
257 Self {
258 data: Arc::new(bytes),
259 layout,
260 dtype,
261 }
262 }
263
264 pub fn zeros(shape: Shape, dtype: DType) -> Self {
266 Self::empty(shape, dtype)
267 }
268
269 pub fn filled_typed<E: bytemuck::Pod + Send + Sync>(
271 shape: Shape,
272 dtype: DType,
273 value: E,
274 ) -> Self {
275 assert_eq!(
276 dtype_size(dtype),
277 core::mem::size_of::<E>(),
278 "filled_typed: dtype size mismatch"
279 );
280 let n = shape.num_elements();
281 let data = alloc::vec![value; n];
282 let bytes = Bytes::from_elems(data);
283 Self {
284 data: Arc::new(bytes),
285 layout: Layout::contiguous(shape),
286 dtype,
287 }
288 }
289
290 pub fn to_contiguous(&self) -> Self {
292 if self.is_contiguous()
298 && self.layout.start_offset() == 0
299 && self.data.len() == self.layout.num_elements() * dtype_size(self.dtype)
300 {
301 return self.clone();
302 }
303
304 match self.dtype {
306 DType::F64 => self.copy_contiguous::<f64>(),
307 DType::F32 => self.copy_contiguous::<f32>(),
308 DType::F16 => self.copy_contiguous::<f16>(),
309 DType::BF16 => self.copy_contiguous::<bf16>(),
310 DType::I64 => self.copy_contiguous::<i64>(),
311 DType::I32 => self.copy_contiguous::<i32>(),
312 DType::I16 => self.copy_contiguous::<i16>(),
313 DType::I8 => self.copy_contiguous::<i8>(),
314 DType::U64 => self.copy_contiguous::<u64>(),
315 DType::U32 => self.copy_contiguous::<u32>(),
316 DType::U16 => self.copy_contiguous::<u16>(),
317 DType::U8 => self.copy_contiguous::<u8>(),
318 DType::Bool(burn_std::BoolStore::Native | burn_std::BoolStore::U8) => {
319 self.copy_contiguous::<u8>()
320 }
321 DType::Bool(burn_std::BoolStore::U32) => {
322 panic!("burn-flex: Bool(U32) storage is not yet supported")
323 }
324 _ => panic!("Unsupported dtype for contiguous copy: {:?}", self.dtype),
325 }
326 }
327
328 fn copy_contiguous<E: Element + bytemuck::Pod>(&self) -> Self {
329 let src: &[E] = bytemuck::cast_slice(&self.data);
330 let n = self.layout.num_elements();
331 let mut dst = Vec::with_capacity(n);
332
333 let collapsed = collapse_for_copy(self.layout.shape(), self.layout.strides());
339 let (shape, strides) = collapsed.as_slices();
340 let offset = self.layout.start_offset() as isize;
341 let all_positive = strides.iter().all(|&s| s >= 0);
342
343 if shape.len() <= 1 && all_positive {
344 let collapsed_numel = if shape.is_empty() { 1 } else { shape[0] };
349 debug_assert_eq!(n, collapsed_numel);
350 unsafe { dst.set_len(n) };
352 if shape.is_empty() {
353 if n > 0 {
354 dst[0] = src[offset as usize];
355 }
356 } else {
357 let len = shape[0];
358 let stride = strides[0];
359 if stride == 1 {
360 dst[..len].copy_from_slice(&src[offset as usize..offset as usize + len]);
361 } else {
362 for (i, slot) in dst.iter_mut().take(len).enumerate() {
363 let idx = (offset + i as isize * stride) as usize;
364 *slot = src[idx];
365 }
366 }
367 }
368 } else if shape.len() == 2 && all_positive {
369 debug_assert_eq!(shape[0] * shape[1], n, "2D strides must cover all elements");
374 unsafe { dst.set_len(n) };
377 copy_2d_tiled(
378 &mut dst, src, offset, shape[0], shape[1], strides[0], strides[1],
379 );
380 } else if all_positive && strides.last().copied() == Some(1) {
381 unsafe { dst.set_len(n) };
384 copy_inner_contiguous_run(&mut dst, src, offset, shape, strides);
385 } else {
386 for idx in crate::strided_index::StridedIter::new(&self.layout) {
389 dst.push(src[idx]);
390 }
391 }
392
393 let bytes = Bytes::from_elems(dst);
394 let layout = Layout::contiguous(self.layout.shape().clone());
395 Self {
396 data: Arc::new(bytes),
397 layout,
398 dtype: self.dtype,
399 }
400 }
401
402 pub fn reshape(&self, new_shape: Shape) -> Self {
404 assert_eq!(
405 self.layout.num_elements(),
406 new_shape.num_elements(),
407 "reshape must preserve total elements"
408 );
409
410 if let Some(new_layout) = self.layout.reshape(new_shape.clone()) {
411 Self {
412 data: Arc::clone(&self.data),
413 layout: new_layout,
414 dtype: self.dtype,
415 }
416 } else {
417 self.to_contiguous().reshape(new_shape)
419 }
420 }
421
422 pub fn transpose(&self, dim1: usize, dim2: usize) -> Self {
424 Self {
425 data: Arc::clone(&self.data),
426 layout: self.layout.transpose(dim1, dim2),
427 dtype: self.dtype,
428 }
429 }
430
431 pub fn narrow(&self, dim: usize, start: usize, len: usize) -> Self {
433 Self {
434 data: Arc::clone(&self.data),
435 layout: self.layout.narrow(dim, start, len),
436 dtype: self.dtype,
437 }
438 }
439
440 pub fn permute(&self, axes: &[usize]) -> Self {
442 Self {
443 data: Arc::clone(&self.data),
444 layout: self.layout.permute(axes),
445 dtype: self.dtype,
446 }
447 }
448}
449
450impl TensorMetadata for FlexTensor {
451 type Device = FlexDevice;
452
453 fn dtype(&self) -> DType {
454 self.dtype
455 }
456
457 fn shape(&self) -> Shape {
458 self.layout.shape().clone()
459 }
460
461 fn rank(&self) -> usize {
462 self.layout.num_dims()
463 }
464
465 fn device(&self) -> Self::Device {
466 FlexDevice
467 }
468
469 fn can_mut(&self) -> bool {
470 self.is_unique()
471 }
472}
473
474const COLLAPSE_MAX_RANK: usize = 8;
477
478#[derive(Debug, Clone, Copy)]
482struct CollapsedLayout {
483 ndim: usize,
484 shape: [usize; COLLAPSE_MAX_RANK],
485 strides: [isize; COLLAPSE_MAX_RANK],
486}
487
488impl CollapsedLayout {
489 #[inline]
490 fn as_slices(&self) -> (&[usize], &[isize]) {
491 (&self.shape[..self.ndim], &self.strides[..self.ndim])
492 }
493}
494
495fn collapse_for_copy(shape: &[usize], strides: &[isize]) -> CollapsedLayout {
519 let mut out = CollapsedLayout {
520 ndim: 0,
521 shape: [0; COLLAPSE_MAX_RANK],
522 strides: [0; COLLAPSE_MAX_RANK],
523 };
524
525 if shape.len() > COLLAPSE_MAX_RANK {
530 out.ndim = shape.len().min(COLLAPSE_MAX_RANK);
531 return out;
532 }
533
534 for (&s, &st) in shape.iter().zip(strides.iter()) {
544 if s == 1 {
545 continue;
546 }
547 let merge = out.ndim > 0
548 && (s as isize)
549 .checked_mul(st)
550 .is_some_and(|run| out.strides[out.ndim - 1] == run);
551 if merge {
552 out.shape[out.ndim - 1] *= s;
553 out.strides[out.ndim - 1] = st;
554 } else {
555 out.shape[out.ndim] = s;
556 out.strides[out.ndim] = st;
557 out.ndim += 1;
558 }
559 }
560
561 out
562}
563
564#[inline]
569fn copy_inner_contiguous_run<E: Copy>(
570 dst: &mut [E],
571 src: &[E],
572 offset: isize,
573 shape: &[usize],
574 strides: &[isize],
575) {
576 debug_assert!(!shape.is_empty());
577 debug_assert_eq!(*strides.last().unwrap(), 1, "innermost stride must be 1");
578
579 let mut expected_stride: isize = 1;
583 let mut split = shape.len(); while split > 0 {
585 let dim_idx = split - 1;
586 if strides[dim_idx] == expected_stride {
587 expected_stride = expected_stride
588 .checked_mul(shape[dim_idx] as isize)
589 .expect("contiguous run length overflows isize");
590 split -= 1;
591 } else {
592 break;
593 }
594 }
595
596 let outer_shape = &shape[..split];
597 let outer_strides = &strides[..split];
598 let inner_run_length: usize = shape[split..].iter().product();
599
600 if outer_shape.is_empty() {
602 let src_index = offset as usize;
603 dst[..inner_run_length].copy_from_slice(&src[src_index..src_index + inner_run_length]);
604 return;
605 }
606
607 if inner_run_length == 0 || outer_shape.contains(&0) {
608 return;
609 }
610
611 let mut counter = [0usize; COLLAPSE_MAX_RANK];
620 let outer_dim_count = outer_shape.len();
621 let mut dst_position = 0usize;
622
623 loop {
624 let mut src_start = offset;
625 for dim_idx in 0..outer_dim_count {
626 src_start += counter[dim_idx] as isize * outer_strides[dim_idx];
627 }
628 let src_index = src_start as usize;
629 dst[dst_position..dst_position + inner_run_length]
630 .copy_from_slice(&src[src_index..src_index + inner_run_length]);
631 dst_position += inner_run_length;
632
633 let mut dim_idx = outer_dim_count;
634 loop {
635 if dim_idx == 0 {
636 return;
637 }
638 dim_idx -= 1;
639 counter[dim_idx] += 1;
640 if counter[dim_idx] < outer_shape[dim_idx] {
641 break;
642 }
643 counter[dim_idx] = 0;
644 }
645 }
646}
647
648#[inline]
653fn copy_2d_tiled<E: Copy>(
654 dst: &mut [E],
655 src: &[E],
656 offset: isize,
657 rows: usize,
658 cols: usize,
659 row_stride: isize,
660 col_stride: isize,
661) {
662 const TILE: usize = 16;
663
664 if row_stride <= col_stride {
665 for col_tile in (0..cols).step_by(TILE) {
667 let col_end = (col_tile + TILE).min(cols);
668 for row_tile in (0..rows).step_by(TILE) {
669 let row_end = (row_tile + TILE).min(rows);
670 for col in col_tile..col_end {
671 let col_base = offset + col as isize * col_stride;
672 for row in row_tile..row_end {
673 let idx = (col_base + row as isize * row_stride) as usize;
674 unsafe {
677 *dst.get_unchecked_mut(row * cols + col) = src[idx];
678 }
679 }
680 }
681 }
682 }
683 } else {
684 for row_tile in (0..rows).step_by(TILE) {
686 let row_end = (row_tile + TILE).min(rows);
687 for col_tile in (0..cols).step_by(TILE) {
688 let col_end = (col_tile + TILE).min(cols);
689 for row in row_tile..row_end {
690 let row_base =
691 offset + row as isize * row_stride + col_tile as isize * col_stride;
692 let dst_base = row * cols + col_tile;
693 for c in 0..(col_end - col_tile) {
694 let idx = (row_base + c as isize * col_stride) as usize;
695 unsafe {
697 *dst.get_unchecked_mut(dst_base + c) = src[idx];
698 }
699 }
700 }
701 }
702 }
703 }
704}
705
706pub(crate) fn dtype_size(dtype: DType) -> usize {
722 let size = dtype.size();
724 assert!(
725 size > 0,
726 "burn-flex: dtype {:?} has zero-byte element size (sub-byte packed \
727 quantization is not yet supported)",
728 dtype
729 );
730 size
731}
732
733#[cfg(test)]
734mod tests {
735 use super::*;
736 use alloc::vec;
737
738 #[test]
739 fn test_from_data_roundtrip() {
740 let data = TensorData::from([1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0]);
741 let tensor = FlexTensor::from_data(data.clone());
742 let result = tensor.into_data();
743 assert_eq!(data.shape, result.shape);
744 assert_eq!(data.dtype, result.dtype);
745 }
746
747 #[test]
748 fn test_collapse_for_copy_squeezes_size1_and_merges_contig() {
749 let shape = vec![1, 244, 224, 48];
751 let strides = vec![2_623_488_isize, 224, 1, 54656];
752 let collapsed = collapse_for_copy(&shape, &strides);
753 let (s, st) = collapsed.as_slices();
754 assert_eq!(s, &[54656, 48]);
755 assert_eq!(st, &[1, 54656]);
756 }
757
758 #[test]
759 fn test_collapse_for_copy_already_contiguous_3d() {
760 let collapsed = collapse_for_copy(&[2, 3, 4], &[12, 4, 1]);
761 let (s, st) = collapsed.as_slices();
762 assert_eq!(s, &[24]);
763 assert_eq!(st, &[1]);
764 }
765
766 #[test]
767 fn test_collapse_for_copy_transpose_2d() {
768 let collapsed = collapse_for_copy(&[5, 3], &[1, 5]);
769 let (s, st) = collapsed.as_slices();
770 assert_eq!(s, &[5, 3]);
771 assert_eq!(st, &[1, 5]);
772 }
773
774 #[test]
775 fn test_collapse_for_copy_all_size1() {
776 let collapsed = collapse_for_copy(&[1, 1, 1], &[0, 0, 0]);
777 let (s, st) = collapsed.as_slices();
778 assert!(s.is_empty());
779 assert!(st.is_empty());
780 }
781
782 #[test]
789 fn test_to_contiguous_zero_sized_narrowed() {
790 let t = FlexTensor::from_data(TensorData::new(
791 (0..6).map(|i| i as f32).collect::<Vec<_>>(),
792 vec![6],
793 ));
794 let empty_view = t.narrow(0, 3, 0);
796 assert_eq!(empty_view.shape().to_vec(), vec![0]);
797 assert_ne!(empty_view.layout().start_offset(), 0);
798
799 let contig = empty_view.to_contiguous();
800 assert_eq!(contig.shape().to_vec(), vec![0]);
801 assert_eq!(contig.layout().start_offset(), 0);
802 assert_eq!(contig.into_data().bytes.len(), 0);
803 }
804
805 #[test]
812 fn test_to_contiguous_prefix_view_shrinks_buffer() {
813 let data: Vec<f32> = (0..40).map(|i| i as f32).collect();
814 let t = FlexTensor::from_data(TensorData::new(data, vec![8, 5]));
815
816 let prefix = t.narrow(0, 0, 5);
817 assert_eq!(prefix.shape().to_vec(), vec![5, 5]);
818 assert_eq!(prefix.layout().strides(), &[5, 1]);
819 assert_eq!(prefix.layout().start_offset(), 0);
820 assert!(prefix.is_contiguous());
821 assert_eq!(prefix.storage::<f32>().len(), 40);
822
823 let contig = prefix.to_contiguous();
824 assert_eq!(contig.storage::<f32>().len(), 25);
825 assert_eq!(contig.layout().num_elements(), 25);
826 assert_eq!(
827 contig.storage::<f32>(),
828 &(0..5)
829 .flat_map(|r| (0..5).map(move |c| (r * 5 + c) as f32))
830 .collect::<Vec<_>>()[..]
831 );
832 }
833
834 #[test]
837 fn test_to_contiguous_4d_permuted_matches_naive() {
838 let dims = [1, 48, 4, 5];
839 let n: usize = dims.iter().product();
840 let data: Vec<f32> = (0..n).map(|i| i as f32).collect();
841 let t = FlexTensor::from_data(TensorData::new(data.clone(), dims.to_vec()));
842 let permuted = t.permute(&[0, 2, 3, 1]);
843 assert!(!permuted.is_contiguous());
844
845 let contig = permuted.to_contiguous();
846 assert!(contig.is_contiguous());
847 assert_eq!(contig.shape().to_vec(), vec![1, 4, 5, 48]);
848
849 let mut expected = Vec::with_capacity(n);
851 for h in 0..4 {
852 for w in 0..5 {
853 for c in 0..48 {
854 let idx = c * 20 + h * 5 + w;
855 expected.push(data[idx]);
856 }
857 }
858 }
859
860 let result_data = contig.into_data();
861 let values = result_data.as_slice::<f32>().unwrap();
862 assert_eq!(values, expected.as_slice());
863 }
864
865 #[test]
870 fn test_to_contiguous_3d_inner_stride1_matches_naive() {
871 let dims = [1, 2, 4, 3, 3]; let n: usize = dims.iter().product();
873 let data: Vec<f32> = (0..n).map(|i| i as f32).collect();
874 let t = FlexTensor::from_data(TensorData::new(data.clone(), dims.to_vec()));
875 let permuted = t.permute(&[0, 2, 1, 3, 4]); assert!(!permuted.is_contiguous());
877
878 let contiguous_data = permuted.to_contiguous();
879 assert!(contiguous_data.is_contiguous());
880 assert_eq!(contiguous_data.shape().to_vec(), vec![1, 4, 2, 3, 3]);
881
882 let mut expected = Vec::with_capacity(n);
885 for c in 0..4 {
886 for g in 0..2 {
887 for h in 0..3 {
888 for w in 0..3 {
889 let idx = g * 36 + c * 9 + h * 3 + w;
890 expected.push(data[idx]);
891 }
892 }
893 }
894 }
895
896 let result_data = contiguous_data.into_data();
898 let values = result_data.as_slice::<f32>().unwrap();
899 assert_eq!(values, expected.as_slice());
900 }
901
902 #[test]
905 fn test_to_contiguous_2d_row_stride_gt_col_stride() {
906 let data: Vec<f32> = (0..18).map(|i| i as f32).collect();
910 let t = FlexTensor::from_data(TensorData::new(data, vec![6, 3]));
911 let stepped = crate::ops::slice::slice(
912 t,
913 &[
914 burn_std::Slice::new(0, Some(6), 2),
915 burn_std::Slice::new(0, None, 1),
916 ],
917 );
918 assert_eq!(stepped.layout().shape().to_vec(), vec![3, 3]);
920 assert_eq!(stepped.layout().strides(), &[6, 1]);
921 assert!(!stepped.layout().is_contiguous());
922
923 let contig = stepped.to_contiguous();
924 assert!(contig.is_contiguous());
925 assert_eq!(contig.shape().to_vec(), vec![3, 3]);
926
927 let result_data = contig.into_data();
928 let values = result_data.as_slice::<f32>().unwrap();
929 let expected = vec![
931 0.0f32, 1.0, 2.0, 6.0, 7.0, 8.0, 12.0, 13.0, 14.0, ];
935 assert_eq!(values, expected.as_slice());
936 }
937
938 #[test]
939 fn test_reshape() {
940 let data = TensorData::new(vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]);
941 let tensor = FlexTensor::from_data(data);
942 let reshaped = tensor.reshape(Shape::from(vec![3, 2]));
943 assert_eq!(reshaped.shape().to_vec(), vec![3, 2]);
944 }
945
946 #[test]
947 fn test_transpose() {
948 let data = TensorData::new(vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]);
949 let tensor = FlexTensor::from_data(data);
950 let transposed = tensor.transpose(0, 1);
951 assert_eq!(transposed.shape().to_vec(), vec![3, 2]);
952 assert!(!transposed.is_contiguous());
953 }
954
955 #[test]
956 fn test_clone_is_cheap() {
957 let data = TensorData::from([1.0f32, 2.0, 3.0, 4.0]);
958 let tensor = FlexTensor::from_data(data);
959
960 assert!(tensor.is_unique());
962
963 let cloned = tensor.clone();
965 assert!(!tensor.is_unique());
966 assert!(!cloned.is_unique());
967
968 assert!(core::ptr::eq(
970 tensor.bytes().as_ptr(),
971 cloned.bytes().as_ptr(),
972 ));
973 }
974
975 #[test]
976 fn test_cow_on_mutation() {
977 let data = TensorData::from([1.0f32, 2.0, 3.0, 4.0]);
978 let tensor = FlexTensor::from_data(data);
979 let mut cloned = tensor.clone();
980
981 assert!(!tensor.is_unique());
983 assert!(!cloned.is_unique());
984
985 let storage: &mut [f32] = cloned.storage_mut();
987 storage[0] = 99.0;
988
989 assert!(tensor.is_unique());
991 assert!(cloned.is_unique());
992
993 assert_ne!(tensor.bytes().as_ptr(), cloned.bytes().as_ptr());
995 assert_eq!(tensor.storage::<f32>()[0], 1.0);
996 assert_eq!(cloned.storage::<f32>()[0], 99.0);
997 }
998
999 #[test]
1000 fn test_into_data_narrowed_at_offset_zero() {
1001 let data = TensorData::new(vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]);
1003 let tensor = FlexTensor::from_data(data);
1004 let narrowed = tensor.narrow(0, 0, 1);
1006 assert!(narrowed.is_contiguous());
1007 assert_eq!(narrowed.layout().start_offset(), 0);
1008
1009 let result = narrowed.into_data();
1010 assert_eq!(result.shape.to_vec(), vec![1, 3]);
1011 assert_eq!(result.bytes.len(), 3 * core::mem::size_of::<f32>());
1013 let values: Vec<f32> = result.to_vec().unwrap();
1014 assert_eq!(values, vec![1.0, 2.0, 3.0]);
1015 }
1016}