1use std::sync::Arc;
20
21use computegraph::graph::{Graph, GraphBuilder};
22use computegraph::types::{OperationRole, ValueRef};
23use computegraph::GraphOperation;
24use tenferro_ops::dim_expr::DimExpr;
25use tenferro_ops::std_tensor_op::StdTensorOp;
26use tenferro_ops::TensorMeta;
27
28use crate::checkpoint::CheckpointNode;
29use crate::error::{Error, ErrorPhase, Result};
30use crate::metadata::{
31 register_scoped_graph_analysis, registered_meta, MetadataScopeChain, RegisteredGraphAnalysis,
32};
33use crate::shape_constraint::{ConstraintScopeChain, ScopedShapeConstraint, ShapeConstraintScope};
34use crate::shape_infer::{infer_extension_output_meta_with_constraints, InferredExtensionMeta};
35use crate::traced::{
36 merge_traced_inputs_map, merge_traced_leaf_metas, next_traced_id, TracedTensor,
37};
38
39type ExpandedOutputMetas = Vec<(tenferro_tensor::DType, Vec<SymDim>)>;
40
41pub use crate::compiler::CompilerOptions;
42#[doc(hidden)]
43pub use crate::shape_infer::{
44 infer_output_dtype, infer_output_extents, infer_output_shapes, promote_dtype,
45 promote_dtype_div_like, promote_dtype_for_binary_op, promote_dtypes,
46};
47pub use tenferro_extension_macros::define_extension_runtime;
48pub use tenferro_ops::ext_op::{
49 ExtensionAlias, ExtensionAliasDeclaration, ExtensionEffect, ExtensionEffectAccess,
50 ExtensionEffectDeclaration, ExtensionOp,
51};
52pub use tenferro_ops::{ExtensionFamilyId, ExtensionShapeContext, SymDim};
53
54pub use crate::extension_cache::{
55 ExtensionCacheKey, ExtensionCacheLimits, ExtensionCacheSelector, ExtensionCacheStore,
56};
57pub use crate::extension_execution_context::ExtensionExecutionContext;
58pub use crate::runtime::{ExtensionModule, ExtensionModuleId, ExtensionModuleRegistrar};
59
60pub fn apply(op: Arc<dyn ExtensionOp>, inputs: &[&TracedTensor]) -> Result<Vec<TracedTensor>> {
117 if inputs.len() != op.input_count() {
118 return Err(Error::invalid_argument(
119 "extension::apply",
120 ErrorPhase::GraphBuild,
121 "inputs",
122 format!(
123 "op family {:?} expects {} inputs, got {}",
124 op.family_id(),
125 op.input_count(),
126 inputs.len()
127 ),
128 ));
129 }
130
131 let append = append_raw_op(StdTensorOp::Extension(op.clone()), inputs)?;
132 let analysis = analyze_extension_graph(append.graph.as_ref())?;
133 let output_metas = append
134 .output_ids
135 .iter()
136 .map(|&output| {
137 let meta = registered_meta(&append.graph.values()[output].key)?;
138 let shape = meta.bound_shape().ok_or_else(|| {
139 Error::invalid_argument(
140 "extension::apply",
141 ErrorPhase::Compile,
142 "output_metadata",
143 format!(
144 "extension family {:?} produced unknown output shape metadata",
145 op.family_id()
146 ),
147 )
148 })?;
149 Ok((meta.dtype, shape))
150 })
151 .collect::<Result<Vec<_>>>()?;
152 traced_outputs_from_analysis(
153 inputs,
154 append.graph,
155 &append.output_ids,
156 output_metas,
157 analysis,
158 )
159}
160
161#[doc(hidden)]
163pub struct RawAppend {
164 pub graph: Arc<Graph<StdTensorOp>>,
165 pub output_ids: Vec<usize>,
166}
167
168#[doc(hidden)]
177pub fn append_raw_op(op: StdTensorOp, inputs: &[&TracedTensor]) -> Result<RawAppend> {
178 let expected = op.input_count();
179 if inputs.len() != expected {
180 return Err(Error::invalid_argument(
181 "extension::append_raw_op",
182 ErrorPhase::GraphBuild,
183 "inputs",
184 format!("op expects {expected} inputs, got {}", inputs.len()),
185 ));
186 }
187 let mut builder = GraphBuilder::<StdTensorOp>::new();
188 for input in inputs {
189 builder.add_parent(input.graph.clone());
190 }
191 let op_inputs: Vec<ValueRef<StdTensorOp>> = inputs
192 .iter()
193 .map(|t| ValueRef::External(t.graph.values()[t.val].key.clone()))
194 .collect();
195 let output_ids = builder.add_operation(op, op_inputs, OperationRole::Primary);
196 builder.set_outputs(output_ids.clone());
197 Ok(RawAppend {
198 graph: Arc::new(builder.build()),
199 output_ids,
200 })
201}
202
203#[doc(hidden)]
205pub fn append_raw_eager_outputs(
206 op: StdTensorOp,
207 inputs: &[&TracedTensor],
208 output_metadata: &[TensorMeta],
209) -> Result<Vec<TracedTensor>> {
210 let append = append_raw_op(op.clone(), inputs)?;
211 if append.output_ids.len() != output_metadata.len() {
212 return Err(Error::Internal(format!(
213 "semantic eager recording expected {} outputs for {op:?}, got {}",
214 output_metadata.len(),
215 append.output_ids.len()
216 )));
217 }
218
219 let inputs_map = merge_traced_inputs_map(inputs.iter().copied());
220 let leaf_metas = merge_traced_leaf_metas(inputs.iter().copied());
221 let mut extra_roots = Vec::new();
222 for input in inputs {
223 extra_roots.extend(input.extra_roots.iter().cloned());
224 }
225 let metadata_scopes =
226 MetadataScopeChain::merge(inputs.iter().map(|input| &input.metadata_scopes));
227
228 Ok(append
229 .output_ids
230 .into_iter()
231 .zip(output_metadata)
232 .map(|(val, meta)| TracedTensor {
233 id: next_traced_id(),
234 rank: meta.rank(),
235 dtype: meta.dtype,
236 graph: Arc::clone(&append.graph),
237 val,
238 data: None,
239 shape_hint: None,
240 inputs_map: Arc::clone(&inputs_map),
241 leaf_metas: Arc::clone(&leaf_metas),
242 extra_roots: extra_roots.clone(),
243 checkpoint_chain: None,
244 metadata_scopes: metadata_scopes.clone(),
245 constraint_scopes: ConstraintScopeChain::empty(),
246 })
247 .collect())
248}
249
250pub(crate) fn analyze_extension_graph(
254 graph: &Graph<StdTensorOp>,
255) -> Result<RegisteredGraphAnalysis> {
256 register_scoped_graph_analysis(graph, std::iter::empty())
257}
258
259#[doc(hidden)]
284pub fn apply_standard_op(op: StdTensorOp, inputs: &[&TracedTensor]) -> Result<Vec<TracedTensor>> {
285 if matches!(op, StdTensorOp::Extension(_)) {
286 return Err(Error::invalid_argument(
287 "extension::apply_standard_op",
288 ErrorPhase::GraphBuild,
289 "op",
290 "Extension ops must be passed to extension::apply",
291 ));
292 }
293 let expected = op.input_count();
294 if inputs.len() != expected {
295 return Err(Error::invalid_argument(
296 "extension::apply_standard_op",
297 ErrorPhase::GraphBuild,
298 "inputs",
299 format!("op expects {expected} inputs, got {}", inputs.len()),
300 ));
301 }
302
303 let append = append_raw_op(op, inputs)?;
304 let analysis = analyze_extension_graph(append.graph.as_ref())?;
305 let output_metas = append
306 .output_ids
307 .iter()
308 .map(|&output| {
309 let meta = registered_meta(&append.graph.values()[output].key)?;
310 let shape = meta.bound_shape().ok_or_else(|| {
311 Error::invalid_argument(
312 "extension::apply_standard_op",
313 ErrorPhase::Compile,
314 "output_metadata",
315 "standard op produced unknown output shape metadata",
316 )
317 })?;
318 Ok((meta.dtype, shape))
319 })
320 .collect::<Result<Vec<_>>>()?;
321 traced_outputs_from_analysis(
322 inputs,
323 append.graph,
324 &append.output_ids,
325 output_metas,
326 analysis,
327 )
328}
329
330#[doc(hidden)]
374pub fn attach_expanded_shape_contract(
375 op: &dyn ExtensionOp,
376 inputs: &[&TracedTensor],
377 output: TracedTensor,
378) -> Result<TracedTensor> {
379 if op.output_count() != 1 {
380 return Err(Error::invalid_argument(
381 "extension::attach_expanded_shape_contract",
382 ErrorPhase::GraphBuild,
383 "outputs",
384 format!(
385 "extension family {:?} contract expects {} outputs, got one expanded output",
386 op.family_id(),
387 op.output_count(),
388 ),
389 ));
390 }
391 let (_, inferred) = infer_expanded_shape_contract(op, inputs)?;
392 attach_inferred_expanded_shape_contract(inputs, vec![output], inferred)?
393 .into_iter()
394 .next()
395 .ok_or_else(|| Error::Internal("expanded shape contract returned no output".into()))
396}
397
398fn infer_expanded_shape_contract(
399 op: &dyn ExtensionOp,
400 inputs: &[&TracedTensor],
401) -> Result<(ExpandedOutputMetas, InferredExtensionMeta)> {
402 if inputs.len() != op.input_count() {
403 return Err(Error::invalid_argument(
404 "extension::infer_expanded_shape_contract",
405 ErrorPhase::GraphBuild,
406 "inputs",
407 format!(
408 "extension family {:?} contract expects {} inputs, got {}",
409 op.family_id(),
410 op.input_count(),
411 inputs.len()
412 ),
413 ));
414 }
415 let input_dtypes: Vec<_> = inputs.iter().map(|input| input.dtype).collect();
416 let input_shapes: Vec<_> = inputs
417 .iter()
418 .enumerate()
419 .map(|(input_idx, input)| DimExpr::input_shape(input_idx, input.rank))
420 .collect();
421 let input_shape_refs: Vec<_> = input_shapes.iter().map(Vec::as_slice).collect();
422 let inferred =
423 infer_extension_output_meta_with_constraints(op, &input_dtypes, &input_shape_refs)?;
424 let input_sym_shapes = inputs
425 .iter()
426 .map(|input| {
427 (0..input.rank)
428 .map(|axis| input.axis_sym_dim(axis))
429 .collect::<Result<Vec<_>>>()
430 })
431 .collect::<Result<Vec<_>>>()?;
432 let input_sym_shape_refs = input_sym_shapes
433 .iter()
434 .map(Vec::as_slice)
435 .collect::<Vec<_>>();
436 let output_metas = inferred
437 .output_metas
438 .iter()
439 .map(|(dtype, shape)| {
440 (
441 *dtype,
442 shape
443 .iter()
444 .map(|dim| SymDim::from_dim_expr(dim, &input_sym_shape_refs))
445 .collect(),
446 )
447 })
448 .collect();
449 Ok((output_metas, inferred))
450}
451
452fn attach_inferred_expanded_shape_contract(
453 inputs: &[&TracedTensor],
454 mut outputs: Vec<TracedTensor>,
455 inferred: InferredExtensionMeta,
456) -> Result<Vec<TracedTensor>> {
457 if inferred.output_metas.len() != outputs.len() {
458 return Err(Error::invalid_argument(
459 "extension::attach_expanded_shape_contract",
460 ErrorPhase::GraphBuild,
461 "outputs",
462 format!(
463 "extension contract inferred {} outputs, but expanded graph produced {}",
464 inferred.output_metas.len(),
465 outputs.len()
466 ),
467 ));
468 }
469 for (output, (dtype, local_shape)) in outputs.iter().zip(inferred.output_metas.iter()) {
470 if output.dtype != *dtype || output.rank != local_shape.len() {
471 return Err(Error::invalid_argument(
472 "extension::attach_expanded_shape_contract",
473 ErrorPhase::GraphBuild,
474 "outputs",
475 format!(
476 "extension contract inferred output {:?} rank {}, but expanded output is {:?} rank {}",
477 dtype,
478 local_shape.len(),
479 output.dtype,
480 output.rank
481 ),
482 ));
483 }
484 }
485 if inferred.constraints.is_empty() {
486 return Ok(outputs);
487 }
488
489 let origins = outputs
490 .iter()
491 .map(|output| output.graph.values()[output.val].key.clone())
492 .collect::<Vec<_>>();
493 let input_keys = inputs
494 .iter()
495 .map(|input| input.graph.values()[input.val].key.clone())
496 .collect::<Vec<_>>();
497 let constraints = inferred
498 .constraints
499 .into_iter()
500 .map(|local| ScopedShapeConstraint {
501 origins: origins.clone(),
502 inputs: input_keys.clone(),
503 local,
504 })
505 .collect();
506 let scope = Arc::new(ShapeConstraintScope::new(constraints));
507 for output in &mut outputs {
508 output.constraint_scopes =
509 ConstraintScopeChain::with_scope(Arc::clone(&scope), [&output.constraint_scopes]);
510 }
511 Ok(outputs)
512}
513
514#[doc(hidden)]
563pub fn apply_expanded_graph_with_shape_contract(
564 op: &dyn ExtensionOp,
565 inputs: &[&TracedTensor],
566 build: impl FnOnce(&mut GraphBuilder<StdTensorOp>, &[ValueRef<StdTensorOp>]) -> Result<Vec<usize>>,
567) -> Result<Vec<TracedTensor>> {
568 let (output_metas, inferred) = infer_expanded_shape_contract(op, inputs)?;
569 let outputs = apply_expanded_graph(inputs, output_metas, build)?;
570 attach_inferred_expanded_shape_contract(inputs, outputs, inferred)
571}
572
573pub fn apply_expanded_graph(
586 inputs: &[&TracedTensor],
587 output_metas: Vec<(tenferro_tensor::DType, Vec<SymDim>)>,
588 build: impl FnOnce(&mut GraphBuilder<StdTensorOp>, &[ValueRef<StdTensorOp>]) -> Result<Vec<usize>>,
589) -> Result<Vec<TracedTensor>> {
590 let mut builder = GraphBuilder::<StdTensorOp>::new();
591 for input in inputs {
592 builder.add_parent(input.graph.clone());
593 }
594 let op_inputs: Vec<ValueRef<StdTensorOp>> = inputs
595 .iter()
596 .map(|t| ValueRef::External(t.graph.values()[t.val].key.clone()))
597 .collect();
598 let outputs = build(&mut builder, &op_inputs)?;
599 if outputs.len() != output_metas.len() {
600 return Err(Error::invalid_argument(
601 "extension::apply_expanded_graph",
602 ErrorPhase::GraphBuild,
603 "outputs",
604 format!(
605 "extension expanded graph returned {} outputs for {} output metadata entries",
606 outputs.len(),
607 output_metas.len()
608 ),
609 ));
610 }
611 builder.set_outputs(outputs.clone());
612 let graph = Arc::new(builder.build());
613 let analysis = register_scoped_graph_analysis(graph.as_ref(), std::iter::empty())?;
614 traced_outputs_from_analysis(inputs, graph, &outputs, output_metas, analysis)
615}
616
617fn traced_outputs_from_analysis(
618 inputs: &[&TracedTensor],
619 graph: Arc<computegraph::graph::Graph<StdTensorOp>>,
620 outputs: &[usize],
621 output_metas: Vec<(tenferro_tensor::DType, Vec<SymDim>)>,
622 analysis: RegisteredGraphAnalysis,
623) -> Result<Vec<TracedTensor>> {
624 let metadata_scope = Arc::new(analysis.metadata);
625 let constraint_scope = Arc::new(analysis.constraints);
626
627 let merged_map = merge_traced_inputs_map(inputs.iter().copied());
628 let merged_leaf_metas = merge_traced_leaf_metas(inputs.iter().copied());
629 let mut extra_roots = Vec::new();
630 let mut checkpoint_chain = None;
631 let metadata_scopes = MetadataScopeChain::with_scope(
632 Arc::clone(&metadata_scope),
633 inputs.iter().map(|input| &input.metadata_scopes),
634 );
635 let constraint_scopes = if constraint_scope.is_empty() {
636 ConstraintScopeChain::merge(inputs.iter().map(|input| &input.constraint_scopes))
637 } else {
638 ConstraintScopeChain::with_scope(
639 constraint_scope,
640 inputs.iter().map(|input| &input.constraint_scopes),
641 )
642 };
643 for input in inputs {
644 extra_roots.extend(input.extra_roots.iter().cloned());
645 checkpoint_chain =
646 CheckpointNode::merge_chains(checkpoint_chain, input.checkpoint_chain.clone());
647 }
648 let all_inputs_concrete = inputs.iter().all(|t| t.shape_hint.is_some());
649 Ok(outputs
650 .iter()
651 .zip(output_metas)
652 .map(|(&val, (dtype, shape))| {
653 let shape_hint = if all_inputs_concrete {
654 Some(shape.clone())
655 } else {
656 None
657 };
658 TracedTensor {
659 id: next_traced_id(),
660 rank: shape.len(),
661 dtype,
662 graph: graph.clone(),
663 val,
664 data: None,
665 shape_hint,
666 inputs_map: merged_map.clone(),
667 leaf_metas: merged_leaf_metas.clone(),
668 extra_roots: extra_roots.clone(),
669 checkpoint_chain: checkpoint_chain.clone(),
670 metadata_scopes: metadata_scopes.clone(),
671 constraint_scopes: constraint_scopes.clone(),
672 }
673 })
674 .collect())
675}
676
677#[cfg(test)]
678mod tests;