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}