Skip to main content

tenferro_runtime/
typed_tensor.rs

1//! Typed tensor operation extension traits.
2//!
3//! Operation families that are no longer part of core, including einsum, live
4//! in their extension crates.
5
6use num_complex::Complex64;
7use tenferro_ops::broadcast::{
8    broadcast_input_plan, broadcast_shape, broadcast_shapes, BroadcastError,
9};
10use tenferro_tensor::validate::matmul_config_for_shapes;
11use tenferro_tensor::{
12    BackendSession, CompareDir, DotGeneralConfig, Error, Result, Tensor, TensorRead, TensorScalar,
13    ValidationError,
14};
15
16use crate::composite;
17use crate::composite::session::{run_session_composite, typed_borrowed};
18use crate::{TypedTensorMaskSessionOpsExt, TypedTensorSessionOpsExt};
19use tenferro_tensor::TypedTensor;
20
21impl<T: TensorScalar> TypedTensorSessionOpsExt<T> for TypedTensor<T> {
22    fn add(
23        &self,
24        rhs: &TypedTensor<T>,
25        session: &mut dyn BackendSession,
26    ) -> Result<TypedTensor<T>> {
27        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
28        let out = session.add_read(lhs.tensor_read(), rhs.tensor_read())?;
29        into_typed_result("add", out)
30    }
31
32    fn mul(
33        &self,
34        rhs: &TypedTensor<T>,
35        session: &mut dyn BackendSession,
36    ) -> Result<TypedTensor<T>> {
37        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
38        let out = session.mul_read(lhs.tensor_read(), rhs.tensor_read())?;
39        into_typed_result("mul", out)
40    }
41
42    fn exp(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
43        let out = session.exp_read(T::tensor_read(self))?;
44        into_typed_result("exp", out)
45    }
46
47    fn reduce_sum(
48        &self,
49        axes: Option<&[usize]>,
50        session: &mut dyn BackendSession,
51    ) -> Result<TypedTensor<T>> {
52        let axes = crate::tensor::all_axes_if_none(self.shape().len(), axes);
53        let out = session.reduce_sum_read(T::tensor_read(self), &axes)?;
54        into_typed_result("reduce_sum", out)
55    }
56
57    fn sub(
58        &self,
59        rhs: &TypedTensor<T>,
60        session: &mut dyn BackendSession,
61    ) -> Result<TypedTensor<T>> {
62        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
63        let out = session.sub_read(lhs.tensor_read(), rhs.tensor_read())?;
64        into_typed_result("sub", out)
65    }
66
67    fn div(
68        &self,
69        rhs: &TypedTensor<T>,
70        session: &mut dyn BackendSession,
71    ) -> Result<TypedTensor<T>> {
72        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
73        let out = session.div_read(lhs.tensor_read(), rhs.tensor_read())?;
74        into_typed_result("div", out)
75    }
76
77    fn rem(
78        &self,
79        rhs: &TypedTensor<T>,
80        session: &mut dyn BackendSession,
81    ) -> Result<TypedTensor<T>> {
82        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
83        let out = session.rem_read(lhs.tensor_read(), rhs.tensor_read())?;
84        into_typed_result("rem", out)
85    }
86
87    fn pow(
88        &self,
89        rhs: &TypedTensor<T>,
90        session: &mut dyn BackendSession,
91    ) -> Result<TypedTensor<T>> {
92        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
93        let out = session.pow_read(lhs.tensor_read(), rhs.tensor_read())?;
94        into_typed_result("pow", out)
95    }
96
97    fn maximum(
98        &self,
99        rhs: &TypedTensor<T>,
100        session: &mut dyn BackendSession,
101    ) -> Result<TypedTensor<T>> {
102        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
103        let out = session.maximum_read(lhs.tensor_read(), rhs.tensor_read())?;
104        into_typed_result("maximum", out)
105    }
106
107    fn minimum(
108        &self,
109        rhs: &TypedTensor<T>,
110        session: &mut dyn BackendSession,
111    ) -> Result<TypedTensor<T>> {
112        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
113        let out = session.minimum_read(lhs.tensor_read(), rhs.tensor_read())?;
114        into_typed_result("minimum", out)
115    }
116
117    fn neg(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
118        let out = session.neg_read(T::tensor_read(self))?;
119        into_typed_result("neg", out)
120    }
121
122    fn abs(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T::Real>> {
123        let out = session.abs_read(T::tensor_read(self))?;
124        into_typed_result::<T::Real>("abs", out)
125    }
126
127    fn sign(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
128        let out = session.sign_read(T::tensor_read(self))?;
129        into_typed_result("sign", out)
130    }
131
132    fn conj(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
133        let out = session.conj_read(T::tensor_read(self))?;
134        into_typed_result("conj", out)
135    }
136
137    fn log(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
138        let out = session.log_read(T::tensor_read(self))?;
139        into_typed_result("log", out)
140    }
141
142    fn expm1(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
143        let out = session.expm1_read(T::tensor_read(self))?;
144        into_typed_result("expm1", out)
145    }
146
147    fn log1p(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
148        let out = session.log1p_read(T::tensor_read(self))?;
149        into_typed_result("log1p", out)
150    }
151
152    fn erf(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
153        let out = session.erf_read(T::tensor_read(self))?;
154        into_typed_result("erf", out)
155    }
156
157    fn sin(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
158        let out = session.sin_read(T::tensor_read(self))?;
159        into_typed_result("sin", out)
160    }
161
162    fn cos(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
163        let out = session.cos_read(T::tensor_read(self))?;
164        into_typed_result("cos", out)
165    }
166
167    fn tanh(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
168        let out = session.tanh_read(T::tensor_read(self))?;
169        into_typed_result("tanh", out)
170    }
171
172    fn sqrt(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
173        let out = session.sqrt_read(T::tensor_read(self))?;
174        into_typed_result("sqrt", out)
175    }
176
177    fn rsqrt(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
178        let out = session.rsqrt_read(T::tensor_read(self))?;
179        into_typed_result("rsqrt", out)
180    }
181
182    fn compare(
183        &self,
184        rhs: &TypedTensor<T>,
185        dir: CompareDir,
186        session: &mut dyn BackendSession,
187    ) -> Result<TypedTensor<bool>> {
188        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
189        let out = session.compare_read(lhs.tensor_read(), rhs.tensor_read(), &dir)?;
190        into_typed_result("compare", out)
191    }
192
193    fn clamp(
194        &self,
195        lower: &TypedTensor<T>,
196        upper: &TypedTensor<T>,
197        session: &mut dyn BackendSession,
198    ) -> Result<TypedTensor<T>> {
199        let (input, lower, upper) = broadcast_ternary_in_read(self, lower, upper, session)?;
200        let out = session.clamp_read(
201            input.tensor_read(),
202            lower.tensor_read(),
203            upper.tensor_read(),
204        )?;
205        into_typed_result("clamp", out)
206    }
207
208    fn matmul(
209        &self,
210        rhs: &TypedTensor<T>,
211        session: &mut dyn BackendSession,
212    ) -> Result<TypedTensor<T>> {
213        let config = matmul_config_for_shapes("matmul", self.shape(), rhs.shape())?;
214        let out = session.dot_general_read(T::tensor_read(self), T::tensor_read(rhs), &config)?;
215        into_typed_result("matmul", out)
216    }
217
218    fn reshape(&self, shape: &[usize], session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
219        let out = session.reshape_read(T::tensor_read(self), shape)?;
220        into_typed_result("reshape", out)
221    }
222
223    fn transpose(
224        &self,
225        perm: &[usize],
226        session: &mut dyn BackendSession,
227    ) -> Result<TypedTensor<T>> {
228        let out = session.transpose_read(T::tensor_read(self), perm)?;
229        into_typed_result("transpose", out)
230    }
231
232    fn broadcast_in_dim(
233        &self,
234        shape: &[usize],
235        dims: &[usize],
236        session: &mut dyn BackendSession,
237    ) -> Result<TypedTensor<T>> {
238        let out = session.broadcast_in_dim_read(T::tensor_read(self), shape, dims)?;
239        into_typed_result("broadcast_in_dim", out)
240    }
241
242    fn reduce_max(
243        &self,
244        axes: Option<&[usize]>,
245        session: &mut dyn BackendSession,
246    ) -> Result<TypedTensor<T>> {
247        let axes = crate::tensor::all_axes_if_none(self.shape().len(), axes);
248        let out = session.reduce_max_read(T::tensor_read(self), &axes)?;
249        into_typed_result("reduce_max", out)
250    }
251
252    fn reduce_min(
253        &self,
254        axes: Option<&[usize]>,
255        session: &mut dyn BackendSession,
256    ) -> Result<TypedTensor<T>> {
257        let axes = crate::tensor::all_axes_if_none(self.shape().len(), axes);
258        let out = session.reduce_min_read(T::tensor_read(self), &axes)?;
259        into_typed_result("reduce_min", out)
260    }
261
262    fn reduce_prod(
263        &self,
264        axes: Option<&[usize]>,
265        session: &mut dyn BackendSession,
266    ) -> Result<TypedTensor<T>> {
267        let axes = crate::tensor::all_axes_if_none(self.shape().len(), axes);
268        let out = session.reduce_prod_read(T::tensor_read(self), &axes)?;
269        into_typed_result("reduce_prod", out)
270    }
271
272    fn reduce_sum_squares(
273        &self,
274        axes: Option<&[usize]>,
275        session: &mut dyn BackendSession,
276    ) -> Result<TypedTensor<T>> {
277        let axes = crate::tensor::all_axes_if_none(self.shape().len(), axes);
278        let out = session.reduce_sum_squares_read(T::tensor_read(self), &axes)?;
279        into_typed_result("reduce_sum_squares", out)
280    }
281
282    fn dot_general(
283        &self,
284        rhs: &TypedTensor<T>,
285        config: DotGeneralConfig,
286        session: &mut dyn BackendSession,
287    ) -> Result<TypedTensor<T>> {
288        let out = session.dot_general_read(T::tensor_read(self), T::tensor_read(rhs), &config)?;
289        into_typed_result("dot_general", out)
290    }
291
292    fn dot_general_with_conj(
293        &self,
294        rhs: &TypedTensor<T>,
295        config: DotGeneralConfig,
296        lhs_conj: bool,
297        rhs_conj: bool,
298        session: &mut dyn BackendSession,
299    ) -> Result<TypedTensor<T>> {
300        let out = session.dot_general_with_conj_read(
301            T::tensor_read(self),
302            T::tensor_read(rhs),
303            &config,
304            lhs_conj,
305            rhs_conj,
306        )?;
307        into_typed_result("dot_general_with_conj", out)
308    }
309
310    fn scale_real(&self, factor: f64, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
311        let scalar = crate::scale::real_scale_scalar(T::dtype(), factor)?;
312        let scalar = session.upload_host_tensor(TensorRead::from_tensor(&scalar))?;
313        let scalar = into_typed_result::<T>("scale_real", scalar)?;
314        TypedTensorSessionOpsExt::mul(self, &scalar, session)
315    }
316
317    fn scale_complex(
318        &self,
319        factor: Complex64,
320        session: &mut dyn BackendSession,
321    ) -> Result<TypedTensor<T>> {
322        let scalar = crate::scale::complex_scale_scalar(T::dtype(), factor)?;
323        let scalar = session.upload_host_tensor(TensorRead::from_tensor(&scalar))?;
324        let scalar = into_typed_result::<T>("scale_complex", scalar)?;
325        TypedTensorSessionOpsExt::mul(self, &scalar, session)
326    }
327
328    fn sigmoid(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
329        let out = run_session_composite(session, |ops| {
330            composite::sigmoid(ops, &typed_borrowed(self))
331        })?;
332        into_typed_result("sigmoid", out)
333    }
334
335    fn silu(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
336        let out =
337            run_session_composite(session, |ops| composite::silu(ops, &typed_borrowed(self)))?;
338        into_typed_result("silu", out)
339    }
340
341    fn softplus(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
342        let out = run_session_composite(session, |ops| {
343            composite::softplus(ops, &typed_borrowed(self))
344        })?;
345        into_typed_result("softplus", out)
346    }
347
348    fn gelu(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
349        let out =
350            run_session_composite(session, |ops| composite::gelu(ops, &typed_borrowed(self)))?;
351        into_typed_result("gelu", out)
352    }
353
354    fn gelu_tanh(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
355        let out = run_session_composite(session, |ops| {
356            composite::gelu_tanh(ops, &typed_borrowed(self))
357        })?;
358        into_typed_result("gelu_tanh", out)
359    }
360
361    fn reduce_mean(
362        &self,
363        axes: Option<&[usize]>,
364        session: &mut dyn BackendSession,
365    ) -> Result<TypedTensor<T>> {
366        let out = run_session_composite(session, |ops| {
367            composite::reduce_mean(ops, &typed_borrowed(self), axes)
368        })?;
369        into_typed_result("reduce_mean", out)
370    }
371
372    fn softmax(&self, axis: usize, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
373        let out = run_session_composite(session, |ops| {
374            composite::softmax(ops, &typed_borrowed(self), axis)
375        })?;
376        into_typed_result("softmax", out)
377    }
378
379    fn log_softmax(&self, axis: usize, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
380        let out = run_session_composite(session, |ops| {
381            composite::log_softmax(ops, &typed_borrowed(self), axis)
382        })?;
383        into_typed_result("log_softmax", out)
384    }
385
386    fn masked_softmax(
387        &self,
388        mask: &TypedTensor<bool>,
389        axis: usize,
390        session: &mut dyn BackendSession,
391    ) -> Result<TypedTensor<T>> {
392        let mask = typed_borrowed(mask);
393        let out = run_session_composite(session, |ops| {
394            composite::masked_softmax(ops, &typed_borrowed(self), &mask, axis)
395        })?;
396        into_typed_result("masked_softmax", out)
397    }
398
399    fn masked_log_softmax(
400        &self,
401        mask: &TypedTensor<bool>,
402        axis: usize,
403        session: &mut dyn BackendSession,
404    ) -> Result<TypedTensor<T>> {
405        let mask = typed_borrowed(mask);
406        let out = run_session_composite(session, |ops| {
407            composite::masked_log_softmax(ops, &typed_borrowed(self), &mask, axis)
408        })?;
409        into_typed_result("masked_log_softmax", out)
410    }
411
412    fn layer_norm(
413        &self,
414        axis: usize,
415        weight: Option<&TypedTensor<T>>,
416        bias: Option<&TypedTensor<T>>,
417        eps: f64,
418        session: &mut dyn BackendSession,
419    ) -> Result<TypedTensor<T>> {
420        let weight = weight.map(typed_borrowed);
421        let bias = bias.map(typed_borrowed);
422        let out = run_session_composite(session, |ops| {
423            composite::layer_norm(
424                ops,
425                &typed_borrowed(self),
426                axis,
427                weight.as_ref(),
428                bias.as_ref(),
429                eps,
430            )
431        })?;
432        into_typed_result("layer_norm", out)
433    }
434
435    fn rms_norm(
436        &self,
437        axis: usize,
438        weight: Option<&TypedTensor<T>>,
439        bias: Option<&TypedTensor<T>>,
440        eps: f64,
441        session: &mut dyn BackendSession,
442    ) -> Result<TypedTensor<T>> {
443        let weight = weight.map(typed_borrowed);
444        let bias = bias.map(typed_borrowed);
445        let out = run_session_composite(session, |ops| {
446            composite::rms_norm(
447                ops,
448                &typed_borrowed(self),
449                axis,
450                weight.as_ref(),
451                bias.as_ref(),
452                eps,
453            )
454        })?;
455        into_typed_result("rms_norm", out)
456    }
457}
458
459impl TypedTensorMaskSessionOpsExt for TypedTensor<bool> {
460    fn where_select<U: TensorScalar>(
461        &self,
462        on_true: &TypedTensor<U>,
463        on_false: &TypedTensor<U>,
464        session: &mut dyn BackendSession,
465    ) -> Result<TypedTensor<U>> {
466        let (condition, on_true, on_false) =
467            broadcast_ternary_in_read(self, on_true, on_false, session)?;
468        let out = session.select_read(
469            condition.tensor_read(),
470            on_true.tensor_read(),
471            on_false.tensor_read(),
472        )?;
473        into_typed_result("where_select", out)
474    }
475}
476
477// INVARIANT: this private adapter keeps borrowed reads borrowed and owns only
478// the explicit fallback tensor; it is never exposed or cloned.
479#[allow(clippy::large_enum_variant)]
480pub(crate) enum ReadInput<'a> {
481    Borrowed(TensorRead<'a>),
482    Owned(Tensor),
483}
484
485impl ReadInput<'_> {
486    pub(crate) fn tensor_read(&self) -> TensorRead<'_> {
487        match self {
488            Self::Borrowed(read) => read.clone(),
489            Self::Owned(tensor) => TensorRead::from_tensor(tensor),
490        }
491    }
492}
493
494pub(crate) fn broadcast_to_in_read<'a>(
495    input: TensorRead<'a>,
496    target_shape: &[usize],
497    session: &mut dyn BackendSession,
498) -> Result<ReadInput<'a>> {
499    if input.shape() == target_shape {
500        return Ok(ReadInput::Borrowed(input));
501    }
502
503    let plan = broadcast_input_plan(input.shape(), target_shape).map_err(broadcast_error)?;
504    let source = if plan.source_shape == input.shape() {
505        ReadInput::Borrowed(input)
506    } else {
507        let reshaped = session.reshape_read(input, &plan.source_shape)?;
508        ReadInput::Owned(reshaped)
509    };
510    let out = session.broadcast_in_dim_read(source.tensor_read(), target_shape, &plan.dims)?;
511    Ok(ReadInput::Owned(out))
512}
513
514fn broadcast_binary_in_read<'a, T: TensorScalar>(
515    lhs: &'a TypedTensor<T>,
516    rhs: &'a TypedTensor<T>,
517    session: &mut dyn BackendSession,
518) -> Result<(ReadInput<'a>, ReadInput<'a>)> {
519    let shape = broadcast_shape(lhs.shape(), rhs.shape()).map_err(broadcast_error)?;
520    Ok((
521        broadcast_to_in_read(T::tensor_read(lhs), &shape, session)?,
522        broadcast_to_in_read(T::tensor_read(rhs), &shape, session)?,
523    ))
524}
525
526fn broadcast_ternary_in_read<'a, C: TensorScalar, T: TensorScalar>(
527    first: &'a TypedTensor<C>,
528    second: &'a TypedTensor<T>,
529    third: &'a TypedTensor<T>,
530    session: &mut dyn BackendSession,
531) -> Result<(ReadInput<'a>, ReadInput<'a>, ReadInput<'a>)> {
532    let shape = broadcast_shapes([first.shape(), second.shape(), third.shape()])
533        .map_err(broadcast_error)?;
534    Ok((
535        broadcast_to_in_read(C::tensor_read(first), &shape, session)?,
536        broadcast_to_in_read(T::tensor_read(second), &shape, session)?,
537        broadcast_to_in_read(T::tensor_read(third), &shape, session)?,
538    ))
539}
540
541pub(crate) fn broadcast_error(err: BroadcastError) -> Error {
542    match err {
543        BroadcastError::IncompatibleBinary { lhs, rhs } => {
544            Error::shape_mismatch("broadcast", lhs, rhs)
545        }
546        BroadcastError::IncompatibleInput { input, output } => {
547            Error::shape_mismatch("broadcast", input, output)
548        }
549        BroadcastError::RankTooLarge { input, output } => {
550            Error::rank_mismatch("broadcast", output.len(), input.len())
551        }
552    }
553}
554
555fn into_typed_result<T: TensorScalar>(op: &'static str, tensor: Tensor) -> Result<TypedTensor<T>> {
556    let actual = tensor.dtype();
557    T::into_typed(tensor).map_err(|_| {
558        Error::validation(
559            op,
560            ValidationError::DTypeMismatch {
561                expected: T::dtype(),
562                actual,
563            },
564        )
565    })
566}