Skip to main content

tenferro_runtime/
tensor.rs

1//! Concrete tensor operation extension trait.
2//!
3//! `tenferro-tensor` owns storage and backend traits. This runtime crate
4//! provides backend-parametric session-explicit operation methods through
5//! [`TensorSessionOpsExt`].
6
7use tenferro_ops::broadcast::{
8    broadcast_error_to_validation, broadcast_input_plan, broadcast_shape, broadcast_shapes,
9};
10use tenferro_tensor::validate::matmul_config_for_shapes;
11use tenferro_tensor::{BackendSession, CompareDir, DType, Error, Result};
12
13use crate::TensorSessionOpsExt;
14use tenferro_tensor::Tensor;
15
16impl TensorSessionOpsExt for Tensor {
17    fn add(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
18        let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
19        session.add(&lhs, &rhs)
20    }
21
22    fn mul(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
23        let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
24        session.mul(&lhs, &rhs)
25    }
26
27    fn exp(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
28        session.exp(self)
29    }
30
31    fn reduce_sum(&self, axes: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
32        session.reduce_sum(self, axes)
33    }
34
35    fn convert(&self, to: DType, session: &mut dyn BackendSession) -> Result<Tensor> {
36        session.convert(self, to)
37    }
38
39    fn cast(&self, to: DType, session: &mut dyn BackendSession) -> Result<Tensor> {
40        session.cast(self, to)
41    }
42
43    fn sub(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
44        let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
45        session.sub(&lhs, &rhs)
46    }
47
48    fn div(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
49        let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
50        session.div(&lhs, &rhs)
51    }
52
53    fn rem(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
54        let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
55        session.rem(&lhs, &rhs)
56    }
57
58    fn pow(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
59        let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
60        session.pow(&lhs, &rhs)
61    }
62
63    fn maximum(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
64        let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
65        session.maximum(&lhs, &rhs)
66    }
67
68    fn minimum(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
69        let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
70        session.minimum(&lhs, &rhs)
71    }
72
73    fn neg(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
74        session.neg(self)
75    }
76
77    fn abs(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
78        session.abs(self)
79    }
80
81    fn sign(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
82        session.sign(self)
83    }
84
85    fn conj(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
86        session.conj(self)
87    }
88
89    fn log(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
90        session.log(self)
91    }
92
93    fn expm1(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
94        session.expm1(self)
95    }
96
97    fn log1p(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
98        session.log1p(self)
99    }
100
101    fn sin(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
102        session.sin(self)
103    }
104
105    fn cos(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
106        session.cos(self)
107    }
108
109    fn tanh(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
110        session.tanh(self)
111    }
112
113    fn sqrt(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
114        session.sqrt(self)
115    }
116
117    fn rsqrt(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
118        session.rsqrt(self)
119    }
120
121    fn compare(
122        &self,
123        rhs: &Tensor,
124        dir: CompareDir,
125        session: &mut dyn BackendSession,
126    ) -> Result<Tensor> {
127        let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
128        session.compare(&lhs, &rhs, &dir)
129    }
130
131    fn where_select(
132        &self,
133        on_true: &Tensor,
134        on_false: &Tensor,
135        session: &mut dyn BackendSession,
136    ) -> Result<Tensor> {
137        let (condition, on_true, on_false) =
138            broadcast_ternary_in(self, on_true, on_false, session)?;
139        session.select(&condition, &on_true, &on_false)
140    }
141
142    fn clamp(
143        &self,
144        lower: &Tensor,
145        upper: &Tensor,
146        session: &mut dyn BackendSession,
147    ) -> Result<Tensor> {
148        let (input, lower, upper) = broadcast_ternary_in(self, lower, upper, session)?;
149        session.clamp(&input, &lower, &upper)
150    }
151
152    fn matmul(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
153        let config = matmul_config_for_shapes("matmul", self.shape(), rhs.shape())?;
154        session.dot_general(self, rhs, &config)
155    }
156
157    fn reshape(&self, shape: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
158        session.reshape(self, shape)
159    }
160
161    fn transpose(&self, perm: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
162        session.transpose(self, perm)
163    }
164}
165
166fn broadcast_to_in(
167    input: &Tensor,
168    target_shape: &[usize],
169    session: &mut dyn BackendSession,
170) -> Result<Tensor> {
171    let input_shape = input.shape();
172    if input_shape == target_shape {
173        return input.duplicate();
174    }
175
176    let plan = broadcast_input_plan(input_shape, target_shape).map_err(broadcast_error)?;
177    let source = if plan.source_shape == input_shape {
178        input.duplicate()?
179    } else {
180        session.reshape(input, &plan.source_shape)?
181    };
182    session.broadcast_in_dim(&source, target_shape, &plan.dims)
183}
184
185fn broadcast_binary_in(
186    lhs: &Tensor,
187    rhs: &Tensor,
188    session: &mut dyn BackendSession,
189) -> Result<(Tensor, Tensor)> {
190    let shape = broadcast_shape(lhs.shape(), rhs.shape()).map_err(broadcast_error)?;
191    Ok((
192        broadcast_to_in(lhs, &shape, session)?,
193        broadcast_to_in(rhs, &shape, session)?,
194    ))
195}
196
197fn broadcast_ternary_in(
198    first: &Tensor,
199    second: &Tensor,
200    third: &Tensor,
201    session: &mut dyn BackendSession,
202) -> Result<(Tensor, Tensor, Tensor)> {
203    let shape = broadcast_shapes([first.shape(), second.shape(), third.shape()])
204        .map_err(broadcast_error)?;
205    Ok((
206        broadcast_to_in(first, &shape, session)?,
207        broadcast_to_in(second, &shape, session)?,
208        broadcast_to_in(third, &shape, session)?,
209    ))
210}
211
212fn broadcast_error(err: tenferro_ops::broadcast::BroadcastError) -> Error {
213    Error::validation("broadcast", broadcast_error_to_validation(err))
214}