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}