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