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}