1use crate::composite;
8use crate::composite::session::{borrowed, run_session_composite};
9use std::borrow::Cow;
10
11use num_complex::Complex64;
12use tenferro_ops::broadcast::{broadcast_error_to_validation, broadcast_shape, broadcast_shapes};
13use tenferro_tensor::validate::matmul_config_for_shapes;
14use tenferro_tensor::{
15 BackendSession, CompareDir, DType, DotGeneralConfig, Error, GatherConfig, PadConfig, Result,
16 ScatterConfig, SliceConfig, TensorRead,
17};
18
19use crate::typed_tensor::{broadcast_to_in_read, ReadInput};
20
21use crate::TensorSessionOpsExt;
22use tenferro_tensor::Tensor;
23
24impl TensorSessionOpsExt for Tensor {
25 fn add(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
26 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
27 session.add_read(lhs.tensor_read(), rhs.tensor_read())
28 }
29
30 fn mul(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
31 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
32 session.mul_read(lhs.tensor_read(), rhs.tensor_read())
33 }
34
35 fn exp(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
36 session.exp_read(TensorRead::from_tensor(self))
37 }
38
39 fn reduce_sum(
40 &self,
41 axes: Option<&[usize]>,
42 session: &mut dyn BackendSession,
43 ) -> Result<Tensor> {
44 let axes = all_axes_if_none(self.shape().len(), axes);
45 session.reduce_sum_read(TensorRead::from_tensor(self), &axes)
46 }
47
48 fn convert(&self, to: DType, session: &mut dyn BackendSession) -> Result<Tensor> {
49 session.convert(self, to)
50 }
51
52 fn cast(&self, to: DType, session: &mut dyn BackendSession) -> Result<Tensor> {
53 session.cast(self, to)
54 }
55
56 fn sub(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
57 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
58 session.sub_read(lhs.tensor_read(), rhs.tensor_read())
59 }
60
61 fn div(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
62 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
63 session.div_read(lhs.tensor_read(), rhs.tensor_read())
64 }
65
66 fn rem(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
67 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
68 session.rem_read(lhs.tensor_read(), rhs.tensor_read())
69 }
70
71 fn pow(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
72 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
73 session.pow_read(lhs.tensor_read(), rhs.tensor_read())
74 }
75
76 fn maximum(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
77 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
78 session.maximum_read(lhs.tensor_read(), rhs.tensor_read())
79 }
80
81 fn minimum(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
82 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
83 session.minimum_read(lhs.tensor_read(), rhs.tensor_read())
84 }
85
86 fn neg(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
87 session.neg_read(TensorRead::from_tensor(self))
88 }
89
90 fn abs(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
91 session.abs_read(TensorRead::from_tensor(self))
92 }
93
94 fn sign(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
95 session.sign_read(TensorRead::from_tensor(self))
96 }
97
98 fn conj(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
99 session.conj_read(TensorRead::from_tensor(self))
100 }
101
102 fn log(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
103 session.log_read(TensorRead::from_tensor(self))
104 }
105
106 fn expm1(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
107 session.expm1_read(TensorRead::from_tensor(self))
108 }
109
110 fn log1p(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
111 session.log1p_read(TensorRead::from_tensor(self))
112 }
113
114 fn erf(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
115 session.erf_read(TensorRead::from_tensor(self))
116 }
117
118 fn sin(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
119 session.sin_read(TensorRead::from_tensor(self))
120 }
121
122 fn cos(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
123 session.cos_read(TensorRead::from_tensor(self))
124 }
125
126 fn tanh(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
127 session.tanh_read(TensorRead::from_tensor(self))
128 }
129
130 fn sqrt(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
131 session.sqrt_read(TensorRead::from_tensor(self))
132 }
133
134 fn rsqrt(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
135 session.rsqrt_read(TensorRead::from_tensor(self))
136 }
137
138 fn compare(
139 &self,
140 rhs: &Tensor,
141 dir: CompareDir,
142 session: &mut dyn BackendSession,
143 ) -> Result<Tensor> {
144 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
145 session.compare_read(lhs.tensor_read(), rhs.tensor_read(), &dir)
146 }
147
148 fn where_select(
149 &self,
150 on_true: &Tensor,
151 on_false: &Tensor,
152 session: &mut dyn BackendSession,
153 ) -> Result<Tensor> {
154 let (condition, on_true, on_false) =
155 broadcast_ternary_in(self, on_true, on_false, session)?;
156 session.select_read(
157 condition.tensor_read(),
158 on_true.tensor_read(),
159 on_false.tensor_read(),
160 )
161 }
162
163 fn clamp(
164 &self,
165 lower: &Tensor,
166 upper: &Tensor,
167 session: &mut dyn BackendSession,
168 ) -> Result<Tensor> {
169 let (input, lower, upper) = broadcast_ternary_in(self, lower, upper, session)?;
170 session.clamp_read(
171 input.tensor_read(),
172 lower.tensor_read(),
173 upper.tensor_read(),
174 )
175 }
176
177 fn matmul(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
178 let config = matmul_config_for_shapes("matmul", self.shape(), rhs.shape())?;
179 session.dot_general_read(
180 TensorRead::from_tensor(self),
181 TensorRead::from_tensor(rhs),
182 &config,
183 )
184 }
185
186 fn reshape(&self, shape: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
187 session.reshape_read(TensorRead::from_tensor(self), shape)
188 }
189
190 fn transpose(&self, perm: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
191 session.transpose_read(TensorRead::from_tensor(self), perm)
192 }
193
194 fn gather(
195 &self,
196 indices: &Tensor,
197 config: GatherConfig,
198 session: &mut dyn BackendSession,
199 ) -> Result<Tensor> {
200 session.gather(self, indices, &config)
201 }
202
203 fn scatter(
204 &self,
205 indices: &Tensor,
206 updates: &Tensor,
207 config: ScatterConfig,
208 session: &mut dyn BackendSession,
209 ) -> Result<Tensor> {
210 session.scatter(self, indices, updates, &config)
211 }
212
213 fn slice(&self, config: SliceConfig, session: &mut dyn BackendSession) -> Result<Tensor> {
214 session.slice(self, &config)
215 }
216
217 fn dynamic_slice(
218 &self,
219 starts: &Tensor,
220 sizes: &[usize],
221 session: &mut dyn BackendSession,
222 ) -> Result<Tensor> {
223 session.dynamic_slice(self, starts, sizes)
224 }
225
226 fn pad(&self, config: PadConfig, session: &mut dyn BackendSession) -> Result<Tensor> {
227 session.pad(self, &config)
228 }
229
230 fn concatenate(
231 inputs: &[&Tensor],
232 axis: usize,
233 session: &mut dyn BackendSession,
234 ) -> Result<Tensor> {
235 session.concatenate(inputs, axis)
236 }
237
238 fn reverse(&self, axes: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
239 session.reverse(self, axes)
240 }
241
242 fn reduce_max(
243 &self,
244 axes: Option<&[usize]>,
245 session: &mut dyn BackendSession,
246 ) -> Result<Tensor> {
247 let axes = all_axes_if_none(self.shape().len(), axes);
248 session.reduce_max_read(TensorRead::from_tensor(self), &axes)
249 }
250
251 fn reduce_min(
252 &self,
253 axes: Option<&[usize]>,
254 session: &mut dyn BackendSession,
255 ) -> Result<Tensor> {
256 let axes = all_axes_if_none(self.shape().len(), axes);
257 session.reduce_min_read(TensorRead::from_tensor(self), &axes)
258 }
259
260 fn reduce_prod(
261 &self,
262 axes: Option<&[usize]>,
263 session: &mut dyn BackendSession,
264 ) -> Result<Tensor> {
265 let axes = all_axes_if_none(self.shape().len(), axes);
266 session.reduce_prod_read(TensorRead::from_tensor(self), &axes)
267 }
268
269 fn reduce_sum_squares(
270 &self,
271 axes: Option<&[usize]>,
272 session: &mut dyn BackendSession,
273 ) -> Result<Tensor> {
274 let axes = all_axes_if_none(self.shape().len(), axes);
275 session.reduce_sum_squares_read(TensorRead::from_tensor(self), &axes)
276 }
277
278 fn broadcast_in_dim(
279 &self,
280 shape: &[usize],
281 dims: &[usize],
282 session: &mut dyn BackendSession,
283 ) -> Result<Tensor> {
284 session.broadcast_in_dim_read(TensorRead::from_tensor(self), shape, dims)
285 }
286
287 fn tril(&self, k: i64, session: &mut dyn BackendSession) -> Result<Tensor> {
288 session.tril(self, k)
289 }
290
291 fn triu(&self, k: i64, session: &mut dyn BackendSession) -> Result<Tensor> {
292 session.triu(self, k)
293 }
294
295 fn extract_diag(
296 &self,
297 axis_a: usize,
298 axis_b: usize,
299 session: &mut dyn BackendSession,
300 ) -> Result<Tensor> {
301 session.extract_diagonal(self, axis_a, axis_b)
302 }
303
304 fn embed_diag(
305 &self,
306 axis_a: usize,
307 axis_b: usize,
308 session: &mut dyn BackendSession,
309 ) -> Result<Tensor> {
310 session.embed_diagonal(self, axis_a, axis_b)
311 }
312
313 fn dot_general(
314 &self,
315 rhs: &Tensor,
316 config: DotGeneralConfig,
317 session: &mut dyn BackendSession,
318 ) -> Result<Tensor> {
319 session.dot_general_read(
320 TensorRead::from_tensor(self),
321 TensorRead::from_tensor(rhs),
322 &config,
323 )
324 }
325
326 fn dot_general_with_conj(
327 &self,
328 rhs: &Tensor,
329 config: DotGeneralConfig,
330 lhs_conj: bool,
331 rhs_conj: bool,
332 session: &mut dyn BackendSession,
333 ) -> Result<Tensor> {
334 session.dot_general_with_conj(self, rhs, &config, lhs_conj, rhs_conj)
335 }
336
337 fn scale_real(&self, factor: f64, session: &mut dyn BackendSession) -> Result<Tensor> {
338 let scalar = crate::scale::real_scale_scalar(self.dtype(), factor)?;
339 let scalar = session.upload_host_tensor(TensorRead::from_tensor(&scalar))?;
340 TensorSessionOpsExt::mul(self, &scalar, session)
341 }
342
343 fn scale_complex(&self, factor: Complex64, session: &mut dyn BackendSession) -> Result<Tensor> {
344 let scalar = crate::scale::complex_scale_scalar(self.dtype(), factor)?;
345 let scalar = session.upload_host_tensor(TensorRead::from_tensor(&scalar))?;
346 TensorSessionOpsExt::mul(self, &scalar, session)
347 }
348
349 fn sigmoid(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
350 run_session_composite(session, |ops| composite::sigmoid(ops, &borrowed(self)))
351 }
352
353 fn silu(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
354 run_session_composite(session, |ops| composite::silu(ops, &borrowed(self)))
355 }
356
357 fn softplus(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
358 run_session_composite(session, |ops| composite::softplus(ops, &borrowed(self)))
359 }
360
361 fn gelu(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
362 run_session_composite(session, |ops| composite::gelu(ops, &borrowed(self)))
363 }
364
365 fn gelu_tanh(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
366 run_session_composite(session, |ops| composite::gelu_tanh(ops, &borrowed(self)))
367 }
368
369 fn reduce_mean(
370 &self,
371 axes: Option<&[usize]>,
372 session: &mut dyn BackendSession,
373 ) -> Result<Tensor> {
374 run_session_composite(session, |ops| {
375 composite::reduce_mean(ops, &borrowed(self), axes)
376 })
377 }
378
379 fn softmax(&self, axis: usize, session: &mut dyn BackendSession) -> Result<Tensor> {
380 run_session_composite(session, |ops| {
381 composite::softmax(ops, &borrowed(self), axis)
382 })
383 }
384
385 fn log_softmax(&self, axis: usize, session: &mut dyn BackendSession) -> Result<Tensor> {
386 run_session_composite(session, |ops| {
387 composite::log_softmax(ops, &borrowed(self), axis)
388 })
389 }
390
391 fn masked_softmax(
392 &self,
393 mask: &Tensor,
394 axis: usize,
395 session: &mut dyn BackendSession,
396 ) -> Result<Tensor> {
397 let mask = borrowed(mask);
398 run_session_composite(session, |ops| {
399 composite::masked_softmax(ops, &borrowed(self), &mask, axis)
400 })
401 }
402
403 fn masked_log_softmax(
404 &self,
405 mask: &Tensor,
406 axis: usize,
407 session: &mut dyn BackendSession,
408 ) -> Result<Tensor> {
409 let mask = borrowed(mask);
410 run_session_composite(session, |ops| {
411 composite::masked_log_softmax(ops, &borrowed(self), &mask, axis)
412 })
413 }
414
415 fn layer_norm(
416 &self,
417 axis: usize,
418 weight: Option<&Tensor>,
419 bias: Option<&Tensor>,
420 eps: f64,
421 session: &mut dyn BackendSession,
422 ) -> Result<Tensor> {
423 let weight = weight.map(borrowed);
424 let bias = bias.map(borrowed);
425 run_session_composite(session, |ops| {
426 composite::layer_norm(
427 ops,
428 &borrowed(self),
429 axis,
430 weight.as_ref(),
431 bias.as_ref(),
432 eps,
433 )
434 })
435 }
436
437 fn rms_norm(
438 &self,
439 axis: usize,
440 weight: Option<&Tensor>,
441 bias: Option<&Tensor>,
442 eps: f64,
443 session: &mut dyn BackendSession,
444 ) -> Result<Tensor> {
445 let weight = weight.map(borrowed);
446 let bias = bias.map(borrowed);
447 run_session_composite(session, |ops| {
448 composite::rms_norm(
449 ops,
450 &borrowed(self),
451 axis,
452 weight.as_ref(),
453 bias.as_ref(),
454 eps,
455 )
456 })
457 }
458
459 fn take_along_axis(
460 &self,
461 indices: &Tensor,
462 axis: usize,
463 session: &mut dyn BackendSession,
464 ) -> Result<Tensor> {
465 let indices = borrowed(indices);
466 run_session_composite(session, |ops| {
467 composite::take_along_axis(ops, &borrowed(self), &indices, axis)
468 })
469 }
470}
471
472fn broadcast_binary_in<'a>(
473 lhs: &'a Tensor,
474 rhs: &'a Tensor,
475 session: &mut dyn BackendSession,
476) -> Result<(ReadInput<'a>, ReadInput<'a>)> {
477 let shape = broadcast_shape(lhs.shape(), rhs.shape()).map_err(broadcast_error)?;
478 Ok((
479 broadcast_to_in_read(TensorRead::from_tensor(lhs), &shape, session)?,
480 broadcast_to_in_read(TensorRead::from_tensor(rhs), &shape, session)?,
481 ))
482}
483
484fn broadcast_ternary_in<'a>(
485 first: &'a Tensor,
486 second: &'a Tensor,
487 third: &'a Tensor,
488 session: &mut dyn BackendSession,
489) -> Result<(ReadInput<'a>, ReadInput<'a>, ReadInput<'a>)> {
490 let shape = broadcast_shapes([first.shape(), second.shape(), third.shape()])
491 .map_err(broadcast_error)?;
492 Ok((
493 broadcast_to_in_read(TensorRead::from_tensor(first), &shape, session)?,
494 broadcast_to_in_read(TensorRead::from_tensor(second), &shape, session)?,
495 broadcast_to_in_read(TensorRead::from_tensor(third), &shape, session)?,
496 ))
497}
498
499fn broadcast_error(err: tenferro_ops::broadcast::BroadcastError) -> Error {
500 Error::validation("broadcast", broadcast_error_to_validation(err))
501}
502
503pub(crate) fn all_axes_if_none(rank: usize, axes: Option<&[usize]>) -> Cow<'_, [usize]> {
507 match axes {
508 Some(axes) => Cow::Borrowed(axes),
509 None => Cow::Owned((0..rank).collect()),
510 }
511}