Skip to main content

tenferro_ops/
std_tensor_op.rs

1use std::hash::{Hash, Hasher};
2use std::sync::Arc;
3
4#[cfg(all(test, feature = "autodiff"))]
5use crate::ad::{ADRuleResult, PrimitiveTransposeInput};
6#[cfg(all(test, feature = "autodiff"))]
7use computegraph::types::{LocalValueId, OperationRole, ValueKey};
8use computegraph::GraphOperation;
9use num_complex::{Complex32, Complex64};
10
11use crate::dim_expr::DimExpr;
12use crate::ext_op::{ext_op_eq, hash_extension, ExtensionOp};
13use crate::input_key::TensorInputKey;
14use tenferro_tensor::{
15    CompareDir, DType, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig, SliceConfig,
16    TensorScalar,
17};
18
19/// Scalar values that can be encoded as tensor constant operations.
20///
21/// # Examples
22///
23/// ```rust
24/// use tenferro_ops::std_tensor_op::ConstantScalar;
25///
26/// assert_eq!(1.0_f64.constant_bytes(), 1.0_f64.to_le_bytes().to_vec());
27/// ```
28pub trait ConstantScalar: TensorScalar + private::Sealed {
29    /// Encode the scalar value as little-endian constant bytes.
30    ///
31    /// # Examples
32    ///
33    /// ```rust
34    /// use tenferro_ops::std_tensor_op::ConstantScalar;
35    ///
36    /// assert_eq!(true.constant_bytes(), vec![1]);
37    /// ```
38    fn constant_bytes(self) -> Vec<u8>;
39}
40
41mod private {
42    pub trait Sealed {}
43
44    impl Sealed for f64 {}
45    impl Sealed for f32 {}
46    impl Sealed for i64 {}
47    impl Sealed for i32 {}
48    impl Sealed for bool {}
49    impl Sealed for num_complex::Complex64 {}
50    impl Sealed for num_complex::Complex32 {}
51}
52
53impl ConstantScalar for f64 {
54    fn constant_bytes(self) -> Vec<u8> {
55        self.to_le_bytes().to_vec()
56    }
57}
58
59impl ConstantScalar for f32 {
60    fn constant_bytes(self) -> Vec<u8> {
61        self.to_le_bytes().to_vec()
62    }
63}
64
65impl ConstantScalar for i64 {
66    fn constant_bytes(self) -> Vec<u8> {
67        self.to_le_bytes().to_vec()
68    }
69}
70
71impl ConstantScalar for i32 {
72    fn constant_bytes(self) -> Vec<u8> {
73        self.to_le_bytes().to_vec()
74    }
75}
76
77impl ConstantScalar for bool {
78    fn constant_bytes(self) -> Vec<u8> {
79        vec![u8::from(self)]
80    }
81}
82
83impl ConstantScalar for Complex64 {
84    fn constant_bytes(self) -> Vec<u8> {
85        let mut bytes = Vec::with_capacity(16);
86        bytes.extend_from_slice(&self.re.to_le_bytes());
87        bytes.extend_from_slice(&self.im.to_le_bytes());
88        bytes
89    }
90}
91
92impl ConstantScalar for Complex32 {
93    fn constant_bytes(self) -> Vec<u8> {
94        let mut bytes = Vec::with_capacity(8);
95        bytes.extend_from_slice(&self.re.to_le_bytes());
96        bytes.extend_from_slice(&self.im.to_le_bytes());
97        bytes
98    }
99}
100
101tenferro_core_ops::define_std_tensor_op!();
102
103impl StdTensorOp {
104    /// Create a scalar constant op from any supported tensor scalar.
105    ///
106    /// # Examples
107    ///
108    /// ```rust
109    /// use num_complex::Complex64;
110    /// use tenferro_ops::std_tensor_op::StdTensorOp;
111    /// use tenferro_tensor::DType;
112    ///
113    /// let real = StdTensorOp::constant(1.5_f64);
114    /// let complex = StdTensorOp::constant(Complex64::new(1.0, -2.0));
115    ///
116    /// assert!(matches!(real, StdTensorOp::Constant { dtype: DType::F64, .. }));
117    /// assert!(matches!(complex, StdTensorOp::Constant { dtype: DType::C64, .. }));
118    /// ```
119    pub fn constant<T: ConstantScalar>(value: T) -> Self {
120        Self::Constant {
121            dtype: T::dtype(),
122            bytes: value.constant_bytes(),
123        }
124    }
125}
126
127impl PartialEq for StdTensorOp {
128    fn eq(&self, other: &Self) -> bool {
129        if std::mem::discriminant(self) != std::mem::discriminant(other) {
130            return false;
131        }
132        match (self, other) {
133            (Self::Add, Self::Add)
134            | (Self::Sub, Self::Sub)
135            | (Self::Mul, Self::Mul)
136            | (Self::Neg, Self::Neg)
137            | (Self::Conj, Self::Conj)
138            | (Self::Div, Self::Div)
139            | (Self::Rem, Self::Rem)
140            | (Self::Abs, Self::Abs)
141            | (Self::Sign, Self::Sign)
142            | (Self::Maximum, Self::Maximum)
143            | (Self::Minimum, Self::Minimum)
144            | (Self::Select, Self::Select)
145            | (Self::Clamp, Self::Clamp)
146            | (Self::Exp, Self::Exp)
147            | (Self::Log, Self::Log)
148            | (Self::Sin, Self::Sin)
149            | (Self::Cos, Self::Cos)
150            | (Self::Tanh, Self::Tanh)
151            | (Self::Sqrt, Self::Sqrt)
152            | (Self::Rsqrt, Self::Rsqrt)
153            | (Self::Pow, Self::Pow)
154            | (Self::Expm1, Self::Expm1)
155            | (Self::Log1p, Self::Log1p)
156            | (Self::Erf, Self::Erf)
157            | (Self::DynamicUpdateSlice, Self::DynamicUpdateSlice) => true,
158            (Self::DotGeneral { config: a }, Self::DotGeneral { config: b }) => a == b,
159            (Self::Transpose { perm: a }, Self::Transpose { perm: b }) => a == b,
160            (Self::Reshape { to_shape: a }, Self::Reshape { to_shape: b }) => a == b,
161            (
162                Self::BroadcastInDim {
163                    shape: sa,
164                    dims: da,
165                },
166                Self::BroadcastInDim {
167                    shape: sb,
168                    dims: db,
169                },
170            ) => sa == sb && da == db,
171            (Self::Convert { from: fa, to: ta }, Self::Convert { from: fb, to: tb }) => {
172                fa == fb && ta == tb
173            }
174            (
175                Self::Constant {
176                    dtype: da,
177                    bytes: ba,
178                },
179                Self::Constant {
180                    dtype: db,
181                    bytes: bb,
182                },
183            ) => da == db && ba == bb,
184            (Self::ReduceSum { axes: a }, Self::ReduceSum { axes: b })
185            | (Self::ReduceSumSquares { axes: a }, Self::ReduceSumSquares { axes: b })
186            | (Self::ReduceProd { axes: a }, Self::ReduceProd { axes: b })
187            | (Self::ReduceMax { axes: a }, Self::ReduceMax { axes: b })
188            | (Self::ReduceMin { axes: a }, Self::ReduceMin { axes: b })
189            | (Self::Reverse { axes: a }, Self::Reverse { axes: b }) => a == b,
190            (Self::Compare(a), Self::Compare(b)) => a == b,
191            (
192                Self::ExtractDiag {
193                    axis_a: aa,
194                    axis_b: ba,
195                },
196                Self::ExtractDiag {
197                    axis_a: ab,
198                    axis_b: bb,
199                },
200            )
201            | (
202                Self::EmbedDiag {
203                    axis_a: aa,
204                    axis_b: ba,
205                },
206                Self::EmbedDiag {
207                    axis_a: ab,
208                    axis_b: bb,
209                },
210            ) => aa == ab && ba == bb,
211            (Self::Tril { k: a }, Self::Tril { k: b })
212            | (Self::Triu { k: a }, Self::Triu { k: b }) => a == b,
213            (Self::Gather(a), Self::Gather(b)) => a == b,
214            (
215                Self::GatherDynamicSliceSizes {
216                    offset_dims: oa,
217                    collapsed_slice_dims: ca,
218                    start_index_map: sa,
219                    index_vector_dim: ia,
220                    slice_sizes: za,
221                },
222                Self::GatherDynamicSliceSizes {
223                    offset_dims: ob,
224                    collapsed_slice_dims: cb,
225                    start_index_map: sb,
226                    index_vector_dim: ib,
227                    slice_sizes: zb,
228                },
229            ) => oa == ob && ca == cb && sa == sb && ia == ib && za == zb,
230            (Self::Scatter(a), Self::Scatter(b)) => a == b,
231            (Self::Slice(a), Self::Slice(b)) => a == b,
232            (Self::DynamicSlice { slice_sizes: a }, Self::DynamicSlice { slice_sizes: b }) => {
233                a == b
234            }
235            (Self::Pad(a), Self::Pad(b)) => a == b,
236            (
237                Self::Concatenate {
238                    axis: a,
239                    input_count: na,
240                },
241                Self::Concatenate {
242                    axis: b,
243                    input_count: nb,
244                },
245            ) => a == b && na == nb,
246            (Self::ShapeOf { axis: a }, Self::ShapeOf { axis: b })
247            | (Self::DynamicTruncate { axis: a }, Self::DynamicTruncate { axis: b })
248            | (Self::PadToMatch { axis: a }, Self::PadToMatch { axis: b }) => a == b,
249            (Self::Extension(a), Self::Extension(b)) => ext_op_eq(a.as_ref(), b.as_ref()),
250            _ => false,
251        }
252    }
253}
254
255impl Eq for StdTensorOp {}
256
257impl Hash for StdTensorOp {
258    fn hash<H: Hasher>(&self, state: &mut H) {
259        std::mem::discriminant(self).hash(state);
260        match self {
261            Self::Add
262            | Self::Sub
263            | Self::Mul
264            | Self::Neg
265            | Self::Conj
266            | Self::Div
267            | Self::Rem
268            | Self::Abs
269            | Self::Sign
270            | Self::Maximum
271            | Self::Minimum
272            | Self::Select
273            | Self::Clamp
274            | Self::Exp
275            | Self::Log
276            | Self::Sin
277            | Self::Cos
278            | Self::Tanh
279            | Self::Sqrt
280            | Self::Rsqrt
281            | Self::Pow
282            | Self::Expm1
283            | Self::Log1p
284            | Self::Erf => {}
285            Self::DotGeneral { config } => {
286                config.hash(state);
287            }
288            Self::Transpose { perm } => perm.hash(state),
289            Self::Reshape { to_shape } => {
290                to_shape.hash(state);
291            }
292            Self::BroadcastInDim { shape, dims } => {
293                shape.hash(state);
294                dims.hash(state);
295            }
296            Self::Convert { from, to } => {
297                from.hash(state);
298                to.hash(state);
299            }
300            Self::Constant { dtype, bytes } => {
301                dtype.hash(state);
302                bytes.hash(state);
303            }
304            Self::ReduceSum { axes } | Self::ReduceSumSquares { axes } => {
305                axes.hash(state);
306            }
307            Self::Compare(dir) => dir.hash(state),
308            Self::ExtractDiag { axis_a, axis_b } | Self::EmbedDiag { axis_a, axis_b } => {
309                axis_a.hash(state);
310                axis_b.hash(state);
311            }
312            Self::Tril { k } | Self::Triu { k } => k.hash(state),
313            Self::Gather(config) => config.hash(state),
314            Self::GatherDynamicSliceSizes {
315                offset_dims,
316                collapsed_slice_dims,
317                start_index_map,
318                index_vector_dim,
319                slice_sizes,
320            } => {
321                offset_dims.hash(state);
322                collapsed_slice_dims.hash(state);
323                start_index_map.hash(state);
324                index_vector_dim.hash(state);
325                slice_sizes.hash(state);
326            }
327            Self::Scatter(config) => config.hash(state),
328            Self::Slice(config) => config.hash(state),
329            Self::DynamicSlice { slice_sizes } => slice_sizes.hash(state),
330            Self::DynamicUpdateSlice => {}
331            Self::Pad(config) => config.hash(state),
332            Self::Concatenate { axis, input_count } => {
333                axis.hash(state);
334                input_count.hash(state);
335            }
336            Self::Reverse { axes } => axes.hash(state),
337            Self::ShapeOf { axis } | Self::DynamicTruncate { axis } | Self::PadToMatch { axis } => {
338                axis.hash(state)
339            }
340            Self::ReduceProd { axes } | Self::ReduceMax { axes } | Self::ReduceMin { axes } => {
341                axes.hash(state);
342            }
343            Self::Extension(op) => hash_extension(op.as_ref(), state),
344        }
345    }
346}
347
348fn n_inputs_from_dim_exprs(min_inputs: usize, exprs: &[&[DimExpr]]) -> usize {
349    let max_idx = exprs
350        .iter()
351        .flat_map(|exprs| exprs.iter())
352        .filter_map(DimExpr::max_input_idx)
353        .max()
354        .map_or(0, |max_idx| max_idx + 1);
355    max_idx.max(min_inputs)
356}
357
358impl GraphOperation for StdTensorOp {
359    // This graph is metadata-only; execution owns tensors in the runtime
360    // submission layer rather than as graph operands.
361    type Operand = ();
362    type Context = ();
363    type InputKey = TensorInputKey;
364
365    fn input_count(&self) -> usize {
366        match self {
367            Self::Add | Self::Sub | Self::Mul | Self::DotGeneral { .. } | Self::Gather(_) => 2,
368            Self::GatherDynamicSliceSizes { slice_sizes, .. } => {
369                n_inputs_from_dim_exprs(2, &[slice_sizes])
370            }
371            Self::Neg
372            | Self::Conj
373            | Self::Transpose { .. }
374            | Self::Convert { .. }
375            | Self::ExtractDiag { .. }
376            | Self::EmbedDiag { .. }
377            | Self::Tril { .. }
378            | Self::Triu { .. }
379            | Self::Slice(_)
380            | Self::Pad(_)
381            | Self::Reverse { .. }
382            | Self::ShapeOf { .. } => 1,
383            Self::DynamicTruncate { .. } | Self::PadToMatch { .. } => 2,
384            Self::Reshape { to_shape } => n_inputs_from_dim_exprs(1, &[to_shape]),
385            Self::BroadcastInDim { shape, .. } => n_inputs_from_dim_exprs(1, &[shape]),
386            Self::ReduceSum { .. }
387            | Self::ReduceSumSquares { .. }
388            | Self::ReduceProd { .. }
389            | Self::ReduceMax { .. }
390            | Self::ReduceMin { .. } => 1,
391            Self::Div
392            | Self::Rem
393            | Self::Maximum
394            | Self::Minimum
395            | Self::Pow
396            | Self::DynamicSlice { .. } => 2,
397            Self::Constant { .. } => 0,
398            Self::Scatter(_) | Self::DynamicUpdateSlice => 3,
399            Self::Concatenate { input_count, .. } => *input_count,
400            Self::Abs
401            | Self::Sign
402            | Self::Exp
403            | Self::Log
404            | Self::Sin
405            | Self::Cos
406            | Self::Tanh
407            | Self::Sqrt
408            | Self::Rsqrt
409            | Self::Expm1
410            | Self::Log1p
411            | Self::Erf => 1,
412            Self::Select | Self::Clamp => 3,
413            Self::Compare(_) => 2,
414            Self::Extension(op) => ExtensionOp::input_count(op.as_ref()),
415        }
416    }
417
418    fn output_count(&self) -> usize {
419        match self {
420            Self::Add
421            | Self::Sub
422            | Self::Mul
423            | Self::Neg
424            | Self::Conj
425            | Self::DotGeneral { .. }
426            | Self::Transpose { .. }
427            | Self::Reshape { .. }
428            | Self::BroadcastInDim { .. }
429            | Self::Convert { .. }
430            | Self::ReduceSum { .. }
431            | Self::ReduceSumSquares { .. }
432            | Self::Div
433            | Self::Rem
434            | Self::Abs
435            | Self::Sign
436            | Self::Maximum
437            | Self::Minimum
438            | Self::Compare(_)
439            | Self::Select
440            | Self::Clamp
441            | Self::Constant { .. }
442            | Self::Exp
443            | Self::Log
444            | Self::Sin
445            | Self::Cos
446            | Self::Tanh
447            | Self::Sqrt
448            | Self::Rsqrt
449            | Self::Pow
450            | Self::Expm1
451            | Self::Log1p
452            | Self::Erf
453            | Self::ExtractDiag { .. }
454            | Self::EmbedDiag { .. }
455            | Self::Tril { .. }
456            | Self::Triu { .. }
457            | Self::Gather(_)
458            | Self::GatherDynamicSliceSizes { .. }
459            | Self::Scatter(_)
460            | Self::Slice(_)
461            | Self::DynamicSlice { .. }
462            | Self::DynamicUpdateSlice
463            | Self::Pad(_)
464            | Self::Reverse { .. }
465            | Self::ShapeOf { .. }
466            | Self::DynamicTruncate { .. }
467            | Self::PadToMatch { .. }
468            | Self::ReduceProd { .. }
469            | Self::ReduceMax { .. }
470            | Self::ReduceMin { .. } => 1,
471            Self::Concatenate { .. } => 1,
472            Self::Extension(op) => ExtensionOp::output_count(op.as_ref()),
473        }
474    }
475}
476
477#[cfg(all(test, feature = "autodiff"))]
478impl StdTensorOp {
479    pub(crate) fn jvp_rule(
480        &self,
481        builder: &mut computegraph::graph::GraphBuilder<Self>,
482        primal_in: &[ValueKey<Self>],
483        primal_out: &[ValueKey<Self>],
484        tangent_in: &[Option<LocalValueId>],
485        ctx: &mut crate::ad::context::ShapeGuardContext,
486    ) -> ADRuleResult<Vec<Option<LocalValueId>>> {
487        crate::ad::linearize(self, builder, primal_in, primal_out, tangent_in, ctx)
488    }
489
490    pub(crate) fn transpose_rule(
491        &self,
492        builder: &mut computegraph::graph::GraphBuilder<Self>,
493        cotangent_out: &[Option<LocalValueId>],
494        inputs: &[computegraph::ValueRef<Self>],
495        mode: &OperationRole,
496        ctx: &mut crate::ad::context::ShapeGuardContext,
497    ) -> ADRuleResult<Vec<Option<LocalValueId>>> {
498        let inputs = inputs
499            .iter()
500            .map(|input| match input {
501                computegraph::ValueRef::Local(local_id) => {
502                    let key = builder.global_key(*local_id).clone();
503                    PrimitiveTransposeInput::Residual(key)
504                }
505                computegraph::ValueRef::External(key) => {
506                    PrimitiveTransposeInput::Residual(key.clone())
507                }
508            })
509            .collect::<Vec<_>>();
510        crate::ad::transpose_rule(self, builder, cotangent_out, inputs.as_slice(), mode, ctx)
511    }
512}