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#[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#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
131pub enum CoreSemanticOpConversionError {
132 #[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#[derive(Clone, Copy)]
365#[non_exhaustive]
366pub enum SemanticOpRef<'a> {
367 Core(&'a CoreSemanticOp),
369 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#[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 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 pub fn inputs(self) -> &'a [ProgramValue] {
406 &self.operation.inputs
407 }
408
409 pub fn outputs(self) -> &'a [ProgramValue] {
411 &self.operation.outputs
412 }
413
414 pub fn effects(self) -> &'a [Effect] {
416 &self.operation.effects
417 }
418
419 pub fn aliases(self) -> &'a [Alias] {
421 &self.operation.aliases
422 }
423
424 pub fn shape_guards(self) -> &'a [ShapeGuard] {
426 &self.operation.shape_guards
427 }
428
429 pub fn provenance(self) -> SemanticProvenanceView<'a> {
431 self.operation.provenance.view()
432 }
433
434 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}