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 tenferro_ops::broadcast::{
7    broadcast_input_plan, broadcast_shape, broadcast_shapes, BroadcastError,
8};
9use tenferro_tensor::validate::matmul_config_for_shapes;
10use tenferro_tensor::{
11    BackendSession, CompareDir, DType, Error, Result, Tensor, TensorRead, TensorScalar,
12    ValidationError,
13};
14
15use crate::{TypedTensorMaskSessionOpsExt, TypedTensorSessionOpsExt};
16use tenferro_tensor::TypedTensor;
17
18impl<T: TensorScalar> TypedTensorSessionOpsExt<T> for TypedTensor<T> {
19    fn add(
20        &self,
21        rhs: &TypedTensor<T>,
22        session: &mut dyn BackendSession,
23    ) -> Result<TypedTensor<T>> {
24        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
25        let out = session.add_read(lhs.tensor_read(), rhs.tensor_read())?;
26        into_typed_result("add", out)
27    }
28
29    fn mul(
30        &self,
31        rhs: &TypedTensor<T>,
32        session: &mut dyn BackendSession,
33    ) -> Result<TypedTensor<T>> {
34        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
35        let out = session.mul_read(lhs.tensor_read(), rhs.tensor_read())?;
36        into_typed_result("mul", out)
37    }
38
39    fn exp(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
40        let out = session.exp_read(T::tensor_read(self))?;
41        into_typed_result("exp", out)
42    }
43
44    fn reduce_sum(
45        &self,
46        axes: &[usize],
47        session: &mut dyn BackendSession,
48    ) -> Result<TypedTensor<T>> {
49        let out = session.reduce_sum_read(T::tensor_read(self), axes)?;
50        into_typed_result("reduce_sum", out)
51    }
52
53    fn sub(
54        &self,
55        rhs: &TypedTensor<T>,
56        session: &mut dyn BackendSession,
57    ) -> Result<TypedTensor<T>> {
58        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
59        let out = session.sub_read(lhs.tensor_read(), rhs.tensor_read())?;
60        into_typed_result("sub", out)
61    }
62
63    fn div(
64        &self,
65        rhs: &TypedTensor<T>,
66        session: &mut dyn BackendSession,
67    ) -> Result<TypedTensor<T>> {
68        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
69        let out = session.div_read(lhs.tensor_read(), rhs.tensor_read())?;
70        into_typed_result("div", out)
71    }
72
73    fn rem(
74        &self,
75        rhs: &TypedTensor<T>,
76        session: &mut dyn BackendSession,
77    ) -> Result<TypedTensor<T>> {
78        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
79        let out = session.rem_read(lhs.tensor_read(), rhs.tensor_read())?;
80        into_typed_result("rem", out)
81    }
82
83    fn pow(
84        &self,
85        rhs: &TypedTensor<T>,
86        session: &mut dyn BackendSession,
87    ) -> Result<TypedTensor<T>> {
88        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
89        let out = session.pow_read(lhs.tensor_read(), rhs.tensor_read())?;
90        into_typed_result("pow", out)
91    }
92
93    fn maximum(
94        &self,
95        rhs: &TypedTensor<T>,
96        session: &mut dyn BackendSession,
97    ) -> Result<TypedTensor<T>> {
98        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
99        let out = session.maximum_read(lhs.tensor_read(), rhs.tensor_read())?;
100        into_typed_result("maximum", out)
101    }
102
103    fn minimum(
104        &self,
105        rhs: &TypedTensor<T>,
106        session: &mut dyn BackendSession,
107    ) -> Result<TypedTensor<T>> {
108        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
109        let out = session.minimum_read(lhs.tensor_read(), rhs.tensor_read())?;
110        into_typed_result("minimum", out)
111    }
112
113    fn neg(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
114        let out = session.neg_read(T::tensor_read(self))?;
115        into_typed_result("neg", out)
116    }
117
118    fn abs(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
119        let out = session.abs_read(T::tensor_read(self))?;
120        into_typed_result("abs", out)
121    }
122
123    fn sign(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
124        let out = session.sign_read(T::tensor_read(self))?;
125        into_typed_result("sign", out)
126    }
127
128    fn conj(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
129        let out = session.conj_read(T::tensor_read(self))?;
130        into_typed_result("conj", out)
131    }
132
133    fn log(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
134        let out = session.log_read(T::tensor_read(self))?;
135        into_typed_result("log", out)
136    }
137
138    fn expm1(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
139        let out = session.expm1_read(T::tensor_read(self))?;
140        into_typed_result("expm1", out)
141    }
142
143    fn log1p(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
144        let out = session.log1p_read(T::tensor_read(self))?;
145        into_typed_result("log1p", out)
146    }
147
148    fn sin(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
149        let out = session.sin_read(T::tensor_read(self))?;
150        into_typed_result("sin", out)
151    }
152
153    fn cos(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
154        let out = session.cos_read(T::tensor_read(self))?;
155        into_typed_result("cos", out)
156    }
157
158    fn tanh(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
159        let out = session.tanh_read(T::tensor_read(self))?;
160        into_typed_result("tanh", out)
161    }
162
163    fn sqrt(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
164        let out = session.sqrt_read(T::tensor_read(self))?;
165        into_typed_result("sqrt", out)
166    }
167
168    fn rsqrt(&self, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
169        let out = session.rsqrt_read(T::tensor_read(self))?;
170        into_typed_result("rsqrt", out)
171    }
172
173    fn compare(
174        &self,
175        rhs: &TypedTensor<T>,
176        dir: CompareDir,
177        session: &mut dyn BackendSession,
178    ) -> Result<TypedTensor<bool>> {
179        let (lhs, rhs) = broadcast_binary_in_read(self, rhs, session)?;
180        let out = session.compare_read(lhs.tensor_read(), rhs.tensor_read(), &dir)?;
181        into_typed_result("compare", out)
182    }
183
184    fn clamp(
185        &self,
186        lower: &TypedTensor<T>,
187        upper: &TypedTensor<T>,
188        session: &mut dyn BackendSession,
189    ) -> Result<TypedTensor<T>> {
190        let (input, lower, upper) = broadcast_ternary_in_read(self, lower, upper, session)?;
191        let out = session.clamp_read(
192            input.tensor_read(),
193            lower.tensor_read(),
194            upper.tensor_read(),
195        )?;
196        into_typed_result("clamp", out)
197    }
198
199    fn matmul(
200        &self,
201        rhs: &TypedTensor<T>,
202        session: &mut dyn BackendSession,
203    ) -> Result<TypedTensor<T>> {
204        let config = matmul_config_for_shapes("matmul", self.shape(), rhs.shape())?;
205        let out = session.dot_general_read(T::tensor_read(self), T::tensor_read(rhs), &config)?;
206        into_typed_result("matmul", out)
207    }
208
209    fn reshape(&self, shape: &[usize], session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
210        let out = session.reshape_read(T::tensor_read(self), shape)?;
211        into_typed_result("reshape", out)
212    }
213
214    fn transpose(
215        &self,
216        perm: &[usize],
217        session: &mut dyn BackendSession,
218    ) -> Result<TypedTensor<T>> {
219        let out = session.transpose_read(T::tensor_read(self), perm)?;
220        into_typed_result("transpose", out)
221    }
222
223    fn broadcast_in_dim(
224        &self,
225        shape: &[usize],
226        dims: &[usize],
227        session: &mut dyn BackendSession,
228    ) -> Result<TypedTensor<T>> {
229        let out = session.broadcast_in_dim_read(T::tensor_read(self), shape, dims)?;
230        into_typed_result("broadcast_in_dim", out)
231    }
232}
233
234impl TypedTensorMaskSessionOpsExt for TypedTensor<bool> {
235    fn where_select<U: TensorScalar>(
236        &self,
237        on_true: &TypedTensor<U>,
238        on_false: &TypedTensor<U>,
239        session: &mut dyn BackendSession,
240    ) -> Result<TypedTensor<U>> {
241        let (condition, on_true, on_false) =
242            broadcast_ternary_in_read(self, on_true, on_false, session)?;
243        let out = session.select_read(
244            condition.tensor_read(),
245            on_true.tensor_read(),
246            on_false.tensor_read(),
247        )?;
248        into_typed_result("where_select", out)
249    }
250}
251
252// INVARIANT: this private adapter keeps borrowed reads borrowed and owns only
253// the explicit fallback tensor; it is never exposed or cloned.
254#[allow(clippy::large_enum_variant)]
255enum ReadInput<'a> {
256    Borrowed(TensorRead<'a>),
257    Owned(Tensor),
258}
259
260impl ReadInput<'_> {
261    fn tensor_read(&self) -> TensorRead<'_> {
262        match self {
263            Self::Borrowed(read) => read.clone(),
264            Self::Owned(tensor) => TensorRead::from_tensor(tensor),
265        }
266    }
267}
268
269fn broadcast_to_in_read<'a, T: TensorScalar>(
270    input: &'a TypedTensor<T>,
271    target_shape: &[usize],
272    session: &mut dyn BackendSession,
273) -> Result<ReadInput<'a>> {
274    if input.shape() == target_shape {
275        return Ok(ReadInput::Borrowed(T::tensor_read(input)));
276    }
277
278    let plan = broadcast_input_plan(input.shape(), target_shape).map_err(broadcast_error)?;
279    let source = if plan.source_shape == input.shape() {
280        ReadInput::Borrowed(T::tensor_read(input))
281    } else {
282        let reshaped = session.reshape_read(T::tensor_read(input), &plan.source_shape)?;
283        ReadInput::Owned(reshaped)
284    };
285    let out = session.broadcast_in_dim_read(source.tensor_read(), target_shape, &plan.dims)?;
286    Ok(ReadInput::Owned(out))
287}
288
289fn broadcast_binary_in_read<'a, T: TensorScalar>(
290    lhs: &'a TypedTensor<T>,
291    rhs: &'a TypedTensor<T>,
292    session: &mut dyn BackendSession,
293) -> Result<(ReadInput<'a>, ReadInput<'a>)> {
294    let shape = broadcast_shape(lhs.shape(), rhs.shape()).map_err(broadcast_error)?;
295    Ok((
296        broadcast_to_in_read(lhs, &shape, session)?,
297        broadcast_to_in_read(rhs, &shape, session)?,
298    ))
299}
300
301fn broadcast_ternary_in_read<'a, C: TensorScalar, T: TensorScalar>(
302    first: &'a TypedTensor<C>,
303    second: &'a TypedTensor<T>,
304    third: &'a TypedTensor<T>,
305    session: &mut dyn BackendSession,
306) -> Result<(ReadInput<'a>, ReadInput<'a>, ReadInput<'a>)> {
307    let shape = broadcast_shapes([first.shape(), second.shape(), third.shape()])
308        .map_err(broadcast_error)?;
309    Ok((
310        broadcast_to_in_read(first, &shape, session)?,
311        broadcast_to_in_read(second, &shape, session)?,
312        broadcast_to_in_read(third, &shape, session)?,
313    ))
314}
315
316fn broadcast_error(err: BroadcastError) -> Error {
317    match err {
318        BroadcastError::IncompatibleBinary { lhs, rhs } => {
319            Error::shape_mismatch("broadcast", lhs, rhs)
320        }
321        BroadcastError::IncompatibleInput { input, output } => {
322            Error::shape_mismatch("broadcast", input, output)
323        }
324        BroadcastError::RankTooLarge { input, output } => {
325            Error::rank_mismatch("broadcast", output.len(), input.len())
326        }
327    }
328}
329
330fn into_typed_result<T: TensorScalar>(op: &'static str, tensor: Tensor) -> Result<TypedTensor<T>> {
331    let actual = tensor.dtype();
332    T::into_typed(tensor).map_err(|_| {
333        Error::validation(
334            op,
335            ValidationError::DTypeMismatch {
336                expected: core_dtype(T::dtype()),
337                actual: core_dtype(actual),
338            },
339        )
340    })
341}
342
343fn core_dtype(dtype: DType) -> tenferro_tensor::core::DType {
344    match dtype {
345        DType::F32 => tenferro_tensor::core::DType::F32,
346        DType::F64 => tenferro_tensor::core::DType::F64,
347        DType::I32 => tenferro_tensor::core::DType::I32,
348        DType::I64 => tenferro_tensor::core::DType::I64,
349        DType::Bool => tenferro_tensor::core::DType::Bool,
350        DType::C32 => tenferro_tensor::core::DType::C32,
351        DType::C64 => tenferro_tensor::core::DType::C64,
352    }
353}