Skip to main content

tenferro_runtime/
typed_session_ext.rs

1//! Backend-explicit session operations on [`TypedTensor`].
2
3use num_complex::Complex64;
4
5use crate::{BackendSession, CompareDir, DotGeneralConfig, TensorScalar, TypedTensor};
6
7/// AD-free tensor operations on [`TypedTensor`], run inside a borrowed backend session.
8///
9/// The methods mirror [`crate::TensorSessionOpsExt`] for a statically known
10/// scalar type, with the session last. Operations whose backend hooks take
11/// owned dtype-erased tensors (indexing, padding, concatenation, triangular
12/// and diagonal masks) are offered on [`crate::Tensor`] only: a typed form
13/// would have to copy the input first. Move a typed tensor into that surface
14/// with `Tensor::from_typed`, which does not copy. Selection with a bool mask
15/// is [`crate::TypedTensorMaskSessionOpsExt::where_select`].
16///
17/// # Examples
18///
19/// ```rust
20/// use tenferro_cpu::CpuBackend;
21/// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
22/// use tenferro_tensor::BackendSessionHost;
23///
24/// let mut backend = CpuBackend::new();
25/// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, -3.0])?;
26/// let y = backend.with_backend_session(|session| x.reduce_max(None, session))??;
27/// assert_eq!(y.host_data()?, &[1.0]);
28/// # Ok::<(), Box<dyn std::error::Error>>(())
29/// ```
30pub trait TypedTensorSessionOpsExt<T: TensorScalar> {
31    /// Elementwise addition with NumPy-style broadcasting inside a session.
32    ///
33    /// # Examples
34    ///
35    /// ```rust
36    /// use tenferro_cpu::CpuBackend;
37    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
38    /// use tenferro_tensor::BackendSessionHost;
39    ///
40    /// let mut backend = CpuBackend::new();
41    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 2.0]).unwrap();
42    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![3.0, 4.0]).unwrap();
43    /// let sum = backend.with_backend_session(|session| a.add(&b, session))??;
44    /// assert_eq!(sum.host_data().unwrap(), &[4.0, 6.0]);
45    /// # Ok::<(), Box<dyn std::error::Error>>(())
46    /// ```
47    ///
48    /// # Errors
49    ///
50    /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` for
51    /// incompatible operands, or [`tenferro_tensor::Error::BackendSource`] for
52    /// a typed backend failure.
53    fn add(
54        &self,
55        rhs: &TypedTensor<T>,
56        session: &mut dyn BackendSession,
57    ) -> tenferro_tensor::Result<TypedTensor<T>>;
58    /// Elementwise multiplication with NumPy-style broadcasting inside a session.
59    ///
60    /// # Examples
61    ///
62    /// ```rust
63    /// use tenferro_cpu::CpuBackend;
64    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
65    /// use tenferro_tensor::BackendSessionHost;
66    ///
67    /// let mut backend = CpuBackend::new();
68    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![1], vec![2.0]).unwrap();
69    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![4], vec![3.0; 4]).unwrap();
70    /// let product = backend.with_backend_session(|session| a.mul(&b, session))??;
71    /// assert_eq!(product.host_data().unwrap(), &[6.0; 4]);
72    /// # Ok::<(), Box<dyn std::error::Error>>(())
73    /// ```
74    ///
75    /// # Errors
76    ///
77    /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` for
78    /// incompatible operands, or [`tenferro_tensor::Error::BackendSource`] for
79    /// a typed backend failure.
80    fn mul(
81        &self,
82        rhs: &TypedTensor<T>,
83        session: &mut dyn BackendSession,
84    ) -> tenferro_tensor::Result<TypedTensor<T>>;
85    /// Elementwise exponential inside a session.
86    ///
87    /// # Examples
88    ///
89    /// ```rust
90    /// use tenferro_cpu::CpuBackend;
91    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
92    /// use tenferro_tensor::BackendSessionHost;
93    ///
94    /// let mut backend = CpuBackend::new();
95    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0, 1.0]).unwrap();
96    /// let y = backend.with_backend_session(|session| x.exp(session))??;
97    /// let y = y.host_data().unwrap();
98    /// assert!((y[0] - 1.0).abs() < 1.0e-12);
99    /// assert!((y[1] - std::f64::consts::E).abs() < 1.0e-12);
100    /// # Ok::<(), Box<dyn std::error::Error>>(())
101    /// ```
102    ///
103    /// # Errors
104    ///
105    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
106    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
107    /// failure.
108    fn exp(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
109    /// Sum over the selected axes inside a session. `None` reduces every
110    /// axis and `Some(&[])` keeps the input shape.
111    ///
112    /// # Examples
113    ///
114    /// ```rust
115    /// use tenferro_cpu::CpuBackend;
116    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
117    /// use tenferro_tensor::BackendSessionHost;
118    ///
119    /// let mut backend = CpuBackend::new();
120    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2, 3], vec![1.0; 6]).unwrap();
121    /// let sums = backend.with_backend_session(|session| x.reduce_sum(Some(&[1]), session))??;
122    /// assert_eq!(sums.host_data().unwrap(), &[3.0, 3.0]);
123    /// # Ok::<(), Box<dyn std::error::Error>>(())
124    /// ```
125    ///
126    /// # Errors
127    ///
128    /// Returns [`tenferro_tensor::Error::Validation`] with `AxisOutOfBounds`
129    /// for an axis outside the input rank or `DuplicateAxis` when `axes`
130    /// repeats an axis, or [`tenferro_tensor::Error::BackendSource`] for a
131    /// typed backend failure.
132    fn reduce_sum(
133        &self,
134        axes: Option<&[usize]>,
135        session: &mut dyn BackendSession,
136    ) -> tenferro_tensor::Result<TypedTensor<T>>;
137    /// Elementwise subtraction with NumPy-style broadcasting inside a session.
138    ///
139    /// # Examples
140    ///
141    /// ```rust
142    /// use tenferro_cpu::CpuBackend;
143    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
144    /// use tenferro_tensor::BackendSessionHost;
145    ///
146    /// let mut backend = CpuBackend::new();
147    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 4.0]).unwrap();
148    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 8.0]).unwrap();
149    /// let y = backend.with_backend_session(|session| a.sub(&b, session))??;
150    /// assert_eq!(y.host_data().unwrap(), &[1.0, -4.0]);
151    /// # Ok::<(), Box<dyn std::error::Error>>(())
152    /// ```
153    ///
154    /// # Errors
155    ///
156    /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` for
157    /// incompatible operands, or [`tenferro_tensor::Error::BackendSource`] for
158    /// a typed backend failure.
159    fn sub(
160        &self,
161        rhs: &TypedTensor<T>,
162        session: &mut dyn BackendSession,
163    ) -> tenferro_tensor::Result<TypedTensor<T>>;
164    /// Elementwise division with NumPy-style broadcasting inside a session.
165    ///
166    /// # Examples
167    ///
168    /// ```rust
169    /// use tenferro_cpu::CpuBackend;
170    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
171    /// use tenferro_tensor::BackendSessionHost;
172    ///
173    /// let mut backend = CpuBackend::new();
174    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![4.0, 8.0]).unwrap();
175    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 4.0]).unwrap();
176    /// let y = backend.with_backend_session(|session| a.div(&b, session))??;
177    /// assert_eq!(y.host_data().unwrap(), &[2.0, 2.0]);
178    /// # Ok::<(), Box<dyn std::error::Error>>(())
179    /// ```
180    ///
181    /// # Errors
182    ///
183    /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` for
184    /// incompatible shapes, a numerical [`tenferro_tensor::Error::Extension`]
185    /// for a detected zero divisor, or [`tenferro_tensor::Error::BackendSource`]
186    /// for a typed backend failure.
187    fn div(
188        &self,
189        rhs: &TypedTensor<T>,
190        session: &mut dyn BackendSession,
191    ) -> tenferro_tensor::Result<TypedTensor<T>>;
192    /// Elementwise remainder with NumPy-style broadcasting inside a session.
193    ///
194    /// # Examples
195    ///
196    /// ```rust
197    /// use tenferro_cpu::CpuBackend;
198    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
199    /// use tenferro_tensor::BackendSessionHost;
200    ///
201    /// let mut backend = CpuBackend::new();
202    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![5.0, 7.0]).unwrap();
203    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 4.0]).unwrap();
204    /// let y = backend.with_backend_session(|session| a.rem(&b, session))??;
205    /// assert_eq!(y.host_data().unwrap(), &[1.0, 3.0]);
206    /// # Ok::<(), Box<dyn std::error::Error>>(())
207    /// ```
208    ///
209    /// # Errors
210    ///
211    /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` for
212    /// incompatible shapes, a numerical [`tenferro_tensor::Error::Extension`]
213    /// for a detected zero divisor, or [`tenferro_tensor::Error::BackendSource`]
214    /// for a typed backend failure.
215    fn rem(
216        &self,
217        rhs: &TypedTensor<T>,
218        session: &mut dyn BackendSession,
219    ) -> tenferro_tensor::Result<TypedTensor<T>>;
220    /// Elementwise power with NumPy-style broadcasting inside a session.
221    ///
222    /// # Examples
223    ///
224    /// ```rust
225    /// use tenferro_cpu::CpuBackend;
226    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
227    /// use tenferro_tensor::BackendSessionHost;
228    ///
229    /// let mut backend = CpuBackend::new();
230    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 3.0]).unwrap();
231    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![3.0, 2.0]).unwrap();
232    /// let y = backend.with_backend_session(|session| a.pow(&b, session))??;
233    /// assert_eq!(y.host_data().unwrap(), &[8.0, 9.0]);
234    /// # Ok::<(), Box<dyn std::error::Error>>(())
235    /// ```
236    ///
237    /// # Errors
238    ///
239    /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` for
240    /// incompatible shapes, a numerical [`tenferro_tensor::Error::Extension`]
241    /// for a detected negative integer exponent, or
242    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
243    fn pow(
244        &self,
245        rhs: &TypedTensor<T>,
246        session: &mut dyn BackendSession,
247    ) -> tenferro_tensor::Result<TypedTensor<T>>;
248    /// Elementwise maximum with NumPy-style broadcasting inside a session.
249    ///
250    /// # Examples
251    ///
252    /// ```rust
253    /// use tenferro_cpu::CpuBackend;
254    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
255    /// use tenferro_tensor::BackendSessionHost;
256    ///
257    /// let mut backend = CpuBackend::new();
258    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 4.0]).unwrap();
259    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 8.0]).unwrap();
260    /// let y = backend.with_backend_session(|session| a.maximum(&b, session))??;
261    /// assert_eq!(y.host_data().unwrap(), &[2.0, 8.0]);
262    /// # Ok::<(), Box<dyn std::error::Error>>(())
263    /// ```
264    ///
265    /// # Errors
266    ///
267    /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` for
268    /// incompatible operands, or [`tenferro_tensor::Error::BackendSource`] for
269    /// a typed backend failure.
270    fn maximum(
271        &self,
272        rhs: &TypedTensor<T>,
273        session: &mut dyn BackendSession,
274    ) -> tenferro_tensor::Result<TypedTensor<T>>;
275    /// Elementwise minimum with NumPy-style broadcasting inside a session.
276    ///
277    /// # Examples
278    ///
279    /// ```rust
280    /// use tenferro_cpu::CpuBackend;
281    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
282    /// use tenferro_tensor::BackendSessionHost;
283    ///
284    /// let mut backend = CpuBackend::new();
285    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 4.0]).unwrap();
286    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 8.0]).unwrap();
287    /// let y = backend.with_backend_session(|session| a.minimum(&b, session))??;
288    /// assert_eq!(y.host_data().unwrap(), &[1.0, 4.0]);
289    /// # Ok::<(), Box<dyn std::error::Error>>(())
290    /// ```
291    ///
292    /// # Errors
293    ///
294    /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` for
295    /// incompatible operands, or [`tenferro_tensor::Error::BackendSource`] for
296    /// a typed backend failure.
297    fn minimum(
298        &self,
299        rhs: &TypedTensor<T>,
300        session: &mut dyn BackendSession,
301    ) -> tenferro_tensor::Result<TypedTensor<T>>;
302    /// Elementwise negation inside a session.
303    ///
304    /// # Examples
305    ///
306    /// ```rust
307    /// use tenferro_cpu::CpuBackend;
308    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
309    /// use tenferro_tensor::BackendSessionHost;
310    ///
311    /// let mut backend = CpuBackend::new();
312    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, -2.0]).unwrap();
313    /// let y = backend.with_backend_session(|session| x.neg(session))??;
314    /// assert_eq!(y.host_data().unwrap(), &[-1.0, 2.0]);
315    /// # Ok::<(), Box<dyn std::error::Error>>(())
316    /// ```
317    ///
318    /// # Errors
319    ///
320    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
321    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
322    /// failure.
323    fn neg(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
324    /// Elementwise absolute value inside a session.
325    ///
326    /// The result has the real counterpart dtype `T::Real`: complex magnitude
327    /// is real, and real or integer inputs keep their own dtype.
328    ///
329    /// # Examples
330    ///
331    /// ```rust
332    /// use num_complex::Complex64;
333    /// use tenferro_cpu::CpuBackend;
334    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
335    /// use tenferro_tensor::BackendSessionHost;
336    ///
337    /// let mut backend = CpuBackend::new();
338    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![-1.0, 2.0])?;
339    /// let y = backend.with_backend_session(|session| x.abs(session))??;
340    /// assert_eq!(y.host_data()?, &[1.0, 2.0]);
341    ///
342    /// let z = TypedTensor::<Complex64>::from_vec_col_major(vec![1], vec![Complex64::new(3.0, 4.0)])?;
343    /// let magnitude: TypedTensor<f64> = backend.with_backend_session(|session| z.abs(session))??;
344    /// assert_eq!(magnitude.host_data()?, &[5.0]);
345    /// # Ok::<(), Box<dyn std::error::Error>>(())
346    /// ```
347    ///
348    /// # Errors
349    ///
350    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
351    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
352    /// failure.
353    fn abs(
354        &self,
355        session: &mut dyn BackendSession,
356    ) -> tenferro_tensor::Result<TypedTensor<T::Real>>;
357    /// Elementwise sign inside a session.
358    ///
359    /// # Examples
360    ///
361    /// ```rust
362    /// use tenferro_cpu::CpuBackend;
363    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
364    /// use tenferro_tensor::BackendSessionHost;
365    ///
366    /// let mut backend = CpuBackend::new();
367    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, -2.0]).unwrap();
368    /// let y = backend.with_backend_session(|session| x.sign(session))??;
369    /// assert_eq!(y.host_data().unwrap(), &[1.0, -1.0]);
370    /// # Ok::<(), Box<dyn std::error::Error>>(())
371    /// ```
372    ///
373    /// # Errors
374    ///
375    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
376    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
377    /// failure.
378    fn sign(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
379    /// Elementwise complex conjugate inside a session.
380    ///
381    /// For real dtypes the conjugate is the identity.
382    ///
383    /// # Examples
384    ///
385    /// ```rust
386    /// use tenferro_cpu::CpuBackend;
387    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
388    /// use tenferro_tensor::BackendSessionHost;
389    ///
390    /// let mut backend = CpuBackend::new();
391    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, -2.0]).unwrap();
392    /// let y = backend.with_backend_session(|session| x.conj(session))??;
393    /// assert_eq!(y.host_data().unwrap(), &[1.0, -2.0]);
394    /// # Ok::<(), Box<dyn std::error::Error>>(())
395    /// ```
396    ///
397    /// # Errors
398    ///
399    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
400    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
401    /// failure.
402    fn conj(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
403    /// Elementwise natural logarithm inside a session.
404    ///
405    /// # Examples
406    ///
407    /// ```rust
408    /// use tenferro_cpu::CpuBackend;
409    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
410    /// use tenferro_tensor::BackendSessionHost;
411    ///
412    /// let mut backend = CpuBackend::new();
413    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, std::f64::consts::E]).unwrap();
414    /// let y = backend.with_backend_session(|session| x.log(session))??;
415    /// let y = y.host_data().unwrap();
416    /// assert!(y[0].abs() < 1.0e-12);
417    /// assert!((y[1] - 1.0).abs() < 1.0e-12);
418    /// # Ok::<(), Box<dyn std::error::Error>>(())
419    /// ```
420    ///
421    /// # Errors
422    ///
423    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
424    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
425    /// failure.
426    fn log(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
427    /// Elementwise `exp(x) - 1` inside a session.
428    ///
429    /// # Examples
430    ///
431    /// ```rust
432    /// use tenferro_cpu::CpuBackend;
433    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
434    /// use tenferro_tensor::BackendSessionHost;
435    ///
436    /// let mut backend = CpuBackend::new();
437    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0, 1.0]).unwrap();
438    /// let y = backend.with_backend_session(|session| x.expm1(session))??;
439    /// let y = y.host_data().unwrap();
440    /// assert!(y[0].abs() < 1.0e-12);
441    /// assert!((y[1] - (std::f64::consts::E - 1.0)).abs() < 1.0e-12);
442    /// # Ok::<(), Box<dyn std::error::Error>>(())
443    /// ```
444    ///
445    /// # Errors
446    ///
447    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
448    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
449    /// failure.
450    fn expm1(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
451    /// Elementwise `log(1 + x)` inside a session.
452    ///
453    /// # Examples
454    ///
455    /// ```rust
456    /// use tenferro_cpu::CpuBackend;
457    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
458    /// use tenferro_tensor::BackendSessionHost;
459    ///
460    /// let mut backend = CpuBackend::new();
461    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0, std::f64::consts::E - 1.0]).unwrap();
462    /// let y = backend.with_backend_session(|session| x.log1p(session))??;
463    /// let y = y.host_data().unwrap();
464    /// assert!(y[0].abs() < 1.0e-12);
465    /// assert!((y[1] - 1.0).abs() < 1.0e-12);
466    /// # Ok::<(), Box<dyn std::error::Error>>(())
467    /// ```
468    ///
469    /// # Errors
470    ///
471    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
472    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
473    /// failure.
474    fn log1p(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
475    /// Elementwise error function `erf(x)` inside a session, for real `f32`/`f64`.
476    ///
477    /// # Examples
478    ///
479    /// ```rust
480    /// use tenferro_cpu::CpuBackend;
481    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
482    /// use tenferro_tensor::BackendSessionHost;
483    ///
484    /// let mut backend = CpuBackend::new();
485    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0, 1.0]).unwrap();
486    /// let y = backend.with_backend_session(|session| x.erf(session))??;
487    /// let y = y.host_data().unwrap();
488    /// assert_eq!(y[0], 0.0);
489    /// assert!((y[1] - 0.842_700_792_949_714_9).abs() < 1.0e-15);
490    /// # Ok::<(), Box<dyn std::error::Error>>(())
491    /// ```
492    ///
493    /// # Errors
494    ///
495    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for a complex or
496    /// integer element type, or [`tenferro_tensor::Error::BackendSource`] for
497    /// a typed backend failure.
498    fn erf(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
499    /// Elementwise sine inside a session.
500    ///
501    /// # Examples
502    ///
503    /// ```rust
504    /// use tenferro_cpu::CpuBackend;
505    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
506    /// use tenferro_tensor::BackendSessionHost;
507    ///
508    /// let mut backend = CpuBackend::new();
509    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0, std::f64::consts::FRAC_PI_2]).unwrap();
510    /// let y = backend.with_backend_session(|session| x.sin(session))??;
511    /// let y = y.host_data().unwrap();
512    /// assert!(y[0].abs() < 1.0e-12);
513    /// assert!((y[1] - 1.0).abs() < 1.0e-12);
514    /// # Ok::<(), Box<dyn std::error::Error>>(())
515    /// ```
516    ///
517    /// # Errors
518    ///
519    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
520    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
521    /// failure.
522    fn sin(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
523    /// Elementwise cosine inside a session.
524    ///
525    /// # Examples
526    ///
527    /// ```rust
528    /// use tenferro_cpu::CpuBackend;
529    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
530    /// use tenferro_tensor::BackendSessionHost;
531    ///
532    /// let mut backend = CpuBackend::new();
533    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0, std::f64::consts::PI]).unwrap();
534    /// let y = backend.with_backend_session(|session| x.cos(session))??;
535    /// let y = y.host_data().unwrap();
536    /// assert!((y[0] - 1.0).abs() < 1.0e-12);
537    /// assert!((y[1] + 1.0).abs() < 1.0e-12);
538    /// # Ok::<(), Box<dyn std::error::Error>>(())
539    /// ```
540    ///
541    /// # Errors
542    ///
543    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
544    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
545    /// failure.
546    fn cos(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
547    /// Elementwise hyperbolic tangent inside a session.
548    ///
549    /// # Examples
550    ///
551    /// ```rust
552    /// use tenferro_cpu::CpuBackend;
553    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
554    /// use tenferro_tensor::BackendSessionHost;
555    ///
556    /// let mut backend = CpuBackend::new();
557    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0, 1.0]).unwrap();
558    /// let y = backend.with_backend_session(|session| x.tanh(session))??;
559    /// let y = y.host_data().unwrap();
560    /// assert!(y[0].abs() < 1.0e-12);
561    /// assert!((y[1] - 0.7615941559557649).abs() < 1.0e-12);
562    /// # Ok::<(), Box<dyn std::error::Error>>(())
563    /// ```
564    ///
565    /// # Errors
566    ///
567    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
568    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
569    /// failure.
570    fn tanh(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
571    /// Elementwise square root inside a session.
572    ///
573    /// # Examples
574    ///
575    /// ```rust
576    /// use tenferro_cpu::CpuBackend;
577    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
578    /// use tenferro_tensor::BackendSessionHost;
579    ///
580    /// let mut backend = CpuBackend::new();
581    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![4.0, 9.0]).unwrap();
582    /// let y = backend.with_backend_session(|session| x.sqrt(session))??;
583    /// assert_eq!(y.host_data().unwrap(), &[2.0, 3.0]);
584    /// # Ok::<(), Box<dyn std::error::Error>>(())
585    /// ```
586    ///
587    /// # Errors
588    ///
589    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
590    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
591    /// failure.
592    fn sqrt(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
593    /// Elementwise reciprocal square root inside a session.
594    ///
595    /// # Examples
596    ///
597    /// ```rust
598    /// use tenferro_cpu::CpuBackend;
599    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
600    /// use tenferro_tensor::BackendSessionHost;
601    ///
602    /// let mut backend = CpuBackend::new();
603    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![4.0, 1.0]).unwrap();
604    /// let y = backend.with_backend_session(|session| x.rsqrt(session))??;
605    /// let y = y.host_data().unwrap();
606    /// assert!((y[0] - 0.5).abs() < 1.0e-12);
607    /// assert!((y[1] - 1.0).abs() < 1.0e-12);
608    /// # Ok::<(), Box<dyn std::error::Error>>(())
609    /// ```
610    ///
611    /// # Errors
612    ///
613    /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
614    /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
615    /// failure.
616    fn rsqrt(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
617    /// Elementwise comparison with NumPy-style broadcasting inside a session.
618    ///
619    /// The result is a bool typed tensor.
620    ///
621    /// # Examples
622    ///
623    /// ```rust
624    /// use tenferro_cpu::CpuBackend;
625    /// use tenferro_runtime::{CompareDir, TypedTensor, TypedTensorSessionOpsExt};
626    /// use tenferro_tensor::BackendSessionHost;
627    ///
628    /// let mut backend = CpuBackend::new();
629    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 4.0]).unwrap();
630    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 8.0]).unwrap();
631    /// let y = backend.with_backend_session(|session| a.compare(&b, CompareDir::Gt, session))??;
632    /// assert_eq!(y.host_data().unwrap(), &[true, false]);
633    /// # Ok::<(), Box<dyn std::error::Error>>(())
634    /// ```
635    ///
636    /// # Errors
637    ///
638    /// Returns [`tenferro_tensor::Error::Validation`] with
639    /// `ShapeMismatch::IncompatibleShapes` when broadcasting the operands is
640    /// impossible, or [`tenferro_tensor::Error::BackendSource`] for a typed
641    /// backend failure.
642    fn compare(
643        &self,
644        rhs: &TypedTensor<T>,
645        dir: CompareDir,
646        session: &mut dyn BackendSession,
647    ) -> tenferro_tensor::Result<TypedTensor<bool>>;
648    /// Clamp values elementwise between lower and upper bounds inside a session.
649    ///
650    /// # Examples
651    ///
652    /// ```rust
653    /// use tenferro_cpu::CpuBackend;
654    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
655    /// use tenferro_tensor::BackendSessionHost;
656    ///
657    /// let mut backend = CpuBackend::new();
658    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![-2.0, 4.0]).unwrap();
659    /// let lower = TypedTensor::<f64>::from_vec_col_major(vec![], vec![0.0]).unwrap();
660    /// let upper = TypedTensor::<f64>::from_vec_col_major(vec![], vec![3.0]).unwrap();
661    /// let y = backend.with_backend_session(|session| x.clamp(&lower, &upper, session))??;
662    /// assert_eq!(y.host_data().unwrap(), &[0.0, 3.0]);
663    /// # Ok::<(), Box<dyn std::error::Error>>(())
664    /// ```
665    ///
666    /// # Errors
667    ///
668    /// Returns [`tenferro_tensor::Error::Validation`] with
669    /// `ShapeMismatch::IncompatibleShapes` when a bound cannot broadcast to
670    /// the input, or [`tenferro_tensor::Error::BackendSource`] for a typed
671    /// backend failure.
672    fn clamp(
673        &self,
674        lower: &TypedTensor<T>,
675        upper: &TypedTensor<T>,
676        session: &mut dyn BackendSession,
677    ) -> tenferro_tensor::Result<TypedTensor<T>>;
678    /// Rank-2 matrix multiplication inside a session.
679    ///
680    /// # Examples
681    ///
682    /// ```rust
683    /// use tenferro_cpu::CpuBackend;
684    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
685    /// use tenferro_tensor::BackendSessionHost;
686    ///
687    /// let mut backend = CpuBackend::new();
688    /// let a = TypedTensor::<f64>::from_vec_col_major(vec![2, 3], vec![1.0; 6]).unwrap();
689    /// let b = TypedTensor::<f64>::from_vec_col_major(vec![3, 2], vec![1.0; 6]).unwrap();
690    /// let c = backend.with_backend_session(|session| a.matmul(&b, session))??;
691    /// assert_eq!(c.shape(), &[2, 2]);
692    /// # Ok::<(), Box<dyn std::error::Error>>(())
693    /// ```
694    ///
695    /// # Errors
696    ///
697    /// Returns [`tenferro_tensor::Error::Validation`] with `RankMismatch` when
698    /// either operand is not rank two or `ShapeMismatch::ContractedDimensions`
699    /// when the inner dimensions differ, or
700    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
701    fn matmul(
702        &self,
703        rhs: &TypedTensor<T>,
704        session: &mut dyn BackendSession,
705    ) -> tenferro_tensor::Result<TypedTensor<T>>;
706    /// Reshape through the backend structural operation inside a session.
707    ///
708    /// # Examples
709    ///
710    /// ```rust
711    /// use tenferro_cpu::CpuBackend;
712    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
713    /// use tenferro_tensor::BackendSessionHost;
714    ///
715    /// let mut backend = CpuBackend::new();
716    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2, 3], vec![1.0; 6]).unwrap();
717    /// let y = backend.with_backend_session(|session| x.reshape(&[3, 2], session))??;
718    /// assert_eq!(y.shape(), &[3, 2]);
719    /// # Ok::<(), Box<dyn std::error::Error>>(())
720    /// ```
721    ///
722    /// # Errors
723    ///
724    /// Returns [`tenferro_tensor::Error::Validation`] with
725    /// `ShapeMismatch::ReshapeElementCount` when the element counts differ,
726    /// `IntegerOverflow` when shape arithmetic overflows, or
727    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
728    fn reshape(
729        &self,
730        shape: &[usize],
731        session: &mut dyn BackendSession,
732    ) -> tenferro_tensor::Result<TypedTensor<T>>;
733    /// Permute axes through the backend structural operation inside a session.
734    ///
735    /// # Examples
736    ///
737    /// ```rust
738    /// use tenferro_cpu::CpuBackend;
739    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
740    /// use tenferro_tensor::BackendSessionHost;
741    ///
742    /// let mut backend = CpuBackend::new();
743    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2, 3], vec![1.0; 6]).unwrap();
744    /// let y = backend.with_backend_session(|session| x.transpose(&[1, 0], session))??;
745    /// assert_eq!(y.shape(), &[3, 2]);
746    /// # Ok::<(), Box<dyn std::error::Error>>(())
747    /// ```
748    ///
749    /// # Errors
750    ///
751    /// Returns [`tenferro_tensor::Error::Validation`] with
752    /// `InvalidPermutationLength` when `perm` has the wrong length,
753    /// `AxisOutOfBounds` for an invalid axis, or `DuplicateAxis` for a
754    /// repeated axis, or [`tenferro_tensor::Error::BackendSource`] for a typed
755    /// backend failure.
756    fn transpose(
757        &self,
758        perm: &[usize],
759        session: &mut dyn BackendSession,
760    ) -> tenferro_tensor::Result<TypedTensor<T>>;
761    /// Broadcast into a larger shape inside a session.
762    ///
763    /// # Examples
764    ///
765    /// ```rust
766    /// use tenferro_cpu::CpuBackend;
767    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
768    /// use tenferro_tensor::BackendSessionHost;
769    ///
770    /// let mut backend = CpuBackend::new();
771    /// let row = TypedTensor::<f64>::from_vec_col_major(vec![3], vec![1.0, 2.0, 3.0]).unwrap();
772    /// let matrix = backend.with_backend_session(|session| row.broadcast_in_dim(&[2, 3], &[1], session))??;
773    /// assert_eq!(matrix.shape(), &[2, 3]);
774    /// # Ok::<(), Box<dyn std::error::Error>>(())
775    /// ```
776    ///
777    /// # Errors
778    ///
779    /// Returns [`tenferro_tensor::Error::Validation`] with `RankMismatch` when
780    /// `dims` does not match the input rank, `AxisOutOfBounds` or
781    /// `DuplicateAxis` for an invalid mapping, or
782    /// `ShapeMismatch::IncompatibleShapes` when known dimensions cannot
783    /// broadcast. [`tenferro_tensor::Error::BackendSource`] reports a typed
784    /// backend failure.
785    fn broadcast_in_dim(
786        &self,
787        shape: &[usize],
788        dims: &[usize],
789        session: &mut dyn BackendSession,
790    ) -> tenferro_tensor::Result<TypedTensor<T>>;
791    /// Take the maximum over the selected axes inside a session. `None` reduces every axis and `Some(&[])` keeps the input shape.
792    ///
793    /// # Examples
794    ///
795    /// ```rust
796    /// use tenferro_cpu::CpuBackend;
797    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
798    /// use tenferro_tensor::BackendSessionHost;
799    ///
800    /// let mut backend = CpuBackend::new();
801    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 5.0, 3.0, 2.0])?;
802    /// let y = backend.with_backend_session(|session| x.reduce_max(Some(&[0]), session))??;
803    /// assert_eq!(y.host_data()?, &[5.0, 3.0]);
804    /// # Ok::<(), Box<dyn std::error::Error>>(())
805    /// ```
806    ///
807    /// # Errors
808    ///
809    /// Returns [`tenferro_tensor::Error::Validation`] for an out-of-range or repeated axis, or
810    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
811    /// An unsupported dtype returns [`tenferro_tensor::Error::Unsupported`].
812    fn reduce_max(
813        &self,
814        axes: Option<&[usize]>,
815        session: &mut dyn BackendSession,
816    ) -> tenferro_tensor::Result<TypedTensor<T>>;
817    /// Take the minimum over the selected axes inside a session. `None` reduces every axis and `Some(&[])` keeps the input shape.
818    ///
819    /// # Examples
820    ///
821    /// ```rust
822    /// use tenferro_cpu::CpuBackend;
823    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
824    /// use tenferro_tensor::BackendSessionHost;
825    ///
826    /// let mut backend = CpuBackend::new();
827    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 5.0, 3.0, 2.0])?;
828    /// let y = backend.with_backend_session(|session| x.reduce_min(Some(&[0]), session))??;
829    /// assert_eq!(y.host_data()?, &[1.0, 2.0]);
830    /// # Ok::<(), Box<dyn std::error::Error>>(())
831    /// ```
832    ///
833    /// # Errors
834    ///
835    /// Returns [`tenferro_tensor::Error::Validation`] for an out-of-range or repeated axis, or
836    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
837    /// An unsupported dtype returns [`tenferro_tensor::Error::Unsupported`].
838    fn reduce_min(
839        &self,
840        axes: Option<&[usize]>,
841        session: &mut dyn BackendSession,
842    ) -> tenferro_tensor::Result<TypedTensor<T>>;
843    /// Multiply over the selected axes inside a session. `None` reduces every axis and `Some(&[])` keeps the input shape.
844    ///
845    /// # Examples
846    ///
847    /// ```rust
848    /// use tenferro_cpu::CpuBackend;
849    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
850    /// use tenferro_tensor::BackendSessionHost;
851    ///
852    /// let mut backend = CpuBackend::new();
853    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 5.0, 3.0, 2.0])?;
854    /// let y = backend.with_backend_session(|session| x.reduce_prod(Some(&[0]), session))??;
855    /// assert_eq!(y.host_data()?, &[5.0, 6.0]);
856    /// # Ok::<(), Box<dyn std::error::Error>>(())
857    /// ```
858    ///
859    /// # Errors
860    ///
861    /// Returns [`tenferro_tensor::Error::Validation`] for an out-of-range or repeated axis, or
862    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
863    /// An unsupported dtype returns [`tenferro_tensor::Error::Unsupported`].
864    fn reduce_prod(
865        &self,
866        axes: Option<&[usize]>,
867        session: &mut dyn BackendSession,
868    ) -> tenferro_tensor::Result<TypedTensor<T>>;
869    /// Sum elementwise squares over the selected axes inside a session (`f32`/`f64`). `None` reduces every axis and `Some(&[])` keeps the input shape.
870    ///
871    /// # Examples
872    ///
873    /// ```rust
874    /// use tenferro_cpu::CpuBackend;
875    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
876    /// use tenferro_tensor::BackendSessionHost;
877    ///
878    /// let mut backend = CpuBackend::new();
879    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 5.0, 3.0, 2.0])?;
880    /// let y = backend.with_backend_session(|session| x.reduce_sum_squares(Some(&[0]), session))??;
881    /// assert_eq!(y.host_data()?, &[26.0, 13.0]);
882    /// # Ok::<(), Box<dyn std::error::Error>>(())
883    /// ```
884    ///
885    /// # Errors
886    ///
887    /// Returns [`tenferro_tensor::Error::Validation`] for an out-of-range or repeated axis, or
888    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
889    /// An unsupported dtype returns [`tenferro_tensor::Error::Unsupported`].
890    fn reduce_sum_squares(
891        &self,
892        axes: Option<&[usize]>,
893        session: &mut dyn BackendSession,
894    ) -> tenferro_tensor::Result<TypedTensor<T>>;
895    /// Contract this tensor with `rhs` inside a session (StableHLO `dot_general`).
896    ///
897    /// The output layout is `[lhs free..., rhs free..., batch...]` (batch axes trail).
898    ///
899    /// # Examples
900    ///
901    /// ```rust
902    /// use tenferro_cpu::CpuBackend;
903    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
904    /// use tenferro_tensor::BackendSessionHost;
905    ///
906    /// let mut backend = CpuBackend::new();
907    /// use tenferro_runtime::DotGeneralConfig;
908    /// let lhs = TypedTensor::<f64>::from_vec_col_major(vec![1, 2], vec![2.0, 3.0])?;
909    /// let rhs = TypedTensor::<f64>::from_vec_col_major(vec![2, 1], vec![4.0, 5.0])?;
910    /// let config = DotGeneralConfig {
911    ///     lhs_contracting_dims: [1].as_slice().into(),
912    ///     rhs_contracting_dims: [0].as_slice().into(),
913    ///     lhs_batch_dims: [].as_slice().into(),
914    ///     rhs_batch_dims: [].as_slice().into(),
915    /// };
916    /// let y = backend.with_backend_session(|session| lhs.dot_general(&rhs, config, session))??;
917    /// assert_eq!(y.host_data()?, &[23.0]);
918    /// # Ok::<(), Box<dyn std::error::Error>>(())
919    /// ```
920    ///
921    /// # Errors
922    ///
923    /// Returns [`tenferro_tensor::Error::Validation`] for incompatible contraction or batch dimensions, or
924    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
925    fn dot_general(
926        &self,
927        rhs: &TypedTensor<T>,
928        config: DotGeneralConfig,
929        session: &mut dyn BackendSession,
930    ) -> tenferro_tensor::Result<TypedTensor<T>>;
931    /// Contract with optional conjugation of either operand inside a session.
932    ///
933    /// # Examples
934    ///
935    /// ```rust
936    /// use tenferro_cpu::CpuBackend;
937    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
938    /// use tenferro_tensor::BackendSessionHost;
939    ///
940    /// let mut backend = CpuBackend::new();
941    /// use num_complex::Complex64;
942    /// use tenferro_runtime::DotGeneralConfig;
943    /// let lhs = TypedTensor::<Complex64>::from_vec_col_major(vec![1, 1], vec![Complex64::new(0.0, 1.0)])?;
944    /// let rhs = TypedTensor::<Complex64>::from_vec_col_major(vec![1, 1], vec![Complex64::new(0.0, 1.0)])?;
945    /// let config = DotGeneralConfig {
946    ///     lhs_contracting_dims: [1].as_slice().into(),
947    ///     rhs_contracting_dims: [0].as_slice().into(),
948    ///     lhs_batch_dims: [].as_slice().into(),
949    ///     rhs_batch_dims: [].as_slice().into(),
950    /// };
951    /// let y = backend.with_backend_session(|session| lhs.dot_general_with_conj(&rhs, config, true, false, session))??;
952    /// assert_eq!(y.host_data()?, &[Complex64::new(1.0, 0.0)]);
953    /// # Ok::<(), Box<dyn std::error::Error>>(())
954    /// ```
955    ///
956    /// # Errors
957    ///
958    /// Returns [`tenferro_tensor::Error::Validation`] for incompatible contraction or batch dimensions, or
959    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
960    fn dot_general_with_conj(
961        &self,
962        rhs: &TypedTensor<T>,
963        config: DotGeneralConfig,
964        lhs_conj: bool,
965        rhs_conj: bool,
966        session: &mut dyn BackendSession,
967    ) -> tenferro_tensor::Result<TypedTensor<T>>;
968    /// Multiply by a real scalar inside a session, with the eager `scale_real` dtype rules.
969    ///
970    /// Integer dtypes round the factor.
971    ///
972    /// # Examples
973    ///
974    /// ```rust
975    /// use tenferro_cpu::CpuBackend;
976    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
977    /// use tenferro_tensor::BackendSessionHost;
978    ///
979    /// let mut backend = CpuBackend::new();
980    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 2.0])?;
981    /// let y = backend.with_backend_session(|session| x.scale_real(2.0, session))??;
982    /// assert_eq!(y.host_data()?, &[2.0, 4.0]);
983    /// # Ok::<(), Box<dyn std::error::Error>>(())
984    /// ```
985    ///
986    /// # Errors
987    ///
988    /// Returns [`tenferro_tensor::Error::Validation`] with `InvalidArgument` for a
989    /// non-finite factor or an integer factor out of range, or
990    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
991    fn scale_real(
992        &self,
993        factor: f64,
994        session: &mut dyn BackendSession,
995    ) -> tenferro_tensor::Result<TypedTensor<T>>;
996    /// Multiply a complex tensor by a complex scalar inside a session.
997    ///
998    /// # Examples
999    ///
1000    /// ```rust
1001    /// use tenferro_cpu::CpuBackend;
1002    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1003    /// use tenferro_tensor::BackendSessionHost;
1004    ///
1005    /// let mut backend = CpuBackend::new();
1006    /// use num_complex::Complex64;
1007    /// let x = TypedTensor::<Complex64>::from_vec_col_major(vec![1], vec![Complex64::new(1.0, 2.0)])?;
1008    /// let y = backend.with_backend_session(|session| x.scale_complex(Complex64::new(0.0, 1.0), session))??;
1009    /// assert_eq!(y.host_data()?, &[Complex64::new(-2.0, 1.0)]);
1010    /// # Ok::<(), Box<dyn std::error::Error>>(())
1011    /// ```
1012    ///
1013    /// # Errors
1014    ///
1015    /// Returns [`tenferro_tensor::Error::Validation`] with `InvalidArgument` when
1016    /// `T` is not complex, or [`tenferro_tensor::Error::BackendSource`] for a typed
1017    /// backend failure.
1018    fn scale_complex(
1019        &self,
1020        factor: Complex64,
1021        session: &mut dyn BackendSession,
1022    ) -> tenferro_tensor::Result<TypedTensor<T>>;
1023    /// Logistic sigmoid `1 / (1 + exp(-x))` inside a session, overflow-free.
1024    ///
1025    /// Evaluated as `1 / (1 + e)` for `x > 0` and `e / (1 + e)` otherwise, with
1026    /// `e = exp(-|x|)`. Real `F32`/`F64` only.
1027    ///
1028    /// # Examples
1029    ///
1030    /// ```rust
1031    /// use tenferro_cpu::CpuBackend;
1032    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1033    /// use tenferro_tensor::BackendSessionHost;
1034    ///
1035    /// let mut backend = CpuBackend::new();
1036    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![3], vec![-700.0, 0.0, 1000.0])?;
1037    /// let y = backend.with_backend_session(|session| x.sigmoid(session))??;
1038    /// let y = y.host_data()?;
1039    /// assert_eq!(y[1], 0.5);
1040    /// assert!(y[0] > 0.0 && y[0] < 1e-300);
1041    /// assert_eq!(y[2], 1.0);
1042    /// # Ok::<(), Box<dyn std::error::Error>>(())
1043    /// ```
1044    ///
1045    /// # Errors
1046    ///
1047    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1048    /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1049    /// for a typed backend failure.
1050    fn sigmoid(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
1051    /// SiLU (swish) `x * sigmoid(x)` inside a session.
1052    ///
1053    /// Real `F32`/`F64` only.
1054    ///
1055    /// # Examples
1056    ///
1057    /// ```rust
1058    /// use tenferro_cpu::CpuBackend;
1059    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1060    /// use tenferro_tensor::BackendSessionHost;
1061    ///
1062    /// let mut backend = CpuBackend::new();
1063    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![3], vec![-1.0, 0.0, 1.0])?;
1064    /// let y = backend.with_backend_session(|session| x.silu(session))??;
1065    /// let y = y.host_data()?;
1066    /// assert_eq!(y[1], 0.0);
1067    /// assert!((y[2] - 1.0 / (1.0 + (-1.0_f64).exp())).abs() < 1e-15);
1068    /// # Ok::<(), Box<dyn std::error::Error>>(())
1069    /// ```
1070    ///
1071    /// # Errors
1072    ///
1073    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1074    /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1075    /// for a typed backend failure.
1076    fn silu(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
1077    /// Softplus `log(1 + exp(x))` inside a session, in the stable form `max(x, 0) + log1p(exp(-|x|))`.
1078    ///
1079    /// Real `F32`/`F64` only; never overflows.
1080    ///
1081    /// # Examples
1082    ///
1083    /// ```rust
1084    /// use tenferro_cpu::CpuBackend;
1085    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1086    /// use tenferro_tensor::BackendSessionHost;
1087    ///
1088    /// let mut backend = CpuBackend::new();
1089    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![3], vec![-1000.0, 0.0, 1000.0])?;
1090    /// let y = backend.with_backend_session(|session| x.softplus(session))??;
1091    /// let y = y.host_data()?;
1092    /// assert_eq!(y[0], 0.0);
1093    /// assert!((y[1] - 2.0_f64.ln()).abs() < 1e-15);
1094    /// assert_eq!(y[2], 1000.0);
1095    /// # Ok::<(), Box<dyn std::error::Error>>(())
1096    /// ```
1097    ///
1098    /// # Errors
1099    ///
1100    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1101    /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1102    /// for a typed backend failure.
1103    fn softplus(&self, session: &mut dyn BackendSession)
1104        -> tenferro_tensor::Result<TypedTensor<T>>;
1105    /// Exact GELU `x/2 * (1 + erf(x / sqrt(2)))` inside a session.
1106    ///
1107    /// Real `F32`/`F64` only (PyTorch `approximate="none"`).
1108    ///
1109    /// # Examples
1110    ///
1111    /// ```rust
1112    /// use tenferro_cpu::CpuBackend;
1113    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1114    /// use tenferro_tensor::BackendSessionHost;
1115    ///
1116    /// let mut backend = CpuBackend::new();
1117    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![3], vec![-1.0, 0.0, 1.0])?;
1118    /// let y = backend.with_backend_session(|session| x.gelu(session))??;
1119    /// let y = y.host_data()?;
1120    /// assert_eq!(y[1], 0.0);
1121    /// assert!((y[2] - 0.841_344_746_068_542_9).abs() < 1e-15);
1122    /// # Ok::<(), Box<dyn std::error::Error>>(())
1123    /// ```
1124    ///
1125    /// # Errors
1126    ///
1127    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1128    /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1129    /// for a typed backend failure.
1130    fn gelu(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<TypedTensor<T>>;
1131    /// GELU tanh approximation inside a session (PyTorch `approximate="tanh"`).
1132    ///
1133    /// `x/2 * (1 + tanh(sqrt(2/pi) * (x + 0.044715 x^3)))`; real `F32`/`F64` only.
1134    ///
1135    /// # Examples
1136    ///
1137    /// ```rust
1138    /// use tenferro_cpu::CpuBackend;
1139    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1140    /// use tenferro_tensor::BackendSessionHost;
1141    ///
1142    /// let mut backend = CpuBackend::new();
1143    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![3], vec![-1.0, 0.0, 1.0])?;
1144    /// let y = backend.with_backend_session(|session| x.gelu_tanh(session))??;
1145    /// let y = y.host_data()?;
1146    /// assert_eq!(y[1], 0.0);
1147    /// assert!((y[2] - 0.841_191_990_608_276_8).abs() < 1e-12);
1148    /// # Ok::<(), Box<dyn std::error::Error>>(())
1149    /// ```
1150    ///
1151    /// # Errors
1152    ///
1153    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1154    /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1155    /// for a typed backend failure.
1156    fn gelu_tanh(
1157        &self,
1158        session: &mut dyn BackendSession,
1159    ) -> tenferro_tensor::Result<TypedTensor<T>>;
1160    /// Arithmetic mean over `axes` inside a session (`None` reduces every axis).
1161    ///
1162    /// Float and complex dtypes. The sum is divided by the element count; a mean
1163    /// over zero elements is `NaN`, and `Some(&[])` is the identity.
1164    ///
1165    /// # Examples
1166    ///
1167    /// ```rust
1168    /// use tenferro_cpu::CpuBackend;
1169    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1170    /// use tenferro_tensor::BackendSessionHost;
1171    ///
1172    /// let mut backend = CpuBackend::new();
1173    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 2.0, 3.0, 4.0])?;
1174    /// let y = backend.with_backend_session(|session| x.reduce_mean(None, session))??;
1175    /// assert_eq!(y.host_data()?, &[2.5]);
1176    /// # Ok::<(), Box<dyn std::error::Error>>(())
1177    /// ```
1178    ///
1179    /// # Errors
1180    ///
1181    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for integer or
1182    /// `Bool` input, [`tenferro_tensor::Error::Validation`] with
1183    /// `AxisOutOfBounds` or `DuplicateAxis` for invalid axes, or
1184    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1185    fn reduce_mean(
1186        &self,
1187        axes: Option<&[usize]>,
1188        session: &mut dyn BackendSession,
1189    ) -> tenferro_tensor::Result<TypedTensor<T>>;
1190    /// Max-subtracted softmax along `axis` inside a session.
1191    ///
1192    /// Real `F32`/`F64` only. A slice that is entirely `-inf` returns zeros
1193    /// instead of `NaN`; a `NaN` or `+inf` entry makes its slice `NaN`.
1194    ///
1195    /// # Examples
1196    ///
1197    /// ```rust
1198    /// use tenferro_cpu::CpuBackend;
1199    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1200    /// use tenferro_tensor::BackendSessionHost;
1201    ///
1202    /// let mut backend = CpuBackend::new();
1203    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 1.0])?;
1204    /// let y = backend.with_backend_session(|session| x.softmax(0, session))??;
1205    /// assert_eq!(y.host_data()?, &[0.5, 0.5]);
1206    /// # Ok::<(), Box<dyn std::error::Error>>(())
1207    /// ```
1208    ///
1209    /// # Errors
1210    ///
1211    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1212    /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1213    /// `AxisOutOfBounds` for an invalid axis, or
1214    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1215    fn softmax(
1216        &self,
1217        axis: usize,
1218        session: &mut dyn BackendSession,
1219    ) -> tenferro_tensor::Result<TypedTensor<T>>;
1220    /// Max-subtracted log-softmax along `axis` inside a session.
1221    ///
1222    /// Real `F32`/`F64` only. A slice that is entirely `-inf` returns `-inf`
1223    /// instead of `NaN`.
1224    ///
1225    /// # Examples
1226    ///
1227    /// ```rust
1228    /// use tenferro_cpu::CpuBackend;
1229    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1230    /// use tenferro_tensor::BackendSessionHost;
1231    ///
1232    /// let mut backend = CpuBackend::new();
1233    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 1.0])?;
1234    /// let y = backend.with_backend_session(|session| x.log_softmax(0, session))??;
1235    /// assert_eq!(y.host_data()?, &[-std::f64::consts::LN_2; 2]);
1236    /// # Ok::<(), Box<dyn std::error::Error>>(())
1237    /// ```
1238    ///
1239    /// # Errors
1240    ///
1241    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1242    /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1243    /// `AxisOutOfBounds` for an invalid axis, or
1244    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1245    fn log_softmax(
1246        &self,
1247        axis: usize,
1248        session: &mut dyn BackendSession,
1249    ) -> tenferro_tensor::Result<TypedTensor<T>>;
1250    /// Softmax along `axis` over the entries where the `Bool` `mask` is true.
1251    ///
1252    /// `mask` broadcasts to the input shape. Masked-out entries are `0` whatever
1253    /// their value; a slice with no unmasked entry is all zeros.
1254    ///
1255    /// # Examples
1256    ///
1257    /// ```rust
1258    /// use tenferro_cpu::CpuBackend;
1259    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1260    /// use tenferro_tensor::BackendSessionHost;
1261    ///
1262    /// let mut backend = CpuBackend::new();
1263    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![3.0, 4.0])?;
1264    /// let mask = TypedTensor::<bool>::from_vec_col_major(vec![2], vec![false, false])?;
1265    /// let y = backend.with_backend_session(|session| x.masked_softmax(&mask, 0, session))??;
1266    /// assert_eq!(y.host_data()?, &[0.0, 0.0]);
1267    /// # Ok::<(), Box<dyn std::error::Error>>(())
1268    /// ```
1269    ///
1270    /// # Errors
1271    ///
1272    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1273    /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1274    /// `DTypeMismatch` for a non-`Bool` mask, `ShapeMismatch` for a mask that
1275    /// does not broadcast to the input, or `AxisOutOfBounds` for an invalid
1276    /// axis, or [`tenferro_tensor::Error::BackendSource`] for a typed backend
1277    /// failure.
1278    fn masked_softmax(
1279        &self,
1280        mask: &TypedTensor<bool>,
1281        axis: usize,
1282        session: &mut dyn BackendSession,
1283    ) -> tenferro_tensor::Result<TypedTensor<T>>;
1284    /// Log-softmax along `axis` over the entries where the `Bool` `mask` is true.
1285    ///
1286    /// Masked-out entries are `-inf`; a slice with no unmasked entry is all `-inf`.
1287    ///
1288    /// # Examples
1289    ///
1290    /// ```rust
1291    /// use tenferro_cpu::CpuBackend;
1292    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1293    /// use tenferro_tensor::BackendSessionHost;
1294    ///
1295    /// let mut backend = CpuBackend::new();
1296    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![3.0, 4.0])?;
1297    /// let mask = TypedTensor::<bool>::from_vec_col_major(vec![2], vec![true, false])?;
1298    /// let y = backend.with_backend_session(|session| x.masked_log_softmax(&mask, 0, session))??;
1299    /// assert_eq!(y.host_data()?, &[0.0, f64::NEG_INFINITY]);
1300    /// # Ok::<(), Box<dyn std::error::Error>>(())
1301    /// ```
1302    ///
1303    /// # Errors
1304    ///
1305    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1306    /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1307    /// `DTypeMismatch` for a non-`Bool` mask, `ShapeMismatch` for a mask that
1308    /// does not broadcast to the input, or `AxisOutOfBounds` for an invalid
1309    /// axis, or [`tenferro_tensor::Error::BackendSource`] for a typed backend
1310    /// failure.
1311    fn masked_log_softmax(
1312        &self,
1313        mask: &TypedTensor<bool>,
1314        axis: usize,
1315        session: &mut dyn BackendSession,
1316    ) -> tenferro_tensor::Result<TypedTensor<T>>;
1317    /// Layer normalization along `axis` with optional affine `weight` / `bias`, inside a session.
1318    ///
1319    /// `(x - mean) / sqrt(var + eps) * weight + bias` with the biased variance;
1320    /// `weight` and `bias` are rank-1 of length `shape[axis]`. Real `F32`/`F64` only.
1321    ///
1322    /// # Examples
1323    ///
1324    /// ```rust
1325    /// use tenferro_cpu::CpuBackend;
1326    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1327    /// use tenferro_tensor::BackendSessionHost;
1328    ///
1329    /// let mut backend = CpuBackend::new();
1330    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 3.0])?;
1331    /// let y = backend.with_backend_session(|session| x.layer_norm(0, None, None, 0.0, session))??;
1332    /// assert_eq!(y.host_data()?, &[-1.0, 1.0]);
1333    /// # Ok::<(), Box<dyn std::error::Error>>(())
1334    /// ```
1335    ///
1336    /// # Errors
1337    ///
1338    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1339    /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1340    /// `AxisOutOfBounds` for an invalid axis, `InvalidArgument` for a negative
1341    /// or non-finite `eps`, or `DTypeMismatch` / `ShapeMismatch` for a weight or
1342    /// bias that is not a same-dtype vector of the axis length, or
1343    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1344    fn layer_norm(
1345        &self,
1346        axis: usize,
1347        weight: Option<&TypedTensor<T>>,
1348        bias: Option<&TypedTensor<T>>,
1349        eps: f64,
1350        session: &mut dyn BackendSession,
1351    ) -> tenferro_tensor::Result<TypedTensor<T>>;
1352    /// RMS normalization along `axis` with optional affine `weight` / `bias`, inside a session.
1353    ///
1354    /// `x / sqrt(mean(x^2) + eps) * weight + bias`; `weight` and `bias` are rank-1
1355    /// of length `shape[axis]`. Real `F32`/`F64` only.
1356    ///
1357    /// # Examples
1358    ///
1359    /// ```rust
1360    /// use tenferro_cpu::CpuBackend;
1361    /// use tenferro_runtime::{TypedTensor, TypedTensorSessionOpsExt};
1362    /// use tenferro_tensor::BackendSessionHost;
1363    ///
1364    /// let mut backend = CpuBackend::new();
1365    /// let x = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0, 0.0])?;
1366    /// let y = backend.with_backend_session(|session| x.rms_norm(0, None, None, 1e-6, session))??;
1367    /// assert_eq!(y.host_data()?, &[0.0, 0.0]);
1368    /// # Ok::<(), Box<dyn std::error::Error>>(())
1369    /// ```
1370    ///
1371    /// # Errors
1372    ///
1373    /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1374    /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1375    /// `AxisOutOfBounds` for an invalid axis, `InvalidArgument` for a negative
1376    /// or non-finite `eps`, or `DTypeMismatch` / `ShapeMismatch` for a weight or
1377    /// bias that is not a same-dtype vector of the axis length, or
1378    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1379    fn rms_norm(
1380        &self,
1381        axis: usize,
1382        weight: Option<&TypedTensor<T>>,
1383        bias: Option<&TypedTensor<T>>,
1384        eps: f64,
1385        session: &mut dyn BackendSession,
1386    ) -> tenferro_tensor::Result<TypedTensor<T>>;
1387}
1388
1389/// Backend-explicit bool-mask session operations for typed tensors.
1390///
1391/// This trait keeps `where_select` available as a method on bool
1392/// `TypedTensor`s while preserving the crate-root extension-trait surface. It
1393/// is public because downstream users call it directly; the implementation
1394/// helper in the private `typed_tensor` module is not a compatibility API.
1395///
1396/// # Examples
1397///
1398/// ```rust
1399/// use tenferro_cpu::CpuBackend;
1400/// use tenferro_runtime::{TypedTensor, TypedTensorMaskSessionOpsExt};
1401/// use tenferro_tensor::BackendSessionHost;
1402///
1403/// let mut backend = CpuBackend::new();
1404/// let condition =
1405///     TypedTensor::<bool>::from_vec_col_major(vec![2], vec![true, false]).unwrap();
1406/// let on_true = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 2.0]).unwrap();
1407/// let on_false = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![3.0, 4.0]).unwrap();
1408/// let selected = backend
1409///     .with_backend_session(|session| condition.where_select(&on_true, &on_false, session))?
1410///     .unwrap();
1411/// assert_eq!(selected.host_data().unwrap(), &[1.0, 4.0]);
1412/// # Ok::<(), Box<dyn std::error::Error>>(())
1413/// ```
1414pub trait TypedTensorMaskSessionOpsExt {
1415    /// Select typed values using this bool tensor as condition.
1416    ///
1417    /// The condition broadcasts against both branches (NumPy rules).
1418    ///
1419    /// # Examples
1420    ///
1421    /// ```rust
1422    /// use tenferro_cpu::CpuBackend;
1423    /// use tenferro_runtime::{TypedTensor, TypedTensorMaskSessionOpsExt};
1424    /// use tenferro_tensor::BackendSessionHost;
1425    ///
1426    /// let mut backend = CpuBackend::new();
1427    /// let mask = TypedTensor::<bool>::from_vec_col_major(vec![], vec![false])?;
1428    /// let x = TypedTensor::<i64>::from_vec_col_major(vec![2], vec![1, 2])?;
1429    /// let y = TypedTensor::<i64>::from_vec_col_major(vec![2], vec![3, 4])?;
1430    /// let picked = backend.with_backend_session(|session| mask.where_select(&x, &y, session))??;
1431    /// assert_eq!(picked.host_data()?, &[3, 4]);
1432    /// # Ok::<(), Box<dyn std::error::Error>>(())
1433    /// ```
1434    ///
1435    /// # Errors
1436    ///
1437    /// Returns [`tenferro_tensor::Error::Validation`] with
1438    /// `ShapeMismatch::IncompatibleShapes` when the condition or either branch
1439    /// cannot broadcast to the other operands, or
1440    /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1441    fn where_select<U: TensorScalar>(
1442        &self,
1443        on_true: &TypedTensor<U>,
1444        on_false: &TypedTensor<U>,
1445        session: &mut dyn BackendSession,
1446    ) -> tenferro_tensor::Result<TypedTensor<U>>;
1447}