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 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
503/// Resolve a reduction-family axis argument: `None` selects every axis.
504///
505/// An explicit axis list stays borrowed; only `None` builds a list.
506pub(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}