# Source code for tensorplay.export._trace Source: https://www.tensorplay.cn/docs/_modules/tensorplay/export/_trace.html ``` """Capture routines for building structured exported programs.""" from __future__ import annotations import inspect from collections.abc import Mapping from typing import Any, Callable from ..graph import GraphCaptureError, GraphModule, Node, Proxy, Tracer from .dynamic_shapes import AdditionalInputs, ConstraintsExceededError, Dim, ShapesCollection, _DimHint from .exported_program import EqualityConstraint, ExportedProgram, ModuleCallEntry, ModuleCallSignature from .graph_signature import ( ConstantArgument, ExportGraphSignature, InputKind, InputSpec, OutputKind, OutputSpec, TensorArgument, ) __all__ = ["ExportTracer", "draft_export", "export", "export_for_training"] _STATE_PREFIXES = {"parameter": "p", "buffer": "b", "constant": "c"} def _qualified(module_name: str, attribute: str) -> str: return f"{module_name}.{attribute}" if module_name else attribute def _collect_attributes(root: Any) -> dict[str, tuple[str, bool]]: """Map qualified attribute paths to ``(kind, persistent)``. Kinds cover parameters, persistent and non-persistent buffers, and plain tensor attributes (recorded as constants). Constant entries carry the value so the tracer can lift them alongside parameters and buffers. """ attributes: dict[str, tuple[str, bool]] = {} named_modules = getattr(root, "named_modules", None) if not callable(named_modules): return attributes import tensorplay as tp for module_name, module in named_modules(remove_duplicate=True): non_persistent = getattr(module, "_non_persistent_buffers_set", set()) for name, value in getattr(module, "_parameters", {}).items(): if value is not None: attributes[_qualified(module_name, name)] = ("parameter", True) for name, value in getattr(module, "_buffers", {}).items(): if value is not None: attributes[_qualified(module_name, name)] = ( "buffer", name not in non_persistent, ) parameters = getattr(module, "_parameters", {}) buffers = getattr(module, "_buffers", {}) children = getattr(module, "_modules", {}) for name, value in vars(module).items(): if name.startswith("_") or name in parameters or name in buffers: continue if name in children or callable(value): continue if isinstance(value, tp.nn.Parameter) or not isinstance(value, tp.Tensor): continue attributes[_qualified(module_name, name)] = ("constant", True) return attributes def _resolve_attribute(root: Any, target: str) -> Any: value = root for atom in target.split("."): value = getattr(value, atom) return value class ExportTracer(Tracer): """Tracer that lifts module state into graph inputs. Parameters, buffers, and constant tensors become placeholders ahead of the user inputs, so the captured graph is functional: it reads no attributes and its only inputs are the flat value list described by the graph signature. Child module forwards are additionally wrapped so every recorded node carries the qualified path of the module that produced it (``nn_module_stack``), and module call boundaries (argument and result nodes) are recorded for later hierarchy reconstruction. """ def __init__(self, concrete_args: dict[str, Any] | None = None) -> None: super().__init__(concrete_args) # qualified attribute path -> (placeholder node, kind, persistent) self.state_targets: dict[str, tuple[Node, str, bool]] = {} self._constant_patches: list[tuple[Any, str, Any]] = [] self._missing_sentinel = object() self._module_stack: tuple[str, ...] = () self._forward_patches: list[tuple[Any, Any]] = [] self.module_calls: list[dict[str, Any]] = [] self._call_keys: set[tuple[Any, ...]] = set() def _register_state(self, root: Any) -> None: for target, (kind, persistent) in _collect_attributes(root).items(): value = _resolve_attribute(root, target) mangled = f"{_STATE_PREFIXES[kind]}_{target.replace('.', '_')}" node = self.graph.create_node("placeholder", mangled, (), {}, name=mangled) node.meta["state_target"] = target node.meta["state_kind"] = kind node.meta["state_persistent"] = persistent self.state_targets[target] = (node, kind, persistent) def _patch_constants(self, root: Any) -> None: for target, (node, kind, _persistent) in self.state_targets.items(): if kind != "constant": continue parent_name, _, leaf = target.rpartition(".") parent = _resolve_attribute(root, parent_name) if parent_name else root previous = getattr(parent, leaf, self._missing_sentinel) self._constant_patches.append((parent, leaf, previous)) setattr(parent, leaf, Proxy(node, self)) def _restore_constants(self) -> None: for parent, leaf, previous in reversed(self._constant_patches): if previous is self._missing_sentinel: delattr(parent, leaf) else: setattr(parent, leaf, previous) self._constant_patches.clear() def _wrap_child_forwards(self, root: Any) -> None: """Route child module calls through stack-tracking wrappers.""" from ..graph._pytree import tree_flatten named_modules = getattr(root, "named_modules", None) if not callable(named_modules): return for module_name, module in named_modules(remove_duplicate=True): if not module_name: continue # the root call is described by the program itself original = getattr(module, "forward", None) if not callable(original): continue def wrapper( *args: Any, _original: Any = original, _qualname: str = module_name, **kwargs: Any, ) -> Any: arg_nodes: list[Any] = [] for value in args: arg_nodes.append( value.node.name if isinstance(value, Proxy) else value ) kwargs_nodes = { key: value.node.name if isinstance(value, Proxy) else value for key, value in kwargs.items() } self._module_stack = (*self._module_stack, _qualname) try: result = _original(*args, **kwargs) finally: self._module_stack = self._module_stack[:-1] result_nodes: list[str] = [] def visit(item: Any) -> None: if isinstance(item, Proxy): result_nodes.append(item.node.name) elif isinstance(item, (tuple, list)): for entry in item: visit(entry) elif isinstance(item, dict): for entry in item.values(): visit(entry) visit(result) _in_spec = tree_flatten(args)[1] _out_spec = tree_flatten(result)[1] key = (_qualname, tuple(map(repr, arg_nodes)), tuple(result_nodes)) if key not in self._call_keys: self._call_keys.add(key) self.module_calls.append( { "fqn": _qualname, "args": arg_nodes, "kwargs": kwargs_nodes, "result": result_nodes, "in_spec": _in_spec, "out_spec": _out_spec, } ) return result try: module.forward = wrapper # type: ignore[method-assign] except Exception: continue self._forward_patches.append((module, original)) def _restore_child_forwards(self) -> None: for module, _original in self._forward_patches: try: del module.forward except AttributeError: pass self._forward_patches.clear() def trace(self, root: Any, sample_inputs: dict[str, Any] | None = None) -> GraphModule: self.root = root if callable(getattr(root, "named_modules", None)): self._register_state(root) self._patch_constants(root) self._wrap_child_forwards(root) try: return super().trace(root, sample_inputs) finally: self._restore_child_forwards() self._restore_constants() return super().trace(root, sample_inputs) def create_proxy( self, kind: str, target: Any, args: tuple[Any, ...], kwargs: dict[str, Any], ) -> Proxy: if kind == "get_attr": entry = self.state_targets.get(target) if entry is not None: return Proxy(entry[0], self) proxy = super().create_proxy(kind, target, args, kwargs) if self._module_stack: proxy.node.meta["nn_module_stack"] = self._module_stack return proxy def _validate_graph(graph_module: GraphModule, attributes: Mapping[str, Any]) -> None: for node in graph_module.graph.nodes: op = node.op if op == "get_attr": # A tensor the function closed over is not an attribute of the # model; the capture holds it itself, under a name of its own. if ( node.target not in attributes and node.target not in graph_module._graph_attrs ): raise GraphCaptureError( f"get_attr target {node.target!r} is not present on the captured model" ) elif op == "call_function": if not callable(node.target): raise GraphCaptureError(f"call_function target {node.target!r} is not callable") elif op == "call_method": if not node.args or not isinstance(node.target, str): raise GraphCaptureError(f"malformed call_method node: {node}") elif op not in {"placeholder", "output", "call_module"}: raise GraphCaptureError(f"unsupported graph node kind: {op!r}") if not graph_module.graph.outputs: raise GraphCaptureError("captured graph has no output") def _normalize_dynamic_shapes( spec: Any, parameter_names: list[str], model: Any, args: tuple[Any, ...], kwargs: Mapping[str, Any], ) -> Any: if spec is None: return {} if isinstance(spec, AdditionalInputs): return spec.dynamic_shapes(model, args, kwargs) if isinstance(spec, ShapesCollection): return spec.dynamic_shapes(model, args, kwargs) if isinstance(spec, (list, tuple)): if len(spec) > len(parameter_names): raise ValueError("dynamic_shapes has more entries than graph inputs") return { name: _validate_dimension_spec(value) for name, value in zip(parameter_names, spec) if value is not None } if not isinstance(spec, dict): raise TypeError("dynamic_shapes must be a mapping or a sequence") normalized: dict[str, Any] = {} for name, dims_spec in spec.items(): if name not in parameter_names: raise ValueError( f"dynamic_shapes key {name!r} does not match any argument; " f"expected one of {parameter_names}" ) if isinstance(dims_spec, dict): entry: dict[int, Any] = {} for index, value in dims_spec.items(): if type(index) is not int or index < 0: raise TypeError(f"dimension index must be a non-negative int, got {index!r}") entry[index] = _validate_dimension_value(value) normalized[name] = entry elif isinstance(dims_spec, (tuple, list)): normalized[name] = _validate_dimension_spec(dims_spec) elif dims_spec is None: normalized[name] = None else: raise TypeError(f"dynamic_shapes[{name!r}] must describe dimensions") return normalized def _validate_dimension_value(value: Any) -> Any: if value is None or isinstance(value, (Dim, _DimHint)): return value if type(value) is int: return value raise TypeError( "dimension spec must be int or Dim (a dim hint or None is also accepted), " f"got {type(value)!r}" ) def _validate_dimension_spec(value: Any) -> Any: if isinstance(value, dict): result: dict[int, Any] = {} for index, item in value.items(): if type(index) is not int or index < 0: raise TypeError(f"dimension index must be a non-negative int, got {index!r}") result[index] = _validate_dimension_value(item) return result if isinstance(value, (tuple, list)): return tuple(_validate_dimension_value(item) for item in value) return _validate_dimension_value(value) def _argument_for_node(node: Any) -> Any: if isinstance(node, Node): return TensorArgument(node.name) if isinstance(node, (str, int, float, bool)) or node is None: return ConstantArgument(f"constant_{abs(hash(repr(node))) % 100000}", node) return ConstantArgument(f"constant_{abs(hash(type(node).__name__)) % 100000}", node) def _flatten_leaves(value: Any) -> list[Any]: leaves: list[Any] = [] def visit(item: Any) -> None: if isinstance(item, Node): leaves.append(item) elif isinstance(item, (tuple, list)): for entry in item: visit(entry) elif isinstance(item, dict): for entry in item.values(): visit(entry) else: leaves.append(item) visit(value) return leaves def _mutation_chain_root(node: Node) -> Node: """Walk in-place op and element-read chains back to the state holder. Two hop rules apply: in-place methods (``add_``) consume the previous value of the object they update, and ``getitem`` reads reach into a container that itself entered the graph as one placeholder. Hopping through both attributes an element update to the container input. """ import operator seen: set[int] = set() current = node while id(current) not in seen: seen.add(id(current)) if ( current.op == "call_method" and isinstance(current.target, str) and current.target.endswith("_") and current.args and isinstance(current.args[0], Node) ): current = current.args[0] continue if ( current.op == "call_function" and current.target is operator.getitem and current.args and isinstance(current.args[0], Node) ): current = current.args[0] continue break return current def _detect_mutations( graph_module: GraphModule, state_targets: Mapping[str, tuple[Node, str, bool]], ) -> list[tuple[Node, OutputKind, str | None]]: """Find in-place updates of lifted state or user inputs. Returns ``(holder, kind, target)`` tuples where ``holder`` is the graph node carrying the final value of the mutated object. """ state_by_placeholder: dict[str, tuple[str, str, bool]] = { node.name: (target, kind, persistent) for target, (node, kind, persistent) in state_targets.items() } holders: dict[str, tuple[OutputKind, str | None]] = {} for node in graph_module.graph.placeholders: entry = state_by_placeholder.get(node.name) if entry is not None: _target, kind, _persistent = entry holders[node.name] = ( ( OutputKind.PARAMETER_MUTATION if kind == "parameter" else OutputKind.BUFFER_MUTATION ), _target, ) else: holders[node.name] = (OutputKind.USER_INPUT_MUTATION, None) mutations: dict[str, tuple[Node, OutputKind, str | None]] = {} for node in graph_module.graph.nodes: if node.op != "call_method" or not isinstance(node.target, str): continue if not node.target.endswith("_") or not node.args or not isinstance(node.args[0], Node): continue root = _mutation_chain_root(node) info = holders.get(root.name) if info is None: continue # only a container element (getitem chain) may update a user input; # a direct in-place call on a whole input tensor is graph-invisible # by construction because capture runs on value copies if info[0] is OutputKind.USER_INPUT_MUTATION and root is node.args[0]: continue target = info[1] if info[1] is not None else root.name mutations[node.name] = (node, info[0], target) return list(mutations.values()) def _state_entry_for_placeholder( state_targets: Mapping[str, tuple[Node, str, bool]], placeholder_name: str, ) -> tuple[str, str, bool] | None: for target, (node, kind, persistent) in state_targets.items(): if node.name == placeholder_name: return target, kind, persistent return None def _state_map_for_placeholders( state_targets: Mapping[str, tuple[Node, str, bool]], ) -> dict[str, tuple[str, str, bool]]: """Placeholder-name-keyed view of the lifted-state table. Signature construction visits placeholders in graph order; a name-keyed map keeps that walk linear instead of scanning the state table per node. """ return { node.name: (target, kind, persistent) for target, (node, kind, persistent) in state_targets.items() } def _rewrite_container_reads( graph_module: GraphModule, mutations: list[tuple[Node, OutputKind, str | None]], ) -> None: """Point element reads after a mutation at the mutated value. A second ``items[0]`` read records its own getitem node; without this rewrite it would observe the pre-mutation element and diverge from eager execution, where both reads return the same object. """ import operator if not mutations: return graph = graph_module.graph order = {node.name: position for position, node in enumerate(graph.nodes)} finals: dict[tuple[str, int], Node] = {} for node, _kind, _target in mutations: current = node container: Node | None = None index: int | None = None while True: if ( current.op == "call_function" and current.target is operator.getitem and current.args and isinstance(current.args[0], Node) ): if index is None and len(current.args) > 1 and isinstance(current.args[1], int): container = current.args[0] index = current.args[1] current = current.args[0] continue break if container is not None and index is not None: finals[(container.name, index)] = node if not finals: return for read in list(graph.nodes): if read.op != "call_function" or read.target is not operator.getitem: continue if len(read.args) != 2 or not isinstance(read.args[0], Node): continue key = (read.args[0].name, read.args[1]) if isinstance(read.args[1], int) else None final = finals.get(key) if key is not None else None if final is None or final is read: continue if order.get(final.name, -1) >= order.get(read.name, 1 << 30): # the read happens before the mutation; it keeps the old value continue read.replace_all_uses_with(final) if not read.users: graph.erase_node(read) def _output_specs( graph_module: GraphModule, mutations: list[tuple[Node, OutputKind, str | None]], ) -> list[OutputSpec]: output = graph_module.graph.output_node leaves = _flatten_leaves(output.args[0]) specs = [ OutputSpec(kind, TensorArgument(node.name), target) for node, kind, target in mutations ] specs.extend( OutputSpec(OutputKind.USER_OUTPUT, _argument_for_node(value)) for value in leaves[len(mutations):] ) return specs def _restructure_output( graph_module: GraphModule, mutations: list[tuple[Node, OutputKind, str | None]], ) -> None: """Rewrite the graph result to ``(mutations..., flattened_user_outputs...)``. A mutated value may also be a user output; the flat contract repeats the node in both roles, exactly as the signature describes. """ if not mutations: return graph = graph_module.graph output = graph.output_node leaves = _flatten_leaves(output.args[0]) mutation_nodes = [node for node, _kind, _target in mutations] graph.output(tuple([*mutation_nodes, *leaves])) # -- runtime assertions for dynamic-shape contracts ------------------------ def _assert_dim_range(tensor: Any, index: int, min: Any, max: Any, name: str) -> Any: size = tuple(tensor.shape)[index] if (min is not None and size < min) or (max is not None and size > max): raise ConstraintsExceededError( f"runtime assertion failed for {name!r}: expected dimension {index} " f"in [{min if min is not None else '-inf'}, {max if max is not None else 'inf'}], " f"got {size}" ) return size def _assert_dims_equal(tensor_a: Any, index_a: int, tensor_b: Any, index_b: int, name: str) -> Any: size_a = tuple(tensor_a.shape)[index_a] size_b = tuple(tensor_b.shape)[index_b] if size_a != size_b: raise ConstraintsExceededError( f"runtime assertion failed for {name!r}: dimensions " f"{index_a} and {index_b} must agree, got {size_a} and {size_b}" ) return size_a def _assert_dim_relation( tensor_root: Any, index_root: int, tensor_derived: Any, index_derived: int, scale: int, offset: int, name: str, ) -> Any: size_root = tuple(tensor_root.shape)[index_root] size_derived = tuple(tensor_derived.shape)[index_derived] expected = scale * size_root + offset if size_derived != expected: raise ConstraintsExceededError( f"runtime assertion failed for {name!r}: expected dimension " f"{index_derived} == {scale} * dim {index_root} + {offset} " f"({expected}), got {size_derived}" ) return size_derived def _apply_dynamic_shape_constraints( graph_module: GraphModule, combined_args: Mapping[str, Any], normalized: Mapping[str, Any], ) -> tuple[dict[str, dict[str, int | None]], list[EqualityConstraint]]: """Validate the spec, insert runtime assertions, and describe shared dims. Returns the per-name range bounds and one equality record per named dimension that appears at more than one input site (shared dims must hold equal sizes across those sites at runtime). """ from .dynamic_shapes import _constraint_program, _process_dynamic_shapes from .dim_constraints import DimConstraints, attach_observed_sizes constraints = _process_dynamic_shapes(combined_args, normalized) if not constraints: return {}, [] # strict checks first: contradictory range declarations for one name are # specification errors regardless of the example inputs asserts, ranges = _constraint_program(constraints) solver = DimConstraints() for constraint, observed in attach_observed_sizes(constraints, combined_args, normalized): solver.add(constraint, observed) if not solver.solve(): raise ConstraintsExceededError( "export-time dimension constraints are inconsistent:\n" + solver.pretty_print() ) placeholders = {node.name: node for node in graph_module.graph.placeholders} identity_to_name = { id(value): name for name, value in combined_args.items() if not isinstance(value, (dict, list, tuple)) } # one equality record per name spanning several distinct input sites site_pairs: dict[str, set[tuple[str, int]]] = {} for constraint in constraints: if constraint.name is None: continue input_name = identity_to_name.get(id(constraint.source)) if input_name is None: continue site_pairs.setdefault(constraint.name, set()).add((input_name, constraint.dim)) equality_constraints = [ EqualityConstraint(tuple(sorted(pairs, key=repr)), name=name) for name, pairs in site_pairs.items() if len(pairs) > 1 ] def node_of(source: Any) -> Node | None: name = identity_to_name.get(id(source)) return placeholders.get(name) if name is not None else None anchors: dict[str, tuple[Node, int]] = {} graph = graph_module.graph output = graph.output_node with graph.inserting_before(output): for constraint in asserts: if constraint.root is None and constraint.name is not None: node = node_of(constraint.source) if node is not None: anchors[constraint.name] = (node, constraint.dim) for constraint in asserts: node = node_of(constraint.source) if node is None: continue if constraint.root is not None: anchor = anchors.get(constraint.root) if anchor is None: continue graph.call_function( _assert_dim_relation, ( anchor[0], anchor[1], node, constraint.dim, constraint.scale, constraint.offset, constraint.name or constraint.root, ), ) continue if constraint.name is not None: anchor = anchors.get(constraint.name) if anchor is not None and anchor[0] is not node: graph.call_function( _assert_dims_equal, ( anchor[0], anchor[1], node, constraint.dim, constraint.name, ), ) continue if constraint.min is None and constraint.max is None: continue graph.call_function( _assert_dim_range, ( node, constraint.dim, constraint.min, constraint.max, constraint.name or f"dim {constraint.dim}", ), ) return ranges, equality_constraints def _graph_signature( graph_module: GraphModule, state_targets: Mapping[str, tuple[Node, str, bool]], mutations: list[tuple[Node, OutputKind, str | None]], ) -> ExportGraphSignature: """Signature over the flat input contract: state first, then user inputs.""" state_by_placeholder = _state_map_for_placeholders(state_targets) inputs: list[InputSpec] = [] for node in graph_module.graph.placeholders: entry = state_by_placeholder.get(node.name) if entry is None: inputs.append(InputSpec(InputKind.USER_INPUT, TensorArgument(node.name), None)) continue target, kind, persistent = entry if kind == "parameter": inputs.append( InputSpec(InputKind.PARAMETER, TensorArgument(node.name), target) ) elif kind == "buffer": inputs.append( InputSpec( InputKind.BUFFER, TensorArgument(node.name), target, persistent=persistent, ) ) else: inputs.append( InputSpec( InputKind.CONSTANT_TENSOR, TensorArgument(node.name), target, persistent=None, ) ) return ExportGraphSignature(inputs, _output_specs(graph_module, mutations)) def _bind_examples(graph_module: GraphModule, args: tuple[Any, ...], kwargs: Mapping[str, Any]) -> dict[str, Any]: signature = graph_module.signature if signature is None: signature = inspect.signature(getattr(graph_module.root, "forward", graph_module.root)) bound = signature.bind_partial(*args, **dict(kwargs)) bound.apply_defaults() return { node.name: bound.arguments[ node.target if isinstance(node.target, str) else node.name ] for node in graph_module.graph.placeholders if ( node.target if isinstance(node.target, str) else node.name ) in bound.arguments } def _flat_signature( graph_module: GraphModule, state_targets: Mapping[str, tuple[Node, str, bool]], example_inputs: Mapping[str, Any] | None = None, ) -> inspect.Signature: """Signature covering every flat input: state placeholders, then user args. User parameters keep their kinds; defaults prefer the export-time binding over the source declaration so omitted call arguments replay the capture. """ example_inputs = example_inputs or {} state_names = {node.name for node, _kind, _persistent in state_targets.values()} user_parameters = ( dict(graph_module.signature.parameters) if graph_module.signature is not None else {} ) parameters: list[inspect.Parameter] = [] for node in graph_module.graph.placeholders: if node.name in state_names: parameters.append( inspect.Parameter(node.name, inspect.Parameter.POSITIONAL_OR_KEYWORD) ) continue original = user_parameters.get(node.name) default = original.default if original is not None else inspect.Parameter.empty if node.name in example_inputs: default = example_inputs[node.name] if original is not None and default is not original.default: original = inspect.Parameter( original.name, original.kind, default=default ) if original is not None: parameters.append(original) elif node.args: parameters.append( inspect.Parameter( node.name, inspect.Parameter.POSITIONAL_OR_KEYWORD, default=node.args[0], ) ) else: parameters.append( inspect.Parameter(node.name, inspect.Parameter.POSITIONAL_OR_KEYWORD) ) return inspect.Signature(parameters) def _capture( model: Callable[..., Any], args: tuple[Any, ...], kwargs: Mapping[str, Any], dynamic_shapes: Any, ) -> ExportedProgram: from ..graph._pytree import tree_flatten from .dynamic_shapes import _combine_args tracer = ExportTracer() target = model.forward if callable(getattr(model, "forward", None)) else model try: bound = inspect.signature(target).bind_partial(*args, **kwargs) bound.apply_defaults() sample_inputs = dict(bound.arguments) except (TypeError, ValueError): sample_inputs = None from tensorplay.compiler import _exporting_context with _exporting_context(): graph_module = tracer.trace(model, sample_inputs=sample_inputs) attributes = _collect_attributes(model) _validate_graph(graph_module, attributes) placeholders = graph_module.graph.placeholders names = [node.name for node in placeholders] examples = _bind_examples(graph_module, args, kwargs) normalized = _normalize_dynamic_shapes(dynamic_shapes, names, model, args, kwargs) meta = graph_module.meta meta["user_signature"] = graph_module.signature meta["state_targets"] = dict(tracer.state_targets) constants: dict[str, Any] = {} for target, (_node, kind, persistent) in tracer.state_targets.items(): if kind == "constant" or (kind == "buffer" and not persistent): constants[target] = _resolve_attribute(model, target) meta["constants"] = constants mutations = _detect_mutations(graph_module, tracer.state_targets) user_out_spec = tree_flatten(graph_module.graph.output_node.args[0])[1] ranges, equality_constraints = _apply_dynamic_shape_constraints( graph_module, _combine_args(model, args, kwargs), normalized ) _rewrite_container_reads(graph_module, mutations) _restructure_output(graph_module, mutations) if mutations or ranges: graph_module.recompile() meta["num_mutations"] = len(mutations) meta["out_spec"] = user_out_spec meta["in_spec"] = tree_flatten((tuple(args), dict(kwargs)))[1] graph_module.signature = _flat_signature( graph_module, tracer.state_targets, examples ) meta["module_calls"] = list(tracer.module_calls) signature = _graph_signature(graph_module, tracer.state_targets, mutations) user_output_args = [ spec.arg for spec in signature.output_specs if spec.kind is OutputKind.USER_OUTPUT ] root_entry = ModuleCallEntry( "", ModuleCallSignature( inputs=[TensorArgument(node.name) for node in placeholders], outputs=user_output_args, in_spec=meta.get("in_spec"), out_spec=meta.get("out_spec"), forward_arg_names=[node.name for node in placeholders], ), ) call_entries: list[ModuleCallEntry] = [] for record in tracer.module_calls: arg_specs = [ TensorArgument(value) if isinstance(value, str) else ConstantArgument("", value) for value in record["args"] ] call_entries.append( ModuleCallEntry( record["fqn"], ModuleCallSignature( inputs=arg_specs, outputs=[TensorArgument(name) for name in record["result"]], in_spec=record.get("in_spec"), out_spec=record.get("out_spec"), forward_arg_names=[ value if isinstance(value, str) else f"arg_{index}" for index, value in enumerate(record["args"]) ], ), ) ) calls = [root_entry, *call_entries] program = ExportedProgram( graph_module=graph_module, graph_signature=signature, example_inputs=examples, dynamic_shapes=normalized, module_call_graph=calls, range_constraints=ranges, equality_constraints=equality_constraints, ) program.validate() return program [docs] def export( model: Callable[..., Any], *args: Any, dynamic_shapes: Any = None, strict: bool = False, preserve_module_call_signature: Any = (), **kwargs: Any, ) -> ExportedProgram: """Capture a callable and return an executable graph program. Args: model: an ``nn.Module`` or plain callable; child modules are inlined. args/kwargs: example inputs binding argument defaults. dynamic_shapes: dimension specification per argument (dict, sequence, :class:`ShapesCollection`, or :class:`AdditionalInputs`). strict: reserved for callers of the strict capture contract; capture validation is identical in both modes. preserve_module_call_signature: submodule paths whose call metadata is recorded in ``module_call_graph`` for module-level tooling. """ if not callable(model): raise TypeError(f"model must be callable, got {type(model).__name__}") if isinstance(model, object) and hasattr(model, "named_modules"): known = {name for name, _ in model.named_modules(remove_duplicate=True)} unknown = [p for p in preserve_module_call_signature if p not in known] if unknown: raise GraphCaptureError( f"preserve_module_call_signature paths {sorted(unknown)} do not " f"exist on the model; known: {sorted(known)}" ) program = _capture(model, args, kwargs, dynamic_shapes) if preserve_module_call_signature: existing = {entry.fqn for entry in program.module_call_graph} program.module_call_graph.extend( ModuleCallEntry(path) for path in preserve_module_call_signature if path not in existing ) del strict return program [docs] def export_for_training( model: Callable[..., Any], *args: Any, dynamic_shapes: Any = None, **kwargs: Any, ) -> ExportedProgram: """Capture a callable while retaining its mutable training state.""" return export(model, *args, dynamic_shapes=dynamic_shapes, **kwargs) def draft_export( model: Callable[..., Any], *args: Any, dynamic_shapes: Any = None, **kwargs: Any, ) -> Any: """Capture with failure reporting; see :mod:`tensorplay.export._draft_export`.""" from ._draft_export import draft_export as _draft_export return _draft_export(model, *args, dynamic_shapes=dynamic_shapes, **kwargs) ```