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::{broadcast_error_to_validation, broadcast_shape, broadcast_shapes};
8use tenferro_tensor::validate::matmul_config_for_shapes;
9use tenferro_tensor::{BackendSession, CompareDir, DType, Error, Result, TensorRead};
10
11use crate::typed_tensor::{broadcast_to_in_read, ReadInput};
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_read(lhs.tensor_read(), rhs.tensor_read())
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_read(lhs.tensor_read(), rhs.tensor_read())
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_read(lhs.tensor_read(), rhs.tensor_read())
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_read(lhs.tensor_read(), rhs.tensor_read())
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_read(lhs.tensor_read(), rhs.tensor_read())
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_read(lhs.tensor_read(), rhs.tensor_read())
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_read(lhs.tensor_read(), rhs.tensor_read())
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_read(lhs.tensor_read(), rhs.tensor_read())
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_read(lhs.tensor_read(), rhs.tensor_read(), &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_read(
140            condition.tensor_read(),
141            on_true.tensor_read(),
142            on_false.tensor_read(),
143        )
144    }
145
146    fn clamp(
147        &self,
148        lower: &Tensor,
149        upper: &Tensor,
150        session: &mut dyn BackendSession,
151    ) -> Result<Tensor> {
152        let (input, lower, upper) = broadcast_ternary_in(self, lower, upper, session)?;
153        session.clamp_read(
154            input.tensor_read(),
155            lower.tensor_read(),
156            upper.tensor_read(),
157        )
158    }
159
160    fn matmul(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
161        let config = matmul_config_for_shapes("matmul", self.shape(), rhs.shape())?;
162        session.dot_general(self, rhs, &config)
163    }
164
165    fn reshape(&self, shape: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
166        session.reshape(self, shape)
167    }
168
169    fn transpose(&self, perm: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
170        session.transpose(self, perm)
171    }
172}
173
174fn broadcast_binary_in<'a>(
175    lhs: &'a Tensor,
176    rhs: &'a Tensor,
177    session: &mut dyn BackendSession,
178) -> Result<(ReadInput<'a>, ReadInput<'a>)> {
179    let shape = broadcast_shape(lhs.shape(), rhs.shape()).map_err(broadcast_error)?;
180    Ok((
181        broadcast_to_in_read(TensorRead::from_tensor(lhs), &shape, session)?,
182        broadcast_to_in_read(TensorRead::from_tensor(rhs), &shape, session)?,
183    ))
184}
185
186fn broadcast_ternary_in<'a>(
187    first: &'a Tensor,
188    second: &'a Tensor,
189    third: &'a Tensor,
190    session: &mut dyn BackendSession,
191) -> Result<(ReadInput<'a>, ReadInput<'a>, ReadInput<'a>)> {
192    let shape = broadcast_shapes([first.shape(), second.shape(), third.shape()])
193        .map_err(broadcast_error)?;
194    Ok((
195        broadcast_to_in_read(TensorRead::from_tensor(first), &shape, session)?,
196        broadcast_to_in_read(TensorRead::from_tensor(second), &shape, session)?,
197        broadcast_to_in_read(TensorRead::from_tensor(third), &shape, session)?,
198    ))
199}
200
201fn broadcast_error(err: tenferro_ops::broadcast::BroadcastError) -> Error {
202    Error::validation("broadcast", broadcast_error_to_validation(err))
203}