Skip to main content

tenferro_ad/eager/
composite.rs

1//! Eager composite operations (activations, normalizations, softmax,
2//! `take_along_axis`) on the borrowed [`EagerSession`].
3//!
4//! The formulation and edge-case policy are shared with the traced and
5//! concrete-session surfaces through `tenferro_runtime::composite`; AD follows
6//! from the recorded primitives.
7
8use tenferro_runtime::composite::{
9    self, scalar_tensor, zero_pad_config, CompositeBinary, CompositeOps, CompositeReduce,
10    CompositeUnary,
11};
12use tenferro_tensor::{CompareDir, DType, GatherConfig};
13
14use super::{EagerSession, EagerTensor};
15use crate::{Error, Result};
16
17/// [`CompositeOps`] over a borrowed eager session.
18struct EagerComposite<'s, 'a> {
19    session: &'s mut EagerSession<'a>,
20}
21
22impl CompositeOps for EagerComposite<'_, '_> {
23    type Value = EagerTensor;
24    type Error = Error;
25
26    fn dtype(&self, value: &EagerTensor) -> DType {
27        value.dtype()
28    }
29
30    fn shape(&self, value: &EagerTensor) -> Result<Vec<usize>> {
31        Ok(value.shape().to_vec())
32    }
33
34    fn scalar(&mut self, dtype: DType, value: f64) -> Result<EagerTensor> {
35        self.session
36            .constant_from_host(scalar_tensor(dtype, value)?)
37    }
38
39    fn unary(&mut self, op: CompositeUnary, value: &EagerTensor) -> Result<EagerTensor> {
40        let session = &mut *self.session;
41        match op {
42            CompositeUnary::Neg => session.neg(value),
43            CompositeUnary::Exp => session.exp(value),
44            CompositeUnary::Log => session.log(value),
45            CompositeUnary::Log1p => session.log1p(value),
46            CompositeUnary::Tanh => session.tanh(value),
47            CompositeUnary::Erf => session.erf(value),
48            CompositeUnary::Rsqrt => session.rsqrt(value),
49        }
50    }
51
52    fn binary(
53        &mut self,
54        op: CompositeBinary,
55        lhs: &EagerTensor,
56        rhs: &EagerTensor,
57    ) -> Result<EagerTensor> {
58        let session = &mut *self.session;
59        match op {
60            CompositeBinary::Add => session.add(lhs, rhs),
61            CompositeBinary::Sub => session.sub(lhs, rhs),
62            CompositeBinary::Mul => session.mul(lhs, rhs),
63            CompositeBinary::Div => session.div(lhs, rhs),
64            CompositeBinary::Maximum => session.maximum(lhs, rhs),
65        }
66    }
67
68    fn compare(
69        &mut self,
70        lhs: &EagerTensor,
71        rhs: &EagerTensor,
72        dir: CompareDir,
73    ) -> Result<EagerTensor> {
74        self.session.compare(lhs, rhs, dir)
75    }
76
77    fn select(
78        &mut self,
79        condition: &EagerTensor,
80        on_true: &EagerTensor,
81        on_false: &EagerTensor,
82    ) -> Result<EagerTensor> {
83        self.session.where_select(condition, on_true, on_false)
84    }
85
86    fn reduce(
87        &mut self,
88        op: CompositeReduce,
89        value: &EagerTensor,
90        axes: &[usize],
91    ) -> Result<EagerTensor> {
92        match op {
93            CompositeReduce::Sum => self.session.reduce_sum(value, Some(axes)),
94            CompositeReduce::Max => self.session.reduce_max(value, Some(axes)),
95            CompositeReduce::SumSquares => self.session.reduce_sum_squares(value, Some(axes)),
96        }
97    }
98
99    fn broadcast_in_dim(
100        &mut self,
101        value: &EagerTensor,
102        shape: &[usize],
103        dims: &[usize],
104    ) -> Result<EagerTensor> {
105        self.session.broadcast_in_dim(value, shape, dims)
106    }
107
108    fn reshape(&mut self, value: &EagerTensor, shape: &[usize]) -> Result<EagerTensor> {
109        self.session.reshape(value, shape.to_vec())
110    }
111
112    fn concatenate(&mut self, values: &[&EagerTensor], axis: usize) -> Result<EagerTensor> {
113        self.session.concatenate(values, axis)
114    }
115
116    fn pad(&mut self, value: &EagerTensor, low: &[usize], high: &[usize]) -> Result<EagerTensor> {
117        self.session.pad(value, zero_pad_config(low, high))
118    }
119
120    fn gather(
121        &mut self,
122        operand: &EagerTensor,
123        indices: &EagerTensor,
124        config: GatherConfig,
125    ) -> Result<EagerTensor> {
126        self.session.gather(operand, indices, config)
127    }
128}
129
130impl<'a> EagerSession<'a> {
131    fn composite(&mut self) -> EagerComposite<'_, 'a> {
132        EagerComposite { session: self }
133    }
134
135    /// Logistic sigmoid `1 / (1 + exp(-x))`, overflow-free.
136    ///
137    /// Evaluated as `1 / (1 + e)` for `x > 0` and `e / (1 + e)` otherwise, with
138    /// `e = exp(-|x|)`; the derivative is finite everywhere (`sigmoid'(0) = 1/4`).
139    /// Real `F32`/`F64` only.
140    ///
141    /// # Examples
142    /// ```rust
143    /// use tenferro_ad::{EagerRuntime, Tensor};
144    /// let ctx = EagerRuntime::new()?;
145    /// let y = ctx.with_eager_session(|s| {
146    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![3], vec![-700.0_f64, 0.0, 1000.0])?)?;
147    ///     s.sigmoid(&x)
148    /// })?;
149    /// let y = y.value()?;
150    /// let y = y.as_slice::<f64>()?;
151    /// assert_eq!(y[1], 0.5);
152    /// assert!(y[0] > 0.0 && y[0] < 1e-300);
153    /// assert_eq!(y[2], 1.0);
154    /// # Ok::<(), tenferro_ad::Error>(())
155    /// ```
156    ///
157    /// # Errors
158    ///
159    /// Returns a typed `UnsupportedDType` error for complex, integer, or `Bool`
160    /// input, [`Error::ContextMismatch`] for a tensor from another runtime, or a
161    /// backend error.
162    pub fn sigmoid(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
163        composite::sigmoid(&mut self.composite(), input)
164    }
165
166    /// SiLU (swish) `x * sigmoid(x)`.
167    ///
168    /// Real `F32`/`F64` only.
169    ///
170    /// # Examples
171    /// ```rust
172    /// use tenferro_ad::{EagerRuntime, Tensor};
173    /// let ctx = EagerRuntime::new()?;
174    /// let y = ctx.with_eager_session(|s| {
175    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![3], vec![-1.0_f64, 0.0, 1.0])?)?;
176    ///     s.silu(&x)
177    /// })?;
178    /// let y = y.value()?;
179    /// let y = y.as_slice::<f64>()?;
180    /// assert_eq!(y[1], 0.0);
181    /// assert!((y[2] - 1.0 / (1.0 + (-1.0_f64).exp())).abs() < 1e-15);
182    /// # Ok::<(), tenferro_ad::Error>(())
183    /// ```
184    ///
185    /// # Errors
186    ///
187    /// Returns a typed `UnsupportedDType` error for complex, integer, or `Bool`
188    /// input, [`Error::ContextMismatch`] for a tensor from another runtime, or a
189    /// backend error.
190    pub fn silu(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
191        composite::silu(&mut self.composite(), input)
192    }
193
194    /// Softplus `log(1 + exp(x))` in the stable form `max(x, 0) + log1p(exp(-|x|))`.
195    ///
196    /// Never overflows; `softplus'(0) = 1/2` and `softplus''(0) = 1/4`. Real `F32`/`F64` only.
197    ///
198    /// # Examples
199    /// ```rust
200    /// use tenferro_ad::{EagerRuntime, Tensor};
201    /// let ctx = EagerRuntime::new()?;
202    /// let y = ctx.with_eager_session(|s| {
203    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![3], vec![-1000.0_f64, 0.0, 1000.0])?)?;
204    ///     s.softplus(&x)
205    /// })?;
206    /// let y = y.value()?;
207    /// let y = y.as_slice::<f64>()?;
208    /// assert_eq!(y[0], 0.0);
209    /// assert!((y[1] - 2.0_f64.ln()).abs() < 1e-15);
210    /// assert_eq!(y[2], 1000.0);
211    /// # Ok::<(), tenferro_ad::Error>(())
212    /// ```
213    ///
214    /// # Errors
215    ///
216    /// Returns a typed `UnsupportedDType` error for complex, integer, or `Bool`
217    /// input, [`Error::ContextMismatch`] for a tensor from another runtime, or a
218    /// backend error.
219    pub fn softplus(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
220        composite::softplus(&mut self.composite(), input)
221    }
222
223    /// Exact GELU `x/2 * (1 + erf(x / sqrt(2)))` (PyTorch `approximate="none"`).
224    ///
225    /// Real `F32`/`F64` only.
226    ///
227    /// # Examples
228    /// ```rust
229    /// use tenferro_ad::{EagerRuntime, Tensor};
230    /// let ctx = EagerRuntime::new()?;
231    /// let y = ctx.with_eager_session(|s| {
232    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![3], vec![-1.0_f64, 0.0, 1.0])?)?;
233    ///     s.gelu(&x)
234    /// })?;
235    /// let y = y.value()?;
236    /// let y = y.as_slice::<f64>()?;
237    /// assert_eq!(y[1], 0.0);
238    /// assert!((y[2] - 0.841_344_746_068_542_9).abs() < 1e-15);
239    /// # Ok::<(), tenferro_ad::Error>(())
240    /// ```
241    ///
242    /// # Errors
243    ///
244    /// Returns a typed `UnsupportedDType` error for complex, integer, or `Bool`
245    /// input, [`Error::ContextMismatch`] for a tensor from another runtime, or a
246    /// backend error.
247    pub fn gelu(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
248        composite::gelu(&mut self.composite(), input)
249    }
250
251    /// GELU tanh approximation (PyTorch `approximate="tanh"`).
252    ///
253    /// `x/2 * (1 + tanh(sqrt(2/pi) * (x + 0.044715 x^3)))`; real `F32`/`F64` only.
254    ///
255    /// # Examples
256    /// ```rust
257    /// use tenferro_ad::{EagerRuntime, Tensor};
258    /// let ctx = EagerRuntime::new()?;
259    /// let y = ctx.with_eager_session(|s| {
260    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![3], vec![-1.0_f64, 0.0, 1.0])?)?;
261    ///     s.gelu_tanh(&x)
262    /// })?;
263    /// let y = y.value()?;
264    /// let y = y.as_slice::<f64>()?;
265    /// assert_eq!(y[1], 0.0);
266    /// assert!((y[2] - 0.841_191_990_608_276_8).abs() < 1e-12);
267    /// # Ok::<(), tenferro_ad::Error>(())
268    /// ```
269    ///
270    /// # Errors
271    ///
272    /// Returns a typed `UnsupportedDType` error for complex, integer, or `Bool`
273    /// input, [`Error::ContextMismatch`] for a tensor from another runtime, or a
274    /// backend error.
275    pub fn gelu_tanh(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
276        composite::gelu_tanh(&mut self.composite(), input)
277    }
278
279    /// Arithmetic mean over `axes` (`None` reduces every axis).
280    ///
281    /// Float and complex dtypes. The sum is divided by the element count; a mean
282    /// over zero elements is `NaN`, and `Some(&[])` is the identity.
283    ///
284    /// # Examples
285    /// ```rust
286    /// use tenferro_ad::{EagerRuntime, Tensor};
287    /// let ctx = EagerRuntime::new()?;
288    /// let y = ctx.with_eager_session(|s| {
289    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?)?;
290    ///     s.reduce_mean(&x, Some(&[1]))
291    /// })?;
292    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[2.0, 3.0]);
293    /// # Ok::<(), tenferro_ad::Error>(())
294    /// ```
295    ///
296    /// # Errors
297    ///
298    /// Returns a typed `UnsupportedDType` error for integer or `Bool` input, an
299    /// `AxisOutOfBounds` / `DuplicateAxis` validation error for invalid axes,
300    /// [`Error::ContextMismatch`] for a tensor from another runtime, or a backend error.
301    pub fn reduce_mean(
302        &mut self,
303        input: &EagerTensor,
304        axes: Option<&[usize]>,
305    ) -> Result<EagerTensor> {
306        composite::reduce_mean(&mut self.composite(), input, axes)
307    }
308
309    /// Max-subtracted softmax along `axis`.
310    ///
311    /// A slice that is entirely `-inf` returns zeros with a finite gradient instead
312    /// of `NaN`; a participating `NaN` or `+inf` makes its slice `NaN`; a
313    /// zero-length `axis` returns an empty result. Real `F32`/`F64` only.
314    ///
315    /// # Examples
316    /// ```rust
317    /// use tenferro_ad::{EagerRuntime, Tensor};
318    /// let ctx = EagerRuntime::new()?;
319    /// let y = ctx.with_eager_session(|s| {
320    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![2], vec![0.0_f64, f64::NEG_INFINITY])?)?;
321    ///     s.softmax(&x, 0)
322    /// })?;
323    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[1.0, 0.0]);
324    /// # Ok::<(), tenferro_ad::Error>(())
325    /// ```
326    ///
327    /// # Errors
328    ///
329    /// Returns a typed `UnsupportedDType` error for non-real input, an
330    /// `AxisOutOfBounds` validation error for an invalid axis,
331    /// [`Error::ContextMismatch`] for a tensor from another runtime, or a backend error.
332    pub fn softmax(&mut self, input: &EagerTensor, axis: usize) -> Result<EagerTensor> {
333        composite::softmax(&mut self.composite(), input, axis)
334    }
335
336    /// Max-subtracted log-softmax along `axis`.
337    ///
338    /// A slice that is entirely `-inf` returns `-inf` with a finite gradient
339    /// instead of `NaN`. Real `F32`/`F64` only.
340    ///
341    /// # Examples
342    /// ```rust
343    /// use tenferro_ad::{EagerRuntime, Tensor};
344    /// let ctx = EagerRuntime::new()?;
345    /// let y = ctx.with_eager_session(|s| {
346    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 1.0])?)?;
347    ///     s.log_softmax(&x, 0)
348    /// })?;
349    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[-std::f64::consts::LN_2; 2]);
350    /// # Ok::<(), tenferro_ad::Error>(())
351    /// ```
352    ///
353    /// # Errors
354    ///
355    /// Returns a typed `UnsupportedDType` error for non-real input, an
356    /// `AxisOutOfBounds` validation error for an invalid axis,
357    /// [`Error::ContextMismatch`] for a tensor from another runtime, or a backend error.
358    pub fn log_softmax(&mut self, input: &EagerTensor, axis: usize) -> Result<EagerTensor> {
359        composite::log_softmax(&mut self.composite(), input, axis)
360    }
361
362    /// Softmax along `axis` over the entries where the `Bool` `mask` is true.
363    ///
364    /// `mask` broadcasts to the input shape. Masked-out entries get probability `0`
365    /// and a zero gradient whatever their value; a slice with no unmasked entry
366    /// returns zeros with a zero gradient.
367    ///
368    /// # Examples
369    /// ```rust
370    /// use tenferro_ad::{EagerRuntime, Tensor};
371    /// let ctx = EagerRuntime::new()?;
372    /// let y = ctx.with_eager_session(|s| {
373    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 1.0, f64::NAN])?)?;
374    ///     let mask = s.constant_from_host(Tensor::from_vec_col_major(vec![3], vec![true, true, false])?)?;
375    ///     s.masked_softmax(&x, &mask, 0)
376    /// })?;
377    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.5, 0.5, 0.0]);
378    /// # Ok::<(), tenferro_ad::Error>(())
379    /// ```
380    ///
381    /// # Errors
382    ///
383    /// Returns a typed `UnsupportedDType` error for non-real input, a
384    /// `DTypeMismatch` validation error for a non-`Bool` mask, `ShapeMismatch` for
385    /// a mask that does not broadcast to the input, `AxisOutOfBounds` for an invalid
386    /// axis, [`Error::ContextMismatch`] for a tensor from another runtime, or a
387    /// backend error.
388    pub fn masked_softmax(
389        &mut self,
390        input: &EagerTensor,
391        mask: &EagerTensor,
392        axis: usize,
393    ) -> Result<EagerTensor> {
394        composite::masked_softmax(&mut self.composite(), input, mask, axis)
395    }
396
397    /// Log-softmax along `axis` over the entries where the `Bool` `mask` is true.
398    ///
399    /// Masked-out entries are `-inf` with a zero gradient; a slice with no unmasked
400    /// entry is all `-inf` with a zero gradient.
401    ///
402    /// # Examples
403    /// ```rust
404    /// use tenferro_ad::{EagerRuntime, Tensor};
405    /// let ctx = EagerRuntime::new()?;
406    /// let y = ctx.with_eager_session(|s| {
407    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0])?)?;
408    ///     let mask = s.constant_from_host(Tensor::from_vec_col_major(vec![2], vec![true, false])?)?;
409    ///     s.masked_log_softmax(&x, &mask, 0)
410    /// })?;
411    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.0, f64::NEG_INFINITY]);
412    /// # Ok::<(), tenferro_ad::Error>(())
413    /// ```
414    ///
415    /// # Errors
416    ///
417    /// Returns a typed `UnsupportedDType` error for non-real input, a
418    /// `DTypeMismatch` validation error for a non-`Bool` mask, `ShapeMismatch` for
419    /// a mask that does not broadcast to the input, `AxisOutOfBounds` for an invalid
420    /// axis, [`Error::ContextMismatch`] for a tensor from another runtime, or a
421    /// backend error.
422    pub fn masked_log_softmax(
423        &mut self,
424        input: &EagerTensor,
425        mask: &EagerTensor,
426        axis: usize,
427    ) -> Result<EagerTensor> {
428        composite::masked_log_softmax(&mut self.composite(), input, mask, axis)
429    }
430
431    /// Layer normalization along `axis` with optional affine `weight` / `bias`.
432    ///
433    /// `(x - mean) / sqrt(var + eps) * weight + bias` with the biased variance of the
434    /// centered values; `weight` and `bias` are rank-1 of length `shape[axis]`. A
435    /// zero-variance slice normalizes to `0` (then `bias`) with a finite gradient
436    /// when `eps > 0`. Real `F32`/`F64` only.
437    ///
438    /// # Examples
439    /// ```rust
440    /// use tenferro_ad::{EagerRuntime, Tensor};
441    /// let ctx = EagerRuntime::new()?;
442    /// let y = ctx.with_eager_session(|s| {
443    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 3.0])?)?;
444    ///     s.layer_norm(&x, 0, None, None, 0.0)
445    /// })?;
446    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[-1.0, 1.0]);
447    /// # Ok::<(), tenferro_ad::Error>(())
448    /// ```
449    ///
450    /// # Errors
451    ///
452    /// Returns a typed `UnsupportedDType` error for non-real input, an
453    /// `AxisOutOfBounds` validation error for an invalid axis, `InvalidArgument` for
454    /// a negative or non-finite `eps`, `DTypeMismatch` / `ShapeMismatch` for a weight
455    /// or bias that is not a same-dtype vector of the axis length,
456    /// [`Error::ContextMismatch`] for a tensor from another runtime, or a backend error.
457    pub fn layer_norm(
458        &mut self,
459        input: &EagerTensor,
460        axis: usize,
461        weight: Option<&EagerTensor>,
462        bias: Option<&EagerTensor>,
463        eps: f64,
464    ) -> Result<EagerTensor> {
465        composite::layer_norm(&mut self.composite(), input, axis, weight, bias, eps)
466    }
467
468    /// RMS normalization along `axis` with optional affine `weight` / `bias`.
469    ///
470    /// `x / sqrt(mean(x^2) + eps) * weight + bias`; `weight` and `bias` are rank-1 of
471    /// length `shape[axis]`. An all-zero slice normalizes to `0` (then `bias`) with
472    /// a finite gradient when `eps > 0`. Real `F32`/`F64` only.
473    ///
474    /// # Examples
475    /// ```rust
476    /// use tenferro_ad::{EagerRuntime, Tensor};
477    /// let ctx = EagerRuntime::new()?;
478    /// let y = ctx.with_eager_session(|s| {
479    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![2], vec![0.0_f64, 0.0])?)?;
480    ///     s.rms_norm(&x, 0, None, None, 1e-6)
481    /// })?;
482    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.0, 0.0]);
483    /// # Ok::<(), tenferro_ad::Error>(())
484    /// ```
485    ///
486    /// # Errors
487    ///
488    /// Returns a typed `UnsupportedDType` error for non-real input, an
489    /// `AxisOutOfBounds` validation error for an invalid axis, `InvalidArgument` for
490    /// a negative or non-finite `eps`, `DTypeMismatch` / `ShapeMismatch` for a weight
491    /// or bias that is not a same-dtype vector of the axis length,
492    /// [`Error::ContextMismatch`] for a tensor from another runtime, or a backend error.
493    pub fn rms_norm(
494        &mut self,
495        input: &EagerTensor,
496        axis: usize,
497        weight: Option<&EagerTensor>,
498        bias: Option<&EagerTensor>,
499        eps: f64,
500    ) -> Result<EagerTensor> {
501        composite::rms_norm(&mut self.composite(), input, axis, weight, bias, eps)
502    }
503
504    /// NumPy-style `take_along_axis` over `gather`.
505    ///
506    /// `out[.., i, ..] = input[.., indices[.., i, ..], ..]` along `axis`. `indices`
507    /// (I32/I64) has the input's rank; every other dimension is either the input's
508    /// extent (batch-varying indices) or `1` (the whole extent is taken). Indices
509    /// must be in bounds. The gradient flows to `input` only.
510    ///
511    /// # Examples
512    /// ```rust
513    /// use tenferro_ad::{EagerRuntime, Tensor};
514    /// let ctx = EagerRuntime::new()?;
515    /// // Per-batch row gather: out[i, j, b] = x[idx[i, b], j, b].
516    /// let y = ctx.with_eager_session(|s| {
517    ///     let x = s.constant_from_host(Tensor::from_vec_col_major(vec![2, 2, 2], (0..8).map(f64::from).collect::<Vec<_>>())?)?;
518    ///     let idx = s.constant_from_host(Tensor::from_vec_col_major(vec![2, 1, 2], vec![1_i64, 0, 0, 0])?)?;
519    ///     s.take_along_axis(&x, &idx, 0)
520    /// })?;
521    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[1.0, 0.0, 3.0, 2.0, 4.0, 4.0, 6.0, 6.0]);
522    /// # Ok::<(), tenferro_ad::Error>(())
523    /// ```
524    ///
525    /// # Errors
526    ///
527    /// Returns a `RankMismatch` / `ShapeMismatch` validation error for incompatible
528    /// index shapes, `AxisOutOfBounds` for an invalid axis, `InvalidArgument` when
529    /// taking from a zero-length axis, a typed `UnsupportedDType` error for a
530    /// non-integer index dtype, [`Error::ContextMismatch`] for a tensor from another
531    /// runtime, or a backend error.
532    pub fn take_along_axis(
533        &mut self,
534        input: &EagerTensor,
535        indices: &EagerTensor,
536        axis: usize,
537    ) -> Result<EagerTensor> {
538        composite::take_along_axis(&mut self.composite(), input, indices, axis)
539    }
540}