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}