Skip to main content

tenferro_runtime/program/
op.rs

1use std::sync::Arc;
2
3use computegraph::GraphOperation;
4use tenferro_ops::dim_expr::DimExpr;
5use tenferro_ops::ext_op::ExtensionOp;
6use tenferro_ops::std_tensor_op::StdTensorOp;
7use tenferro_tensor::{
8    CompareDir, DType, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig, SliceConfig,
9};
10
11use super::metadata::SemanticProvenance;
12use super::{
13    Alias, Effect, ProgramValue, SemanticPlacementConstraint, SemanticProvenanceView, ShapeGuard,
14};
15
16/// Closed backend-neutral vocabulary of core semantic tensor operations.
17#[derive(Clone, Debug, PartialEq)]
18#[non_exhaustive]
19pub enum CoreSemanticOp {
20    Add,
21    Sub,
22    Mul,
23    Neg,
24    Conj,
25    DotGeneral {
26        config: DotGeneralConfig,
27    },
28    Transpose {
29        perm: Vec<usize>,
30    },
31    Reshape {
32        to_shape: Vec<DimExpr>,
33    },
34    BroadcastInDim {
35        shape: Vec<DimExpr>,
36        dims: Vec<usize>,
37    },
38    Convert {
39        from: DType,
40        to: DType,
41    },
42    Constant {
43        dtype: DType,
44        bytes: Vec<u8>,
45    },
46    ReduceSum {
47        axes: Vec<usize>,
48    },
49    ReduceSumSquares {
50        axes: Vec<usize>,
51    },
52    Div,
53    Rem,
54    Abs,
55    Sign,
56    Maximum,
57    Minimum,
58    Compare(CompareDir),
59    Select,
60    Clamp,
61    Exp,
62    Log,
63    Sin,
64    Cos,
65    Tanh,
66    Sqrt,
67    Rsqrt,
68    Pow,
69    Expm1,
70    Log1p,
71    Erf,
72    ExtractDiag {
73        axis_a: usize,
74        axis_b: usize,
75    },
76    EmbedDiag {
77        axis_a: usize,
78        axis_b: usize,
79    },
80    Tril {
81        k: i64,
82    },
83    Triu {
84        k: i64,
85    },
86    Gather(GatherConfig),
87    GatherDynamicSliceSizes {
88        offset_dims: Vec<usize>,
89        collapsed_slice_dims: Vec<usize>,
90        start_index_map: Vec<usize>,
91        index_vector_dim: usize,
92        slice_sizes: Vec<DimExpr>,
93    },
94    Scatter(ScatterConfig),
95    Slice(SliceConfig),
96    DynamicSlice {
97        slice_sizes: Vec<usize>,
98    },
99    DynamicUpdateSlice,
100    Pad(PadConfig),
101    Concatenate {
102        axis: usize,
103        input_count: usize,
104    },
105    Reverse {
106        axes: Vec<usize>,
107    },
108    ShapeOf {
109        axis: usize,
110    },
111    DynamicTruncate {
112        axis: usize,
113    },
114    PadToMatch {
115        axis: usize,
116    },
117    ReduceProd {
118        axes: Vec<usize>,
119    },
120    ReduceMax {
121        axes: Vec<usize>,
122    },
123    ReduceMin {
124        axes: Vec<usize>,
125    },
126}
127
128/// Failure to convert a standard-operation carrier into the closed core
129/// semantic vocabulary.
130#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
131pub enum CoreSemanticOpConversionError {
132    /// Extension payloads must enter through
133    /// [`SemanticProgramBuilder::add_extension`](super::SemanticProgramBuilder::add_extension).
134    #[error("extension operations are not core semantic operations")]
135    ExtensionCarrier,
136}
137
138impl CoreSemanticOp {
139    pub(crate) fn input_count(&self) -> usize {
140        let standard = StdTensorOp::from(self);
141        GraphOperation::input_count(&standard)
142    }
143
144    pub(crate) fn output_count(&self) -> usize {
145        let standard = StdTensorOp::from(self);
146        GraphOperation::output_count(&standard)
147    }
148}
149
150impl TryFrom<&StdTensorOp> for CoreSemanticOp {
151    type Error = CoreSemanticOpConversionError;
152
153    fn try_from(op: &StdTensorOp) -> Result<Self, Self::Error> {
154        Ok(match op {
155            StdTensorOp::Add => Self::Add,
156            StdTensorOp::Sub => Self::Sub,
157            StdTensorOp::Mul => Self::Mul,
158            StdTensorOp::Neg => Self::Neg,
159            StdTensorOp::Conj => Self::Conj,
160            StdTensorOp::DotGeneral { config } => Self::DotGeneral {
161                config: config.clone(),
162            },
163            StdTensorOp::Transpose { perm } => Self::Transpose { perm: perm.clone() },
164            StdTensorOp::Reshape { to_shape } => Self::Reshape {
165                to_shape: to_shape.clone(),
166            },
167            StdTensorOp::BroadcastInDim { shape, dims } => Self::BroadcastInDim {
168                shape: shape.clone(),
169                dims: dims.clone(),
170            },
171            StdTensorOp::Convert { from, to } => Self::Convert {
172                from: *from,
173                to: *to,
174            },
175            StdTensorOp::Constant { dtype, bytes } => Self::Constant {
176                dtype: *dtype,
177                bytes: bytes.clone(),
178            },
179            StdTensorOp::ReduceSum { axes } => Self::ReduceSum { axes: axes.clone() },
180            StdTensorOp::ReduceSumSquares { axes } => Self::ReduceSumSquares { axes: axes.clone() },
181            StdTensorOp::Div => Self::Div,
182            StdTensorOp::Rem => Self::Rem,
183            StdTensorOp::Abs => Self::Abs,
184            StdTensorOp::Sign => Self::Sign,
185            StdTensorOp::Maximum => Self::Maximum,
186            StdTensorOp::Minimum => Self::Minimum,
187            StdTensorOp::Compare(direction) => Self::Compare(direction.clone()),
188            StdTensorOp::Select => Self::Select,
189            StdTensorOp::Clamp => Self::Clamp,
190            StdTensorOp::Exp => Self::Exp,
191            StdTensorOp::Log => Self::Log,
192            StdTensorOp::Sin => Self::Sin,
193            StdTensorOp::Cos => Self::Cos,
194            StdTensorOp::Tanh => Self::Tanh,
195            StdTensorOp::Sqrt => Self::Sqrt,
196            StdTensorOp::Rsqrt => Self::Rsqrt,
197            StdTensorOp::Pow => Self::Pow,
198            StdTensorOp::Expm1 => Self::Expm1,
199            StdTensorOp::Log1p => Self::Log1p,
200            StdTensorOp::Erf => Self::Erf,
201            StdTensorOp::ExtractDiag { axis_a, axis_b } => Self::ExtractDiag {
202                axis_a: *axis_a,
203                axis_b: *axis_b,
204            },
205            StdTensorOp::EmbedDiag { axis_a, axis_b } => Self::EmbedDiag {
206                axis_a: *axis_a,
207                axis_b: *axis_b,
208            },
209            StdTensorOp::Tril { k } => Self::Tril { k: *k },
210            StdTensorOp::Triu { k } => Self::Triu { k: *k },
211            StdTensorOp::Gather(config) => Self::Gather(config.clone()),
212            StdTensorOp::GatherDynamicSliceSizes {
213                offset_dims,
214                collapsed_slice_dims,
215                start_index_map,
216                index_vector_dim,
217                slice_sizes,
218            } => Self::GatherDynamicSliceSizes {
219                offset_dims: offset_dims.clone(),
220                collapsed_slice_dims: collapsed_slice_dims.clone(),
221                start_index_map: start_index_map.clone(),
222                index_vector_dim: *index_vector_dim,
223                slice_sizes: slice_sizes.clone(),
224            },
225            StdTensorOp::Scatter(config) => Self::Scatter(config.clone()),
226            StdTensorOp::Slice(config) => Self::Slice(config.clone()),
227            StdTensorOp::DynamicSlice { slice_sizes } => Self::DynamicSlice {
228                slice_sizes: slice_sizes.clone(),
229            },
230            StdTensorOp::DynamicUpdateSlice => Self::DynamicUpdateSlice,
231            StdTensorOp::Pad(config) => Self::Pad(config.clone()),
232            StdTensorOp::Concatenate { axis, input_count } => Self::Concatenate {
233                axis: *axis,
234                input_count: *input_count,
235            },
236            StdTensorOp::Reverse { axes } => Self::Reverse { axes: axes.clone() },
237            StdTensorOp::ShapeOf { axis } => Self::ShapeOf { axis: *axis },
238            StdTensorOp::DynamicTruncate { axis } => Self::DynamicTruncate { axis: *axis },
239            StdTensorOp::PadToMatch { axis } => Self::PadToMatch { axis: *axis },
240            StdTensorOp::ReduceProd { axes } => Self::ReduceProd { axes: axes.clone() },
241            StdTensorOp::ReduceMax { axes } => Self::ReduceMax { axes: axes.clone() },
242            StdTensorOp::ReduceMin { axes } => Self::ReduceMin { axes: axes.clone() },
243            StdTensorOp::Extension(_) => {
244                return Err(CoreSemanticOpConversionError::ExtensionCarrier);
245            }
246        })
247    }
248}
249
250impl From<&CoreSemanticOp> for StdTensorOp {
251    fn from(op: &CoreSemanticOp) -> Self {
252        match op {
253            CoreSemanticOp::Add => Self::Add,
254            CoreSemanticOp::Sub => Self::Sub,
255            CoreSemanticOp::Mul => Self::Mul,
256            CoreSemanticOp::Neg => Self::Neg,
257            CoreSemanticOp::Conj => Self::Conj,
258            CoreSemanticOp::DotGeneral { config } => Self::DotGeneral {
259                config: config.clone(),
260            },
261            CoreSemanticOp::Transpose { perm } => Self::Transpose { perm: perm.clone() },
262            CoreSemanticOp::Reshape { to_shape } => Self::Reshape {
263                to_shape: to_shape.clone(),
264            },
265            CoreSemanticOp::BroadcastInDim { shape, dims } => Self::BroadcastInDim {
266                shape: shape.clone(),
267                dims: dims.clone(),
268            },
269            CoreSemanticOp::Convert { from, to } => Self::Convert {
270                from: *from,
271                to: *to,
272            },
273            CoreSemanticOp::Constant { dtype, bytes } => Self::Constant {
274                dtype: *dtype,
275                bytes: bytes.clone(),
276            },
277            CoreSemanticOp::ReduceSum { axes } => Self::ReduceSum { axes: axes.clone() },
278            CoreSemanticOp::ReduceSumSquares { axes } => {
279                Self::ReduceSumSquares { axes: axes.clone() }
280            }
281            CoreSemanticOp::Div => Self::Div,
282            CoreSemanticOp::Rem => Self::Rem,
283            CoreSemanticOp::Abs => Self::Abs,
284            CoreSemanticOp::Sign => Self::Sign,
285            CoreSemanticOp::Maximum => Self::Maximum,
286            CoreSemanticOp::Minimum => Self::Minimum,
287            CoreSemanticOp::Compare(direction) => Self::Compare(direction.clone()),
288            CoreSemanticOp::Select => Self::Select,
289            CoreSemanticOp::Clamp => Self::Clamp,
290            CoreSemanticOp::Exp => Self::Exp,
291            CoreSemanticOp::Log => Self::Log,
292            CoreSemanticOp::Sin => Self::Sin,
293            CoreSemanticOp::Cos => Self::Cos,
294            CoreSemanticOp::Tanh => Self::Tanh,
295            CoreSemanticOp::Sqrt => Self::Sqrt,
296            CoreSemanticOp::Rsqrt => Self::Rsqrt,
297            CoreSemanticOp::Pow => Self::Pow,
298            CoreSemanticOp::Expm1 => Self::Expm1,
299            CoreSemanticOp::Log1p => Self::Log1p,
300            CoreSemanticOp::Erf => Self::Erf,
301            CoreSemanticOp::ExtractDiag { axis_a, axis_b } => Self::ExtractDiag {
302                axis_a: *axis_a,
303                axis_b: *axis_b,
304            },
305            CoreSemanticOp::EmbedDiag { axis_a, axis_b } => Self::EmbedDiag {
306                axis_a: *axis_a,
307                axis_b: *axis_b,
308            },
309            CoreSemanticOp::Tril { k } => Self::Tril { k: *k },
310            CoreSemanticOp::Triu { k } => Self::Triu { k: *k },
311            CoreSemanticOp::Gather(config) => Self::Gather(config.clone()),
312            CoreSemanticOp::GatherDynamicSliceSizes {
313                offset_dims,
314                collapsed_slice_dims,
315                start_index_map,
316                index_vector_dim,
317                slice_sizes,
318            } => Self::GatherDynamicSliceSizes {
319                offset_dims: offset_dims.clone(),
320                collapsed_slice_dims: collapsed_slice_dims.clone(),
321                start_index_map: start_index_map.clone(),
322                index_vector_dim: *index_vector_dim,
323                slice_sizes: slice_sizes.clone(),
324            },
325            CoreSemanticOp::Scatter(config) => Self::Scatter(config.clone()),
326            CoreSemanticOp::Slice(config) => Self::Slice(config.clone()),
327            CoreSemanticOp::DynamicSlice { slice_sizes } => Self::DynamicSlice {
328                slice_sizes: slice_sizes.clone(),
329            },
330            CoreSemanticOp::DynamicUpdateSlice => Self::DynamicUpdateSlice,
331            CoreSemanticOp::Pad(config) => Self::Pad(config.clone()),
332            CoreSemanticOp::Concatenate { axis, input_count } => Self::Concatenate {
333                axis: *axis,
334                input_count: *input_count,
335            },
336            CoreSemanticOp::Reverse { axes } => Self::Reverse { axes: axes.clone() },
337            CoreSemanticOp::ShapeOf { axis } => Self::ShapeOf { axis: *axis },
338            CoreSemanticOp::DynamicTruncate { axis } => Self::DynamicTruncate { axis: *axis },
339            CoreSemanticOp::PadToMatch { axis } => Self::PadToMatch { axis: *axis },
340            CoreSemanticOp::ReduceProd { axes } => Self::ReduceProd { axes: axes.clone() },
341            CoreSemanticOp::ReduceMax { axes } => Self::ReduceMax { axes: axes.clone() },
342            CoreSemanticOp::ReduceMin { axes } => Self::ReduceMin { axes: axes.clone() },
343        }
344    }
345}
346
347pub(crate) enum SemanticOp {
348    Core(CoreSemanticOp),
349    Extension(Arc<dyn ExtensionOp>),
350}
351
352pub(crate) struct SemanticOperation {
353    pub(crate) op: SemanticOp,
354    pub(crate) inputs: Box<[ProgramValue]>,
355    pub(crate) outputs: Box<[ProgramValue]>,
356    pub(crate) effects: Box<[Effect]>,
357    pub(crate) aliases: Box<[Alias]>,
358    pub(crate) shape_guards: Box<[ShapeGuard]>,
359    pub(crate) placement: SemanticPlacementConstraint,
360    pub(crate) provenance: SemanticProvenance,
361}
362
363/// Borrowed semantic operation payload.
364#[derive(Clone, Copy)]
365#[non_exhaustive]
366pub enum SemanticOpRef<'a> {
367    /// Closed core operation.
368    Core(&'a CoreSemanticOp),
369    /// Extension semantic payload.
370    Extension(&'a dyn ExtensionOp),
371}
372
373impl std::fmt::Debug for SemanticOpRef<'_> {
374    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
375        match self {
376            Self::Core(_) => formatter.write_str("SemanticOpRef::Core(<bounded>)"),
377            Self::Extension(op) => formatter
378                .debug_tuple("SemanticOpRef::Extension")
379                .field(&op.family_id())
380                .finish(),
381        }
382    }
383}
384
385/// Allocation-free immutable view of one semantic operation.
386#[derive(Clone, Copy)]
387pub struct SemanticOperationView<'a> {
388    operation: &'a SemanticOperation,
389}
390
391impl<'a> SemanticOperationView<'a> {
392    pub(crate) const fn new(operation: &'a SemanticOperation) -> Self {
393        Self { operation }
394    }
395
396    /// Borrow the semantic operation payload.
397    pub fn op(self) -> SemanticOpRef<'a> {
398        match &self.operation.op {
399            SemanticOp::Core(op) => SemanticOpRef::Core(op),
400            SemanticOp::Extension(op) => SemanticOpRef::Extension(op.as_ref()),
401        }
402    }
403
404    /// Borrow ordered SSA inputs.
405    pub fn inputs(self) -> &'a [ProgramValue] {
406        &self.operation.inputs
407    }
408
409    /// Borrow ordered SSA outputs.
410    pub fn outputs(self) -> &'a [ProgramValue] {
411        &self.operation.outputs
412    }
413
414    /// Borrow ordered observable effects.
415    pub fn effects(self) -> &'a [Effect] {
416        &self.operation.effects
417    }
418
419    /// Borrow output alias declarations.
420    pub fn aliases(self) -> &'a [Alias] {
421        &self.operation.aliases
422    }
423
424    /// Borrow operation-local symbolic guards.
425    pub fn shape_guards(self) -> &'a [ShapeGuard] {
426        &self.operation.shape_guards
427    }
428
429    /// Return bounded diagnostic provenance without source identities.
430    pub fn provenance(self) -> SemanticProvenanceView<'a> {
431        self.operation.provenance.view()
432    }
433
434    /// Return the unresolved placement constraint.
435    pub fn placement(self) -> SemanticPlacementConstraint {
436        self.operation.placement
437    }
438}
439
440impl std::fmt::Debug for SemanticOperationView<'_> {
441    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
442        formatter
443            .debug_struct("SemanticOperationView")
444            .field("op", &self.op())
445            .field("inputs", &self.inputs().len())
446            .field("outputs", &self.outputs().len())
447            .field("effects", &self.effects().len())
448            .field("aliases", &self.aliases().len())
449            .field("shape_guards", &self.shape_guards().len())
450            .finish()
451    }
452}