Skip to main content

tenferro_ad/
shape_packing.rs

1use std::ops::Range;
2
3use tenferro_tensor::{GatherConfig, SliceConfig, Tensor, TypedTensor};
4
5use crate::eager::{EagerSession, EagerTensor};
6use crate::error::{Error, Result};
7
8fn normalize_existing_axis(op: &'static str, axis: isize, rank: usize) -> Result<usize> {
9    let normalized = if axis >= 0 {
10        axis as usize
11    } else {
12        rank.checked_sub(axis.unsigned_abs()).ok_or_else(|| {
13            tenferro_tensor::Error::axis_out_of_bounds(op, axis.unsigned_abs(), rank)
14        })?
15    };
16    if normalized >= rank {
17        return Err(
18            tenferro_tensor::Error::axis_out_of_bounds(op, axis.unsigned_abs(), rank).into(),
19        );
20    }
21    Ok(normalized)
22}
23
24fn normalize_insert_axis(op: &'static str, axis: isize, rank: usize) -> Result<usize> {
25    let insert_rank = rank
26        .checked_add(1)
27        .ok_or_else(|| tenferro_tensor::Error::axis_out_of_bounds(op, axis.unsigned_abs(), rank))?;
28    let normalized = if axis >= 0 {
29        axis as usize
30    } else {
31        insert_rank
32            .checked_sub(axis.unsigned_abs())
33            .ok_or_else(|| {
34                tenferro_tensor::Error::axis_out_of_bounds(op, axis.unsigned_abs(), insert_rank)
35            })?
36    };
37    if normalized > rank {
38        return Err(tenferro_tensor::Error::axis_out_of_bounds(
39            op,
40            axis.unsigned_abs(),
41            insert_rank,
42        )
43        .into());
44    }
45    Ok(normalized)
46}
47
48fn index_select_config(
49    shape: &[usize],
50    axis: isize,
51    positions: &[usize],
52) -> Result<(Tensor, GatherConfig)> {
53    let axis = normalize_existing_axis("index_select", axis, shape.len())?;
54    let axis_extent = shape[axis];
55    for &position in positions {
56        if position >= axis_extent {
57            return Err(tenferro_tensor::Error::invalid_argument(
58                "index_select",
59                "position",
60                format!(
61                    "position {position} out of bounds for axis {axis} with extent {axis_extent}"
62                ),
63            )
64            .into());
65        }
66    }
67
68    let mut slice_sizes = shape.to_vec();
69    slice_sizes[axis] = 1;
70
71    let offset_dims = (0..shape.len()).filter(|&dim| dim != axis).collect();
72    let index_data = positions
73        .iter()
74        .map(|&position| {
75            i64::try_from(position).map_err(|_| {
76                tenferro_tensor::Error::invalid_argument(
77                    "index_select",
78                    "position",
79                    format!("position {position} cannot be represented as i64"),
80                )
81            })
82        })
83        .collect::<tenferro_tensor::Result<Vec<_>>>()?;
84    let indices = Tensor::from_typed::<i64>(TypedTensor::from_vec_col_major(
85        vec![positions.len(), 1],
86        index_data,
87    )?);
88
89    let config = GatherConfig {
90        offset_dims,
91        collapsed_slice_dims: vec![axis],
92        start_index_map: vec![axis],
93        index_vector_dim: 1,
94        slice_sizes,
95    };
96
97    Ok((indices, config))
98}
99
100fn validate_stack_shapes(op: &'static str, shapes: &[&[usize]]) -> Result<()> {
101    let Some(first) = shapes.first() else {
102        return Err(tenferro_tensor::Error::invalid_argument(
103            op,
104            "inputs",
105            "stack requires at least one input",
106        )
107        .into());
108    };
109    for shape in shapes.iter().skip(1) {
110        if *shape != *first {
111            return Err(tenferro_tensor::Error::shape_mismatch(op, *first, *shape).into());
112        }
113    }
114    Ok(())
115}
116
117#[derive(Clone, Debug)]
118enum AxisSelection {
119    Slice {
120        axis: usize,
121        range: Range<usize>,
122        step: usize,
123    },
124    Take {
125        axis: usize,
126        indices: Vec<usize>,
127    },
128}
129
130fn validate_axis_selection(
131    op: &'static str,
132    rank: usize,
133    seen: &mut [bool],
134    axis: usize,
135) -> Result<()> {
136    if axis >= rank {
137        return Err(tenferro_tensor::Error::axis_out_of_bounds(op, axis, rank).into());
138    }
139    if seen[axis] {
140        return Err(tenferro_tensor::Error::duplicate_axis(op, axis, "selection").into());
141    }
142    seen[axis] = true;
143    Ok(())
144}
145
146fn apply_slice_axis_config(
147    op: &'static str,
148    shape: &[usize],
149    selections: &[AxisSelection],
150) -> Result<Option<SliceConfig>> {
151    let mut starts = vec![0; shape.len()];
152    let mut limits = shape.to_vec();
153    let mut strides = vec![1; shape.len()];
154    let mut has_slice = false;
155    for selection in selections {
156        let AxisSelection::Slice { axis, range, step } = selection else {
157            continue;
158        };
159        if *step == 0 {
160            return Err(tenferro_tensor::Error::invalid_argument(
161                op,
162                "step",
163                format!("axis {axis} has zero step"),
164            )
165            .into());
166        }
167        let extent = shape[*axis];
168        if range.start > range.end || range.end > extent {
169            return Err(tenferro_tensor::Error::invalid_argument(
170                op,
171                "range",
172                format!(
173                    "axis {axis} range {}..{} is out of bounds for extent {extent}",
174                    range.start, range.end
175                ),
176            )
177            .into());
178        }
179        starts[*axis] = range.start;
180        limits[*axis] = range.end;
181        strides[*axis] = *step;
182        has_slice = true;
183    }
184    Ok(has_slice.then_some(SliceConfig {
185        starts,
186        limits,
187        strides,
188    }))
189}
190
191/// Rank-preserving eager tensor slicing builder.
192///
193/// Unspecified axes are kept whole. Range selections become one `Slice`
194/// operation; host-known position selections become `Gather`/`index_select`
195/// operations.
196///
197/// # Examples
198///
199/// ```rust
200/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
201///
202/// let ctx = EagerRuntime::new()?;
203/// let x = EagerTensor::from_tensor_in(
204///     Tensor::from_vec_col_major(vec![3, 4], vec![0.0_f64; 12]).unwrap(),
205///     ctx,
206/// ).unwrap();
207/// let y = x.runtime().with_eager_session(|s| x.slice_builder().axis(0, 0..2).axis_step(1, 0..4, 2).apply(s))?;
208/// assert_eq!(y.shape(), &[2, 2]);
209/// # Ok::<(), tenferro_ad::Error>(())
210/// ```
211#[derive(Clone, Debug)]
212pub struct EagerSliceBuilder<'a> {
213    tensor: &'a EagerTensor,
214    selections: Vec<AxisSelection>,
215}
216
217impl<'a> EagerSliceBuilder<'a> {
218    fn new(tensor: &'a EagerTensor) -> Self {
219        Self {
220            tensor,
221            selections: Vec::new(),
222        }
223    }
224
225    /// Add an exclusive-end range selection for one axis.
226    ///
227    /// # Examples
228    ///
229    /// ```rust
230    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
231    ///
232    /// let ctx = EagerRuntime::new()?;
233    /// let x = EagerTensor::from_tensor_in(
234    ///     Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0]).unwrap(),
235    ///     ctx,
236    /// ).unwrap();
237    /// let y = x.runtime().with_eager_session(|s| x.slice_builder().axis(0, 1..3).apply(s))?;
238    /// assert_eq!(y.shape(), &[2]);
239    /// # Ok::<(), tenferro_ad::Error>(())
240    /// ```
241    pub fn axis(mut self, axis: usize, range: Range<usize>) -> Self {
242        self.selections.push(AxisSelection::Slice {
243            axis,
244            range,
245            step: 1,
246        });
247        self
248    }
249
250    /// Add an exclusive-end strided range selection for one axis.
251    ///
252    /// # Examples
253    ///
254    /// ```rust
255    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
256    ///
257    /// let ctx = EagerRuntime::new()?;
258    /// let x = EagerTensor::from_tensor_in(
259    ///     Tensor::from_vec_col_major(vec![5], vec![1.0_f64, 2.0, 3.0, 4.0, 5.0]).unwrap(),
260    ///     ctx,
261    /// ).unwrap();
262    /// let y = x.runtime().with_eager_session(|s| x.slice_builder().axis_step(0, 0..5, 2).apply(s))?;
263    /// assert_eq!(y.shape(), &[3]);
264    /// # Ok::<(), tenferro_ad::Error>(())
265    /// ```
266    pub fn axis_step(mut self, axis: usize, range: Range<usize>, step: usize) -> Self {
267        self.selections
268            .push(AxisSelection::Slice { axis, range, step });
269        self
270    }
271
272    /// Add a host-known position selection for one axis.
273    ///
274    /// # Examples
275    ///
276    /// ```rust
277    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
278    ///
279    /// let ctx = EagerRuntime::new()?;
280    /// let x = EagerTensor::from_tensor_in(
281    ///     Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0]).unwrap(),
282    ///     ctx,
283    /// ).unwrap();
284    /// let y = x.runtime().with_eager_session(|s| x.slice_builder().take_axis(0, &[2, 0]).apply(s))?;
285    /// assert_eq!(y.shape(), &[2]);
286    /// # Ok::<(), tenferro_ad::Error>(())
287    /// ```
288    pub fn take_axis(mut self, axis: usize, indices: &[usize]) -> Self {
289        self.selections.push(AxisSelection::Take {
290            axis,
291            indices: indices.to_vec(),
292        });
293        self
294    }
295
296    /// Build and apply the requested slice/take operations.
297    ///
298    /// # Examples
299    ///
300    /// ```rust
301    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
302    ///
303    /// let ctx = EagerRuntime::new()?;
304    /// let x = EagerTensor::from_tensor_in(
305    ///     Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0]).unwrap(),
306    ///     ctx,
307    /// ).unwrap();
308    /// let y = x.runtime().with_eager_session(|s| x.slice_builder().axis(0, 1..4).apply(s))?;
309    /// assert_eq!(y.shape(), &[3]);
310    /// # Ok::<(), tenferro_ad::Error>(())
311    /// ```
312    /// # Errors
313    ///
314    /// Returns [`tenferro_tensor::ValidationError::AxisOutOfBounds`] or
315    /// `DuplicateAxis` when selections address an invalid/repeated axis,
316    /// `InvalidArgument` for zero steps or out-of-bounds ranges, or a typed
317    /// backend/runtime-state error while applying the selections.
318    pub fn apply(self, session: &mut EagerSession<'_>) -> Result<EagerTensor> {
319        session.ensure_runtime(self.tensor)?;
320        let shape = self.tensor.shape().to_vec();
321        let mut seen = vec![false; shape.len()];
322        for selection in &self.selections {
323            let axis = match selection {
324                AxisSelection::Slice { axis, .. } | AxisSelection::Take { axis, .. } => *axis,
325            };
326            validate_axis_selection("slice_builder", shape.len(), &mut seen, axis)?;
327        }
328
329        let mut output = self.tensor.clone();
330        if let Some(config) = apply_slice_axis_config("slice_builder", &shape, &self.selections)? {
331            output = session.slice(&output, config)?;
332        }
333        for selection in self.selections {
334            if let AxisSelection::Take { axis, indices } = selection {
335                output = session.take_axis(&output, axis, &indices)?;
336            }
337        }
338        Ok(output)
339    }
340}
341
342impl EagerTensor {
343    /// Start a rank-preserving slicing builder for this tensor.
344    ///
345    /// # Examples
346    ///
347    /// ```rust
348    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
349    ///
350    /// let ctx = EagerRuntime::new()?;
351    /// let x = EagerTensor::from_tensor_in(
352    ///     Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0]).unwrap(),
353    ///     ctx,
354    /// ).unwrap();
355    /// let y = x.runtime().with_eager_session(|s| x.slice_builder().axis(0, 0..2).apply(s))?;
356    /// assert_eq!(y.shape(), &[2]);
357    /// # Ok::<(), tenferro_ad::Error>(())
358    /// ```
359    pub fn slice_builder(&self) -> EagerSliceBuilder<'_> {
360        EagerSliceBuilder::new(self)
361    }
362}
363
364impl EagerSession<'_> {
365    /// Select positions from one axis using a borrowed session.
366    ///
367    /// # Examples
368    /// ```rust
369    /// use tenferro_ad::{EagerRuntime, Tensor};
370    /// let ctx = EagerRuntime::new()?;
371    /// let result = ctx.with_eager_session(|s| {
372    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0])?)?;
373    ///     s.index_select(&x, -1, &[2, 0])
374    /// })?;
375    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[3.0, 1.0]);
376    /// # Ok::<(), tenferro_ad::Error>(())
377    /// ```
378    /// # Errors
379    /// Returns a typed foreign-runtime, invalid-axis/index, or backend error.
380    pub fn index_select(
381        &mut self,
382        tensor: &EagerTensor,
383        axis: isize,
384        positions: &[usize],
385    ) -> Result<EagerTensor> {
386        self.ensure_runtime(tensor)?;
387        let (indices, config) = index_select_config(tensor.shape(), axis, positions)?;
388        let indices = self.constant_from_host(indices)?;
389        self.gather(tensor, &indices, config)
390    }
391
392    /// Select entries from an axis by host-known positions.
393    ///
394    /// # Examples
395    /// ```rust
396    /// use tenferro_ad::{EagerRuntime, Tensor};
397    /// let ctx = EagerRuntime::new()?;
398    /// let result = ctx.with_eager_session(|s| {
399    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?)?;
400    ///     s.take_axis(&x, 0, &[1])
401    /// })?;
402    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[2.0]);
403    /// # Ok::<(), tenferro_ad::Error>(())
404    /// ```
405    /// # Errors
406    /// Returns a typed foreign-runtime, invalid-axis/index, or backend error.
407    pub fn take_axis(
408        &mut self,
409        tensor: &EagerTensor,
410        axis: usize,
411        positions: &[usize],
412    ) -> Result<EagerTensor> {
413        self.ensure_runtime(tensor)?;
414        let axis = isize::try_from(axis).map_err(|_| {
415            Error::TensorRuntime(tenferro_tensor::Error::invalid_argument(
416                "take_axis",
417                "axis",
418                format!("{axis} cannot be represented as isize"),
419            ))
420        })?;
421        self.index_select(tensor, axis, positions)
422    }
423
424    /// Select matrix rows by host-known positions.
425    ///
426    /// # Examples
427    /// ```rust
428    /// use tenferro_ad::{EagerRuntime, Tensor};
429    /// let ctx = EagerRuntime::new()?;
430    /// let result = ctx.with_eager_session(|s| {
431    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2, 1], vec![1.0_f64, 2.0])?)?;
432    ///     s.take_rows(&x, &[1])
433    /// })?;
434    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[2.0]);
435    /// # Ok::<(), tenferro_ad::Error>(())
436    /// ```
437    /// # Errors
438    /// Returns a typed foreign-runtime, invalid-row, or backend error.
439    pub fn take_rows(&mut self, tensor: &EagerTensor, rows: &[usize]) -> Result<EagerTensor> {
440        self.take_axis(tensor, 0, rows)
441    }
442
443    /// Select matrix columns by host-known positions.
444    ///
445    /// # Examples
446    /// ```rust
447    /// use tenferro_ad::{EagerRuntime, Tensor};
448    /// let ctx = EagerRuntime::new()?;
449    /// let result = ctx.with_eager_session(|s| {
450    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1, 2], vec![1.0_f64, 2.0])?)?;
451    ///     s.take_cols(&x, &[1])
452    /// })?;
453    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[2.0]);
454    /// # Ok::<(), tenferro_ad::Error>(())
455    /// ```
456    /// # Errors
457    /// Returns a typed foreign-runtime, invalid-column, or backend error.
458    pub fn take_cols(&mut self, tensor: &EagerTensor, cols: &[usize]) -> Result<EagerTensor> {
459        self.take_axis(tensor, 1, cols)
460    }
461
462    /// Select a matrix block by host-known row and column positions.
463    ///
464    /// # Examples
465    /// ```rust
466    /// use tenferro_ad::{EagerRuntime, Tensor};
467    /// let ctx = EagerRuntime::new()?;
468    /// let result = ctx.with_eager_session(|s| {
469    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?)?;
470    ///     s.take_block(&x, &[1], &[0])
471    /// })?;
472    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[2.0]);
473    /// # Ok::<(), tenferro_ad::Error>(())
474    /// ```
475    /// # Errors
476    /// Returns a typed foreign-runtime, invalid-row/column, or backend error.
477    pub fn take_block(
478        &mut self,
479        tensor: &EagerTensor,
480        rows: &[usize],
481        cols: &[usize],
482    ) -> Result<EagerTensor> {
483        let rows = self.take_rows(tensor, rows)?;
484        self.take_cols(&rows, cols)
485    }
486
487    /// Slice one axis using an exclusive-end range in this borrowed session.
488    ///
489    /// # Examples
490    /// ```rust
491    /// use tenferro_ad::{EagerRuntime, Tensor};
492    /// let ctx = EagerRuntime::new()?;
493    /// let result = ctx.with_eager_session(|s| {
494    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0])?)?;
495    ///     s.slice_axis(&x, 0, 1..3)
496    /// })?;
497    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[2.0, 3.0]);
498    /// # Ok::<(), tenferro_ad::Error>(())
499    /// ```
500    /// # Errors
501    /// Returns a typed foreign-runtime, invalid-axis/range, or backend error.
502    pub fn slice_axis(
503        &mut self,
504        tensor: &EagerTensor,
505        axis: usize,
506        range: Range<usize>,
507    ) -> Result<EagerTensor> {
508        tensor.slice_builder().axis(axis, range).apply(self)
509    }
510
511    /// Stack eager tensors along a new axis in this borrowed session.
512    ///
513    /// # Examples
514    /// ```rust
515    /// use tenferro_ad::{EagerRuntime, Tensor};
516    /// let ctx = EagerRuntime::new()?;
517    /// let result = ctx.with_eager_session(|s| {
518    ///     let a = s.constant_from(Tensor::from_vec_col_major(vec![], vec![1.0_f64])?)?;
519    ///     let b = s.constant_from(Tensor::from_vec_col_major(vec![], vec![2.0_f64])?)?;
520    ///     s.stack(&[&a, &b], -1)
521    /// })?;
522    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[1.0, 2.0]);
523    /// # Ok::<(), tenferro_ad::Error>(())
524    /// ```
525    /// # Errors
526    /// Returns a typed empty-input, invalid-axis/shape, foreign-runtime, or backend error.
527    pub fn stack(&mut self, tensors: &[&EagerTensor], dim: isize) -> Result<EagerTensor> {
528        let first = tensors.first().copied().ok_or_else(|| {
529            Error::TensorRuntime(tenferro_tensor::Error::invalid_argument(
530                "stack",
531                "inputs",
532                "stack requires at least one input",
533            ))
534        })?;
535        let shapes = tensors
536            .iter()
537            .map(|tensor| tensor.shape())
538            .collect::<Vec<_>>();
539        validate_stack_shapes("stack", &shapes)?;
540        let axis = normalize_insert_axis("stack", dim, first.shape().len())?;
541        let mut expanded_shape = first.shape().to_vec();
542        expanded_shape.insert(axis, 1);
543        let expanded = tensors
544            .iter()
545            .map(|tensor| self.reshape(tensor, &expanded_shape))
546            .collect::<Result<Vec<_>>>()?;
547        let refs = expanded.iter().collect::<Vec<_>>();
548        self.concatenate(&refs, axis)
549    }
550}
551
552#[cfg(test)]
553mod tests {
554    use super::{normalize_existing_axis, normalize_insert_axis};
555
556    #[test]
557    fn axis_normalization_handles_ranks_larger_than_isize_max() {
558        assert_eq!(normalize_existing_axis("test", 0, usize::MAX).unwrap(), 0);
559        assert_eq!(
560            normalize_existing_axis("test", -1, usize::MAX).unwrap(),
561            usize::MAX - 1
562        );
563        assert_eq!(
564            normalize_insert_axis("test", -1, usize::MAX - 1).unwrap(),
565            usize::MAX - 1
566        );
567        assert!(normalize_insert_axis("test", -1, usize::MAX).is_err());
568    }
569}