1use std::sync::{Arc, RwLock};
2
3use crate::{Bool, DTypeKind, Device, Dim, Dims, Error, Float, Int, Layout, Shape, Storage, Tensor, TensorMeta, ViewOp};
4
5pub trait ShapeDTypeKind<D: Device>: DTypeKind<D> {
6 fn contiguous_dispatch(s: &Self::Storage, l: &Layout) -> crate::Result<Self::Storage>;
7 fn cat_dispatch(srcs: &[(&Self::Storage, &Layout)], dim: usize) -> crate::Result<(Self::Storage, Shape)>;
8 fn view_dispatch(src: &Self::Storage, src_l: &Layout, dst_l: &Layout, view: ViewOp) -> crate::Result<Option<Self::Storage>>;
9}
10
11impl<D: Device> ShapeDTypeKind<D> for Float {
12 #[inline]
13 fn contiguous_dispatch(s: &Self::Storage, l: &Layout) -> crate::Result<Self::Storage> {
14 D::f_contiguous(s, l)
15 }
16
17 #[inline]
18 fn cat_dispatch(srcs: &[(&Self::Storage, &Layout)], dim: usize) -> crate::Result<(Self::Storage, Shape)> {
19 D::f_cat(srcs, dim)
20 }
21
22 fn view_dispatch(src: &Self::Storage, src_l: &Layout, dst_l: &Layout, view: ViewOp) -> crate::Result<Option<Self::Storage>> {
23 D::f_view(src, src_l, dst_l, view)
24 }
25}
26
27impl<D: Device> ShapeDTypeKind<D> for Int {
28 #[inline]
29 fn contiguous_dispatch(s: &Self::Storage, l: &Layout) -> crate::Result<Self::Storage> {
30 D::i_contiguous(s, l)
31 }
32
33 #[inline]
34 fn cat_dispatch(srcs: &[(&Self::Storage, &Layout)], dim: usize) -> crate::Result<(Self::Storage, Shape)> {
35 D::i_cat(srcs, dim)
36 }
37
38 fn view_dispatch(src: &Self::Storage, src_l: &Layout, dst_l: &Layout, view: ViewOp) -> crate::Result<Option<Self::Storage>> {
39 D::i_view(src, src_l, dst_l, view)
40 }
41}
42
43impl<D: Device> ShapeDTypeKind<D> for Bool {
44 #[inline]
45 fn contiguous_dispatch(s: &Self::Storage, l: &Layout) -> crate::Result<Self::Storage> {
46 D::b_contiguous(s, l)
47 }
48
49 #[inline]
50 fn cat_dispatch(srcs: &[(&Self::Storage, &Layout)], dim: usize) -> crate::Result<(Self::Storage, Shape)> {
51 D::b_cat(srcs, dim)
52 }
53
54 fn view_dispatch(src: &Self::Storage, src_l: &Layout, dst_l: &Layout, view: ViewOp) -> crate::Result<Option<Self::Storage>> {
55 D::b_view(src, src_l, dst_l, view)
56 }
57}
58
59impl<D: Device, K: ShapeDTypeKind<D>> Tensor<D, K> {
60 fn resolve_view_storage(&self, view: ViewOp, layout: &Layout) -> crate::Result<Option<Arc<RwLock<K::Storage>>>> {
65 let Some(lock) = self.0.storage.as_ref() else {
66 return Ok(None);
67 };
68 let new_storage = {
69 let guard = lock.read().expect("storage read lock");
70 K::view_dispatch(&*guard, self.layout(), layout, view)?
71 };
72 match new_storage {
73 Some(s) => Ok(Some(Arc::new(RwLock::new(s)))),
74 None => Ok(self.0.storage.clone()),
75 }
76 }
77
78 pub fn copy(&self) -> crate::Result<Self> {
80 let storage = K::contiguous_dispatch(&*self.storage_read()?, self.layout())?;
81 assert_eq!(self.dtype(), storage.dtype());
82 Ok(Self::from_storage(storage, self.shape().clone(), K::Meta::on_copy(self)))
83 }
84
85 pub fn copy_(&mut self, src: &Self) -> crate::Result<()> {
87 if self.shape() != src.shape() {
88 return Err(Error::ShapeMismatchBinaryOp { lhs: self.shape().clone(), rhs: src.shape().clone(), op: "copy_" });
89 }
90
91 let src_storage = K::contiguous_dispatch(&*src.storage_read()?, src.layout())?;
92
93 let this = std::sync::Arc::get_mut(&mut self.0).expect("copy_: tensor is shared (cloned elsewhere); cannot update layout");
96
97 match &this.storage {
98 Some(lock) => {
99 *lock.write().expect("storage write lock") = src_storage;
101 }
102 None => {
103 this.storage = Some(std::sync::Arc::new(std::sync::RwLock::new(src_storage)));
105 }
106 }
107
108 this.layout = Layout::contiguous(src.shape().clone());
109 Ok(())
110 }
111
112 pub fn reshape<S: Into<Shape>>(&self, shape: S) -> crate::Result<Self> {
113 let shape = shape.into();
114 if shape.element_count() != self.element_count() {
115 return Err(Error::ElementCountMismatchInReshape { origin: self.shape().clone(), target: shape });
116 }
117 let meta = K::Meta::on_reshape(self);
118 if self.is_contiguous() {
119 let layout = Layout::contiguous_with_offset(shape, self.layout().start_offset());
120 let storage = self.resolve_view_storage(ViewOp::Reshape, &layout)?;
121 Ok(self.share_storage(layout, meta, storage))
122 } else {
123 let storage = K::contiguous_dispatch(&*self.storage_read()?, self.layout())?;
124 Ok(Self::from_storage(storage, shape, meta))
125 }
126 }
127
128 pub fn transpose<D1: Dim, D2: Dim>(&self, dim1: D1, dim2: D2) -> crate::Result<Self> {
129 let dim1 = dim1.to_index(self.shape(), "transpose")?;
130 let dim2 = dim2.to_index(self.shape(), "transpose")?;
131 if dim1 == dim2 {
132 return Ok(self.clone());
133 }
134 let layout = self.layout().transpose(dim1, dim2)?;
135 let storage = self.resolve_view_storage(ViewOp::Transpose(dim1, dim2), &layout)?;
136 Ok(self.share_storage(layout, K::Meta::on_transpose(self, dim1, dim2), storage))
137 }
138
139 pub fn transpose_last(&self) -> crate::Result<Self> {
140 self.transpose(crate::D::Minus1, crate::D::Minus2)
141 }
142
143 pub fn permute<Ds: Dims>(&self, dims: Ds) -> crate::Result<Self> {
144 let dims = dims.to_indexes(self.shape(), "permute")?;
145 let layout = self.layout().permute(&dims)?;
146 let storage = self.resolve_view_storage(ViewOp::Permute(dims.clone()), &layout)?;
147 Ok(self.share_storage(layout, K::Meta::on_permute(self, dims), storage))
148 }
149
150 pub fn narrow<Dm: Dim>(&self, dim: Dm, start: usize, len: usize) -> crate::Result<Self> {
151 let dim = dim.to_index(self.shape(), "narrow")?;
152 let dims = self.dims();
153 if start.saturating_add(len) > dims[dim] {
154 return Err(Error::NarrowInvalidArgs { shape: self.shape().clone(), dim, start, len, msg: "start + len > dim_len" });
155 }
156 if start == 0 && dims[dim] == len {
157 return Ok(self.clone());
158 }
159 let layout = self.layout().narrow(dim, start, len)?;
160 let storage = self.resolve_view_storage(ViewOp::Narrow(dim, start, len), &layout)?;
161 Ok(self.share_storage(layout, K::Meta::on_narrow(self, dim, start, len), storage))
162 }
163
164 pub fn slice<Dm: Dim>(&self, dim: Dm, start: usize, end: usize, step: usize) -> crate::Result<Self> {
170 let dim = dim.to_index(self.shape(), "slice")?;
171 let layout = self.layout().slice(dim, start, end, step)?;
172 let meta = K::Meta::on_slice(self, dim, start, end, step);
173 let storage = self.resolve_view_storage(ViewOp::Slice(dim, start, end, step), &layout)?;
174 Ok(self.share_storage(layout, meta, storage))
175 }
176
177 pub fn broadcast_as<S: Into<Shape>>(&self, shape: S) -> crate::Result<Self> {
178 let layout = self.layout().broadcast_as(shape)?;
179 let storage = self.resolve_view_storage(ViewOp::Broadcast, &layout)?;
180 Ok(self.share_storage(layout, K::Meta::on_broadcast(self), storage))
181 }
182
183 pub fn squeeze<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
184 let dim = dim.to_index(self.shape(), "squeeze")?;
185 let dims = self.dims();
186 if dims[dim] != 1 {
187 return Err(Error::SqueezeDimNot1 { shape: self.shape().clone(), dim });
188 }
189 let mut new_dims = dims.to_vec();
190 let mut strides = self.stride().to_vec();
191 new_dims.remove(dim);
192 strides.remove(dim);
193 let layout = Layout::new(new_dims, strides, self.layout().start_offset());
194 let storage = self.resolve_view_storage(ViewOp::Squeeze, &layout)?;
195 Ok(self.share_storage(layout, K::Meta::on_reshape(self), storage))
196 }
197
198 pub fn unsqueeze<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
199 let dim = dim.to_index_plus_one(self.shape(), "unsqueeze")?;
200 let mut new_dims = self.dims().to_vec();
201 let mut strides = self.stride().to_vec();
202 new_dims.insert(dim, 1);
203 let stride = if dim < strides.len() { strides[dim] } else { 1 };
204 strides.insert(dim, stride);
205 let layout = Layout::new(new_dims, strides, self.layout().start_offset());
206 let storage = self.resolve_view_storage(ViewOp::Unsqueeze, &layout)?;
207 Ok(self.share_storage(layout, K::Meta::on_reshape(self), storage))
208 }
209
210 pub fn contiguous(&self) -> crate::Result<Self> {
212 if self.is_contiguous() {
213 return Ok(self.clone());
214 }
215 let storage = K::contiguous_dispatch(&*self.storage_read()?, self.layout())?;
216 assert_eq!(self.dtype(), storage.dtype());
217 Ok(Self::from_storage(storage, self.shape().clone(), K::Meta::on_copy(self)))
218 }
219
220 pub fn cat<A: AsRef<Self>, Dm: Dim>(arrs: &[A], dim: Dm) -> crate::Result<Self> {
222 if arrs.is_empty() {
223 return Err(Error::OpRequiresAtLeastOneTensor { op: "cat" });
224 }
225 let first = arrs[0].as_ref();
226 let dim = dim.to_index(first.shape(), "cat")?;
227
228 let guards: Vec<_> = arrs.iter().map(|a| a.as_ref().storage_read()).collect::<crate::Result<_>>()?;
229 let views: Vec<(&K::Storage, &Layout)> = guards.iter().zip(arrs.iter()).map(|(g, a)| (&**g, a.as_ref().layout())).collect();
230 let (storage, shape) = K::cat_dispatch(&views, dim)?;
231 drop(guards);
232
233 let meta = K::Meta::on_cat(arrs, dim);
234 Ok(Self::from_storage(storage, shape, meta))
235 }
236
237 pub fn flatten<D1: Dim, D2: Dim>(&self, start_dim: D1, end_dim: D2) -> crate::Result<Self> {
239 let start = start_dim.to_index(self.shape(), "flatten")?;
240 let end = end_dim.to_index(self.shape(), "flatten")?;
241 if start > end {
242 return Err(Error::NarrowInvalidArgs {
243 shape: self.shape().clone(),
244 dim: start,
245 start,
246 len: end,
247 msg: "flatten: start_dim > end_dim",
248 });
249 }
250 let merged: usize = self.dims()[start..=end].iter().product();
251 let mut new_dims = self.dims().to_vec();
252 new_dims.splice(start..=end, [merged]);
253 self.reshape(Shape::from(new_dims))
254 }
255
256 pub fn flatten_from<Dm: Dim>(&self, start_dim: Dm) -> crate::Result<Self> {
257 self.flatten(start_dim, self.rank() - 1)
258 }
259
260 pub fn flatten_to<Dm: Dim>(&self, end_dim: Dm) -> crate::Result<Self> {
261 self.flatten(0usize, end_dim)
262 }
263
264 pub fn flatten_all(&self) -> crate::Result<Self> {
265 self.reshape(Shape::from(self.element_count()))
266 }
267
268 pub fn stack<A: AsRef<Self>, Dm: Dim>(args: &[A], dim: Dm) -> crate::Result<Self> {
270 if args.is_empty() {
271 return Err(Error::OpRequiresAtLeastOneTensor { op: "stack" });
272 }
273 let first = args[0].as_ref();
274 let dim = dim.to_index_plus_one(first.shape(), "stack")?;
275 let unsqueezed: crate::Result<Vec<Self>> = args.iter().map(|a| a.as_ref().unsqueeze(dim)).collect();
276 Self::cat(&unsqueezed?, dim)
277 }
278
279 pub fn split<Dm: Dim>(&self, dim: Dm) -> crate::Result<Vec<Self>> {
281 let dim = dim.to_index(self.shape(), "split")?;
282 let n = self.dims()[dim];
283 (0..n).map(|i| self.narrow(dim, i, 1)).collect()
284 }
285
286 pub fn chunk<Dm: Dim>(&self, chunks: usize, dim: Dm) -> crate::Result<Vec<Self>> {
288 let dim = dim.to_index(self.shape(), "chunk")?;
289 let n = self.dims()[dim];
290 if chunks == 0 {
291 return Err(Error::Msg("chunk: chunks cannot be 0".into()));
292 }
293 let size = (n + chunks - 1) / chunks; let mut result = Vec::new();
295 let mut start = 0;
296 while start < n {
297 let len = size.min(n - start);
298 result.push(self.narrow(dim, start, len)?);
299 start += len;
300 }
301 Ok(result)
302 }
303
304 pub fn repeat_dim<Dm: Dim>(&self, dim: Dm, times: usize) -> crate::Result<Self> {
306 let dim = dim.to_index(self.shape(), "repeat_dim")?;
307 let copies: Vec<&Self> = (0..times).map(|_| self).collect();
308 Self::cat(&copies, dim)
309 }
310}