Skip to main content

tenferro_runtime/
extension.rs

1//! Public surface for out-of-tree extension primitives.
2//!
3//! This module exposes the Stage 6 `ExtensionOp` mechanism through the
4//! runtime crate. External crates implement
5//! [`ExtensionOp`] and build traced graphs containing the extension via
6//! [`apply`].
7//!
8//! See `docs/spec/extension-op.md` for the normative contract.
9//!
10//! # Examples
11//!
12//! ```rust
13//! use tenferro_runtime::extension::{apply, ExtensionOp};
14//!
15//! // Construct an `Arc<dyn ExtensionOp>` and call `apply(op, &[input])`
16//! // to lower it into a `TracedTensor`.
17//! ```
18
19use 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
60/// Apply an extension op in the traced graph.
61///
62/// The `op` value is cloned into a `StdTensorOp::Extension(Arc<dyn ExtensionOp>)`
63/// carrier. The returned vector contains one [`TracedTensor`] per declared
64/// output slot of the extension. Output shapes are inferred via
65/// [`ExtensionOp::infer_output_meta`] using the input shape hints.
66///
67/// `inputs.len()` must equal `op.input_count()`, and each input's
68/// `shape_hint` must be present (i.e. the extension must be used on
69/// tensors whose rank is known at graph-build time). For symbolic-shape
70/// composition, pass concrete tensors to [`crate::Runtime::run_compiled`] at
71/// evaluation time.
72///
73/// # Examples
74///
75/// ```rust
76/// # use std::any::Any;
77/// use std::sync::Arc;
78/// use tenferro_runtime::extension::{apply, ExtensionOp, ExtensionShapeContext};
79/// use tenferro_runtime::{DType, SymDim, TracedTensor};
80///
81/// # #[derive(Clone, Debug)]
82/// # struct IdentityExt;
83/// # impl ExtensionOp for IdentityExt {
84/// #     fn family_id(&self) -> &'static str { "example.identity.v1" }
85/// #     fn payload_hash(&self, _hasher: &mut dyn std::hash::Hasher) {}
86/// #     fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
87/// #         other.as_any().downcast_ref::<IdentityExt>().is_some()
88/// #     }
89/// #     fn clone_arc(&self) -> Arc<dyn ExtensionOp> { Arc::new(self.clone()) }
90/// #     fn as_any(&self) -> &dyn Any { self }
91/// #     fn input_count(&self) -> usize { 1 }
92/// #     fn output_count(&self) -> usize { 1 }
93/// #     fn infer_output_meta(
94/// #         &self,
95/// #         ctx: &mut ExtensionShapeContext<'_>,
96/// #     ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
97/// #         Ok(vec![(ctx.input_dtype(0)?, ctx.input_shape(0)?.to_vec())])
98/// #     }
99/// # }
100/// let op: Arc<dyn ExtensionOp> = Arc::new(IdentityExt);
101/// let a = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
102/// let outputs = apply(op, &[&a])?;
103/// assert_eq!(outputs.len(), 1);
104/// # Ok::<(), tenferro_runtime::Error>(())
105/// ```
106///
107/// # Errors
108///
109/// Returns [`Error::Validation`] with `ValidationError::InvalidArgument` when
110/// the extension receives the wrong number of traced inputs or produces an
111/// unknown output shape. Canonical metadata inference failures, including a
112/// returned metadata count that differs from [`ExtensionOp::output_count`],
113/// are returned as [`Error::TensorRuntime`] containing the typed tensor
114/// validation source, while poisoned metadata state is retained as
115/// [`Error::RuntimeStateSource`].
116pub 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/// Raw result of appending one op to a traced/eager graph without analysis.
162#[doc(hidden)]
163pub struct RawAppend {
164    pub graph: Arc<Graph<StdTensorOp>>,
165    pub output_ids: Vec<usize>,
166}
167
168/// Append one op to a traced graph without running metadata analysis.
169///
170/// This is the O(inputs)/op half of [`apply`]: it builds only the raw
171/// `Graph<StdTensorOp>` carrier (parent edges + op + declared outputs).
172/// Analysis (metadata registration, `infer_output_meta`, constraint scopes)
173/// is deferred via [`analyze_extension_graph`]. The traced path runs both
174/// immediately; the eager-AD path appends now and analyzes at the first AD
175/// request.
176#[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/// Append an eager semantic op and retain its private runtime carrier state.
204#[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
250/// Run the deferred analysis half of an append once: register metadata and
251/// derive constraint scopes for every live value. Idempotent per graph value
252/// key, so repeated calls (one per first AD request) are safe.
253pub(crate) fn analyze_extension_graph(
254    graph: &Graph<StdTensorOp>,
255) -> Result<RegisteredGraphAnalysis> {
256    register_scoped_graph_analysis(graph, std::iter::empty())
257}
258
259/// Apply a core standard op in the traced graph.
260///
261/// This is an internal crate-boundary helper used by eager AD recording to keep
262/// a semantic traced graph beside the existing eager trace. Extension ops
263/// must use [`apply`] instead.
264///
265/// # Examples
266///
267/// ```rust
268/// use tenferro_ops::std_tensor_op::StdTensorOp;
269/// use tenferro_runtime::{extension, TracedTensor};
270///
271/// let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
272/// let outputs = extension::apply_standard_op(StdTensorOp::Neg, &[&x])?;
273/// assert_eq!(outputs.len(), 1);
274/// # Ok::<(), tenferro_runtime::Error>(())
275/// ```
276///
277/// # Errors
278///
279/// Returns [`Error::Validation`] with `InvalidArgument` when the op is an
280/// extension op, receives the wrong number of traced inputs, or produces
281/// metadata without a known output bound. Metadata-analysis and registry
282/// failures are returned as typed runtime errors with their source preserved.
283#[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/// Attach an extension's inferred shape contract to an equivalent expanded output.
331///
332/// Standard extension crates use this when a traced fast path lowers an extension
333/// directly to core operations. The extension remains the single source of truth
334/// for metadata equalities while the executable graph keeps the core-operation
335/// fast path.
336///
337/// # Examples
338///
339/// ```rust
340/// # use std::{any::Any, sync::Arc};
341/// use tenferro_runtime::extension::{
342///     attach_expanded_shape_contract, ExtensionOp, ExtensionShapeContext,
343/// };
344/// use tenferro_runtime::{DType, SymDim, TracedTensor};
345///
346/// # #[derive(Clone, Debug)]
347/// # struct SameShapeAdd;
348/// # impl ExtensionOp for SameShapeAdd {
349/// #     fn family_id(&self) -> &'static str { "example.same-shape-add.v1" }
350/// #     fn payload_hash(&self, _hasher: &mut dyn std::hash::Hasher) {}
351/// #     fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
352/// #         other.as_any().downcast_ref::<Self>().is_some()
353/// #     }
354/// #     fn clone_arc(&self) -> Arc<dyn ExtensionOp> { Arc::new(self.clone()) }
355/// #     fn as_any(&self) -> &dyn Any { self }
356/// #     fn input_count(&self) -> usize { 2 }
357/// #     fn output_count(&self) -> usize { 1 }
358/// #     fn infer_output_meta(
359/// #         &self,
360/// #         ctx: &mut ExtensionShapeContext<'_>,
361/// #     ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
362/// #         ctx.require_same_shape(0, 1)?;
363/// #         Ok(vec![(ctx.input_dtype(0)?, ctx.input_shape(0)?.to_vec())])
364/// #     }
365/// # }
366/// let lhs = TracedTensor::input_symbolic_shape(DType::F64, 1)?;
367/// let rhs = TracedTensor::input_symbolic_shape(DType::F64, 1)?;
368/// let expanded = (&lhs + &rhs)?;
369/// let output = attach_expanded_shape_contract(&SameShapeAdd, &[&lhs, &rhs], expanded)?;
370/// assert_eq!(output.rank, 1);
371/// # Ok::<(), tenferro_runtime::Error>(())
372/// ```
373#[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/// Apply an expanded core graph while retaining one extension metadata contract.
515///
516/// Metadata inference runs exactly once. Its output metadata builds the traced
517/// outputs and its equality constraints are attached to those same outputs.
518///
519/// # Examples
520///
521/// ```rust
522/// # use std::{any::Any, sync::Arc};
523/// use computegraph::types::OperationRole;
524/// use tenferro_ops::std_tensor_op::StdTensorOp;
525/// use tenferro_runtime::extension::{
526///     apply_expanded_graph_with_shape_contract, ExtensionOp, ExtensionShapeContext,
527/// };
528/// use tenferro_runtime::{DType, SymDim, TracedTensor};
529///
530/// # #[derive(Clone, Debug)]
531/// # struct SameShapeAdd;
532/// # impl ExtensionOp for SameShapeAdd {
533/// #     fn family_id(&self) -> &'static str { "example.expanded-add.v1" }
534/// #     fn payload_hash(&self, _hasher: &mut dyn std::hash::Hasher) {}
535/// #     fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
536/// #         other.as_any().downcast_ref::<Self>().is_some()
537/// #     }
538/// #     fn clone_arc(&self) -> Arc<dyn ExtensionOp> { Arc::new(self.clone()) }
539/// #     fn as_any(&self) -> &dyn Any { self }
540/// #     fn input_count(&self) -> usize { 2 }
541/// #     fn output_count(&self) -> usize { 1 }
542/// #     fn infer_output_meta(
543/// #         &self,
544/// #         ctx: &mut ExtensionShapeContext<'_>,
545/// #     ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
546/// #         ctx.require_same_shape(0, 1)?;
547/// #         Ok(vec![(ctx.input_dtype(0)?, ctx.input_shape(0)?.to_vec())])
548/// #     }
549/// # }
550/// let lhs = TracedTensor::input_symbolic_shape(DType::F64, 1)?;
551/// let rhs = TracedTensor::input_symbolic_shape(DType::F64, 1)?;
552/// let outputs = apply_expanded_graph_with_shape_contract(
553///     &SameShapeAdd,
554///     &[&lhs, &rhs],
555///     |builder, inputs| {
556///         Ok(builder.add_operation(StdTensorOp::Add, inputs.to_vec(), OperationRole::Primary))
557///     },
558/// )?;
559/// assert_eq!(outputs[0].rank, 1);
560/// # Ok::<(), tenferro_runtime::Error>(())
561/// ```
562#[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
573/// Apply an extension-provided lowering as ordinary traced graph operations.
574///
575/// This is for extension crates whose operation can be expanded at graph-build
576/// time. It preserves the same parent graph and metadata merging behavior as
577/// [`apply`], but does not insert a `StdTensorOp::Extension` carrier.
578///
579/// # Errors
580///
581/// Returns [`Error::Validation`] with `InvalidArgument` when lowering produces
582/// an invalid output count or unknown output metadata, [`Error::Internal`] for
583/// an invalid graph reference, and [`Error::RuntimeStateSource`] when metadata
584/// registration cannot retain the lowered graph state.
585pub 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;