Skip to main content

luma_tensor/ops/
shape.rs

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    /// Resolve the storage for a view of `self` under `layout`.
61    ///
62    /// Compute devices return `self.0.storage.clone()` (alias); a tracing device
63    /// returns a freshly-built storage so the view is a distinct graph value.
64    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    /// Deep-copy the tensor data to independent storage (records `Op::Copy` for `Float`).
79    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    /// Copy the data from `src` into `self` **in-place**, preserving [`TensorId`].
86    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        // Requires exclusive TensorImpl access (normal when tensor is reached via
94        // a mutable visitor or held uniquely).
95        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                // Overwrite through the existing RwLock — keeps the same Arc.
100                *lock.write().expect("storage write lock") = src_storage;
101            }
102            None => {
103                // Phantom tensor — create storage for the first time.
104                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    /// Slice along `dim` with `start`, `end`, and `step`.
165    ///
166    /// Returns a view (no copy) when possible. For `step == 1` this is
167    /// equivalent to `narrow`; for `step > 1` the backward pass is not yet
168    /// supported on autograd tensors.
169    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    /// Materialize a contiguous copy (records `Op::Copy` for `Float`).
211    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    /// Concatenate tensors along `dim` (materializes new storage).
221    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    /// Flatten dims `start_dim..=end_dim` into one.
238    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    /// Stack tensors along a new axis at `dim` (unsqueeze each + cat).
269    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    /// Split into individual slices along `dim` (one per index).
280    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    /// Split into `chunks` roughly equal pieces along `dim`.
287    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; // ceiling division
294        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    /// Tile `self` by `times` along `dim`.
305    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}