1#[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#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
38pub enum DTypePolicy {
39 SameAny,
40 SameNumeric,
41 SameFloat,
42 AbsToReal,
44 SameFloatOrComplex,
45 CompareToBool,
46 BoolSelect,
47 Convert,
48 Shape,
49 Constant,
50}
51
52#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
64pub struct PrimitiveOpDescriptor {
65 pub kind: PrimitiveOpKind,
67 pub name: &'static str,
69 pub category: OpCategory,
71 pub dtype_policy: DTypePolicy,
73 pub min_inputs: u8,
75 pub max_inputs: u8,
77 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 #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
152 pub enum PrimitiveOpKind {
153 $( $variant, )*
154 }
155
156 impl PrimitiveOpKind {
157 pub const COUNT: usize = [$(PrimitiveOpKind::$variant),*].len();
167
168 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 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
223pub 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 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 Div,
280 Rem,
281 Abs,
282 Sign,
283 Maximum,
284 Minimum,
285 Compare(CompareDir),
286 Select,
287 Clamp,
288
289 Exp,
291 Log,
292 Sin,
293 Cos,
294 Tanh,
295 Sqrt,
296 Rsqrt,
297 Pow,
298 Expm1,
299 Log1p,
300 Erf,
301
302 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 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 ReduceProd {
353 axes: Vec<usize>,
354 },
355 ReduceMax {
356 axes: Vec<usize>,
357 },
358 ReduceMin {
359 axes: Vec<usize>,
360 },
361
362 Extension(Arc<dyn ExtensionOp>),
368 }
369
370 impl StdTensorOp {
371 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 #[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 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}