tenferro_runtime/session_ext.rs
1//! Backend-explicit session operations on [`Tensor`].
2
3use num_complex::Complex64;
4
5use crate::{
6 BackendSession, CompareDir, DType, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig,
7 SliceConfig, Tensor,
8};
9
10/// AD-free tensor operations on [`Tensor`], run inside a borrowed backend session.
11///
12/// Every method takes the receiver first and the `session` last, with the
13/// arguments and config types of the eager `EagerSession` method of the same
14/// name, so AD-free and eager code differ only in where the session comes
15/// from. Enter the session with `BackendSessionHost::with_backend_session`;
16/// it returns `Result<R, SessionEntryError>` around the operation's own
17/// `Result`, so a single call is written `??`. Group several operations in one
18/// session instead of entering one per operation.
19///
20/// # Examples
21///
22/// ```rust
23/// use tenferro_cpu::CpuBackend;
24/// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
25/// use tenferro_tensor::BackendSessionHost;
26///
27/// let mut backend = CpuBackend::new();
28/// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?;
29/// let y = backend.with_backend_session(|session| {
30/// let lower = x.tril(0, session)?;
31/// lower.reduce_sum(None, session)
32/// })??;
33/// assert_eq!(y.as_slice::<f64>()?, &[7.0]);
34/// # Ok::<(), Box<dyn std::error::Error>>(())
35/// ```
36pub trait TensorSessionOpsExt {
37 /// Elementwise addition with NumPy-style broadcasting inside a session.
38 ///
39 /// The broadcast (reshape + `broadcast_in_dim`, or a copy when shapes
40 /// already match) and the add itself all run in the caller's `session`;
41 /// this op never enters a session of its own.
42 ///
43 /// # Examples
44 ///
45 /// ```rust
46 /// use tenferro_cpu::CpuBackend;
47 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
48 /// use tenferro_tensor::BackendSessionHost;
49 ///
50 /// let mut backend = CpuBackend::new();
51 /// let a = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
52 /// let b = Tensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0]).unwrap();
53 /// let sum = backend.with_backend_session(|session| a.add(&b, session))??;
54 /// assert_eq!(sum.as_slice::<f64>().unwrap(), &[4.0, 6.0]);
55 /// # Ok::<(), Box<dyn std::error::Error>>(())
56 /// ```
57 ///
58 /// # Errors
59 ///
60 /// Returns [`tenferro_tensor::Error::Validation`] with a
61 /// [`ShapeMismatch`](tenferro_tensor::ValidationError::ShapeMismatch) or
62 /// `DTypeMismatch` payload when operands are incompatible, or
63 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
64 fn add(
65 &self,
66 rhs: &Tensor,
67 session: &mut dyn BackendSession,
68 ) -> tenferro_tensor::Result<Tensor>;
69 /// Elementwise multiplication with NumPy-style broadcasting inside a session.
70 ///
71 /// Like [`Self::add`], broadcast and multiply run in the one `session`.
72 ///
73 /// # Examples
74 ///
75 /// ```rust
76 /// use tenferro_cpu::CpuBackend;
77 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
78 /// use tenferro_tensor::BackendSessionHost;
79 ///
80 /// let mut backend = CpuBackend::new();
81 /// let a = Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap();
82 /// let b = Tensor::from_vec_col_major(vec![4], vec![3.0_f64; 4]).unwrap();
83 /// let product = backend.with_backend_session(|session| a.mul(&b, session))??;
84 /// assert_eq!(product.as_slice::<f64>().unwrap(), &[6.0; 4]);
85 /// # Ok::<(), Box<dyn std::error::Error>>(())
86 /// ```
87 ///
88 /// # Errors
89 ///
90 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
91 /// `DTypeMismatch` for incompatible operands, or
92 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
93 fn mul(
94 &self,
95 rhs: &Tensor,
96 session: &mut dyn BackendSession,
97 ) -> tenferro_tensor::Result<Tensor>;
98 /// Elementwise exponential inside a session.
99 ///
100 /// # Examples
101 ///
102 /// ```rust
103 /// use tenferro_cpu::CpuBackend;
104 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
105 /// use tenferro_tensor::BackendSessionHost;
106 ///
107 /// let mut backend = CpuBackend::new();
108 /// let x = Tensor::from_vec_col_major(vec![2], vec![0.0_f64, 1.0]).unwrap();
109 /// let y = backend.with_backend_session(|session| x.exp(session))??;
110 /// let y = y.as_slice::<f64>().unwrap();
111 /// assert!((y[0] - 1.0).abs() < 1.0e-12);
112 /// assert!((y[1] - std::f64::consts::E).abs() < 1.0e-12);
113 /// # Ok::<(), Box<dyn std::error::Error>>(())
114 /// ```
115 ///
116 /// # Errors
117 ///
118 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
119 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
120 /// failure.
121 fn exp(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
122 /// Sum over the selected axes inside a session. `None` reduces every
123 /// axis and `Some(&[])` keeps the input shape, as in the eager and traced
124 /// reduction family.
125 ///
126 /// # Examples
127 ///
128 /// ```rust
129 /// use tenferro_cpu::CpuBackend;
130 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
131 /// use tenferro_tensor::BackendSessionHost;
132 ///
133 /// let mut backend = CpuBackend::new();
134 /// let x = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
135 /// let sums = backend.with_backend_session(|session| x.reduce_sum(Some(&[1]), session))??;
136 /// assert_eq!(sums.as_slice::<f64>().unwrap(), &[3.0, 3.0]);
137 /// let total = backend.with_backend_session(|session| x.reduce_sum(None, session))??;
138 /// assert_eq!(total.as_slice::<f64>().unwrap(), &[6.0]);
139 /// # Ok::<(), Box<dyn std::error::Error>>(())
140 /// ```
141 ///
142 /// # Errors
143 ///
144 /// Returns [`tenferro_tensor::Error::Validation`] with `AxisOutOfBounds`
145 /// or `DuplicateAxis` for invalid reductions, or
146 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
147 fn reduce_sum(
148 &self,
149 axes: Option<&[usize]>,
150 session: &mut dyn BackendSession,
151 ) -> tenferro_tensor::Result<Tensor>;
152 /// Convert to a different dtype using the checked conversion lattice inside a session.
153 ///
154 /// # Examples
155 ///
156 /// ```rust
157 /// use tenferro_cpu::CpuBackend;
158 /// use tenferro_runtime::{DType, Tensor, TensorSessionOpsExt};
159 /// use tenferro_tensor::BackendSessionHost;
160 ///
161 /// let mut backend = CpuBackend::new();
162 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
163 /// let y = backend.with_backend_session(|session| x.convert(DType::C64, session))??;
164 /// assert_eq!(y.dtype(), DType::C64);
165 /// # Ok::<(), Box<dyn std::error::Error>>(())
166 /// ```
167 ///
168 /// # Errors
169 ///
170 /// Returns [`tenferro_tensor::Error::UnsupportedDTypeConversion`] when the
171 /// conversion is outside the checked lattice,
172 /// [`tenferro_tensor::Error::Validation`] with `DTypeMismatch` or
173 /// `InvalidArgument` for invalid tensor metadata, or
174 /// [`tenferro_tensor::Error::BackendSource`] when the backend reports a
175 /// typed failure.
176 fn convert(
177 &self,
178 to: DType,
179 session: &mut dyn BackendSession,
180 ) -> tenferro_tensor::Result<Tensor>;
181 /// Cast to a different dtype using explicit lossy projection inside a session.
182 ///
183 /// # Examples
184 ///
185 /// ```rust
186 /// use tenferro_cpu::CpuBackend;
187 /// use tenferro_runtime::{DType, Tensor, TensorSessionOpsExt};
188 /// use tenferro_tensor::BackendSessionHost;
189 ///
190 /// let mut backend = CpuBackend::new();
191 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.2_f64, -2.8]).unwrap();
192 /// let y = backend.with_backend_session(|session| x.cast(DType::I32, session))??;
193 /// assert_eq!(y.as_slice::<i32>().unwrap(), &[1, -2]);
194 /// # Ok::<(), Box<dyn std::error::Error>>(())
195 /// ```
196 ///
197 /// # Errors
198 ///
199 /// Returns [`tenferro_tensor::Error::UnsupportedDTypeConversion`] when the
200 /// requested cast is unsupported, [`tenferro_tensor::Error::Validation`]
201 /// with `DTypeMismatch` or `InvalidArgument` for invalid tensor metadata,
202 /// or [`tenferro_tensor::Error::BackendSource`] for a typed backend
203 /// failure.
204 fn cast(&self, to: DType, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
205 /// Elementwise subtraction with NumPy-style broadcasting inside a session.
206 ///
207 /// Like [`Self::add`], the broadcast and the subtraction run in the one
208 /// `session`.
209 ///
210 /// # Examples
211 ///
212 /// ```rust
213 /// use tenferro_cpu::CpuBackend;
214 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
215 /// use tenferro_tensor::BackendSessionHost;
216 ///
217 /// let mut backend = CpuBackend::new();
218 /// let a = Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 4.0]).unwrap();
219 /// let b = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 8.0]).unwrap();
220 /// let y = backend.with_backend_session(|session| a.sub(&b, session))??;
221 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[1.0, -4.0]);
222 /// # Ok::<(), Box<dyn std::error::Error>>(())
223 /// ```
224 ///
225 /// # Errors
226 ///
227 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
228 /// `DTypeMismatch` for incompatible operands, or
229 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
230 fn sub(
231 &self,
232 rhs: &Tensor,
233 session: &mut dyn BackendSession,
234 ) -> tenferro_tensor::Result<Tensor>;
235 /// Elementwise division with NumPy-style broadcasting inside a session.
236 ///
237 /// # Examples
238 ///
239 /// ```rust
240 /// use tenferro_cpu::CpuBackend;
241 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
242 /// use tenferro_tensor::BackendSessionHost;
243 ///
244 /// let mut backend = CpuBackend::new();
245 /// let a = Tensor::from_vec_col_major(vec![2], vec![4.0_f64, 8.0]).unwrap();
246 /// let b = Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 4.0]).unwrap();
247 /// let y = backend.with_backend_session(|session| a.div(&b, session))??;
248 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[2.0, 2.0]);
249 /// # Ok::<(), Box<dyn std::error::Error>>(())
250 /// ```
251 ///
252 /// # Errors
253 ///
254 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
255 /// `DTypeMismatch` for shape/dtype incompatibility,
256 /// [`tenferro_tensor::Error::Extension`] with a numerical classification
257 /// for a detected zero divisor, or
258 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
259 fn div(
260 &self,
261 rhs: &Tensor,
262 session: &mut dyn BackendSession,
263 ) -> tenferro_tensor::Result<Tensor>;
264 /// Elementwise remainder with NumPy-style broadcasting inside a session.
265 ///
266 /// # Examples
267 ///
268 /// ```rust
269 /// use tenferro_cpu::CpuBackend;
270 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
271 /// use tenferro_tensor::BackendSessionHost;
272 ///
273 /// let mut backend = CpuBackend::new();
274 /// let a = Tensor::from_vec_col_major(vec![2], vec![5.0_f64, 7.0]).unwrap();
275 /// let b = Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 4.0]).unwrap();
276 /// let y = backend.with_backend_session(|session| a.rem(&b, session))??;
277 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[1.0, 3.0]);
278 /// # Ok::<(), Box<dyn std::error::Error>>(())
279 /// ```
280 ///
281 /// # Errors
282 ///
283 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
284 /// `DTypeMismatch` for shape/dtype incompatibility, a numerical
285 /// [`tenferro_tensor::Error::Extension`] for a detected zero divisor, or
286 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
287 fn rem(
288 &self,
289 rhs: &Tensor,
290 session: &mut dyn BackendSession,
291 ) -> tenferro_tensor::Result<Tensor>;
292 /// Elementwise power with NumPy-style broadcasting inside a session.
293 ///
294 /// # Examples
295 ///
296 /// ```rust
297 /// use tenferro_cpu::CpuBackend;
298 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
299 /// use tenferro_tensor::BackendSessionHost;
300 ///
301 /// let mut backend = CpuBackend::new();
302 /// let a = Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0]).unwrap();
303 /// let b = Tensor::from_vec_col_major(vec![2], vec![3.0_f64, 2.0]).unwrap();
304 /// let y = backend.with_backend_session(|session| a.pow(&b, session))??;
305 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[8.0, 9.0]);
306 /// # Ok::<(), Box<dyn std::error::Error>>(())
307 /// ```
308 ///
309 /// # Errors
310 ///
311 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
312 /// `DTypeMismatch` for incompatible metadata, a numerical
313 /// [`tenferro_tensor::Error::Extension`] for a detected negative integer
314 /// exponent, or [`tenferro_tensor::Error::BackendSource`] for a typed
315 /// backend failure.
316 fn pow(
317 &self,
318 rhs: &Tensor,
319 session: &mut dyn BackendSession,
320 ) -> tenferro_tensor::Result<Tensor>;
321 /// Elementwise maximum with NumPy-style broadcasting inside a session.
322 ///
323 /// # Examples
324 ///
325 /// ```rust
326 /// use tenferro_cpu::CpuBackend;
327 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
328 /// use tenferro_tensor::BackendSessionHost;
329 ///
330 /// let mut backend = CpuBackend::new();
331 /// let a = Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 4.0]).unwrap();
332 /// let b = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 8.0]).unwrap();
333 /// let y = backend.with_backend_session(|session| a.maximum(&b, session))??;
334 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[2.0, 8.0]);
335 /// # Ok::<(), Box<dyn std::error::Error>>(())
336 /// ```
337 ///
338 /// # Errors
339 ///
340 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
341 /// `DTypeMismatch` for incompatible operands, or
342 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
343 fn maximum(
344 &self,
345 rhs: &Tensor,
346 session: &mut dyn BackendSession,
347 ) -> tenferro_tensor::Result<Tensor>;
348 /// Elementwise minimum with NumPy-style broadcasting inside a session.
349 ///
350 /// # Examples
351 ///
352 /// ```rust
353 /// use tenferro_cpu::CpuBackend;
354 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
355 /// use tenferro_tensor::BackendSessionHost;
356 ///
357 /// let mut backend = CpuBackend::new();
358 /// let a = Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 4.0]).unwrap();
359 /// let b = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 8.0]).unwrap();
360 /// let y = backend.with_backend_session(|session| a.minimum(&b, session))??;
361 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[1.0, 4.0]);
362 /// # Ok::<(), Box<dyn std::error::Error>>(())
363 /// ```
364 ///
365 /// # Errors
366 ///
367 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
368 /// `DTypeMismatch` for incompatible operands, or
369 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
370 fn minimum(
371 &self,
372 rhs: &Tensor,
373 session: &mut dyn BackendSession,
374 ) -> tenferro_tensor::Result<Tensor>;
375 /// Elementwise negation inside a session.
376 ///
377 /// # Examples
378 ///
379 /// ```rust
380 /// use tenferro_cpu::CpuBackend;
381 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
382 /// use tenferro_tensor::BackendSessionHost;
383 ///
384 /// let mut backend = CpuBackend::new();
385 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, -2.0]).unwrap();
386 /// let y = backend.with_backend_session(|session| x.neg(session))??;
387 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[-1.0, 2.0]);
388 /// # Ok::<(), Box<dyn std::error::Error>>(())
389 /// ```
390 ///
391 /// # Errors
392 ///
393 /// Returns [`tenferro_tensor::Error::Unsupported`] when the dtype is not
394 /// supported by the operation, or [`tenferro_tensor::Error::BackendSource`]
395 /// for a typed backend failure.
396 fn neg(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
397 /// Elementwise absolute value inside a session.
398 ///
399 /// # Examples
400 ///
401 /// ```rust
402 /// use tenferro_cpu::CpuBackend;
403 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
404 /// use tenferro_tensor::BackendSessionHost;
405 ///
406 /// let mut backend = CpuBackend::new();
407 /// let x = Tensor::from_vec_col_major(vec![2], vec![-1.0_f64, 2.0]).unwrap();
408 /// let y = backend.with_backend_session(|session| x.abs(session))??;
409 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[1.0, 2.0]);
410 /// # Ok::<(), Box<dyn std::error::Error>>(())
411 /// ```
412 ///
413 /// # Errors
414 ///
415 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
416 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
417 /// failure.
418 fn abs(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
419 /// Elementwise sign inside a session.
420 ///
421 /// # Examples
422 ///
423 /// ```rust
424 /// use tenferro_cpu::CpuBackend;
425 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
426 /// use tenferro_tensor::BackendSessionHost;
427 ///
428 /// let mut backend = CpuBackend::new();
429 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, -2.0]).unwrap();
430 /// let y = backend.with_backend_session(|session| x.sign(session))??;
431 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[1.0, -1.0]);
432 /// # Ok::<(), Box<dyn std::error::Error>>(())
433 /// ```
434 ///
435 /// # Errors
436 ///
437 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
438 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
439 /// failure.
440 fn sign(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
441 /// Elementwise complex conjugate inside a session.
442 ///
443 /// For real dtypes the conjugate is the identity.
444 ///
445 /// # Examples
446 ///
447 /// ```rust
448 /// use tenferro_cpu::CpuBackend;
449 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
450 /// use tenferro_tensor::BackendSessionHost;
451 ///
452 /// let mut backend = CpuBackend::new();
453 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, -2.0]).unwrap();
454 /// let y = backend.with_backend_session(|session| x.conj(session))??;
455 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[1.0, -2.0]);
456 /// # Ok::<(), Box<dyn std::error::Error>>(())
457 /// ```
458 ///
459 /// # Errors
460 ///
461 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
462 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
463 /// failure.
464 fn conj(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
465 /// Elementwise natural logarithm inside a session.
466 ///
467 /// # Examples
468 ///
469 /// ```rust
470 /// use tenferro_cpu::CpuBackend;
471 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
472 /// use tenferro_tensor::BackendSessionHost;
473 ///
474 /// let mut backend = CpuBackend::new();
475 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, std::f64::consts::E]).unwrap();
476 /// let y = backend.with_backend_session(|session| x.log(session))??;
477 /// let y = y.as_slice::<f64>().unwrap();
478 /// assert!(y[0].abs() < 1.0e-12);
479 /// assert!((y[1] - 1.0).abs() < 1.0e-12);
480 /// # Ok::<(), Box<dyn std::error::Error>>(())
481 /// ```
482 ///
483 /// # Errors
484 ///
485 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
486 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
487 /// failure.
488 fn log(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
489 /// Elementwise `exp(x) - 1` inside a session.
490 ///
491 /// # Examples
492 ///
493 /// ```rust
494 /// use tenferro_cpu::CpuBackend;
495 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
496 /// use tenferro_tensor::BackendSessionHost;
497 ///
498 /// let mut backend = CpuBackend::new();
499 /// let x = Tensor::from_vec_col_major(vec![2], vec![0.0_f64, 1.0]).unwrap();
500 /// let y = backend.with_backend_session(|session| x.expm1(session))??;
501 /// let y = y.as_slice::<f64>().unwrap();
502 /// assert!(y[0].abs() < 1.0e-12);
503 /// assert!((y[1] - (std::f64::consts::E - 1.0)).abs() < 1.0e-12);
504 /// # Ok::<(), Box<dyn std::error::Error>>(())
505 /// ```
506 ///
507 /// # Errors
508 ///
509 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
510 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
511 /// failure.
512 fn expm1(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
513 /// Elementwise `log(1 + x)` inside a session.
514 ///
515 /// # Examples
516 ///
517 /// ```rust
518 /// use tenferro_cpu::CpuBackend;
519 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
520 /// use tenferro_tensor::BackendSessionHost;
521 ///
522 /// let mut backend = CpuBackend::new();
523 /// let x = Tensor::from_vec_col_major(vec![2], vec![0.0_f64, std::f64::consts::E - 1.0]).unwrap();
524 /// let y = backend.with_backend_session(|session| x.log1p(session))??;
525 /// let y = y.as_slice::<f64>().unwrap();
526 /// assert!(y[0].abs() < 1.0e-12);
527 /// assert!((y[1] - 1.0).abs() < 1.0e-12);
528 /// # Ok::<(), Box<dyn std::error::Error>>(())
529 /// ```
530 ///
531 /// # Errors
532 ///
533 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
534 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
535 /// failure.
536 fn log1p(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
537 /// Elementwise error function `erf(x)` inside a session, for real `F32`/`F64`.
538 ///
539 /// # Examples
540 ///
541 /// ```rust
542 /// use tenferro_cpu::CpuBackend;
543 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
544 /// use tenferro_tensor::BackendSessionHost;
545 ///
546 /// let mut backend = CpuBackend::new();
547 /// let x = Tensor::from_vec_col_major(vec![2], vec![0.0_f64, 1.0]).unwrap();
548 /// let y = backend.with_backend_session(|session| x.erf(session))??;
549 /// let y = y.as_slice::<f64>().unwrap();
550 /// assert_eq!(y[0], 0.0);
551 /// assert!((y[1] - 0.842_700_792_949_714_9).abs() < 1.0e-15);
552 /// # Ok::<(), Box<dyn std::error::Error>>(())
553 /// ```
554 ///
555 /// # Errors
556 ///
557 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
558 /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
559 /// for a typed backend failure.
560 fn erf(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
561 /// Elementwise sine inside a session.
562 ///
563 /// # Examples
564 ///
565 /// ```rust
566 /// use tenferro_cpu::CpuBackend;
567 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
568 /// use tenferro_tensor::BackendSessionHost;
569 ///
570 /// let mut backend = CpuBackend::new();
571 /// let x = Tensor::from_vec_col_major(vec![2], vec![0.0_f64, std::f64::consts::FRAC_PI_2]).unwrap();
572 /// let y = backend.with_backend_session(|session| x.sin(session))??;
573 /// let y = y.as_slice::<f64>().unwrap();
574 /// assert!(y[0].abs() < 1.0e-12);
575 /// assert!((y[1] - 1.0).abs() < 1.0e-12);
576 /// # Ok::<(), Box<dyn std::error::Error>>(())
577 /// ```
578 ///
579 /// # Errors
580 ///
581 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
582 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
583 /// failure.
584 fn sin(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
585 /// Elementwise cosine inside a session.
586 ///
587 /// # Examples
588 ///
589 /// ```rust
590 /// use tenferro_cpu::CpuBackend;
591 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
592 /// use tenferro_tensor::BackendSessionHost;
593 ///
594 /// let mut backend = CpuBackend::new();
595 /// let x = Tensor::from_vec_col_major(vec![2], vec![0.0_f64, std::f64::consts::PI]).unwrap();
596 /// let y = backend.with_backend_session(|session| x.cos(session))??;
597 /// let y = y.as_slice::<f64>().unwrap();
598 /// assert!((y[0] - 1.0).abs() < 1.0e-12);
599 /// assert!((y[1] + 1.0).abs() < 1.0e-12);
600 /// # Ok::<(), Box<dyn std::error::Error>>(())
601 /// ```
602 ///
603 /// # Errors
604 ///
605 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
606 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
607 /// failure.
608 fn cos(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
609 /// Elementwise hyperbolic tangent inside a session.
610 ///
611 /// # Examples
612 ///
613 /// ```rust
614 /// use tenferro_cpu::CpuBackend;
615 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
616 /// use tenferro_tensor::BackendSessionHost;
617 ///
618 /// let mut backend = CpuBackend::new();
619 /// let x = Tensor::from_vec_col_major(vec![2], vec![0.0_f64, 1.0]).unwrap();
620 /// let y = backend.with_backend_session(|session| x.tanh(session))??;
621 /// let y = y.as_slice::<f64>().unwrap();
622 /// assert!(y[0].abs() < 1.0e-12);
623 /// assert!((y[1] - 0.7615941559557649).abs() < 1.0e-12);
624 /// # Ok::<(), Box<dyn std::error::Error>>(())
625 /// ```
626 ///
627 /// # Errors
628 ///
629 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
630 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
631 /// failure.
632 fn tanh(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
633 /// Elementwise square root inside a session.
634 ///
635 /// # Examples
636 ///
637 /// ```rust
638 /// use tenferro_cpu::CpuBackend;
639 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
640 /// use tenferro_tensor::BackendSessionHost;
641 ///
642 /// let mut backend = CpuBackend::new();
643 /// let x = Tensor::from_vec_col_major(vec![2], vec![4.0_f64, 9.0]).unwrap();
644 /// let y = backend.with_backend_session(|session| x.sqrt(session))??;
645 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[2.0, 3.0]);
646 /// # Ok::<(), Box<dyn std::error::Error>>(())
647 /// ```
648 ///
649 /// # Errors
650 ///
651 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
652 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
653 /// failure.
654 fn sqrt(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
655 /// Elementwise reciprocal square root inside a session.
656 ///
657 /// # Examples
658 ///
659 /// ```rust
660 /// use tenferro_cpu::CpuBackend;
661 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
662 /// use tenferro_tensor::BackendSessionHost;
663 ///
664 /// let mut backend = CpuBackend::new();
665 /// let x = Tensor::from_vec_col_major(vec![2], vec![4.0_f64, 1.0]).unwrap();
666 /// let y = backend.with_backend_session(|session| x.rsqrt(session))??;
667 /// let y = y.as_slice::<f64>().unwrap();
668 /// assert!((y[0] - 0.5).abs() < 1.0e-12);
669 /// assert!((y[1] - 1.0).abs() < 1.0e-12);
670 /// # Ok::<(), Box<dyn std::error::Error>>(())
671 /// ```
672 ///
673 /// # Errors
674 ///
675 /// Returns [`tenferro_tensor::Error::Unsupported`] for an unsupported
676 /// dtype or [`tenferro_tensor::Error::BackendSource`] for a typed backend
677 /// failure.
678 fn rsqrt(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
679 /// Elementwise comparison with NumPy-style broadcasting inside a session.
680 ///
681 /// The result is a bool tensor.
682 ///
683 /// # Examples
684 ///
685 /// ```rust
686 /// use tenferro_cpu::CpuBackend;
687 /// use tenferro_runtime::{CompareDir, Tensor, TensorSessionOpsExt};
688 /// use tenferro_tensor::BackendSessionHost;
689 ///
690 /// let mut backend = CpuBackend::new();
691 /// let a = Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 4.0]).unwrap();
692 /// let b = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 8.0]).unwrap();
693 /// let y = backend.with_backend_session(|session| a.compare(&b, CompareDir::Gt, session))??;
694 /// assert_eq!(y.as_slice::<bool>().unwrap(), &[true, false]);
695 /// # Ok::<(), Box<dyn std::error::Error>>(())
696 /// ```
697 ///
698 /// # Errors
699 ///
700 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
701 /// `DTypeMismatch` for incompatible shape/dtype metadata, or
702 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
703 fn compare(
704 &self,
705 rhs: &Tensor,
706 dir: CompareDir,
707 session: &mut dyn BackendSession,
708 ) -> tenferro_tensor::Result<Tensor>;
709 /// Select values from `on_true` or `on_false` using this tensor as condition inside a session.
710 ///
711 /// # Examples
712 ///
713 /// ```rust
714 /// use tenferro_cpu::CpuBackend;
715 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
716 /// use tenferro_tensor::BackendSessionHost;
717 ///
718 /// let mut backend = CpuBackend::new();
719 /// let condition = Tensor::from_vec_col_major(vec![2], vec![true, false]).unwrap();
720 /// let on_true = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
721 /// let on_false = Tensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0]).unwrap();
722 /// let y = backend.with_backend_session(|session| condition.where_select(&on_true, &on_false, session))??;
723 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[1.0, 4.0]);
724 /// # Ok::<(), Box<dyn std::error::Error>>(())
725 /// ```
726 ///
727 /// # Errors
728 ///
729 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
730 /// `DTypeMismatch` when the condition and branches are incompatible, or
731 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
732 fn where_select(
733 &self,
734 on_true: &Tensor,
735 on_false: &Tensor,
736 session: &mut dyn BackendSession,
737 ) -> tenferro_tensor::Result<Tensor>;
738 /// Clamp values elementwise between lower and upper bounds inside a session.
739 ///
740 /// # Examples
741 ///
742 /// ```rust
743 /// use tenferro_cpu::CpuBackend;
744 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
745 /// use tenferro_tensor::BackendSessionHost;
746 ///
747 /// let mut backend = CpuBackend::new();
748 /// let x = Tensor::from_vec_col_major(vec![2], vec![-2.0_f64, 4.0]).unwrap();
749 /// let lower = Tensor::from_vec_col_major(vec![], vec![0.0_f64]).unwrap();
750 /// let upper = Tensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
751 /// let y = backend.with_backend_session(|session| x.clamp(&lower, &upper, session))??;
752 /// assert_eq!(y.as_slice::<f64>().unwrap(), &[0.0, 3.0]);
753 /// # Ok::<(), Box<dyn std::error::Error>>(())
754 /// ```
755 ///
756 /// # Errors
757 ///
758 /// Returns [`tenferro_tensor::Error::Validation`] with `ShapeMismatch` or
759 /// `DTypeMismatch` when bounds are incompatible with the input, or
760 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
761 fn clamp(
762 &self,
763 lower: &Tensor,
764 upper: &Tensor,
765 session: &mut dyn BackendSession,
766 ) -> tenferro_tensor::Result<Tensor>;
767 /// Rank-2 matrix multiplication inside a session.
768 ///
769 /// # Examples
770 ///
771 /// ```rust
772 /// use tenferro_cpu::CpuBackend;
773 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
774 /// use tenferro_tensor::BackendSessionHost;
775 ///
776 /// let mut backend = CpuBackend::new();
777 /// let a = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
778 /// let b = Tensor::from_vec_col_major(vec![3, 2], vec![1.0_f64; 6]).unwrap();
779 /// let c = backend.with_backend_session(|session| a.matmul(&b, session))??;
780 /// assert_eq!(c.shape(), &[2, 2]);
781 /// # Ok::<(), Box<dyn std::error::Error>>(())
782 /// ```
783 ///
784 /// # Errors
785 ///
786 /// Returns [`tenferro_tensor::Error::Validation`] with `RankMismatch`,
787 /// `ShapeMismatch`, or `DTypeMismatch` for incompatible matrices, or
788 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
789 fn matmul(
790 &self,
791 rhs: &Tensor,
792 session: &mut dyn BackendSession,
793 ) -> tenferro_tensor::Result<Tensor>;
794 /// Reshape without changing element order inside a session.
795 ///
796 /// # Examples
797 ///
798 /// ```rust
799 /// use tenferro_cpu::CpuBackend;
800 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
801 /// use tenferro_tensor::BackendSessionHost;
802 ///
803 /// let mut backend = CpuBackend::new();
804 /// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0]).unwrap();
805 /// let y = backend.with_backend_session(|session| x.reshape(&[4], session))??;
806 /// assert_eq!(y.shape(), &[4]);
807 /// # Ok::<(), Box<dyn std::error::Error>>(())
808 /// ```
809 ///
810 /// # Errors
811 ///
812 /// Returns [`tenferro_tensor::Error::Validation`] with
813 /// `ShapeMismatch`, `RankMismatch`, or `InvalidArgument` when element
814 /// counts or ranks are invalid, or
815 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
816 fn reshape(
817 &self,
818 shape: &[usize],
819 session: &mut dyn BackendSession,
820 ) -> tenferro_tensor::Result<Tensor>;
821 /// Permute axes inside a session.
822 ///
823 /// # Examples
824 ///
825 /// ```rust
826 /// use tenferro_cpu::CpuBackend;
827 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
828 /// use tenferro_tensor::BackendSessionHost;
829 ///
830 /// let mut backend = CpuBackend::new();
831 /// let x = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
832 /// let y = backend.with_backend_session(|session| x.transpose(&[1, 0], session))??;
833 /// assert_eq!(y.shape(), &[3, 2]);
834 /// # Ok::<(), Box<dyn std::error::Error>>(())
835 /// ```
836 ///
837 /// # Errors
838 ///
839 /// Returns [`tenferro_tensor::Error::Validation`] with
840 /// `InvalidPermutationLength`, `AxisOutOfBounds`, or `DuplicateAxis` for
841 /// an invalid permutation, or
842 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
843 fn transpose(
844 &self,
845 perm: &[usize],
846 session: &mut dyn BackendSession,
847 ) -> tenferro_tensor::Result<Tensor>;
848 /// Gather slices of this tensor at `indices` inside a session (StableHLO `gather`).
849 ///
850 /// # Examples
851 ///
852 /// ```rust
853 /// use tenferro_cpu::CpuBackend;
854 /// use tenferro_runtime::{GatherConfig, Tensor, TensorSessionOpsExt};
855 /// use tenferro_tensor::BackendSessionHost;
856 ///
857 /// let mut backend = CpuBackend::new();
858 /// let x = Tensor::from_vec_col_major(vec![3], vec![10.0_f64, 20.0, 30.0])?;
859 /// let indices = Tensor::from_vec_col_major(vec![2, 1], vec![2_i64, 0])?;
860 /// let config = GatherConfig {
861 /// offset_dims: vec![],
862 /// collapsed_slice_dims: vec![0],
863 /// start_index_map: vec![0],
864 /// index_vector_dim: 1,
865 /// slice_sizes: vec![1],
866 /// };
867 /// let y = backend.with_backend_session(|session| x.gather(&indices, config, session))??;
868 /// assert_eq!(y.as_slice::<f64>()?, &[30.0, 10.0]);
869 /// # Ok::<(), Box<dyn std::error::Error>>(())
870 /// ```
871 ///
872 /// # Errors
873 ///
874 /// Returns [`tenferro_tensor::Error::Validation`] for an invalid configuration, index dtype or out-of-range index, or
875 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
876 fn gather(
877 &self,
878 indices: &Tensor,
879 config: GatherConfig,
880 session: &mut dyn BackendSession,
881 ) -> tenferro_tensor::Result<Tensor>;
882 /// Scatter `updates` into a copy of this tensor at `indices` inside a session (StableHLO `scatter`).
883 ///
884 /// # Examples
885 ///
886 /// ```rust
887 /// use tenferro_cpu::CpuBackend;
888 /// use tenferro_runtime::{ScatterConfig, Tensor, TensorSessionOpsExt};
889 /// use tenferro_tensor::BackendSessionHost;
890 ///
891 /// let mut backend = CpuBackend::new();
892 /// let x = Tensor::from_vec_col_major(vec![4], vec![0.0_f64; 4])?;
893 /// let indices = Tensor::from_vec_col_major(vec![2, 1], vec![1_i64, 3])?;
894 /// let updates = Tensor::from_vec_col_major(vec![2], vec![5.0_f64, 7.0])?;
895 /// let config = ScatterConfig {
896 /// update_window_dims: vec![],
897 /// inserted_window_dims: vec![0],
898 /// scatter_dims_to_operand_dims: vec![0],
899 /// index_vector_dim: 1,
900 /// };
901 /// let y = backend.with_backend_session(|session| x.scatter(&indices, &updates, config, session))??;
902 /// assert_eq!(y.as_slice::<f64>()?, &[0.0, 5.0, 0.0, 7.0]);
903 /// # Ok::<(), Box<dyn std::error::Error>>(())
904 /// ```
905 ///
906 /// # Errors
907 ///
908 /// Returns [`tenferro_tensor::Error::Validation`] for an invalid configuration, index dtype or update shape, or
909 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
910 fn scatter(
911 &self,
912 indices: &Tensor,
913 updates: &Tensor,
914 config: ScatterConfig,
915 session: &mut dyn BackendSession,
916 ) -> tenferro_tensor::Result<Tensor>;
917 /// Slice this tensor with explicit start, limit and stride per axis inside a session.
918 ///
919 /// # Examples
920 ///
921 /// ```rust
922 /// use tenferro_cpu::CpuBackend;
923 /// use tenferro_runtime::{SliceConfig, Tensor, TensorSessionOpsExt};
924 /// use tenferro_tensor::BackendSessionHost;
925 ///
926 /// let mut backend = CpuBackend::new();
927 /// let x = Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0])?;
928 /// let config = SliceConfig { starts: vec![1], limits: vec![3], strides: vec![1] };
929 /// let y = backend.with_backend_session(|session| x.slice(config, session))??;
930 /// assert_eq!(y.as_slice::<f64>()?, &[2.0, 3.0]);
931 /// # Ok::<(), Box<dyn std::error::Error>>(())
932 /// ```
933 ///
934 /// # Errors
935 ///
936 /// Returns [`tenferro_tensor::Error::Validation`] for bounds or strides that do not fit the input, or
937 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
938 fn slice(
939 &self,
940 config: SliceConfig,
941 session: &mut dyn BackendSession,
942 ) -> tenferro_tensor::Result<Tensor>;
943 /// Slice this tensor at runtime `starts` (an integer tensor) with static `sizes` inside a session.
944 ///
945 /// # Examples
946 ///
947 /// ```rust
948 /// use tenferro_cpu::CpuBackend;
949 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
950 /// use tenferro_tensor::BackendSessionHost;
951 ///
952 /// let mut backend = CpuBackend::new();
953 /// let x = Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0])?;
954 /// let starts = Tensor::from_vec_col_major(vec![1], vec![1_i64])?;
955 /// let y = backend.with_backend_session(|session| x.dynamic_slice(&starts, &[2], session))??;
956 /// assert_eq!(y.as_slice::<f64>()?, &[2.0, 3.0]);
957 /// # Ok::<(), Box<dyn std::error::Error>>(())
958 /// ```
959 ///
960 /// # Errors
961 ///
962 /// Returns [`tenferro_tensor::Error::Validation`] for a start-index dtype or rank mismatch or sizes larger than the input, or
963 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
964 fn dynamic_slice(
965 &self,
966 starts: &Tensor,
967 sizes: &[usize],
968 session: &mut dyn BackendSession,
969 ) -> tenferro_tensor::Result<Tensor>;
970 /// Pad this tensor with zeros (edge and interior padding per axis) inside a session.
971 ///
972 /// # Examples
973 ///
974 /// ```rust
975 /// use tenferro_cpu::CpuBackend;
976 /// use tenferro_runtime::{PadConfig, Tensor, TensorSessionOpsExt};
977 /// use tenferro_tensor::BackendSessionHost;
978 ///
979 /// let mut backend = CpuBackend::new();
980 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
981 /// let config = PadConfig {
982 /// edge_padding_low: vec![1],
983 /// edge_padding_high: vec![1],
984 /// interior_padding: vec![1],
985 /// };
986 /// let y = backend.with_backend_session(|session| x.pad(config, session))??;
987 /// assert_eq!(y.as_slice::<f64>()?, &[0.0, 1.0, 0.0, 2.0, 0.0]);
988 /// # Ok::<(), Box<dyn std::error::Error>>(())
989 /// ```
990 ///
991 /// # Errors
992 ///
993 /// Returns [`tenferro_tensor::Error::Validation`] for a padding configuration whose length does not match the input rank, or
994 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
995 fn pad(
996 &self,
997 config: PadConfig,
998 session: &mut dyn BackendSession,
999 ) -> tenferro_tensor::Result<Tensor>;
1000 /// Concatenate tensors along `axis` inside a session.
1001 ///
1002 /// This has no receiver, like the eager `EagerSession::concatenate`: call it as
1003 /// `Tensor::concatenate(&[&a, &b], axis, session)`.
1004 ///
1005 /// # Examples
1006 ///
1007 /// ```rust
1008 /// use tenferro_cpu::CpuBackend;
1009 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1010 /// use tenferro_tensor::BackendSessionHost;
1011 ///
1012 /// let mut backend = CpuBackend::new();
1013 /// let a = Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?;
1014 /// let b = Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?;
1015 /// let y = backend.with_backend_session(|session| Tensor::concatenate(&[&a, &b], 0, session))??;
1016 /// assert_eq!(y.as_slice::<f64>()?, &[1.0, 2.0]);
1017 /// # Ok::<(), Box<dyn std::error::Error>>(())
1018 /// ```
1019 ///
1020 /// # Errors
1021 ///
1022 /// Returns [`tenferro_tensor::Error::Validation`] for an empty input list, an axis out of range, or mismatched shapes or dtypes, or
1023 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1024 fn concatenate(
1025 inputs: &[&Tensor],
1026 axis: usize,
1027 session: &mut dyn BackendSession,
1028 ) -> tenferro_tensor::Result<Tensor>
1029 where
1030 Self: Sized;
1031 /// Reverse the elements along `axes` inside a session.
1032 ///
1033 /// # Examples
1034 ///
1035 /// ```rust
1036 /// use tenferro_cpu::CpuBackend;
1037 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1038 /// use tenferro_tensor::BackendSessionHost;
1039 ///
1040 /// let mut backend = CpuBackend::new();
1041 /// let x = Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0])?;
1042 /// let y = backend.with_backend_session(|session| x.reverse(&[0], session))??;
1043 /// assert_eq!(y.as_slice::<f64>()?, &[3.0, 2.0, 1.0]);
1044 /// # Ok::<(), Box<dyn std::error::Error>>(())
1045 /// ```
1046 ///
1047 /// # Errors
1048 ///
1049 /// Returns [`tenferro_tensor::Error::Validation`] for an out-of-range or repeated axis, or
1050 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1051 fn reverse(
1052 &self,
1053 axes: &[usize],
1054 session: &mut dyn BackendSession,
1055 ) -> tenferro_tensor::Result<Tensor>;
1056 /// Take the maximum over the selected axes inside a session. `None` reduces every axis and `Some(&[])` keeps the input shape.
1057 ///
1058 /// # Examples
1059 ///
1060 /// ```rust
1061 /// use tenferro_cpu::CpuBackend;
1062 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1063 /// use tenferro_tensor::BackendSessionHost;
1064 ///
1065 /// let mut backend = CpuBackend::new();
1066 /// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 5.0, 3.0, 2.0])?;
1067 /// let y = backend.with_backend_session(|session| x.reduce_max(Some(&[0]), session))??;
1068 /// assert_eq!(y.as_slice::<f64>()?, &[5.0, 3.0]);
1069 /// # Ok::<(), Box<dyn std::error::Error>>(())
1070 /// ```
1071 ///
1072 /// # Errors
1073 ///
1074 /// Returns [`tenferro_tensor::Error::Validation`] for an out-of-range or repeated axis, or
1075 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1076 /// An unsupported dtype returns [`tenferro_tensor::Error::Unsupported`].
1077 fn reduce_max(
1078 &self,
1079 axes: Option<&[usize]>,
1080 session: &mut dyn BackendSession,
1081 ) -> tenferro_tensor::Result<Tensor>;
1082 /// Take the minimum over the selected axes inside a session. `None` reduces every axis and `Some(&[])` keeps the input shape.
1083 ///
1084 /// # Examples
1085 ///
1086 /// ```rust
1087 /// use tenferro_cpu::CpuBackend;
1088 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1089 /// use tenferro_tensor::BackendSessionHost;
1090 ///
1091 /// let mut backend = CpuBackend::new();
1092 /// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 5.0, 3.0, 2.0])?;
1093 /// let y = backend.with_backend_session(|session| x.reduce_min(Some(&[0]), session))??;
1094 /// assert_eq!(y.as_slice::<f64>()?, &[1.0, 2.0]);
1095 /// # Ok::<(), Box<dyn std::error::Error>>(())
1096 /// ```
1097 ///
1098 /// # Errors
1099 ///
1100 /// Returns [`tenferro_tensor::Error::Validation`] for an out-of-range or repeated axis, or
1101 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1102 /// An unsupported dtype returns [`tenferro_tensor::Error::Unsupported`].
1103 fn reduce_min(
1104 &self,
1105 axes: Option<&[usize]>,
1106 session: &mut dyn BackendSession,
1107 ) -> tenferro_tensor::Result<Tensor>;
1108 /// Multiply over the selected axes inside a session. `None` reduces every axis and `Some(&[])` keeps the input shape.
1109 ///
1110 /// # Examples
1111 ///
1112 /// ```rust
1113 /// use tenferro_cpu::CpuBackend;
1114 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1115 /// use tenferro_tensor::BackendSessionHost;
1116 ///
1117 /// let mut backend = CpuBackend::new();
1118 /// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 5.0, 3.0, 2.0])?;
1119 /// let y = backend.with_backend_session(|session| x.reduce_prod(Some(&[0]), session))??;
1120 /// assert_eq!(y.as_slice::<f64>()?, &[5.0, 6.0]);
1121 /// # Ok::<(), Box<dyn std::error::Error>>(())
1122 /// ```
1123 ///
1124 /// # Errors
1125 ///
1126 /// Returns [`tenferro_tensor::Error::Validation`] for an out-of-range or repeated axis, or
1127 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1128 /// An unsupported dtype returns [`tenferro_tensor::Error::Unsupported`].
1129 fn reduce_prod(
1130 &self,
1131 axes: Option<&[usize]>,
1132 session: &mut dyn BackendSession,
1133 ) -> tenferro_tensor::Result<Tensor>;
1134 /// Sum elementwise squares over the selected axes inside a session (`f32`/`f64`). `None` reduces every axis and `Some(&[])` keeps the input shape.
1135 ///
1136 /// # Examples
1137 ///
1138 /// ```rust
1139 /// use tenferro_cpu::CpuBackend;
1140 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1141 /// use tenferro_tensor::BackendSessionHost;
1142 ///
1143 /// let mut backend = CpuBackend::new();
1144 /// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 5.0, 3.0, 2.0])?;
1145 /// let y = backend.with_backend_session(|session| x.reduce_sum_squares(Some(&[0]), session))??;
1146 /// assert_eq!(y.as_slice::<f64>()?, &[26.0, 13.0]);
1147 /// # Ok::<(), Box<dyn std::error::Error>>(())
1148 /// ```
1149 ///
1150 /// # Errors
1151 ///
1152 /// Returns [`tenferro_tensor::Error::Validation`] for an out-of-range or repeated axis, or
1153 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1154 /// An unsupported dtype returns [`tenferro_tensor::Error::Unsupported`].
1155 fn reduce_sum_squares(
1156 &self,
1157 axes: Option<&[usize]>,
1158 session: &mut dyn BackendSession,
1159 ) -> tenferro_tensor::Result<Tensor>;
1160 /// Broadcast this tensor into `shape`, mapping input axis `i` to output axis `dims[i]`, inside a session.
1161 ///
1162 /// # Examples
1163 ///
1164 /// ```rust
1165 /// use tenferro_cpu::CpuBackend;
1166 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1167 /// use tenferro_tensor::BackendSessionHost;
1168 ///
1169 /// let mut backend = CpuBackend::new();
1170 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
1171 /// let y = backend.with_backend_session(|session| x.broadcast_in_dim(&[2, 2], &[0], session))??;
1172 /// assert_eq!(y.as_slice::<f64>()?, &[1.0, 2.0, 1.0, 2.0]);
1173 /// # Ok::<(), Box<dyn std::error::Error>>(())
1174 /// ```
1175 ///
1176 /// # Errors
1177 ///
1178 /// Returns [`tenferro_tensor::Error::Validation`] for a dimension mapping that does not fit the input or target shape, or
1179 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1180 fn broadcast_in_dim(
1181 &self,
1182 shape: &[usize],
1183 dims: &[usize],
1184 session: &mut dyn BackendSession,
1185 ) -> tenferro_tensor::Result<Tensor>;
1186 /// Keep the lower triangle (on and below diagonal `k`) of the trailing matrix axes inside a session.
1187 ///
1188 /// # Examples
1189 ///
1190 /// ```rust
1191 /// use tenferro_cpu::CpuBackend;
1192 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1193 /// use tenferro_tensor::BackendSessionHost;
1194 ///
1195 /// let mut backend = CpuBackend::new();
1196 /// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?;
1197 /// let y = backend.with_backend_session(|session| x.tril(0, session))??;
1198 /// assert_eq!(y.as_slice::<f64>()?, &[1.0, 2.0, 0.0, 4.0]);
1199 /// # Ok::<(), Box<dyn std::error::Error>>(())
1200 /// ```
1201 ///
1202 /// # Errors
1203 ///
1204 /// Returns [`tenferro_tensor::Error::Validation`] for an input of rank below 2, or
1205 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1206 fn tril(&self, k: i64, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1207 /// Keep the upper triangle (on and above diagonal `k`) of the trailing matrix axes inside a session.
1208 ///
1209 /// # Examples
1210 ///
1211 /// ```rust
1212 /// use tenferro_cpu::CpuBackend;
1213 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1214 /// use tenferro_tensor::BackendSessionHost;
1215 ///
1216 /// let mut backend = CpuBackend::new();
1217 /// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?;
1218 /// let y = backend.with_backend_session(|session| x.triu(0, session))??;
1219 /// assert_eq!(y.as_slice::<f64>()?, &[1.0, 0.0, 3.0, 4.0]);
1220 /// # Ok::<(), Box<dyn std::error::Error>>(())
1221 /// ```
1222 ///
1223 /// # Errors
1224 ///
1225 /// Returns [`tenferro_tensor::Error::Validation`] for an input of rank below 2, or
1226 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1227 fn triu(&self, k: i64, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1228 /// Extract the diagonal along `axis_a` and `axis_b` inside a session.
1229 ///
1230 /// # Examples
1231 ///
1232 /// ```rust
1233 /// use tenferro_cpu::CpuBackend;
1234 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1235 /// use tenferro_tensor::BackendSessionHost;
1236 ///
1237 /// let mut backend = CpuBackend::new();
1238 /// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?;
1239 /// let y = backend.with_backend_session(|session| x.extract_diag(0, 1, session))??;
1240 /// assert_eq!(y.as_slice::<f64>()?, &[1.0, 4.0]);
1241 /// # Ok::<(), Box<dyn std::error::Error>>(())
1242 /// ```
1243 ///
1244 /// # Errors
1245 ///
1246 /// Returns [`tenferro_tensor::Error::Validation`] for equal or out-of-range axes, or axes of different extent, or
1247 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1248 fn extract_diag(
1249 &self,
1250 axis_a: usize,
1251 axis_b: usize,
1252 session: &mut dyn BackendSession,
1253 ) -> tenferro_tensor::Result<Tensor>;
1254 /// Embed this tensor along the diagonal of `axis_a` and `axis_b` inside a session.
1255 ///
1256 /// # Examples
1257 ///
1258 /// ```rust
1259 /// use tenferro_cpu::CpuBackend;
1260 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1261 /// use tenferro_tensor::BackendSessionHost;
1262 ///
1263 /// let mut backend = CpuBackend::new();
1264 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
1265 /// let y = backend.with_backend_session(|session| x.embed_diag(0, 1, session))??;
1266 /// assert_eq!(y.as_slice::<f64>()?, &[1.0, 0.0, 0.0, 2.0]);
1267 /// # Ok::<(), Box<dyn std::error::Error>>(())
1268 /// ```
1269 ///
1270 /// # Errors
1271 ///
1272 /// Returns [`tenferro_tensor::Error::Validation`] for invalid diagonal axes, or
1273 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1274 fn embed_diag(
1275 &self,
1276 axis_a: usize,
1277 axis_b: usize,
1278 session: &mut dyn BackendSession,
1279 ) -> tenferro_tensor::Result<Tensor>;
1280 /// Contract this tensor with `rhs` inside a session (StableHLO `dot_general`).
1281 ///
1282 /// The output layout is `[lhs free..., rhs free..., batch...]` (batch axes trail).
1283 ///
1284 /// # Examples
1285 ///
1286 /// ```rust
1287 /// use tenferro_cpu::CpuBackend;
1288 /// use tenferro_runtime::{DotGeneralConfig, Tensor, TensorSessionOpsExt};
1289 /// use tenferro_tensor::BackendSessionHost;
1290 ///
1291 /// let mut backend = CpuBackend::new();
1292 /// let lhs = Tensor::from_vec_col_major(vec![1, 2], vec![2.0_f64, 3.0])?;
1293 /// let rhs = Tensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 5.0])?;
1294 /// let config = DotGeneralConfig {
1295 /// lhs_contracting_dims: [1].as_slice().into(),
1296 /// rhs_contracting_dims: [0].as_slice().into(),
1297 /// lhs_batch_dims: [].as_slice().into(),
1298 /// rhs_batch_dims: [].as_slice().into(),
1299 /// };
1300 /// let y = backend.with_backend_session(|session| lhs.dot_general(&rhs, config, session))??;
1301 /// assert_eq!(y.as_slice::<f64>()?, &[23.0]);
1302 /// # Ok::<(), Box<dyn std::error::Error>>(())
1303 /// ```
1304 ///
1305 /// # Errors
1306 ///
1307 /// Returns [`tenferro_tensor::Error::Validation`] for incompatible contraction or batch dimensions or dtypes, or
1308 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1309 fn dot_general(
1310 &self,
1311 rhs: &Tensor,
1312 config: DotGeneralConfig,
1313 session: &mut dyn BackendSession,
1314 ) -> tenferro_tensor::Result<Tensor>;
1315 /// Contract with optional conjugation of either operand inside a session.
1316 ///
1317 /// # Examples
1318 ///
1319 /// ```rust
1320 /// use tenferro_cpu::CpuBackend;
1321 /// use tenferro_runtime::{DotGeneralConfig, Tensor, TensorSessionOpsExt};
1322 /// use tenferro_tensor::BackendSessionHost;
1323 ///
1324 /// let mut backend = CpuBackend::new();
1325 /// use num_complex::Complex64;
1326 /// let lhs = Tensor::from_vec_col_major(vec![1, 1], vec![Complex64::new(0.0, 1.0)])?;
1327 /// let rhs = Tensor::from_vec_col_major(vec![1, 1], vec![Complex64::new(0.0, 1.0)])?;
1328 /// let config = DotGeneralConfig {
1329 /// lhs_contracting_dims: [1].as_slice().into(),
1330 /// rhs_contracting_dims: [0].as_slice().into(),
1331 /// lhs_batch_dims: [].as_slice().into(),
1332 /// rhs_batch_dims: [].as_slice().into(),
1333 /// };
1334 /// let y = backend.with_backend_session(|session| lhs.dot_general_with_conj(&rhs, config, true, false, session))??;
1335 /// assert_eq!(y.as_slice::<Complex64>()?, &[Complex64::new(1.0, 0.0)]);
1336 /// # Ok::<(), Box<dyn std::error::Error>>(())
1337 /// ```
1338 ///
1339 /// # Errors
1340 ///
1341 /// Returns [`tenferro_tensor::Error::Validation`] for incompatible contraction or batch dimensions or dtypes, or
1342 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1343 fn dot_general_with_conj(
1344 &self,
1345 rhs: &Tensor,
1346 config: DotGeneralConfig,
1347 lhs_conj: bool,
1348 rhs_conj: bool,
1349 session: &mut dyn BackendSession,
1350 ) -> tenferro_tensor::Result<Tensor>;
1351 /// Multiply by a real scalar inside a session, with the eager `scale_real` dtype rules.
1352 ///
1353 /// Integer dtypes round the factor.
1354 ///
1355 /// # Examples
1356 ///
1357 /// ```rust
1358 /// use tenferro_cpu::CpuBackend;
1359 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1360 /// use tenferro_tensor::BackendSessionHost;
1361 ///
1362 /// let mut backend = CpuBackend::new();
1363 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
1364 /// let y = backend.with_backend_session(|session| x.scale_real(2.0, session))??;
1365 /// assert_eq!(y.as_slice::<f64>()?, &[2.0, 4.0]);
1366 /// # Ok::<(), Box<dyn std::error::Error>>(())
1367 /// ```
1368 ///
1369 /// # Errors
1370 ///
1371 /// Returns [`tenferro_tensor::Error::Validation`] with `InvalidArgument` for a
1372 /// non-finite factor, an integer factor out of range, or an external dtype,
1373 /// or [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1374 fn scale_real(
1375 &self,
1376 factor: f64,
1377 session: &mut dyn BackendSession,
1378 ) -> tenferro_tensor::Result<Tensor>;
1379 /// Multiply a complex tensor by a complex scalar inside a session.
1380 ///
1381 /// # Examples
1382 ///
1383 /// ```rust
1384 /// use tenferro_cpu::CpuBackend;
1385 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1386 /// use tenferro_tensor::BackendSessionHost;
1387 ///
1388 /// let mut backend = CpuBackend::new();
1389 /// use num_complex::Complex64;
1390 /// let x = Tensor::from_vec_col_major(vec![1], vec![Complex64::new(1.0, 2.0)])?;
1391 /// let y = backend.with_backend_session(|session| x.scale_complex(Complex64::new(0.0, 1.0), session))??;
1392 /// assert_eq!(y.as_slice::<Complex64>()?, &[Complex64::new(-2.0, 1.0)]);
1393 /// # Ok::<(), Box<dyn std::error::Error>>(())
1394 /// ```
1395 ///
1396 /// # Errors
1397 ///
1398 /// Returns [`tenferro_tensor::Error::Validation`] with `InvalidArgument` when
1399 /// the input dtype is not complex, or [`tenferro_tensor::Error::BackendSource`]
1400 /// for a typed backend failure.
1401 fn scale_complex(
1402 &self,
1403 factor: Complex64,
1404 session: &mut dyn BackendSession,
1405 ) -> tenferro_tensor::Result<Tensor>;
1406 /// Logistic sigmoid `1 / (1 + exp(-x))` inside a session, overflow-free.
1407 ///
1408 /// Evaluated as `1 / (1 + e)` for `x > 0` and `e / (1 + e)` otherwise, with
1409 /// `e = exp(-|x|)`. Real `F32`/`F64` only.
1410 ///
1411 /// # Examples
1412 ///
1413 /// ```rust
1414 /// use tenferro_cpu::CpuBackend;
1415 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1416 /// use tenferro_tensor::BackendSessionHost;
1417 ///
1418 /// let mut backend = CpuBackend::new();
1419 /// let x = Tensor::from_vec_col_major(vec![3], vec![-700.0_f64, 0.0, 1000.0])?;
1420 /// let y = backend.with_backend_session(|session| x.sigmoid(session))??;
1421 /// let y = y.as_slice::<f64>()?;
1422 /// assert_eq!(y[1], 0.5);
1423 /// assert!(y[0] > 0.0 && y[0] < 1e-300);
1424 /// assert_eq!(y[2], 1.0);
1425 /// # Ok::<(), Box<dyn std::error::Error>>(())
1426 /// ```
1427 ///
1428 /// # Errors
1429 ///
1430 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1431 /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1432 /// for a typed backend failure.
1433 fn sigmoid(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1434 /// SiLU (swish) `x * sigmoid(x)` inside a session.
1435 ///
1436 /// Real `F32`/`F64` only.
1437 ///
1438 /// # Examples
1439 ///
1440 /// ```rust
1441 /// use tenferro_cpu::CpuBackend;
1442 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1443 /// use tenferro_tensor::BackendSessionHost;
1444 ///
1445 /// let mut backend = CpuBackend::new();
1446 /// let x = Tensor::from_vec_col_major(vec![3], vec![-1.0_f64, 0.0, 1.0])?;
1447 /// let y = backend.with_backend_session(|session| x.silu(session))??;
1448 /// let y = y.as_slice::<f64>()?;
1449 /// assert_eq!(y[1], 0.0);
1450 /// assert!((y[2] - 1.0 / (1.0 + (-1.0_f64).exp())).abs() < 1e-15);
1451 /// # Ok::<(), Box<dyn std::error::Error>>(())
1452 /// ```
1453 ///
1454 /// # Errors
1455 ///
1456 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1457 /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1458 /// for a typed backend failure.
1459 fn silu(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1460 /// Softplus `log(1 + exp(x))` inside a session, in the stable form `max(x, 0) + log1p(exp(-|x|))`.
1461 ///
1462 /// Real `F32`/`F64` only; never overflows.
1463 ///
1464 /// # Examples
1465 ///
1466 /// ```rust
1467 /// use tenferro_cpu::CpuBackend;
1468 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1469 /// use tenferro_tensor::BackendSessionHost;
1470 ///
1471 /// let mut backend = CpuBackend::new();
1472 /// let x = Tensor::from_vec_col_major(vec![3], vec![-1000.0_f64, 0.0, 1000.0])?;
1473 /// let y = backend.with_backend_session(|session| x.softplus(session))??;
1474 /// let y = y.as_slice::<f64>()?;
1475 /// assert_eq!(y[0], 0.0);
1476 /// assert!((y[1] - 2.0_f64.ln()).abs() < 1e-15);
1477 /// assert_eq!(y[2], 1000.0);
1478 /// # Ok::<(), Box<dyn std::error::Error>>(())
1479 /// ```
1480 ///
1481 /// # Errors
1482 ///
1483 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1484 /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1485 /// for a typed backend failure.
1486 fn softplus(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1487 /// Exact GELU `x/2 * (1 + erf(x / sqrt(2)))` inside a session.
1488 ///
1489 /// Real `F32`/`F64` only (PyTorch `approximate="none"`).
1490 ///
1491 /// # Examples
1492 ///
1493 /// ```rust
1494 /// use tenferro_cpu::CpuBackend;
1495 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1496 /// use tenferro_tensor::BackendSessionHost;
1497 ///
1498 /// let mut backend = CpuBackend::new();
1499 /// let x = Tensor::from_vec_col_major(vec![3], vec![-1.0_f64, 0.0, 1.0])?;
1500 /// let y = backend.with_backend_session(|session| x.gelu(session))??;
1501 /// let y = y.as_slice::<f64>()?;
1502 /// assert_eq!(y[1], 0.0);
1503 /// assert!((y[2] - 0.841_344_746_068_542_9).abs() < 1e-15);
1504 /// # Ok::<(), Box<dyn std::error::Error>>(())
1505 /// ```
1506 ///
1507 /// # Errors
1508 ///
1509 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1510 /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1511 /// for a typed backend failure.
1512 fn gelu(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1513 /// GELU tanh approximation inside a session (PyTorch `approximate="tanh"`).
1514 ///
1515 /// `x/2 * (1 + tanh(sqrt(2/pi) * (x + 0.044715 x^3)))`; real `F32`/`F64` only.
1516 ///
1517 /// # Examples
1518 ///
1519 /// ```rust
1520 /// use tenferro_cpu::CpuBackend;
1521 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1522 /// use tenferro_tensor::BackendSessionHost;
1523 ///
1524 /// let mut backend = CpuBackend::new();
1525 /// let x = Tensor::from_vec_col_major(vec![3], vec![-1.0_f64, 0.0, 1.0])?;
1526 /// let y = backend.with_backend_session(|session| x.gelu_tanh(session))??;
1527 /// let y = y.as_slice::<f64>()?;
1528 /// assert_eq!(y[1], 0.0);
1529 /// assert!((y[2] - 0.841_191_990_608_276_8).abs() < 1e-12);
1530 /// # Ok::<(), Box<dyn std::error::Error>>(())
1531 /// ```
1532 ///
1533 /// # Errors
1534 ///
1535 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1536 /// integer, or `Bool` input, or [`tenferro_tensor::Error::BackendSource`]
1537 /// for a typed backend failure.
1538 fn gelu_tanh(&self, session: &mut dyn BackendSession) -> tenferro_tensor::Result<Tensor>;
1539 /// Arithmetic mean over `axes` inside a session (`None` reduces every axis).
1540 ///
1541 /// Float and complex dtypes. The sum is divided by the element count; a mean
1542 /// over zero elements is `NaN`, and `Some(&[])` is the identity.
1543 ///
1544 /// # Examples
1545 ///
1546 /// ```rust
1547 /// use tenferro_cpu::CpuBackend;
1548 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1549 /// use tenferro_tensor::BackendSessionHost;
1550 ///
1551 /// let mut backend = CpuBackend::new();
1552 /// let x = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?;
1553 /// let y = backend.with_backend_session(|session| x.reduce_mean(Some(&[1]), session))??;
1554 /// assert_eq!(y.as_slice::<f64>()?, &[2.0, 3.0]);
1555 /// # Ok::<(), Box<dyn std::error::Error>>(())
1556 /// ```
1557 ///
1558 /// # Errors
1559 ///
1560 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for integer or
1561 /// `Bool` input, [`tenferro_tensor::Error::Validation`] with
1562 /// `AxisOutOfBounds` or `DuplicateAxis` for invalid axes, or
1563 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1564 fn reduce_mean(
1565 &self,
1566 axes: Option<&[usize]>,
1567 session: &mut dyn BackendSession,
1568 ) -> tenferro_tensor::Result<Tensor>;
1569 /// Max-subtracted softmax along `axis` inside a session.
1570 ///
1571 /// Real `F32`/`F64` only. A slice that is entirely `-inf` returns zeros
1572 /// instead of `NaN`; a `NaN` or `+inf` entry makes its slice `NaN`.
1573 ///
1574 /// # Examples
1575 ///
1576 /// ```rust
1577 /// use tenferro_cpu::CpuBackend;
1578 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1579 /// use tenferro_tensor::BackendSessionHost;
1580 ///
1581 /// let mut backend = CpuBackend::new();
1582 /// let x = Tensor::from_vec_col_major(vec![2], vec![0.0_f64, f64::NEG_INFINITY])?;
1583 /// let y = backend.with_backend_session(|session| x.softmax(0, session))??;
1584 /// assert_eq!(y.as_slice::<f64>()?, &[1.0, 0.0]);
1585 /// # Ok::<(), Box<dyn std::error::Error>>(())
1586 /// ```
1587 ///
1588 /// # Errors
1589 ///
1590 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1591 /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1592 /// `AxisOutOfBounds` for an invalid axis, or
1593 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1594 fn softmax(
1595 &self,
1596 axis: usize,
1597 session: &mut dyn BackendSession,
1598 ) -> tenferro_tensor::Result<Tensor>;
1599 /// Max-subtracted log-softmax along `axis` inside a session.
1600 ///
1601 /// Real `F32`/`F64` only. A slice that is entirely `-inf` returns `-inf`
1602 /// instead of `NaN`.
1603 ///
1604 /// # Examples
1605 ///
1606 /// ```rust
1607 /// use tenferro_cpu::CpuBackend;
1608 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1609 /// use tenferro_tensor::BackendSessionHost;
1610 ///
1611 /// let mut backend = CpuBackend::new();
1612 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 1.0])?;
1613 /// let y = backend.with_backend_session(|session| x.log_softmax(0, session))??;
1614 /// assert_eq!(y.as_slice::<f64>()?, &[-std::f64::consts::LN_2; 2]);
1615 /// # Ok::<(), Box<dyn std::error::Error>>(())
1616 /// ```
1617 ///
1618 /// # Errors
1619 ///
1620 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1621 /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1622 /// `AxisOutOfBounds` for an invalid axis, or
1623 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1624 fn log_softmax(
1625 &self,
1626 axis: usize,
1627 session: &mut dyn BackendSession,
1628 ) -> tenferro_tensor::Result<Tensor>;
1629 /// Softmax along `axis` over the entries where the `Bool` `mask` is true.
1630 ///
1631 /// `mask` broadcasts to the input shape. Masked-out entries are `0` whatever
1632 /// their value; a slice with no unmasked entry is all zeros.
1633 ///
1634 /// # Examples
1635 ///
1636 /// ```rust
1637 /// use tenferro_cpu::CpuBackend;
1638 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1639 /// use tenferro_tensor::BackendSessionHost;
1640 ///
1641 /// let mut backend = CpuBackend::new();
1642 /// let x = Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 1.0, f64::NAN])?;
1643 /// let mask = Tensor::from_vec_col_major(vec![3], vec![true, true, false])?;
1644 /// let y = backend.with_backend_session(|session| x.masked_softmax(&mask, 0, session))??;
1645 /// assert_eq!(y.as_slice::<f64>()?, &[0.5, 0.5, 0.0]);
1646 /// # Ok::<(), Box<dyn std::error::Error>>(())
1647 /// ```
1648 ///
1649 /// # Errors
1650 ///
1651 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1652 /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1653 /// `DTypeMismatch` for a non-`Bool` mask, `ShapeMismatch` for a mask that
1654 /// does not broadcast to the input, or `AxisOutOfBounds` for an invalid
1655 /// axis, or [`tenferro_tensor::Error::BackendSource`] for a typed backend
1656 /// failure.
1657 fn masked_softmax(
1658 &self,
1659 mask: &Tensor,
1660 axis: usize,
1661 session: &mut dyn BackendSession,
1662 ) -> tenferro_tensor::Result<Tensor>;
1663 /// Log-softmax along `axis` over the entries where the `Bool` `mask` is true.
1664 ///
1665 /// Masked-out entries are `-inf`; a slice with no unmasked entry is all `-inf`.
1666 ///
1667 /// # Examples
1668 ///
1669 /// ```rust
1670 /// use tenferro_cpu::CpuBackend;
1671 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1672 /// use tenferro_tensor::BackendSessionHost;
1673 ///
1674 /// let mut backend = CpuBackend::new();
1675 /// let x = Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 1.0, 5.0])?;
1676 /// let mask = Tensor::from_vec_col_major(vec![3], vec![true, true, false])?;
1677 /// let y = backend.with_backend_session(|session| x.masked_log_softmax(&mask, 0, session))??;
1678 /// let ln_half = -std::f64::consts::LN_2;
1679 /// assert_eq!(y.as_slice::<f64>()?, &[ln_half, ln_half, f64::NEG_INFINITY]);
1680 /// # Ok::<(), Box<dyn std::error::Error>>(())
1681 /// ```
1682 ///
1683 /// # Errors
1684 ///
1685 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1686 /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1687 /// `DTypeMismatch` for a non-`Bool` mask, `ShapeMismatch` for a mask that
1688 /// does not broadcast to the input, or `AxisOutOfBounds` for an invalid
1689 /// axis, or [`tenferro_tensor::Error::BackendSource`] for a typed backend
1690 /// failure.
1691 fn masked_log_softmax(
1692 &self,
1693 mask: &Tensor,
1694 axis: usize,
1695 session: &mut dyn BackendSession,
1696 ) -> tenferro_tensor::Result<Tensor>;
1697 /// Layer normalization along `axis` with optional affine `weight` / `bias`, inside a session.
1698 ///
1699 /// `(x - mean) / sqrt(var + eps) * weight + bias` with the biased variance;
1700 /// `weight` and `bias` are rank-1 of length `shape[axis]`. Real `F32`/`F64` only.
1701 ///
1702 /// # Examples
1703 ///
1704 /// ```rust
1705 /// use tenferro_cpu::CpuBackend;
1706 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1707 /// use tenferro_tensor::BackendSessionHost;
1708 ///
1709 /// let mut backend = CpuBackend::new();
1710 /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 3.0])?;
1711 /// let bias = Tensor::from_vec_col_major(vec![2], vec![10.0_f64, 10.0])?;
1712 /// let y = backend.with_backend_session(|session| x.layer_norm(0, None, Some(&bias), 0.0, session))??;
1713 /// assert_eq!(y.as_slice::<f64>()?, &[9.0, 11.0]);
1714 /// # Ok::<(), Box<dyn std::error::Error>>(())
1715 /// ```
1716 ///
1717 /// # Errors
1718 ///
1719 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1720 /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1721 /// `AxisOutOfBounds` for an invalid axis, `InvalidArgument` for a negative
1722 /// or non-finite `eps`, or `DTypeMismatch` / `ShapeMismatch` for a weight or
1723 /// bias that is not a same-dtype vector of the axis length, or
1724 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1725 fn layer_norm(
1726 &self,
1727 axis: usize,
1728 weight: Option<&Tensor>,
1729 bias: Option<&Tensor>,
1730 eps: f64,
1731 session: &mut dyn BackendSession,
1732 ) -> tenferro_tensor::Result<Tensor>;
1733 /// RMS normalization along `axis` with optional affine `weight` / `bias`, inside a session.
1734 ///
1735 /// `x / sqrt(mean(x^2) + eps) * weight + bias`; `weight` and `bias` are rank-1
1736 /// of length `shape[axis]`. Real `F32`/`F64` only.
1737 ///
1738 /// # Examples
1739 ///
1740 /// ```rust
1741 /// use tenferro_cpu::CpuBackend;
1742 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1743 /// use tenferro_tensor::BackendSessionHost;
1744 ///
1745 /// let mut backend = CpuBackend::new();
1746 /// let x = Tensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0])?;
1747 /// let weight = Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 1.0])?;
1748 /// let y = backend.with_backend_session(|session| x.rms_norm(0, Some(&weight), None, 0.0, session))??;
1749 /// let y = y.as_slice::<f64>()?;
1750 /// let rms = 12.5_f64.sqrt();
1751 /// assert!((y[0] - 6.0 / rms).abs() < 1e-15 && (y[1] - 4.0 / rms).abs() < 1e-15);
1752 /// # Ok::<(), Box<dyn std::error::Error>>(())
1753 /// ```
1754 ///
1755 /// # Errors
1756 ///
1757 /// Returns [`tenferro_tensor::Error::UnsupportedDType`] for complex,
1758 /// integer, or `Bool` input, [`tenferro_tensor::Error::Validation`] with
1759 /// `AxisOutOfBounds` for an invalid axis, `InvalidArgument` for a negative
1760 /// or non-finite `eps`, or `DTypeMismatch` / `ShapeMismatch` for a weight or
1761 /// bias that is not a same-dtype vector of the axis length, or
1762 /// [`tenferro_tensor::Error::BackendSource`] for a typed backend failure.
1763 fn rms_norm(
1764 &self,
1765 axis: usize,
1766 weight: Option<&Tensor>,
1767 bias: Option<&Tensor>,
1768 eps: f64,
1769 session: &mut dyn BackendSession,
1770 ) -> tenferro_tensor::Result<Tensor>;
1771 /// NumPy-style `take_along_axis` over `gather`, inside a session.
1772 ///
1773 /// `out[.., i, ..] = self[.., indices[.., i, ..], ..]` along `axis`. `indices`
1774 /// (I32/I64) has the input's rank; every other dimension is either the
1775 /// input's extent (batch-varying indices) or `1` (the whole extent is taken).
1776 /// Indices must be in bounds.
1777 ///
1778 /// # Examples
1779 ///
1780 /// ```rust
1781 /// use tenferro_cpu::CpuBackend;
1782 /// use tenferro_runtime::{Tensor, TensorSessionOpsExt};
1783 /// use tenferro_tensor::BackendSessionHost;
1784 ///
1785 /// let mut backend = CpuBackend::new();
1786 /// // Per-batch row gather: out[i, j, b] = x[idx[i, b], j, b].
1787 /// let x = Tensor::from_vec_col_major(vec![2, 2, 2], (0..8).map(f64::from).collect::<Vec<_>>())?;
1788 /// let idx = Tensor::from_vec_col_major(vec![2, 1, 2], vec![1_i64, 0, 0, 0])?;
1789 /// let y = backend.with_backend_session(|session| x.take_along_axis(&idx, 0, session))??;
1790 /// assert_eq!(y.as_slice::<f64>()?, &[1.0, 0.0, 3.0, 2.0, 4.0, 4.0, 6.0, 6.0]);
1791 /// # Ok::<(), Box<dyn std::error::Error>>(())
1792 /// ```
1793 ///
1794 /// # Errors
1795 ///
1796 /// Returns [`tenferro_tensor::Error::Validation`] with `RankMismatch` or
1797 /// `ShapeMismatch` for incompatible index shapes, `AxisOutOfBounds` for an
1798 /// invalid axis, or `InvalidArgument` when taking from a zero-length axis;
1799 /// [`tenferro_tensor::Error::UnsupportedDType`] for a non-integer index
1800 /// dtype; or [`tenferro_tensor::Error::BackendSource`] for a typed backend
1801 /// failure.
1802 fn take_along_axis(
1803 &self,
1804 indices: &Tensor,
1805 axis: usize,
1806 session: &mut dyn BackendSession,
1807 ) -> tenferro_tensor::Result<Tensor>;
1808}