1use crate::prelude_dev::*;
4use core::mem::ManuallyDrop;
5
6impl<T, const N: usize> PackableArrayAPI<T, N> for DataOwned<Vec<T>> {
11 type Array = [T; N];
12 type ArrayVec = DataOwned<Vec<[T; N]>>;
13}
14
15impl<T> PackArrayAPI<T> for DataOwned<Vec<T>> {
16 fn pack_array_f<const N: usize>(self) -> Result<<Self as PackableArrayAPI<T, N>>::ArrayVec> {
17 let raw = self.into_raw();
18 let raw = raw.pack_array_f::<N>()?;
19 Ok(DataOwned::from(raw))
20 }
21}
22
23impl<T, const N: usize> UnpackArrayAPI for DataOwned<Vec<[T; N]>> {
24 type Output = DataOwned<Vec<T>>;
25
26 fn unpack_array(self) -> Self::Output {
27 let raw = self.into_raw();
28 let raw = raw.unpack_array();
29 DataOwned::from(raw)
30 }
31}
32
33impl<'l, T, const N: usize> PackableArrayAPI<T, N> for DataRef<'l, Vec<T>> {
36 type Array = [T; N];
37 type ArrayVec = DataRef<'l, Vec<[T; N]>>;
38}
39
40impl<'l, T> PackArrayAPI<T> for DataRef<'l, Vec<T>> {
41 fn pack_array_f<const N: usize>(self) -> Result<<Self as PackableArrayAPI<T, N>>::ArrayVec> {
42 let raw = self.raw().as_slice().pack_array_f::<N>()?;
43 let vec = unsafe { Vec::from_raw_parts(raw.as_ptr() as *mut [T; N], raw.len(), raw.len()) };
44 Ok(DataRef::from_manually_drop(ManuallyDrop::new(vec)))
45 }
46}
47
48impl<'l, T, const N: usize> UnpackArrayAPI for DataRef<'l, Vec<[T; N]>> {
49 type Output = DataRef<'l, Vec<T>>;
50
51 fn unpack_array(self) -> Self::Output {
52 let raw = self.raw().as_slice().unpack_array();
53 let vec = unsafe { Vec::from_raw_parts(raw.as_ptr() as *mut T, raw.len(), raw.len()) };
54 DataRef::from_manually_drop(ManuallyDrop::new(vec))
55 }
56}
57
58impl<'l, T, const N: usize> PackableArrayAPI<T, N> for DataMut<'l, Vec<T>> {
61 type Array = [T; N];
62 type ArrayVec = DataMut<'l, Vec<[T; N]>>;
63}
64
65impl<'l, T> PackArrayAPI<T> for DataMut<'l, Vec<T>> {
66 fn pack_array_f<const N: usize>(self) -> Result<<Self as PackableArrayAPI<T, N>>::ArrayVec> {
67 let raw = self.raw().as_slice().pack_array_f::<N>()?;
68 let vec = unsafe { Vec::from_raw_parts(raw.as_ptr() as *mut [T; N], raw.len(), raw.len()) };
69 Ok(DataMut::from_manually_drop(ManuallyDrop::new(vec)))
70 }
71}
72
73impl<'l, T, const N: usize> UnpackArrayAPI for DataMut<'l, Vec<[T; N]>> {
74 type Output = DataMut<'l, Vec<T>>;
75
76 fn unpack_array(self) -> Self::Output {
77 let raw = self.raw().as_slice().unpack_array();
78 let vec = unsafe { Vec::from_raw_parts(raw.as_ptr() as *mut T, raw.len(), raw.len()) };
79 DataMut::from_manually_drop(ManuallyDrop::new(vec))
80 }
81}
82
83impl<'l, T, const N: usize> PackableArrayAPI<T, N> for DataCow<'l, Vec<T>> {
86 type Array = [T; N];
87 type ArrayVec = DataCow<'l, Vec<[T; N]>>;
88}
89
90impl<'l, T> PackArrayAPI<T> for DataCow<'l, Vec<T>> {
91 fn pack_array_f<const N: usize>(self) -> Result<<Self as PackableArrayAPI<T, N>>::ArrayVec> {
92 match self {
93 DataCow::Owned(data) => Ok(DataCow::Owned(data.pack_array_f::<N>()?)),
94 DataCow::Ref(data) => Ok(DataCow::Ref(data.pack_array_f::<N>()?)),
95 }
96 }
97}
98
99impl<'l, T, const N: usize> UnpackArrayAPI for DataCow<'l, Vec<[T; N]>> {
100 type Output = DataCow<'l, Vec<T>>;
101
102 fn unpack_array(self) -> Self::Output {
103 match self {
104 DataCow::Owned(data) => DataCow::Owned(data.unpack_array()),
105 DataCow::Ref(data) => DataCow::Ref(data.unpack_array()),
106 }
107 }
108}
109
110impl<R, T, B, D> TensorAny<R, T, B, D>
115where
116 R: DataAPI<Data = B::Raw>,
117 B: DeviceAPI<T>,
118 D: DimAPI + DimSmallerOneAPI,
119 D::SmallerOne: DimAPI,
120{
121 #[substitute_item(
122 ArrayData [<R as PackableArrayAPI<T, N>>::ArrayVec];
123 ArrayType [<R as PackableArrayAPI<T, N>>::Array];
124 )]
125 #[allow(clippy::type_complexity)]
126 pub fn into_pack_array_f<const N: usize>(
127 self,
128 axis: isize,
129 ) -> Result<TensorAny<ArrayData, ArrayType, B, D::SmallerOne>>
130 where
131 B: DeviceAPI<ArrayType>,
132 R: PackableArrayAPI<T, N> + PackArrayAPI<T>,
133 ArrayData: DataAPI<Data = <B as DeviceRawAPI<ArrayType>>::Raw>,
134 {
135 let axis = rstsr_check_axis!(axis, self.ndim())?;
138 rstsr_assert_eq!(self.layout().stride()[axis], 1, InvalidLayout, "The axis must be contiguous")?;
139 rstsr_assert_eq!(self.layout().shape()[axis], N, InvalidLayout, "The axis length must be a exactly {N}")?;
140 rstsr_assert!(self.layout().offset() % N == 0, InvalidLayout, "The offset must be a multiple of {N}")?;
141
142 let (storage, layout) = self.into_raw_parts();
143 let (data, device) = storage.into_raw_parts();
144 let data = data.pack_array_f::<N>()?;
145 let storage = Storage::new(data, device);
146 let layout = layout.dim_chop(axis as isize)?;
147 let stride = layout
148 .stride()
149 .as_ref()
150 .iter()
151 .map(|&s| s / N as isize)
152 .collect_vec()
153 .try_into()
154 .unwrap_or_else(|_| panic!("stride conversion failed"));
155 let new_offset = layout.offset() / N;
156 let new_layout = unsafe { Layout::new_unchecked(layout.shape().clone(), stride, new_offset) };
157 let tensor = unsafe { TensorAny::new_unchecked(storage, new_layout) };
158 Ok(tensor)
159 }
160
161 #[substitute_item(
162 ArrayData [<R as PackableArrayAPI<T, N>>::ArrayVec];
163 ArrayType [<R as PackableArrayAPI<T, N>>::Array];
164 )]
165 #[allow(clippy::type_complexity)]
166 pub fn into_pack_array<const N: usize>(self, axis: isize) -> TensorAny<ArrayData, ArrayType, B, D::SmallerOne>
167 where
168 B: DeviceAPI<ArrayType>,
169 R: PackableArrayAPI<T, N> + PackArrayAPI<T>,
170 ArrayData: DataAPI<Data = <B as DeviceRawAPI<ArrayType>>::Raw>,
171 {
172 self.into_pack_array_f::<N>(axis).rstsr_unwrap()
173 }
174}
175
176impl<R, T, B, D, const N: usize> TensorAny<R, [T; N], B, D>
181where
182 R: DataAPI<Data = <B as DeviceRawAPI<[T; N]>>::Raw>,
183 B: DeviceAPI<T> + DeviceAPI<[T; N]>,
184 D: DimAPI + DimLargerOneAPI,
185 D::LargerOne: DimAPI,
186{
187 #[substitute_item(ROut [<R as UnpackArrayAPI>::Output])]
188 pub fn into_unpack_array_f(self, axis: isize) -> Result<TensorAny<ROut, T, B, D::LargerOne>>
189 where
190 R: UnpackArrayAPI,
191 ROut: DataAPI<Data = <B as DeviceRawAPI<T>>::Raw>,
192 B: DeviceAPI<T>,
193 {
194 let axis = rstsr_check_axis_insert!(axis, self.ndim())?;
196
197 let (storage, layout) = self.into_raw_parts();
198 let (data, device) = storage.into_raw_parts();
199 let data = data.unpack_array();
200 let storage = Storage::new(data, device);
201
202 let mut shape = layout.shape().as_ref().to_vec();
203 let mut stride = layout.stride().as_ref().to_vec();
204 let mut offset = layout.offset();
205
206 shape.insert(axis, N);
207 stride.iter_mut().map(|s| *s *= N as isize).count();
208 stride.insert(axis, 1);
209 offset *= N;
210 let layout = unsafe { Layout::new_unchecked(shape, stride, offset) };
211 let layout = layout.into_dim().rstsr_unwrap();
212 let tensor = unsafe { TensorAny::new_unchecked(storage, layout) };
213 Ok(tensor)
214 }
215
216 #[substitute_item(ROut [<R as UnpackArrayAPI>::Output])]
217 pub fn into_unpack_array(self, axis: isize) -> TensorAny<ROut, T, B, D::LargerOne>
218 where
219 R: UnpackArrayAPI,
220 ROut: DataAPI<Data = <B as DeviceRawAPI<T>>::Raw>,
221 B: DeviceAPI<T>,
222 {
223 self.into_unpack_array_f(axis).rstsr_unwrap()
224 }
225}
226
227#[cfg(test)]
230mod test {
231 use super::*;
232
233 #[test]
234 fn test_pack_array_owned() {
235 let device = DeviceCpuSerial::default();
236 let a = asarray((vec![1, 2, 3, 4, 5, 6], [3, 2].c(), &device));
237 let b = a.into_pack_array_f::<2>(-1).unwrap();
238 println!("{b:?}");
239 assert_eq!(b.raw(), &vec![[1, 2], [3, 4], [5, 6]]);
240
241 let c = b.into_unpack_array(-1);
242 println!("{c:?}");
243 assert_eq!(c.raw(), &vec![1, 2, 3, 4, 5, 6]);
244 }
245
246 #[test]
247 fn test_pack_array_ref() {
248 let device = DeviceCpuSerial::default();
249 let vec = vec![1, 2, 3, 4, 5, 6];
250 let a = asarray((&vec, [3, 2].c(), &device));
251 let b = a.into_pack_array_f::<2>(-1).unwrap();
252 println!("{b:?}");
253 assert_eq!(b.raw(), &vec![[1, 2], [3, 4], [5, 6]]);
254
255 let c = b.into_unpack_array(-1);
256 println!("{c:?}");
257 assert_eq!(c.raw(), &vec![1, 2, 3, 4, 5, 6]);
258 }
259}