Skip to main content

tenferro_runtime/traced/
composite_ops.rs

1//! Traced composite operations (activations, normalizations, softmax,
2//! `take_along_axis`), built from primitives by [`crate::composite`].
3
4use tenferro_ops::std_tensor_op::StdTensorOp;
5use tenferro_tensor::{CompareDir, DType, GatherConfig};
6
7use super::{apply_nullary, TracedTensor};
8use crate::composite::{
9    self, scalar_bytes, zero_pad_config, CompositeBinary, CompositeOps, CompositeReduce,
10    CompositeUnary,
11};
12use crate::error::{Error, Result};
13
14/// [`CompositeOps`] over traced graph construction.
15pub(crate) struct TracedComposite;
16
17impl CompositeOps for TracedComposite {
18    type Value = TracedTensor;
19    type Error = Error;
20
21    fn dtype(&self, value: &TracedTensor) -> DType {
22        value.dtype()
23    }
24
25    fn shape(&self, value: &TracedTensor) -> Result<Vec<usize>> {
26        value.concrete_shape()
27    }
28
29    fn scalar(&mut self, dtype: DType, value: f64) -> Result<TracedTensor> {
30        let bytes = scalar_bytes(dtype, value)?;
31        apply_nullary(
32            StdTensorOp::Constant { dtype, bytes },
33            0,
34            dtype,
35            Some(vec![]),
36        )
37    }
38
39    fn unary(&mut self, op: CompositeUnary, value: &TracedTensor) -> Result<TracedTensor> {
40        match op {
41            CompositeUnary::Neg => value.neg(),
42            CompositeUnary::Exp => value.exp(),
43            CompositeUnary::Log => value.log(),
44            CompositeUnary::Log1p => value.log1p(),
45            CompositeUnary::Tanh => value.tanh(),
46            CompositeUnary::Erf => value.erf(),
47            CompositeUnary::Rsqrt => value.rsqrt(),
48        }
49    }
50
51    fn binary(
52        &mut self,
53        op: CompositeBinary,
54        lhs: &TracedTensor,
55        rhs: &TracedTensor,
56    ) -> Result<TracedTensor> {
57        match op {
58            CompositeBinary::Add => lhs.add(rhs),
59            CompositeBinary::Sub => lhs.sub(rhs),
60            CompositeBinary::Mul => lhs.mul(rhs),
61            CompositeBinary::Div => lhs.div(rhs),
62            CompositeBinary::Maximum => lhs.maximum(rhs),
63        }
64    }
65
66    fn compare(
67        &mut self,
68        lhs: &TracedTensor,
69        rhs: &TracedTensor,
70        dir: CompareDir,
71    ) -> Result<TracedTensor> {
72        lhs.compare(rhs, dir)
73    }
74
75    fn select(
76        &mut self,
77        condition: &TracedTensor,
78        on_true: &TracedTensor,
79        on_false: &TracedTensor,
80    ) -> Result<TracedTensor> {
81        TracedTensor::where_select(condition, on_true, on_false)
82    }
83
84    fn reduce(
85        &mut self,
86        op: CompositeReduce,
87        value: &TracedTensor,
88        axes: &[usize],
89    ) -> Result<TracedTensor> {
90        match op {
91            CompositeReduce::Sum => value.reduce_sum(Some(axes)),
92            CompositeReduce::Max => value.reduce_max(Some(axes)),
93            CompositeReduce::SumSquares => value.reduce_sum_squares(Some(axes)),
94        }
95    }
96
97    fn broadcast_in_dim(
98        &mut self,
99        value: &TracedTensor,
100        shape: &[usize],
101        dims: &[usize],
102    ) -> Result<TracedTensor> {
103        value.broadcast_in_dim(shape, dims)
104    }
105
106    fn reshape(&mut self, value: &TracedTensor, shape: &[usize]) -> Result<TracedTensor> {
107        value.reshape(shape.to_vec())
108    }
109
110    fn concatenate(&mut self, values: &[&TracedTensor], axis: usize) -> Result<TracedTensor> {
111        TracedTensor::concatenate(values, axis)
112    }
113
114    fn pad(&mut self, value: &TracedTensor, low: &[usize], high: &[usize]) -> Result<TracedTensor> {
115        value.pad(zero_pad_config(low, high))
116    }
117
118    fn gather(
119        &mut self,
120        operand: &TracedTensor,
121        indices: &TracedTensor,
122        config: GatherConfig,
123    ) -> Result<TracedTensor> {
124        operand.gather(indices, config)
125    }
126}
127
128impl TracedTensor {
129    /// Logistic sigmoid `1 / (1 + exp(-x))`, for real `F32`/`F64` tensors.
130    ///
131    /// Evaluated as `1 / (1 + e)` for `x > 0` and `e / (1 + e)` otherwise,
132    /// with `e = exp(-|x|)`, so no intermediate overflows and the derivative
133    /// is finite everywhere (`sigmoid'(0) = 1/4`).
134    ///
135    /// # Examples
136    ///
137    /// ```rust
138    /// # use tenferro_runtime::TracedTensor;
139    /// let x = TracedTensor::from_vec_col_major(vec![3], vec![-1.0_f64, 0.0, 1.0])?;
140    /// let y = x.sigmoid()?;
141    /// assert_eq!(y.try_concrete_shape(), Some(vec![3]));
142    /// # Ok::<(), tenferro_runtime::Error>(())
143    /// ```
144    ///
145    /// # Errors
146    ///
147    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
148    /// complex, integer, or `Bool` input, [`Error::Validation`] with
149    /// `InvalidArgument` when the input shape is symbolic, or
150    /// [`Error::RuntimeStateSource`] when graph metadata registration fails.
151    pub fn sigmoid(&self) -> Result<TracedTensor> {
152        composite::sigmoid(&mut TracedComposite, self)
153    }
154
155    /// SiLU (swish) `x * sigmoid(x)`, for real `F32`/`F64` tensors.
156    ///
157    /// # Examples
158    ///
159    /// ```rust
160    /// # use tenferro_runtime::TracedTensor;
161    /// let x = TracedTensor::from_vec_col_major(vec![2], vec![-1.0_f64, 1.0])?;
162    /// let y = x.silu()?;
163    /// assert_eq!(y.dtype(), tenferro_runtime::DType::F64);
164    /// # Ok::<(), tenferro_runtime::Error>(())
165    /// ```
166    ///
167    /// # Errors
168    ///
169    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
170    /// complex, integer, or `Bool` input, [`Error::Validation`] with
171    /// `InvalidArgument` when the input shape is symbolic, or
172    /// [`Error::RuntimeStateSource`] when graph metadata registration fails.
173    pub fn silu(&self) -> Result<TracedTensor> {
174        composite::silu(&mut TracedComposite, self)
175    }
176
177    /// Softplus `log(1 + exp(x))` in the stable form `max(x, 0) + log1p(exp(-|x|))`.
178    ///
179    /// Never overflows; `softplus'(0) = 1/2` and `softplus''(0) = 1/4`.
180    ///
181    /// # Examples
182    ///
183    /// ```rust
184    /// # use tenferro_runtime::TracedTensor;
185    /// let x = TracedTensor::from_vec_col_major(vec![2], vec![-1000.0_f64, 1000.0])?;
186    /// let y = x.softplus()?;
187    /// assert_eq!(y.try_concrete_shape(), Some(vec![2]));
188    /// # Ok::<(), tenferro_runtime::Error>(())
189    /// ```
190    ///
191    /// # Errors
192    ///
193    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
194    /// complex, integer, or `Bool` input, [`Error::Validation`] with
195    /// `InvalidArgument` when the input shape is symbolic, or
196    /// [`Error::RuntimeStateSource`] when graph metadata registration fails.
197    pub fn softplus(&self) -> Result<TracedTensor> {
198        composite::softplus(&mut TracedComposite, self)
199    }
200
201    /// Exact GELU `x/2 * (1 + erf(x / sqrt(2)))` (PyTorch `approximate="none"`).
202    ///
203    /// # Examples
204    ///
205    /// ```rust
206    /// # use tenferro_runtime::TracedTensor;
207    /// let x = TracedTensor::from_vec_col_major(vec![2], vec![-1.0_f64, 1.0])?;
208    /// let y = x.gelu()?;
209    /// assert_eq!(y.try_concrete_shape(), Some(vec![2]));
210    /// # Ok::<(), tenferro_runtime::Error>(())
211    /// ```
212    ///
213    /// # Errors
214    ///
215    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
216    /// complex, integer, or `Bool` input, [`Error::Validation`] with
217    /// `InvalidArgument` when the input shape is symbolic, or
218    /// [`Error::RuntimeStateSource`] when graph metadata registration fails.
219    pub fn gelu(&self) -> Result<TracedTensor> {
220        composite::gelu(&mut TracedComposite, self)
221    }
222
223    /// GELU tanh approximation (PyTorch `approximate="tanh"`).
224    ///
225    /// `x/2 * (1 + tanh(sqrt(2/pi) * (x + 0.044715 x^3)))`.
226    ///
227    /// # Examples
228    ///
229    /// ```rust
230    /// # use tenferro_runtime::TracedTensor;
231    /// let x = TracedTensor::from_vec_col_major(vec![2], vec![-1.0_f64, 1.0])?;
232    /// let y = x.gelu_tanh()?;
233    /// assert_eq!(y.try_concrete_shape(), Some(vec![2]));
234    /// # Ok::<(), tenferro_runtime::Error>(())
235    /// ```
236    ///
237    /// # Errors
238    ///
239    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
240    /// complex, integer, or `Bool` input, [`Error::Validation`] with
241    /// `InvalidArgument` when the input shape is symbolic, or
242    /// [`Error::RuntimeStateSource`] when graph metadata registration fails.
243    pub fn gelu_tanh(&self) -> Result<TracedTensor> {
244        composite::gelu_tanh(&mut TracedComposite, self)
245    }
246
247    /// Arithmetic mean over `axes` (`None` reduces every axis).
248    ///
249    /// Defined for float and complex tensors. A mean over zero elements is
250    /// `NaN`; `Some(&[])` is the identity.
251    ///
252    /// # Examples
253    ///
254    /// ```rust
255    /// # use tenferro_runtime::TracedTensor;
256    /// let x = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?;
257    /// let y = x.reduce_mean(Some(&[0]))?;
258    /// assert_eq!(y.try_concrete_shape(), Some(vec![2]));
259    /// # Ok::<(), tenferro_runtime::Error>(())
260    /// ```
261    ///
262    /// # Errors
263    ///
264    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
265    /// integer or `Bool` input or `AxisOutOfBounds` / `DuplicateAxis` for
266    /// invalid axes, [`Error::Validation`] with `InvalidArgument` when the
267    /// input shape is symbolic, or [`Error::RuntimeStateSource`] when graph
268    /// metadata registration fails.
269    pub fn reduce_mean(&self, axes: Option<&[usize]>) -> Result<TracedTensor> {
270        composite::reduce_mean(&mut TracedComposite, self, axes)
271    }
272
273    /// Max-subtracted softmax along `axis`, for real `F32`/`F64` tensors.
274    ///
275    /// A slice that is entirely `-inf` returns zeros (with a finite gradient)
276    /// instead of `NaN`; a participating `NaN` or `+inf` makes its slice
277    /// `NaN`; a zero-length `axis` returns an empty result.
278    ///
279    /// # Examples
280    ///
281    /// ```rust
282    /// # use tenferro_runtime::TracedTensor;
283    /// let x = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64, 2.0, 3.0, 4.0, 5.0, 6.0])?;
284    /// let y = x.softmax(1)?;
285    /// assert_eq!(y.try_concrete_shape(), Some(vec![2, 3]));
286    /// # Ok::<(), tenferro_runtime::Error>(())
287    /// ```
288    ///
289    /// # Errors
290    ///
291    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
292    /// non-real input or `AxisOutOfBounds` for an invalid axis,
293    /// [`Error::Validation`] with `InvalidArgument` when the input shape is
294    /// symbolic, or [`Error::RuntimeStateSource`] when graph metadata
295    /// registration fails.
296    pub fn softmax(&self, axis: usize) -> Result<TracedTensor> {
297        composite::softmax(&mut TracedComposite, self, axis)
298    }
299
300    /// Max-subtracted log-softmax along `axis`, for real `F32`/`F64` tensors.
301    ///
302    /// A slice that is entirely `-inf` returns `-inf` (with a finite
303    /// gradient) instead of `NaN`.
304    ///
305    /// # Examples
306    ///
307    /// ```rust
308    /// # use tenferro_runtime::TracedTensor;
309    /// let x = TracedTensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0])?;
310    /// let y = x.log_softmax(0)?;
311    /// assert_eq!(y.try_concrete_shape(), Some(vec![3]));
312    /// # Ok::<(), tenferro_runtime::Error>(())
313    /// ```
314    ///
315    /// # Errors
316    ///
317    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
318    /// non-real input or `AxisOutOfBounds` for an invalid axis,
319    /// [`Error::Validation`] with `InvalidArgument` when the input shape is
320    /// symbolic, or [`Error::RuntimeStateSource`] when graph metadata
321    /// registration fails.
322    pub fn log_softmax(&self, axis: usize) -> Result<TracedTensor> {
323        composite::log_softmax(&mut TracedComposite, self, axis)
324    }
325
326    /// Softmax along `axis` over the entries where the `Bool` `mask` is true.
327    ///
328    /// `mask` broadcasts to this tensor's shape. Masked-out entries get
329    /// probability `0` and a zero gradient, whatever their value; a slice
330    /// with no unmasked entry returns zeros with a zero gradient.
331    ///
332    /// # Examples
333    ///
334    /// ```rust
335    /// # use tenferro_runtime::TracedTensor;
336    /// let x = TracedTensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0])?;
337    /// let mask = TracedTensor::from_vec_col_major(vec![3], vec![true, true, false])?;
338    /// let y = x.masked_softmax(&mask, 0)?;
339    /// assert_eq!(y.try_concrete_shape(), Some(vec![3]));
340    /// # Ok::<(), tenferro_runtime::Error>(())
341    /// ```
342    ///
343    /// # Errors
344    ///
345    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
346    /// non-real input, `DTypeMismatch` for a non-`Bool` mask, `ShapeMismatch`
347    /// for a mask that does not broadcast to the input, or `AxisOutOfBounds`;
348    /// [`Error::Validation`] with `InvalidArgument` for symbolic shapes; or
349    /// [`Error::RuntimeStateSource`] when graph metadata registration fails.
350    pub fn masked_softmax(&self, mask: &TracedTensor, axis: usize) -> Result<TracedTensor> {
351        composite::masked_softmax(&mut TracedComposite, self, mask, axis)
352    }
353
354    /// Log-softmax along `axis` over the entries where the `Bool` `mask` is true.
355    ///
356    /// Masked-out entries are `-inf` with a zero gradient; a slice with no
357    /// unmasked entry is all `-inf` with a zero gradient.
358    ///
359    /// # Examples
360    ///
361    /// ```rust
362    /// # use tenferro_runtime::TracedTensor;
363    /// let x = TracedTensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0])?;
364    /// let mask = TracedTensor::from_vec_col_major(vec![3], vec![true, false, true])?;
365    /// let y = x.masked_log_softmax(&mask, 0)?;
366    /// assert_eq!(y.try_concrete_shape(), Some(vec![3]));
367    /// # Ok::<(), tenferro_runtime::Error>(())
368    /// ```
369    ///
370    /// # Errors
371    ///
372    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
373    /// non-real input, `DTypeMismatch` for a non-`Bool` mask, `ShapeMismatch`
374    /// for a mask that does not broadcast to the input, or `AxisOutOfBounds`;
375    /// [`Error::Validation`] with `InvalidArgument` for symbolic shapes; or
376    /// [`Error::RuntimeStateSource`] when graph metadata registration fails.
377    pub fn masked_log_softmax(&self, mask: &TracedTensor, axis: usize) -> Result<TracedTensor> {
378        composite::masked_log_softmax(&mut TracedComposite, self, mask, axis)
379    }
380
381    /// Layer normalization along `axis` with optional affine `weight` / `bias`.
382    ///
383    /// `(x - mean) / sqrt(var + eps) * weight + bias`, with the biased
384    /// variance of the centered values. `weight` and `bias` are rank-1 of
385    /// length `shape[axis]`.
386    ///
387    /// # Examples
388    ///
389    /// ```rust
390    /// # use tenferro_runtime::TracedTensor;
391    /// let x = TracedTensor::from_vec_col_major(vec![4, 2], vec![1.0_f64, 2.0, 3.0, 4.0, 0.0, 0.0, 1.0, 1.0])?;
392    /// let w = TracedTensor::from_vec_col_major(vec![4], vec![1.0_f64, 1.0, 2.0, 2.0])?;
393    /// let y = x.layer_norm(0, Some(&w), None, 1e-5)?;
394    /// assert_eq!(y.try_concrete_shape(), Some(vec![4, 2]));
395    /// # Ok::<(), tenferro_runtime::Error>(())
396    /// ```
397    ///
398    /// # Errors
399    ///
400    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
401    /// non-real input, `AxisOutOfBounds`, `InvalidArgument` for a negative or
402    /// non-finite `eps`, or `DTypeMismatch` / `ShapeMismatch` for a bad
403    /// weight or bias; [`Error::Validation`] with `InvalidArgument` for
404    /// symbolic shapes; or [`Error::RuntimeStateSource`] when graph metadata
405    /// registration fails.
406    pub fn layer_norm(
407        &self,
408        axis: usize,
409        weight: Option<&TracedTensor>,
410        bias: Option<&TracedTensor>,
411        eps: f64,
412    ) -> Result<TracedTensor> {
413        composite::layer_norm(&mut TracedComposite, self, axis, weight, bias, eps)
414    }
415
416    /// RMS normalization along `axis` with optional affine `weight` / `bias`.
417    ///
418    /// `x / sqrt(mean(x^2) + eps) * weight + bias`; `weight` and `bias` are
419    /// rank-1 of length `shape[axis]`.
420    ///
421    /// # Examples
422    ///
423    /// ```rust
424    /// # use tenferro_runtime::TracedTensor;
425    /// let x = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?;
426    /// let y = x.rms_norm(0, None, None, 1e-6)?;
427    /// assert_eq!(y.try_concrete_shape(), Some(vec![2, 2]));
428    /// # Ok::<(), tenferro_runtime::Error>(())
429    /// ```
430    ///
431    /// # Errors
432    ///
433    /// Returns [`Error::TensorRuntime`] wrapping `UnsupportedDType` for
434    /// non-real input, `AxisOutOfBounds`, `InvalidArgument` for a negative or
435    /// non-finite `eps`, or `DTypeMismatch` / `ShapeMismatch` for a bad
436    /// weight or bias; [`Error::Validation`] with `InvalidArgument` for
437    /// symbolic shapes; or [`Error::RuntimeStateSource`] when graph metadata
438    /// registration fails.
439    pub fn rms_norm(
440        &self,
441        axis: usize,
442        weight: Option<&TracedTensor>,
443        bias: Option<&TracedTensor>,
444        eps: f64,
445    ) -> Result<TracedTensor> {
446        composite::rms_norm(&mut TracedComposite, self, axis, weight, bias, eps)
447    }
448
449    /// NumPy-style `take_along_axis` over `gather`.
450    ///
451    /// `out[.., i, ..] = self[.., indices[.., i, ..], ..]` along `axis`.
452    /// `indices` (I32/I64) has this tensor's rank; each other dimension is
453    /// either this tensor's extent (batch-varying indices) or `1` (the whole
454    /// extent is taken). Indices must be in bounds.
455    ///
456    /// # Examples
457    ///
458    /// ```rust
459    /// # use tenferro_runtime::TracedTensor;
460    /// // Gather whole rows of each matrix in a [2, 2, batch=2] stack.
461    /// let x = TracedTensor::from_vec_col_major(vec![2, 2, 2], (0..8).map(f64::from).collect::<Vec<_>>())?;
462    /// let rows = TracedTensor::from_vec_col_major(vec![2, 1, 2], vec![1_i64, 0, 0, 0])?;
463    /// let y = x.take_along_axis(&rows, 0)?;
464    /// assert_eq!(y.try_concrete_shape(), Some(vec![2, 2, 2]));
465    /// # Ok::<(), tenferro_runtime::Error>(())
466    /// ```
467    ///
468    /// # Errors
469    ///
470    /// Returns [`Error::TensorRuntime`] wrapping `RankMismatch` /
471    /// `ShapeMismatch` for incompatible index shapes, `AxisOutOfBounds`,
472    /// `UnsupportedDType` for a non-integer index dtype, or `InvalidArgument`
473    /// when taking from a zero-length axis; [`Error::Validation`] with
474    /// `InvalidArgument` for symbolic shapes; or
475    /// [`Error::RuntimeStateSource`] when graph metadata registration fails.
476    pub fn take_along_axis(&self, indices: &TracedTensor, axis: usize) -> Result<TracedTensor> {
477        composite::take_along_axis(&mut TracedComposite, self, indices, axis)
478    }
479}