1use 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#[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}