Skip to main content

tenferro_core_ops/
catalog.rs

1/// High-level category for a core primitive operation.
2///
3/// # Examples
4///
5/// ```rust
6/// use tenferro_core_ops::{descriptor, OpCategory, PrimitiveOpKind};
7///
8/// assert_eq!(
9///     descriptor(PrimitiveOpKind::ShapeOf).category,
10///     OpCategory::Host
11/// );
12/// ```
13#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
14pub enum OpCategory {
15    Elementwise,
16    Analytic,
17    Structural,
18    Reduction,
19    Contraction,
20    Indexing,
21    Dynamic,
22    Host,
23}
24
25/// Dtype compatibility policy for a core primitive operation.
26///
27/// # Examples
28///
29/// ```rust
30/// use tenferro_core_ops::{descriptor, DTypePolicy, PrimitiveOpKind};
31///
32/// assert_eq!(
33///     descriptor(PrimitiveOpKind::Compare).dtype_policy,
34///     DTypePolicy::CompareToBool
35/// );
36/// ```
37#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
38pub enum DTypePolicy {
39    SameAny,
40    SameNumeric,
41    SameFloat,
42    /// Preserve real numeric dtype and map complex magnitude to the matching real dtype.
43    AbsToReal,
44    SameFloatOrComplex,
45    CompareToBool,
46    BoolSelect,
47    Convert,
48    Shape,
49    Constant,
50}
51
52/// Static metadata for one core primitive operation.
53///
54/// # Examples
55///
56/// ```rust
57/// use tenferro_core_ops::{descriptor, PrimitiveOpKind};
58///
59/// let add = descriptor(PrimitiveOpKind::Add);
60/// assert_eq!(add.min_inputs, 2);
61/// assert_eq!(add.max_inputs, 2);
62/// ```
63#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
64pub struct PrimitiveOpDescriptor {
65    /// Catalog key for this operation.
66    pub kind: PrimitiveOpKind,
67    /// Stable snake-case operation name for diagnostics and descriptors.
68    pub name: &'static str,
69    /// Broad execution category.
70    pub category: OpCategory,
71    /// Dtype compatibility policy.
72    pub dtype_policy: DTypePolicy,
73    /// Minimum number of inputs accepted by the op.
74    pub min_inputs: u8,
75    /// Maximum number of inputs accepted by the op.
76    pub max_inputs: u8,
77    /// Whether this op is executed by host/runtime logic rather than a tensor backend.
78    pub host_only: bool,
79}
80
81macro_rules! primitive_ops {
82    ($macro:ident) => {
83        $macro! {
84            Add, "add", Elementwise, SameNumeric, 2, 2, false;
85            Sub, "sub", Elementwise, SameNumeric, 2, 2, false;
86            Mul, "mul", Elementwise, SameNumeric, 2, 2, false;
87            Neg, "neg", Elementwise, SameNumeric, 1, 1, false;
88            Conj, "conj", Elementwise, SameFloatOrComplex, 1, 1, false;
89            Div, "div", Elementwise, SameNumeric, 2, 2, false;
90            Rem, "rem", Elementwise, SameNumeric, 2, 2, false;
91            Abs, "abs", Elementwise, AbsToReal, 1, 1, false;
92            Sign, "sign", Elementwise, SameNumeric, 1, 1, false;
93            Maximum, "maximum", Elementwise, SameNumeric, 2, 2, false;
94            Minimum, "minimum", Elementwise, SameNumeric, 2, 2, false;
95            Compare, "compare", Elementwise, CompareToBool, 2, 2, false;
96            Select, "select", Elementwise, BoolSelect, 3, 3, false;
97            Clamp, "clamp", Elementwise, SameFloat, 3, 3, false;
98            Exp, "exp", Analytic, SameFloatOrComplex, 1, 1, false;
99            Log, "log", Analytic, SameFloatOrComplex, 1, 1, false;
100            Sin, "sin", Analytic, SameFloatOrComplex, 1, 1, false;
101            Cos, "cos", Analytic, SameFloatOrComplex, 1, 1, false;
102            Tanh, "tanh", Analytic, SameFloatOrComplex, 1, 1, false;
103            Sqrt, "sqrt", Analytic, SameFloatOrComplex, 1, 1, false;
104            Rsqrt, "rsqrt", Analytic, SameFloatOrComplex, 1, 1, false;
105            Pow, "pow", Analytic, SameNumeric, 2, 2, false;
106            Expm1, "expm1", Analytic, SameFloatOrComplex, 1, 1, false;
107            Log1p, "log1p", Analytic, SameFloatOrComplex, 1, 1, false;
108            Erf, "erf", Analytic, SameFloat, 1, 1, false;
109            DotGeneral, "dot_general", Contraction, SameFloatOrComplex, 2, 2, false;
110            ReduceSum, "reduce_sum", Reduction, SameNumeric, 1, 1, false;
111            ReduceSumSquares, "reduce_sum_squares", Reduction, SameFloat, 1, 1, false;
112            ReduceProd, "reduce_prod", Reduction, SameNumeric, 1, 1, false;
113            ReduceMax, "reduce_max", Reduction, SameNumeric, 1, 1, false;
114            ReduceMin, "reduce_min", Reduction, SameNumeric, 1, 1, false;
115            Transpose, "transpose", Structural, SameAny, 1, 1, false;
116            Reshape, "reshape", Structural, SameAny, 1, 1, false;
117            BroadcastInDim, "broadcast_in_dim", Structural, SameAny, 1, 1, false;
118            Convert, "convert", Structural, Convert, 1, 1, false;
119            ExtractDiag, "extract_diag", Structural, SameAny, 1, 1, false;
120            EmbedDiag, "embed_diag", Structural, SameAny, 1, 1, false;
121            Tril, "tril", Structural, SameAny, 1, 1, false;
122            Triu, "triu", Structural, SameAny, 1, 1, false;
123            Gather, "gather", Indexing, SameAny, 2, 2, false;
124            GatherDynamicSliceSizes, "gather_dynamic_slice_sizes", Indexing, SameAny, 2, 2, false;
125            Scatter, "scatter", Indexing, SameAny, 3, 3, false;
126            Slice, "slice", Indexing, SameAny, 1, 1, false;
127            DynamicSlice, "dynamic_slice", Indexing, SameAny, 2, 2, false;
128            DynamicUpdateSlice, "dynamic_update_slice", Indexing, SameAny, 3, 3, false;
129            Pad, "pad", Indexing, SameAny, 1, 1, false;
130            Concatenate, "concatenate", Indexing, SameAny, 1, u8::MAX, false;
131            Reverse, "reverse", Indexing, SameAny, 1, 1, false;
132            ShapeOf, "shape_of", Host, Shape, 1, 1, true;
133            DynamicTruncate, "dynamic_truncate", Dynamic, SameAny, 2, 2, true;
134            PadToMatch, "pad_to_match", Dynamic, SameAny, 2, 2, true;
135            Constant, "constant", Host, Constant, 0, 0, true;
136        }
137    };
138}
139
140macro_rules! define_kind {
141    ($( $variant:ident, $name:literal, $category:ident, $policy:ident, $min:expr, $max:expr, $host:expr; )*) => {
142        /// Catalog key for a core primitive operation.
143        ///
144        /// # Examples
145        ///
146        /// ```rust
147        /// use tenferro_core_ops::{descriptor, PrimitiveOpKind};
148        ///
149        /// assert_eq!(descriptor(PrimitiveOpKind::Add).name, "add");
150        /// ```
151        #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
152        pub enum PrimitiveOpKind {
153            $( $variant, )*
154        }
155
156        impl PrimitiveOpKind {
157            /// Number of primitive operation kinds in the catalog.
158            ///
159            /// # Examples
160            ///
161            /// ```rust
162            /// use tenferro_core_ops::PrimitiveOpKind;
163            ///
164            /// assert!(PrimitiveOpKind::COUNT > 0);
165            /// ```
166            pub const COUNT: usize = [$(PrimitiveOpKind::$variant),*].len();
167
168            /// Return this kind's dense catalog index.
169            ///
170            /// # Examples
171            ///
172            /// ```rust
173            /// use tenferro_core_ops::PrimitiveOpKind;
174            ///
175            /// assert_eq!(PrimitiveOpKind::Add.as_index(), 0);
176            /// ```
177            pub const fn as_index(self) -> usize {
178                self as usize
179            }
180        }
181    };
182}
183
184primitive_ops!(define_kind);
185
186macro_rules! define_descriptors {
187    ($( $variant:ident, $name:literal, $category:ident, $policy:ident, $min:expr, $max:expr, $host:expr; )*) => {
188        const DESCRIPTORS: &[PrimitiveOpDescriptor] = &[
189            $(
190                PrimitiveOpDescriptor {
191                    kind: PrimitiveOpKind::$variant,
192                    name: $name,
193                    category: OpCategory::$category,
194                    dtype_policy: DTypePolicy::$policy,
195                    min_inputs: $min,
196                    max_inputs: $max,
197                    host_only: $host,
198                },
199            )*
200        ];
201
202        /// Return the descriptor for a primitive operation kind.
203        ///
204        /// # Examples
205        ///
206        /// ```rust
207        /// use tenferro_core_ops::{descriptor, PrimitiveOpKind};
208        ///
209        /// assert_eq!(descriptor(PrimitiveOpKind::Add).name, "add");
210        /// ```
211        pub fn descriptor(kind: PrimitiveOpKind) -> &'static PrimitiveOpDescriptor {
212            match kind {
213                $(
214                    PrimitiveOpKind::$variant => &DESCRIPTORS[PrimitiveOpKind::$variant as usize],
215                )*
216            }
217        }
218    };
219}
220
221primitive_ops!(define_descriptors);
222
223/// Return all core primitive operation descriptors in catalog order.
224///
225/// # Examples
226///
227/// ```rust
228/// use tenferro_core_ops::all_primitive_descriptors;
229///
230/// assert!(all_primitive_descriptors()
231///     .iter()
232///     .any(|descriptor| descriptor.name == "add"));
233/// ```
234pub fn all_primitive_descriptors() -> &'static [PrimitiveOpDescriptor] {
235    DESCRIPTORS
236}
237
238#[doc(hidden)]
239#[macro_export]
240macro_rules! define_std_tensor_op {
241    () => {
242        #[derive(Clone, Debug)]
243        pub enum StdTensorOp {
244            // Semiring arithmetic core
245            Add,
246            Sub,
247            Mul,
248            Neg,
249            Conj,
250            DotGeneral {
251                config: DotGeneralConfig,
252            },
253            Transpose {
254                perm: Vec<usize>,
255            },
256            Reshape {
257                to_shape: Vec<DimExpr>,
258            },
259            BroadcastInDim {
260                shape: Vec<DimExpr>,
261                dims: Vec<usize>,
262            },
263            Convert {
264                from: DType,
265                to: DType,
266            },
267            Constant {
268                dtype: DType,
269                bytes: Vec<u8>,
270            },
271            ReduceSum {
272                axes: Vec<usize>,
273            },
274            ReduceSumSquares {
275                axes: Vec<usize>,
276            },
277
278            // Elementwise (non-semiring)
279            Div,
280            Rem,
281            Abs,
282            Sign,
283            Maximum,
284            Minimum,
285            Compare(CompareDir),
286            Select,
287            Clamp,
288
289            // Analytic
290            Exp,
291            Log,
292            Sin,
293            Cos,
294            Tanh,
295            Sqrt,
296            Rsqrt,
297            Pow,
298            Expm1,
299            Log1p,
300            Erf,
301
302            // Diagonal extraction / embedding (AD-closed pair)
303            ExtractDiag {
304                axis_a: usize,
305                axis_b: usize,
306            },
307            EmbedDiag {
308                axis_a: usize,
309                axis_b: usize,
310            },
311            Tril {
312                k: i64,
313            },
314            Triu {
315                k: i64,
316            },
317
318            // Indexing
319            Gather(GatherConfig),
320            GatherDynamicSliceSizes {
321                offset_dims: Vec<usize>,
322                collapsed_slice_dims: Vec<usize>,
323                start_index_map: Vec<usize>,
324                index_vector_dim: usize,
325                slice_sizes: Vec<DimExpr>,
326            },
327            Scatter(ScatterConfig),
328            Slice(SliceConfig),
329            DynamicSlice {
330                slice_sizes: Vec<usize>,
331            },
332            DynamicUpdateSlice,
333            Pad(PadConfig),
334            Concatenate {
335                axis: usize,
336                input_count: usize,
337            },
338            Reverse {
339                axes: Vec<usize>,
340            },
341            ShapeOf {
342                axis: usize,
343            },
344            DynamicTruncate {
345                axis: usize,
346            },
347            PadToMatch {
348                axis: usize,
349            },
350
351            // Reductions
352            ReduceProd {
353                axes: Vec<usize>,
354            },
355            ReduceMax {
356                axes: Vec<usize>,
357            },
358            ReduceMin {
359                axes: Vec<usize>,
360            },
361
362            /// Out-of-tree extension carrier.
363            ///
364            /// See [`crate::ext_op`] and `docs/spec/extension-op.md`. Identity,
365            /// hashing, equality, arity, shape inference, and AD rules are delegated
366            /// to the inner [`ExtensionOp`] trait object.
367            Extension(Arc<dyn ExtensionOp>),
368        }
369
370        impl StdTensorOp {
371            /// Return the core primitive catalog kind for this graph operation.
372            ///
373            /// Extension operations do not claim a core primitive kind; they are
374            /// dispatched through their extension family id instead.
375            ///
376            /// # Examples
377            ///
378            /// ```rust
379            /// use tenferro_core_ops::PrimitiveOpKind;
380            /// use tenferro_ops::std_tensor_op::StdTensorOp;
381            ///
382            /// assert_eq!(StdTensorOp::Add.primitive_kind(), Some(PrimitiveOpKind::Add));
383            /// ```
384            pub fn primitive_kind(&self) -> Option<$crate::PrimitiveOpKind> {
385                let kind = match self {
386                    Self::Add => $crate::PrimitiveOpKind::Add,
387                    Self::Sub => $crate::PrimitiveOpKind::Sub,
388                    Self::Mul => $crate::PrimitiveOpKind::Mul,
389                    Self::Neg => $crate::PrimitiveOpKind::Neg,
390                    Self::Conj => $crate::PrimitiveOpKind::Conj,
391                    Self::DotGeneral { .. } => $crate::PrimitiveOpKind::DotGeneral,
392                    Self::Transpose { .. } => $crate::PrimitiveOpKind::Transpose,
393                    Self::Reshape { .. } => $crate::PrimitiveOpKind::Reshape,
394                    Self::BroadcastInDim { .. } => $crate::PrimitiveOpKind::BroadcastInDim,
395                    Self::Convert { .. } => $crate::PrimitiveOpKind::Convert,
396                    Self::Constant { .. } => $crate::PrimitiveOpKind::Constant,
397                    Self::ReduceSum { .. } => $crate::PrimitiveOpKind::ReduceSum,
398                    Self::ReduceSumSquares { .. } => $crate::PrimitiveOpKind::ReduceSumSquares,
399                    Self::Div => $crate::PrimitiveOpKind::Div,
400                    Self::Rem => $crate::PrimitiveOpKind::Rem,
401                    Self::Abs => $crate::PrimitiveOpKind::Abs,
402                    Self::Sign => $crate::PrimitiveOpKind::Sign,
403                    Self::Maximum => $crate::PrimitiveOpKind::Maximum,
404                    Self::Minimum => $crate::PrimitiveOpKind::Minimum,
405                    Self::Compare(_) => $crate::PrimitiveOpKind::Compare,
406                    Self::Select => $crate::PrimitiveOpKind::Select,
407                    Self::Clamp => $crate::PrimitiveOpKind::Clamp,
408                    Self::Exp => $crate::PrimitiveOpKind::Exp,
409                    Self::Log => $crate::PrimitiveOpKind::Log,
410                    Self::Sin => $crate::PrimitiveOpKind::Sin,
411                    Self::Cos => $crate::PrimitiveOpKind::Cos,
412                    Self::Tanh => $crate::PrimitiveOpKind::Tanh,
413                    Self::Sqrt => $crate::PrimitiveOpKind::Sqrt,
414                    Self::Rsqrt => $crate::PrimitiveOpKind::Rsqrt,
415                    Self::Pow => $crate::PrimitiveOpKind::Pow,
416                    Self::Expm1 => $crate::PrimitiveOpKind::Expm1,
417                    Self::Log1p => $crate::PrimitiveOpKind::Log1p,
418                    Self::Erf => $crate::PrimitiveOpKind::Erf,
419                    Self::ExtractDiag { .. } => $crate::PrimitiveOpKind::ExtractDiag,
420                    Self::EmbedDiag { .. } => $crate::PrimitiveOpKind::EmbedDiag,
421                    Self::Tril { .. } => $crate::PrimitiveOpKind::Tril,
422                    Self::Triu { .. } => $crate::PrimitiveOpKind::Triu,
423                    Self::Gather(_) => $crate::PrimitiveOpKind::Gather,
424                    Self::GatherDynamicSliceSizes { .. } => {
425                        $crate::PrimitiveOpKind::GatherDynamicSliceSizes
426                    }
427                    Self::Scatter(_) => $crate::PrimitiveOpKind::Scatter,
428                    Self::Slice(_) => $crate::PrimitiveOpKind::Slice,
429                    Self::DynamicSlice { .. } => $crate::PrimitiveOpKind::DynamicSlice,
430                    Self::DynamicUpdateSlice => $crate::PrimitiveOpKind::DynamicUpdateSlice,
431                    Self::Pad(_) => $crate::PrimitiveOpKind::Pad,
432                    Self::Concatenate { .. } => $crate::PrimitiveOpKind::Concatenate,
433                    Self::Reverse { .. } => $crate::PrimitiveOpKind::Reverse,
434                    Self::ShapeOf { .. } => $crate::PrimitiveOpKind::ShapeOf,
435                    Self::DynamicTruncate { .. } => $crate::PrimitiveOpKind::DynamicTruncate,
436                    Self::PadToMatch { .. } => $crate::PrimitiveOpKind::PadToMatch,
437                    Self::ReduceProd { .. } => $crate::PrimitiveOpKind::ReduceProd,
438                    Self::ReduceMax { .. } => $crate::PrimitiveOpKind::ReduceMax,
439                    Self::ReduceMin { .. } => $crate::PrimitiveOpKind::ReduceMin,
440                    Self::Extension(_) => return None,
441                };
442                Some(kind)
443            }
444
445            #[cfg(test)]
446            pub(crate) fn sample_from_kind(kind: $crate::PrimitiveOpKind) -> Self {
447                match kind {
448                    $crate::PrimitiveOpKind::Add => Self::Add,
449                    $crate::PrimitiveOpKind::Sub => Self::Sub,
450                    $crate::PrimitiveOpKind::Mul => Self::Mul,
451                    $crate::PrimitiveOpKind::Neg => Self::Neg,
452                    $crate::PrimitiveOpKind::Conj => Self::Conj,
453                    $crate::PrimitiveOpKind::DotGeneral => Self::DotGeneral {
454                        config: DotGeneralConfig {
455                            lhs_contracting_dims: [0].as_slice().into(),
456                            rhs_contracting_dims: [0].as_slice().into(),
457                            lhs_batch_dims: [].as_slice().into(),
458                            rhs_batch_dims: [].as_slice().into(),
459                        },
460                    },
461                    $crate::PrimitiveOpKind::Transpose => Self::Transpose { perm: vec![0] },
462                    $crate::PrimitiveOpKind::Reshape => Self::Reshape {
463                        to_shape: vec![DimExpr::Const(1)],
464                    },
465                    $crate::PrimitiveOpKind::BroadcastInDim => Self::BroadcastInDim {
466                        shape: vec![DimExpr::Const(1)],
467                        dims: vec![0],
468                    },
469                    $crate::PrimitiveOpKind::Convert => Self::Convert {
470                        from: DType::F32,
471                        to: DType::F64,
472                    },
473                    $crate::PrimitiveOpKind::Constant => Self::Constant {
474                        dtype: DType::F64,
475                        bytes: 0.0_f64.to_le_bytes().to_vec(),
476                    },
477                    $crate::PrimitiveOpKind::ReduceSum => Self::ReduceSum { axes: vec![0] },
478                    $crate::PrimitiveOpKind::ReduceSumSquares => {
479                        Self::ReduceSumSquares { axes: vec![0] }
480                    }
481                    $crate::PrimitiveOpKind::Div => Self::Div,
482                    $crate::PrimitiveOpKind::Rem => Self::Rem,
483                    $crate::PrimitiveOpKind::Abs => Self::Abs,
484                    $crate::PrimitiveOpKind::Sign => Self::Sign,
485                    $crate::PrimitiveOpKind::Maximum => Self::Maximum,
486                    $crate::PrimitiveOpKind::Minimum => Self::Minimum,
487                    $crate::PrimitiveOpKind::Compare => Self::Compare(CompareDir::Eq),
488                    $crate::PrimitiveOpKind::Select => Self::Select,
489                    $crate::PrimitiveOpKind::Clamp => Self::Clamp,
490                    $crate::PrimitiveOpKind::Exp => Self::Exp,
491                    $crate::PrimitiveOpKind::Log => Self::Log,
492                    $crate::PrimitiveOpKind::Sin => Self::Sin,
493                    $crate::PrimitiveOpKind::Cos => Self::Cos,
494                    $crate::PrimitiveOpKind::Tanh => Self::Tanh,
495                    $crate::PrimitiveOpKind::Sqrt => Self::Sqrt,
496                    $crate::PrimitiveOpKind::Rsqrt => Self::Rsqrt,
497                    $crate::PrimitiveOpKind::Pow => Self::Pow,
498                    $crate::PrimitiveOpKind::Expm1 => Self::Expm1,
499                    $crate::PrimitiveOpKind::Log1p => Self::Log1p,
500                    $crate::PrimitiveOpKind::Erf => Self::Erf,
501                    $crate::PrimitiveOpKind::ExtractDiag => Self::ExtractDiag {
502                        axis_a: 0,
503                        axis_b: 1,
504                    },
505                    $crate::PrimitiveOpKind::EmbedDiag => Self::EmbedDiag {
506                        axis_a: 0,
507                        axis_b: 1,
508                    },
509                    $crate::PrimitiveOpKind::Tril => Self::Tril { k: 0 },
510                    $crate::PrimitiveOpKind::Triu => Self::Triu { k: 0 },
511                    $crate::PrimitiveOpKind::Gather => Self::Gather(GatherConfig {
512                        offset_dims: vec![],
513                        collapsed_slice_dims: vec![0],
514                        start_index_map: vec![0],
515                        index_vector_dim: 1,
516                        slice_sizes: vec![1],
517                    }),
518                    $crate::PrimitiveOpKind::GatherDynamicSliceSizes => {
519                        Self::GatherDynamicSliceSizes {
520                            offset_dims: vec![],
521                            collapsed_slice_dims: vec![0],
522                            start_index_map: vec![0],
523                            index_vector_dim: 1,
524                            slice_sizes: vec![DimExpr::Const(1)],
525                        }
526                    }
527                    $crate::PrimitiveOpKind::Scatter => Self::Scatter(ScatterConfig {
528                        update_window_dims: vec![],
529                        inserted_window_dims: vec![0],
530                        scatter_dims_to_operand_dims: vec![0],
531                        index_vector_dim: 1,
532                    }),
533                    $crate::PrimitiveOpKind::Slice => Self::Slice(SliceConfig {
534                        starts: vec![0],
535                        limits: vec![1],
536                        strides: vec![1],
537                    }),
538                    $crate::PrimitiveOpKind::DynamicSlice => Self::DynamicSlice {
539                        slice_sizes: vec![1],
540                    },
541                    $crate::PrimitiveOpKind::DynamicUpdateSlice => Self::DynamicUpdateSlice,
542                    $crate::PrimitiveOpKind::Pad => Self::Pad(PadConfig {
543                        edge_padding_low: vec![0],
544                        edge_padding_high: vec![0],
545                        interior_padding: vec![0],
546                    }),
547                    $crate::PrimitiveOpKind::Concatenate => Self::Concatenate {
548                        axis: 0,
549                        input_count: 1,
550                    },
551                    $crate::PrimitiveOpKind::Reverse => Self::Reverse { axes: vec![0] },
552                    $crate::PrimitiveOpKind::ShapeOf => Self::ShapeOf { axis: 0 },
553                    $crate::PrimitiveOpKind::DynamicTruncate => Self::DynamicTruncate { axis: 0 },
554                    $crate::PrimitiveOpKind::PadToMatch => Self::PadToMatch { axis: 0 },
555                    $crate::PrimitiveOpKind::ReduceProd => Self::ReduceProd { axes: vec![0] },
556                    $crate::PrimitiveOpKind::ReduceMax => Self::ReduceMax { axes: vec![0] },
557                    $crate::PrimitiveOpKind::ReduceMin => Self::ReduceMin { axes: vec![0] },
558                }
559            }
560        }
561    };
562}
563
564#[doc(hidden)]
565#[macro_export]
566macro_rules! define_elementwise_fusion_op {
567    () => {
568        /// Elementwise op kinds supported by backend fusion implementations.
569        #[doc(hidden)]
570        #[derive(Clone, Copy, Debug, Hash, PartialEq, Eq)]
571        pub enum ElementwiseFusionOp {
572            Add,
573            Multiply,
574            Negate,
575            Conj,
576            Divide,
577            Remainder,
578            Abs,
579            Maximum,
580            Minimum,
581            Clamp,
582            Exp,
583            Log,
584            Sin,
585            Cos,
586            Tanh,
587            Sqrt,
588            Rsqrt,
589            Pow,
590            Expm1,
591            Log1p,
592            Erf,
593        }
594
595        #[cfg(test)]
596        impl ElementwiseFusionOp {
597            pub(crate) fn iter() -> impl Iterator<Item = Self> {
598                [
599                    Self::Add,
600                    Self::Multiply,
601                    Self::Negate,
602                    Self::Conj,
603                    Self::Divide,
604                    Self::Remainder,
605                    Self::Abs,
606                    Self::Maximum,
607                    Self::Minimum,
608                    Self::Clamp,
609                    Self::Exp,
610                    Self::Log,
611                    Self::Sin,
612                    Self::Cos,
613                    Self::Tanh,
614                    Self::Sqrt,
615                    Self::Rsqrt,
616                    Self::Pow,
617                    Self::Expm1,
618                    Self::Log1p,
619                    Self::Erf,
620                ]
621                .into_iter()
622            }
623
624            pub(crate) fn from_primitive_kind(kind: $crate::PrimitiveOpKind) -> Option<Self> {
625                match kind {
626                    $crate::PrimitiveOpKind::Add => Some(Self::Add),
627                    $crate::PrimitiveOpKind::Mul => Some(Self::Multiply),
628                    $crate::PrimitiveOpKind::Neg => Some(Self::Negate),
629                    $crate::PrimitiveOpKind::Conj => Some(Self::Conj),
630                    $crate::PrimitiveOpKind::Div => Some(Self::Divide),
631                    $crate::PrimitiveOpKind::Rem => Some(Self::Remainder),
632                    $crate::PrimitiveOpKind::Abs => Some(Self::Abs),
633                    $crate::PrimitiveOpKind::Maximum => Some(Self::Maximum),
634                    $crate::PrimitiveOpKind::Minimum => Some(Self::Minimum),
635                    $crate::PrimitiveOpKind::Clamp => Some(Self::Clamp),
636                    $crate::PrimitiveOpKind::Exp => Some(Self::Exp),
637                    $crate::PrimitiveOpKind::Log => Some(Self::Log),
638                    $crate::PrimitiveOpKind::Sin => Some(Self::Sin),
639                    $crate::PrimitiveOpKind::Cos => Some(Self::Cos),
640                    $crate::PrimitiveOpKind::Tanh => Some(Self::Tanh),
641                    $crate::PrimitiveOpKind::Sqrt => Some(Self::Sqrt),
642                    $crate::PrimitiveOpKind::Rsqrt => Some(Self::Rsqrt),
643                    $crate::PrimitiveOpKind::Pow => Some(Self::Pow),
644                    $crate::PrimitiveOpKind::Expm1 => Some(Self::Expm1),
645                    $crate::PrimitiveOpKind::Log1p => Some(Self::Log1p),
646                    $crate::PrimitiveOpKind::Erf => Some(Self::Erf),
647                    _ => None,
648                }
649            }
650
651            pub(crate) fn primitive_kind(self) -> $crate::PrimitiveOpKind {
652                match self {
653                    Self::Add => $crate::PrimitiveOpKind::Add,
654                    Self::Multiply => $crate::PrimitiveOpKind::Mul,
655                    Self::Negate => $crate::PrimitiveOpKind::Neg,
656                    Self::Conj => $crate::PrimitiveOpKind::Conj,
657                    Self::Divide => $crate::PrimitiveOpKind::Div,
658                    Self::Remainder => $crate::PrimitiveOpKind::Rem,
659                    Self::Abs => $crate::PrimitiveOpKind::Abs,
660                    Self::Maximum => $crate::PrimitiveOpKind::Maximum,
661                    Self::Minimum => $crate::PrimitiveOpKind::Minimum,
662                    Self::Clamp => $crate::PrimitiveOpKind::Clamp,
663                    Self::Exp => $crate::PrimitiveOpKind::Exp,
664                    Self::Log => $crate::PrimitiveOpKind::Log,
665                    Self::Sin => $crate::PrimitiveOpKind::Sin,
666                    Self::Cos => $crate::PrimitiveOpKind::Cos,
667                    Self::Tanh => $crate::PrimitiveOpKind::Tanh,
668                    Self::Sqrt => $crate::PrimitiveOpKind::Sqrt,
669                    Self::Rsqrt => $crate::PrimitiveOpKind::Rsqrt,
670                    Self::Pow => $crate::PrimitiveOpKind::Pow,
671                    Self::Expm1 => $crate::PrimitiveOpKind::Expm1,
672                    Self::Log1p => $crate::PrimitiveOpKind::Log1p,
673                    Self::Erf => $crate::PrimitiveOpKind::Erf,
674                }
675            }
676        }
677    };
678}
679
680#[doc(hidden)]
681#[macro_export]
682macro_rules! define_exec_op {
683    () => {
684        #[derive(Clone, Debug)]
685        pub enum ExecOp {
686            Transpose {
687                perm: Vec<usize>,
688            },
689            Reshape {
690                shape: Vec<DimExpr>,
691            },
692            BroadcastInDim {
693                shape: Vec<DimExpr>,
694                dims: Vec<usize>,
695            },
696            Convert {
697                to: DType,
698            },
699            Constant {
700                dtype: DType,
701                bytes: Vec<u8>,
702            },
703            DotGeneral(DotGeneralConfig),
704            DotGeneralWithConj {
705                config: DotGeneralConfig,
706                lhs_conj: bool,
707                rhs_conj: bool,
708            },
709            ReduceSum {
710                axes: Vec<usize>,
711            },
712            ReduceSumSquares {
713                axes: Vec<usize>,
714            },
715            ExtractDiag {
716                axis_a: usize,
717                axis_b: usize,
718            },
719            EmbedDiag {
720                axis_a: usize,
721                axis_b: usize,
722            },
723            Tril {
724                k: i64,
725            },
726            Triu {
727                k: i64,
728            },
729            Add,
730            Subtract,
731            Multiply,
732            Negate,
733            Conj,
734            Divide,
735            Remainder,
736            Abs,
737            Sign,
738            Maximum,
739            Minimum,
740            Compare(CompareDir),
741            Select,
742            Clamp,
743            Exp,
744            Log,
745            Sin,
746            Cos,
747            Tanh,
748            Sqrt,
749            Rsqrt,
750            Pow,
751            Expm1,
752            Log1p,
753            Erf,
754            Gather(GatherConfig),
755            GatherDynamicSliceSizes {
756                offset_dims: Vec<usize>,
757                collapsed_slice_dims: Vec<usize>,
758                start_index_map: Vec<usize>,
759                index_vector_dim: usize,
760                slice_sizes: Vec<DimExpr>,
761            },
762            Scatter(ScatterConfig),
763            Slice(SliceConfig),
764            DynamicSlice {
765                slice_sizes: Vec<usize>,
766            },
767            DynamicUpdateSlice,
768            Pad(PadConfig),
769            Concatenate {
770                axis: usize,
771            },
772            Reverse {
773                axes: Vec<usize>,
774            },
775            ShapeOf {
776                axis: usize,
777            },
778            DynamicTruncate {
779                axis: usize,
780            },
781            PadToMatch {
782                axis: usize,
783            },
784            ReduceProd {
785                axes: Vec<usize>,
786            },
787            ReduceMax {
788                axes: Vec<usize>,
789            },
790            ReduceMin {
791                axes: Vec<usize>,
792            },
793            /// Out-of-tree extension carrier in the execution IR.
794            ///
795            /// Payload and dispatch are defined by the inner [`ExtensionOp`]. The
796            /// execution pipeline treats extensions as single-instruction FFI
797            /// boundaries (spec Section 8): no elementwise fusion, and dispatch is
798            /// routed through the executor's registered extension runtime.
799            Extension(Arc<dyn ExtensionOp>),
800        }
801
802        impl ExecOp {
803            pub(crate) fn primitive_kind(&self) -> Option<$crate::PrimitiveOpKind> {
804                let kind = match self {
805                    Self::Transpose { .. } => $crate::PrimitiveOpKind::Transpose,
806                    Self::Reshape { .. } => $crate::PrimitiveOpKind::Reshape,
807                    Self::BroadcastInDim { .. } => $crate::PrimitiveOpKind::BroadcastInDim,
808                    Self::Convert { .. } => $crate::PrimitiveOpKind::Convert,
809                    Self::Constant { .. } => $crate::PrimitiveOpKind::Constant,
810                    Self::DotGeneral(_) | Self::DotGeneralWithConj { .. } => {
811                        $crate::PrimitiveOpKind::DotGeneral
812                    }
813                    Self::ReduceSum { .. } => $crate::PrimitiveOpKind::ReduceSum,
814                    Self::ReduceSumSquares { .. } => $crate::PrimitiveOpKind::ReduceSumSquares,
815                    Self::ExtractDiag { .. } => $crate::PrimitiveOpKind::ExtractDiag,
816                    Self::EmbedDiag { .. } => $crate::PrimitiveOpKind::EmbedDiag,
817                    Self::Tril { .. } => $crate::PrimitiveOpKind::Tril,
818                    Self::Triu { .. } => $crate::PrimitiveOpKind::Triu,
819                    Self::Add => $crate::PrimitiveOpKind::Add,
820                    Self::Subtract => $crate::PrimitiveOpKind::Sub,
821                    Self::Multiply => $crate::PrimitiveOpKind::Mul,
822                    Self::Negate => $crate::PrimitiveOpKind::Neg,
823                    Self::Conj => $crate::PrimitiveOpKind::Conj,
824                    Self::Divide => $crate::PrimitiveOpKind::Div,
825                    Self::Remainder => $crate::PrimitiveOpKind::Rem,
826                    Self::Abs => $crate::PrimitiveOpKind::Abs,
827                    Self::Sign => $crate::PrimitiveOpKind::Sign,
828                    Self::Maximum => $crate::PrimitiveOpKind::Maximum,
829                    Self::Minimum => $crate::PrimitiveOpKind::Minimum,
830                    Self::Compare(_) => $crate::PrimitiveOpKind::Compare,
831                    Self::Select => $crate::PrimitiveOpKind::Select,
832                    Self::Clamp => $crate::PrimitiveOpKind::Clamp,
833                    Self::Exp => $crate::PrimitiveOpKind::Exp,
834                    Self::Log => $crate::PrimitiveOpKind::Log,
835                    Self::Sin => $crate::PrimitiveOpKind::Sin,
836                    Self::Cos => $crate::PrimitiveOpKind::Cos,
837                    Self::Tanh => $crate::PrimitiveOpKind::Tanh,
838                    Self::Sqrt => $crate::PrimitiveOpKind::Sqrt,
839                    Self::Rsqrt => $crate::PrimitiveOpKind::Rsqrt,
840                    Self::Pow => $crate::PrimitiveOpKind::Pow,
841                    Self::Expm1 => $crate::PrimitiveOpKind::Expm1,
842                    Self::Log1p => $crate::PrimitiveOpKind::Log1p,
843                    Self::Erf => $crate::PrimitiveOpKind::Erf,
844                    Self::Gather(_) => $crate::PrimitiveOpKind::Gather,
845                    Self::GatherDynamicSliceSizes { .. } => {
846                        $crate::PrimitiveOpKind::GatherDynamicSliceSizes
847                    }
848                    Self::Scatter(_) => $crate::PrimitiveOpKind::Scatter,
849                    Self::Slice(_) => $crate::PrimitiveOpKind::Slice,
850                    Self::DynamicSlice { .. } => $crate::PrimitiveOpKind::DynamicSlice,
851                    Self::DynamicUpdateSlice => $crate::PrimitiveOpKind::DynamicUpdateSlice,
852                    Self::Pad(_) => $crate::PrimitiveOpKind::Pad,
853                    Self::Concatenate { .. } => $crate::PrimitiveOpKind::Concatenate,
854                    Self::Reverse { .. } => $crate::PrimitiveOpKind::Reverse,
855                    Self::ShapeOf { .. } => $crate::PrimitiveOpKind::ShapeOf,
856                    Self::DynamicTruncate { .. } => $crate::PrimitiveOpKind::DynamicTruncate,
857                    Self::PadToMatch { .. } => $crate::PrimitiveOpKind::PadToMatch,
858                    Self::ReduceProd { .. } => $crate::PrimitiveOpKind::ReduceProd,
859                    Self::ReduceMax { .. } => $crate::PrimitiveOpKind::ReduceMax,
860                    Self::ReduceMin { .. } => $crate::PrimitiveOpKind::ReduceMin,
861                    Self::Extension(_) => return None,
862                };
863                Some(kind)
864            }
865
866            pub(crate) fn from_std_tensor_op(
867                op: &tenferro_ops::std_tensor_op::StdTensorOp,
868            ) -> Self {
869                match op {
870                    tenferro_ops::std_tensor_op::StdTensorOp::Add => Self::Add,
871                    tenferro_ops::std_tensor_op::StdTensorOp::Sub => Self::Subtract,
872                    tenferro_ops::std_tensor_op::StdTensorOp::Mul => Self::Multiply,
873                    tenferro_ops::std_tensor_op::StdTensorOp::Neg => Self::Negate,
874                    tenferro_ops::std_tensor_op::StdTensorOp::Conj => Self::Conj,
875                    tenferro_ops::std_tensor_op::StdTensorOp::Div => Self::Divide,
876                    tenferro_ops::std_tensor_op::StdTensorOp::Rem => Self::Remainder,
877                    tenferro_ops::std_tensor_op::StdTensorOp::Abs => Self::Abs,
878                    tenferro_ops::std_tensor_op::StdTensorOp::Sign => Self::Sign,
879                    tenferro_ops::std_tensor_op::StdTensorOp::Maximum => Self::Maximum,
880                    tenferro_ops::std_tensor_op::StdTensorOp::Minimum => Self::Minimum,
881                    tenferro_ops::std_tensor_op::StdTensorOp::Compare(dir) => {
882                        Self::Compare(dir.clone())
883                    }
884                    tenferro_ops::std_tensor_op::StdTensorOp::Select => Self::Select,
885                    tenferro_ops::std_tensor_op::StdTensorOp::Clamp => Self::Clamp,
886                    tenferro_ops::std_tensor_op::StdTensorOp::Exp => Self::Exp,
887                    tenferro_ops::std_tensor_op::StdTensorOp::Log => Self::Log,
888                    tenferro_ops::std_tensor_op::StdTensorOp::Sin => Self::Sin,
889                    tenferro_ops::std_tensor_op::StdTensorOp::Cos => Self::Cos,
890                    tenferro_ops::std_tensor_op::StdTensorOp::Tanh => Self::Tanh,
891                    tenferro_ops::std_tensor_op::StdTensorOp::Sqrt => Self::Sqrt,
892                    tenferro_ops::std_tensor_op::StdTensorOp::Rsqrt => Self::Rsqrt,
893                    tenferro_ops::std_tensor_op::StdTensorOp::Pow => Self::Pow,
894                    tenferro_ops::std_tensor_op::StdTensorOp::Expm1 => Self::Expm1,
895                    tenferro_ops::std_tensor_op::StdTensorOp::Log1p => Self::Log1p,
896                    tenferro_ops::std_tensor_op::StdTensorOp::Erf => Self::Erf,
897                    tenferro_ops::std_tensor_op::StdTensorOp::Transpose { perm } => {
898                        Self::Transpose { perm: perm.clone() }
899                    }
900                    tenferro_ops::std_tensor_op::StdTensorOp::Reshape { to_shape } => {
901                        Self::Reshape {
902                            shape: to_shape.clone(),
903                        }
904                    }
905                    tenferro_ops::std_tensor_op::StdTensorOp::BroadcastInDim { shape, dims } => {
906                        Self::BroadcastInDim {
907                            shape: shape.clone(),
908                            dims: dims.clone(),
909                        }
910                    }
911                    tenferro_ops::std_tensor_op::StdTensorOp::Convert { to, .. } => {
912                        Self::Convert { to: *to }
913                    }
914                    tenferro_ops::std_tensor_op::StdTensorOp::Constant { dtype, bytes } => {
915                        Self::Constant {
916                            dtype: *dtype,
917                            bytes: bytes.clone(),
918                        }
919                    }
920                    tenferro_ops::std_tensor_op::StdTensorOp::DotGeneral { config } => {
921                        Self::DotGeneral(config.clone())
922                    }
923                    tenferro_ops::std_tensor_op::StdTensorOp::ReduceSum { axes } => {
924                        Self::ReduceSum { axes: axes.clone() }
925                    }
926                    tenferro_ops::std_tensor_op::StdTensorOp::ReduceSumSquares { axes } => {
927                        Self::ReduceSumSquares { axes: axes.clone() }
928                    }
929                    tenferro_ops::std_tensor_op::StdTensorOp::ReduceProd { axes } => {
930                        Self::ReduceProd { axes: axes.clone() }
931                    }
932                    tenferro_ops::std_tensor_op::StdTensorOp::ReduceMax { axes } => {
933                        Self::ReduceMax { axes: axes.clone() }
934                    }
935                    tenferro_ops::std_tensor_op::StdTensorOp::ReduceMin { axes } => {
936                        Self::ReduceMin { axes: axes.clone() }
937                    }
938                    tenferro_ops::std_tensor_op::StdTensorOp::ExtractDiag { axis_a, axis_b } => {
939                        Self::ExtractDiag {
940                            axis_a: *axis_a,
941                            axis_b: *axis_b,
942                        }
943                    }
944                    tenferro_ops::std_tensor_op::StdTensorOp::EmbedDiag { axis_a, axis_b } => {
945                        Self::EmbedDiag {
946                            axis_a: *axis_a,
947                            axis_b: *axis_b,
948                        }
949                    }
950                    tenferro_ops::std_tensor_op::StdTensorOp::Tril { k } => Self::Tril { k: *k },
951                    tenferro_ops::std_tensor_op::StdTensorOp::Triu { k } => Self::Triu { k: *k },
952                    tenferro_ops::std_tensor_op::StdTensorOp::Gather(config) => {
953                        Self::Gather(config.clone())
954                    }
955                    tenferro_ops::std_tensor_op::StdTensorOp::GatherDynamicSliceSizes {
956                        offset_dims,
957                        collapsed_slice_dims,
958                        start_index_map,
959                        index_vector_dim,
960                        slice_sizes,
961                    } => Self::GatherDynamicSliceSizes {
962                        offset_dims: offset_dims.clone(),
963                        collapsed_slice_dims: collapsed_slice_dims.clone(),
964                        start_index_map: start_index_map.clone(),
965                        index_vector_dim: *index_vector_dim,
966                        slice_sizes: slice_sizes.clone(),
967                    },
968                    tenferro_ops::std_tensor_op::StdTensorOp::Scatter(config) => {
969                        Self::Scatter(config.clone())
970                    }
971                    tenferro_ops::std_tensor_op::StdTensorOp::Slice(config) => {
972                        Self::Slice(config.clone())
973                    }
974                    tenferro_ops::std_tensor_op::StdTensorOp::DynamicSlice { slice_sizes } => {
975                        Self::DynamicSlice {
976                            slice_sizes: slice_sizes.clone(),
977                        }
978                    }
979                    tenferro_ops::std_tensor_op::StdTensorOp::DynamicUpdateSlice => {
980                        Self::DynamicUpdateSlice
981                    }
982                    tenferro_ops::std_tensor_op::StdTensorOp::Pad(config) => {
983                        Self::Pad(config.clone())
984                    }
985                    tenferro_ops::std_tensor_op::StdTensorOp::Concatenate { axis, .. } => {
986                        Self::Concatenate { axis: *axis }
987                    }
988                    tenferro_ops::std_tensor_op::StdTensorOp::Reverse { axes } => {
989                        Self::Reverse { axes: axes.clone() }
990                    }
991                    tenferro_ops::std_tensor_op::StdTensorOp::ShapeOf { axis } => {
992                        Self::ShapeOf { axis: *axis }
993                    }
994                    tenferro_ops::std_tensor_op::StdTensorOp::DynamicTruncate { axis } => {
995                        Self::DynamicTruncate { axis: *axis }
996                    }
997                    tenferro_ops::std_tensor_op::StdTensorOp::PadToMatch { axis } => {
998                        Self::PadToMatch { axis: *axis }
999                    }
1000                    tenferro_ops::std_tensor_op::StdTensorOp::Extension(op) => {
1001                        Self::Extension(op.clone())
1002                    }
1003                }
1004            }
1005
1006            pub(crate) fn elementwise_fusion_op(&self) -> Option<ElementwiseFusionOp> {
1007                match self {
1008                    Self::Add => Some(ElementwiseFusionOp::Add),
1009                    Self::Multiply => Some(ElementwiseFusionOp::Multiply),
1010                    Self::Negate => Some(ElementwiseFusionOp::Negate),
1011                    Self::Conj => Some(ElementwiseFusionOp::Conj),
1012                    Self::Divide => Some(ElementwiseFusionOp::Divide),
1013                    Self::Abs => Some(ElementwiseFusionOp::Abs),
1014                    Self::Maximum => Some(ElementwiseFusionOp::Maximum),
1015                    Self::Minimum => Some(ElementwiseFusionOp::Minimum),
1016                    Self::Clamp => Some(ElementwiseFusionOp::Clamp),
1017                    Self::Exp => Some(ElementwiseFusionOp::Exp),
1018                    Self::Log => Some(ElementwiseFusionOp::Log),
1019                    Self::Sin => Some(ElementwiseFusionOp::Sin),
1020                    Self::Cos => Some(ElementwiseFusionOp::Cos),
1021                    Self::Tanh => Some(ElementwiseFusionOp::Tanh),
1022                    Self::Sqrt => Some(ElementwiseFusionOp::Sqrt),
1023                    Self::Rsqrt => Some(ElementwiseFusionOp::Rsqrt),
1024                    Self::Pow => Some(ElementwiseFusionOp::Pow),
1025                    Self::Expm1 => Some(ElementwiseFusionOp::Expm1),
1026                    Self::Log1p => Some(ElementwiseFusionOp::Log1p),
1027                    Self::Erf => Some(ElementwiseFusionOp::Erf),
1028                    _ => None,
1029                }
1030            }
1031
1032            #[cfg(test)]
1033            pub(crate) fn input_arity_bounds(&self) -> Option<(u8, u8)> {
1034                self.primitive_kind().map(|kind| {
1035                    let descriptor = $crate::descriptor(kind);
1036                    (descriptor.min_inputs, descriptor.max_inputs)
1037                })
1038            }
1039
1040            #[cfg(test)]
1041            pub(crate) fn sample_from_kind(kind: $crate::PrimitiveOpKind) -> Self {
1042                match kind {
1043                    $crate::PrimitiveOpKind::Transpose => Self::Transpose { perm: vec![0] },
1044                    $crate::PrimitiveOpKind::Reshape => Self::Reshape {
1045                        shape: vec![DimExpr::Const(1)],
1046                    },
1047                    $crate::PrimitiveOpKind::BroadcastInDim => Self::BroadcastInDim {
1048                        shape: vec![DimExpr::Const(1)],
1049                        dims: vec![0],
1050                    },
1051                    $crate::PrimitiveOpKind::Convert => Self::Convert { to: DType::F64 },
1052                    $crate::PrimitiveOpKind::Constant => Self::Constant {
1053                        dtype: DType::F64,
1054                        bytes: 0.0_f64.to_le_bytes().to_vec(),
1055                    },
1056                    $crate::PrimitiveOpKind::DotGeneral => Self::DotGeneral(DotGeneralConfig {
1057                        lhs_contracting_dims: [0].as_slice().into(),
1058                        rhs_contracting_dims: [0].as_slice().into(),
1059                        lhs_batch_dims: [].as_slice().into(),
1060                        rhs_batch_dims: [].as_slice().into(),
1061                    }),
1062                    $crate::PrimitiveOpKind::ReduceSum => Self::ReduceSum { axes: vec![0] },
1063                    $crate::PrimitiveOpKind::ReduceSumSquares => {
1064                        Self::ReduceSumSquares { axes: vec![0] }
1065                    }
1066                    $crate::PrimitiveOpKind::ExtractDiag => Self::ExtractDiag {
1067                        axis_a: 0,
1068                        axis_b: 1,
1069                    },
1070                    $crate::PrimitiveOpKind::EmbedDiag => Self::EmbedDiag {
1071                        axis_a: 0,
1072                        axis_b: 1,
1073                    },
1074                    $crate::PrimitiveOpKind::Tril => Self::Tril { k: 0 },
1075                    $crate::PrimitiveOpKind::Triu => Self::Triu { k: 0 },
1076                    $crate::PrimitiveOpKind::Add => Self::Add,
1077                    $crate::PrimitiveOpKind::Sub => Self::Subtract,
1078                    $crate::PrimitiveOpKind::Mul => Self::Multiply,
1079                    $crate::PrimitiveOpKind::Neg => Self::Negate,
1080                    $crate::PrimitiveOpKind::Conj => Self::Conj,
1081                    $crate::PrimitiveOpKind::Div => Self::Divide,
1082                    $crate::PrimitiveOpKind::Rem => Self::Remainder,
1083                    $crate::PrimitiveOpKind::Abs => Self::Abs,
1084                    $crate::PrimitiveOpKind::Sign => Self::Sign,
1085                    $crate::PrimitiveOpKind::Maximum => Self::Maximum,
1086                    $crate::PrimitiveOpKind::Minimum => Self::Minimum,
1087                    $crate::PrimitiveOpKind::Compare => Self::Compare(CompareDir::Eq),
1088                    $crate::PrimitiveOpKind::Select => Self::Select,
1089                    $crate::PrimitiveOpKind::Clamp => Self::Clamp,
1090                    $crate::PrimitiveOpKind::Exp => Self::Exp,
1091                    $crate::PrimitiveOpKind::Log => Self::Log,
1092                    $crate::PrimitiveOpKind::Sin => Self::Sin,
1093                    $crate::PrimitiveOpKind::Cos => Self::Cos,
1094                    $crate::PrimitiveOpKind::Tanh => Self::Tanh,
1095                    $crate::PrimitiveOpKind::Sqrt => Self::Sqrt,
1096                    $crate::PrimitiveOpKind::Rsqrt => Self::Rsqrt,
1097                    $crate::PrimitiveOpKind::Pow => Self::Pow,
1098                    $crate::PrimitiveOpKind::Expm1 => Self::Expm1,
1099                    $crate::PrimitiveOpKind::Log1p => Self::Log1p,
1100                    $crate::PrimitiveOpKind::Erf => Self::Erf,
1101                    $crate::PrimitiveOpKind::Gather => Self::Gather(GatherConfig {
1102                        offset_dims: vec![],
1103                        collapsed_slice_dims: vec![0],
1104                        start_index_map: vec![0],
1105                        index_vector_dim: 1,
1106                        slice_sizes: vec![1],
1107                    }),
1108                    $crate::PrimitiveOpKind::GatherDynamicSliceSizes => {
1109                        Self::GatherDynamicSliceSizes {
1110                            offset_dims: vec![],
1111                            collapsed_slice_dims: vec![0],
1112                            start_index_map: vec![0],
1113                            index_vector_dim: 1,
1114                            slice_sizes: vec![DimExpr::Const(1)],
1115                        }
1116                    }
1117                    $crate::PrimitiveOpKind::Scatter => Self::Scatter(ScatterConfig {
1118                        update_window_dims: vec![],
1119                        inserted_window_dims: vec![0],
1120                        scatter_dims_to_operand_dims: vec![0],
1121                        index_vector_dim: 1,
1122                    }),
1123                    $crate::PrimitiveOpKind::Slice => Self::Slice(SliceConfig {
1124                        starts: vec![0],
1125                        limits: vec![1],
1126                        strides: vec![1],
1127                    }),
1128                    $crate::PrimitiveOpKind::DynamicSlice => Self::DynamicSlice {
1129                        slice_sizes: vec![1],
1130                    },
1131                    $crate::PrimitiveOpKind::DynamicUpdateSlice => Self::DynamicUpdateSlice,
1132                    $crate::PrimitiveOpKind::Pad => Self::Pad(PadConfig {
1133                        edge_padding_low: vec![0],
1134                        edge_padding_high: vec![0],
1135                        interior_padding: vec![0],
1136                    }),
1137                    $crate::PrimitiveOpKind::Concatenate => Self::Concatenate { axis: 0 },
1138                    $crate::PrimitiveOpKind::Reverse => Self::Reverse { axes: vec![0] },
1139                    $crate::PrimitiveOpKind::ShapeOf => Self::ShapeOf { axis: 0 },
1140                    $crate::PrimitiveOpKind::DynamicTruncate => Self::DynamicTruncate { axis: 0 },
1141                    $crate::PrimitiveOpKind::PadToMatch => Self::PadToMatch { axis: 0 },
1142                    $crate::PrimitiveOpKind::ReduceProd => Self::ReduceProd { axes: vec![0] },
1143                    $crate::PrimitiveOpKind::ReduceMax => Self::ReduceMax { axes: vec![0] },
1144                    $crate::PrimitiveOpKind::ReduceMin => Self::ReduceMin { axes: vec![0] },
1145                }
1146            }
1147        }
1148    };
1149}