latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
Source code for tensorplay.export.graph_signature
"""Structured input and output descriptions for exported graphs."""
from __future__ import annotations
import dataclasses
from collections.abc import Collection, Mapping
from enum import Enum, auto
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from ..graph.node import Node
__all__ = [
"ArgumentSpec",
"ConstantArgument",
"CustomObjArgument",
"ExportBackwardSignature",
"ExportGraphSignature",
"GraphSignature",
"InputKind",
"InputSpec",
"OutputKind",
"OutputSpec",
"SymBoolArgument",
"SymFloatArgument",
"SymIntArgument",
"TensorArgument",
"TokenArgument",
]
[docs]
@dataclasses.dataclass
class TensorArgument:
name: str
[docs]
@dataclasses.dataclass
class TokenArgument:
name: str
[docs]
@dataclasses.dataclass
class SymIntArgument:
name: str
[docs]
@dataclasses.dataclass
class SymFloatArgument:
name: str
[docs]
@dataclasses.dataclass
class SymBoolArgument:
name: str
[docs]
@dataclasses.dataclass
class CustomObjArgument:
name: str
class_fqn: str
fake_val: Any = None
[docs]
@dataclasses.dataclass
class ConstantArgument:
name: str
value: int | float | bool | str | None
ArgumentSpec = (
TensorArgument
| SymIntArgument
| SymFloatArgument
| SymBoolArgument
| ConstantArgument
| CustomObjArgument
| TokenArgument
)
[docs]
class InputKind(Enum):
USER_INPUT = auto()
PARAMETER = auto()
BUFFER = auto()
CONSTANT_TENSOR = auto()
CUSTOM_OBJ = auto()
TOKEN = auto()
[docs]
class OutputKind(Enum):
USER_OUTPUT = auto()
LOSS_OUTPUT = auto()
BUFFER_MUTATION = auto()
PARAMETER_MUTATION = auto()
GRADIENT_TO_PARAMETER = auto()
GRADIENT_TO_USER_INPUT = auto()
USER_INPUT_MUTATION = auto()
TOKEN = auto()
_ARGUMENT_TYPES = (
TensorArgument,
SymIntArgument,
SymFloatArgument,
SymBoolArgument,
ConstantArgument,
CustomObjArgument,
TokenArgument,
)
[docs]
@dataclasses.dataclass
class InputSpec:
kind: InputKind
arg: ArgumentSpec
target: str | None = None
persistent: bool | None = None
def __post_init__(self) -> None:
if not isinstance(self.kind, InputKind):
raise TypeError(f"input kind must be InputKind, got {type(self.kind).__name__}")
if not isinstance(self.arg, _ARGUMENT_TYPES):
raise TypeError(f"invalid input argument type: {type(self.arg).__name__}")
if self.kind is InputKind.BUFFER and self.persistent is None:
raise ValueError("buffer input requires a persistence flag")
def __str__(self) -> str:
target = "" if self.target is None else f" target={self.target!r}"
persistent = "" if self.persistent is None else f" persistent={self.persistent}"
return f"{self.arg.name}: {self.kind.name}{target}{persistent}"
[docs]
@dataclasses.dataclass
class OutputSpec:
kind: OutputKind
arg: ArgumentSpec
target: str | None = None
def __post_init__(self) -> None:
if not isinstance(self.kind, OutputKind):
raise TypeError(f"output kind must be OutputKind, got {type(self.kind).__name__}")
if not isinstance(self.arg, _ARGUMENT_TYPES):
raise TypeError(f"invalid output argument type: {type(self.arg).__name__}")
def __str__(self) -> str:
target = "" if self.target is None else f" target={self.target!r}"
return f"{self.arg.name}: {self.kind.name}{target}"
[docs]
@dataclasses.dataclass
class ExportBackwardSignature:
gradients_to_parameters: dict[str, str]
gradients_to_user_inputs: dict[str, str]
loss_output: str
[docs]
@dataclasses.dataclass
class ExportGraphSignature:
"""Describe lifted state, user values, mutations, and graph outputs."""
input_specs: list[InputSpec]
output_specs: list[OutputSpec]
def __post_init__(self) -> None:
self.input_specs = list(self.input_specs)
self.output_specs = list(self.output_specs)
@property
def parameters(self) -> Collection[str]:
return tuple(
spec.target
for spec in self.input_specs
if spec.kind is InputKind.PARAMETER and isinstance(spec.target, str)
)
@property
def buffers(self) -> Collection[str]:
return tuple(
spec.target
for spec in self.input_specs
if spec.kind is InputKind.BUFFER and isinstance(spec.target, str)
)
@property
def non_persistent_buffers(self) -> Collection[str]:
return tuple(
spec.target
for spec in self.input_specs
if spec.kind is InputKind.BUFFER
and spec.persistent is False
and isinstance(spec.target, str)
)
@property
def lifted_tensor_constants(self) -> Collection[str]:
return tuple(
spec.target
for spec in self.input_specs
if spec.kind is InputKind.CONSTANT_TENSOR and isinstance(spec.target, str)
)
@property
def lifted_custom_objs(self) -> Collection[str]:
return tuple(
spec.target
for spec in self.input_specs
if spec.kind is InputKind.CUSTOM_OBJ and isinstance(spec.target, str)
)
# neutral aliases: tensor constants and custom objects live in the same
# input-spec table; callers may prefer the shorter names
@property
def constants(self) -> Collection[str]:
return self.lifted_tensor_constants
@property
def tensor_constants(self) -> Collection[str]:
return self.lifted_tensor_constants
@property
def custom_objs(self) -> Collection[str]:
return self.lifted_custom_objs
[docs]
def is_param(self, name: str) -> bool:
"""Whether ``name`` is a placeholder carrying a lifted parameter."""
return any(
spec.kind is InputKind.PARAMETER
and isinstance(spec.arg, TensorArgument)
and spec.arg.name == name
for spec in self.input_specs
)
[docs]
def is_buffer(self, name: str) -> bool:
"""Whether ``name`` is a placeholder carrying a lifted buffer."""
return any(
spec.kind is InputKind.BUFFER
and isinstance(spec.arg, TensorArgument)
and spec.arg.name == name
for spec in self.input_specs
)
[docs]
def get_param_to_buffer(self) -> Mapping[str, str]:
"""Map parameter targets to the buffer targets holding their gradients.
Gradients are declared as ``GRADIENT_TO_PARAMETER`` outputs; a gradient
for a parameter whose optimizer state lives in a buffer binds the two
targets under the parameter's FQN.
"""
result: dict[str, str] = {}
gradient_targets = {
spec.target
for spec in self.output_specs
if spec.kind is OutputKind.GRADIENT_TO_PARAMETER and isinstance(spec.target, str)
}
if not gradient_targets:
return result
buffer_targets = {
spec.target
for spec in self.input_specs
if spec.kind is InputKind.BUFFER and isinstance(spec.target, str)
}
for target in gradient_targets:
if target in buffer_targets:
result[target] = target
return result
@property
def user_inputs(self) -> Collection[Any]:
values: list[Any] = []
for spec in self.input_specs:
if spec.kind is not InputKind.USER_INPUT:
continue
if isinstance(spec.arg, ConstantArgument):
values.append(spec.arg.value)
else:
values.append(spec.arg.name)
return tuple(values)
@property
def user_outputs(self) -> Collection[Any]:
values: list[Any] = []
for spec in self.output_specs:
if spec.kind not in (OutputKind.USER_OUTPUT, OutputKind.LOSS_OUTPUT):
continue
if isinstance(spec.arg, ConstantArgument):
values.append(spec.arg.value)
else:
values.append(spec.arg.name)
return tuple(values)
@property
def inputs_to_parameters(self) -> Mapping[str, str]:
return {
spec.arg.name: spec.target
for spec in self.input_specs
if spec.kind is InputKind.PARAMETER
and isinstance(spec.arg, TensorArgument)
and isinstance(spec.target, str)
}
@property
def inputs_to_buffers(self) -> Mapping[str, str]:
return {
spec.arg.name: spec.target
for spec in self.input_specs
if spec.kind is InputKind.BUFFER
and isinstance(spec.arg, TensorArgument)
and isinstance(spec.target, str)
}
@property
def inputs_to_lifted_tensor_constants(self) -> Mapping[str, str]:
return {
spec.arg.name: spec.target
for spec in self.input_specs
if spec.kind is InputKind.CONSTANT_TENSOR
and isinstance(spec.arg, TensorArgument)
and isinstance(spec.target, str)
}
@property
def buffers_to_mutate(self) -> Mapping[str, str]:
return {
spec.arg.name: spec.target
for spec in self.output_specs
if spec.kind is OutputKind.BUFFER_MUTATION
and isinstance(spec.arg, TensorArgument)
and isinstance(spec.target, str)
}
@property
def parameters_to_mutate(self) -> Mapping[str, str]:
return {
spec.arg.name: spec.target
for spec in self.output_specs
if spec.kind is OutputKind.PARAMETER_MUTATION
and isinstance(spec.arg, TensorArgument)
and isinstance(spec.target, str)
}
@property
def user_inputs_to_mutate(self) -> Mapping[str, str]:
return {
spec.arg.name: spec.target
for spec in self.output_specs
if spec.kind is OutputKind.USER_INPUT_MUTATION
and isinstance(spec.arg, TensorArgument)
and isinstance(spec.target, str)
}
@property
def backward_signature(self) -> ExportBackwardSignature | None:
loss_output: str | None = None
gradients_to_parameters: dict[str, str] = {}
gradients_to_user_inputs: dict[str, str] = {}
for spec in self.output_specs:
if spec.kind is OutputKind.LOSS_OUTPUT:
if loss_output is not None or not isinstance(spec.arg, TensorArgument):
raise ValueError("loss output must be one tensor argument")
loss_output = spec.arg.name
elif spec.kind is OutputKind.GRADIENT_TO_PARAMETER:
if isinstance(spec.arg, TensorArgument) and isinstance(spec.target, str):
gradients_to_parameters[spec.arg.name] = spec.target
elif spec.kind is OutputKind.GRADIENT_TO_USER_INPUT:
if isinstance(spec.arg, TensorArgument) and isinstance(spec.target, str):
gradients_to_user_inputs[spec.arg.name] = spec.target
if loss_output is None:
return None
return ExportBackwardSignature(
gradients_to_parameters=gradients_to_parameters,
gradients_to_user_inputs=gradients_to_user_inputs,
loss_output=loss_output,
)
@property
def input_tokens(self) -> Collection[str]:
return tuple(
spec.arg.name
for spec in self.input_specs
if spec.kind is InputKind.TOKEN and isinstance(spec.arg, TokenArgument)
)
@property
def output_tokens(self) -> Collection[str]:
return tuple(
spec.arg.name
for spec in self.output_specs
if spec.kind is OutputKind.TOKEN and isinstance(spec.arg, TokenArgument)
)
@property
def assertion_dep_token(self) -> Mapping[int, str] | None:
"""Position of the assertion dependency token output, if present."""
tokens = self.output_tokens
if not tokens:
return None
index = len(self.user_outputs) + len(self.buffers_to_mutate)
return {index: tokens[0]}
[docs]
def replace_all_uses(self, old: str, new: str) -> None:
"""Rename a graph value across every input and output spec."""
if not isinstance(old, str) or not isinstance(new, str):
raise TypeError("replace_all_uses expects string names")
for spec in (*self.output_specs, *self.input_specs):
if spec.arg.name == old:
spec.arg.name = new
[docs]
def get_replace_hook(self, replace_inputs: bool = False):
"""Build a rename hook suitable for graph rewriting passes."""
def hook(old: Any, new: Any, user: Any) -> None:
if getattr(user, "op", None) == "output":
self.replace_all_uses(old.name, new)
if replace_inputs and getattr(old, "op", None) == "placeholder":
self.replace_all_uses(old.name, new)
return hook
[docs]
def clone(self) -> "ExportGraphSignature":
"""Deep copy: specs and argument records are duplicated, not shared."""
return dataclasses.replace(
self,
input_specs=[
dataclasses.replace(spec, arg=dataclasses.replace(spec.arg))
for spec in self.input_specs
],
output_specs=[
dataclasses.replace(spec, arg=dataclasses.replace(spec.arg))
for spec in self.output_specs
],
)
def __deepcopy__(self, memo: dict[int, Any]) -> "ExportGraphSignature":
return self.clone()
def __str__(self) -> str:
inputs = "\n".join(str(spec) for spec in self.input_specs)
outputs = "\n".join(str(spec) for spec in self.output_specs)
return f"\n# inputs\n{inputs}\n\n# outputs\n{outputs}\n"
[docs]
@dataclasses.dataclass(frozen=True)
class GraphSignature:
"""Compact signature retained for callers that build signatures directly."""
parameters: tuple[str, ...]
buffers: tuple[str, ...]
non_persistent_buffers: tuple[str, ...]
user_inputs: tuple[str, ...]
def to_export_signature(self) -> ExportGraphSignature:
inputs: list[InputSpec] = [
InputSpec(InputKind.PARAMETER, TensorArgument(name), name)
for name in self.parameters
]
inputs.extend(
InputSpec(
InputKind.BUFFER,
TensorArgument(name),
name,
persistent=name not in self.non_persistent_buffers,
)
for name in self.buffers
)
inputs.extend(
InputSpec(InputKind.USER_INPUT, TensorArgument(name), None)
for name in self.user_inputs
)
return ExportGraphSignature(inputs, [])
def _immutable_dict(items) -> Mapping[str, str]:
"""A mapping that rejects addition, deletion, and item assignment."""
from types import MappingProxyType
return MappingProxyType(dict(items))
def _make_argument_spec(value: Any, token_names: Any = ()) -> ArgumentSpec:
"""Classify one flattened graph value for signature bookkeeping."""
token_names = set(token_names)
if isinstance(value, Node):
if value.name in token_names:
return TokenArgument(value.name)
meta = getattr(value, "meta", {})
val = meta.get("val")
if val is not None and not hasattr(val, "shape"):
if isinstance(val, (int, float, bool, str)) or val is None:
return ConstantArgument(value.name, val)
fqn = getattr(val, "constant_name", None) or f"{type(val).__module__}.{type(val).__qualname__}"
return CustomObjArgument(value.name, fqn)
return TensorArgument(value.name)
if isinstance(value, (int, float, bool, str)) or value is None:
return ConstantArgument("", value)
raise TypeError(
f"expected a graph node or a scalar constant, got {type(value).__name__}"
)
def _convert_to_export_graph_signature(
graph: Any,
*,
user_inputs: Any,
inputs_to_parameters: Mapping[str, str],
inputs_to_buffers: Mapping[str, str],
user_outputs: Any,
buffer_mutations: Mapping[str, str] | None = None,
parameter_mutations: Mapping[str, str] | None = None,
user_input_mutations: Mapping[str, str] | None = None,
input_tokens: Any = (),
output_tokens: Any = (),
non_persistent_buffers: Any = (),
) -> ExportGraphSignature:
"""Build an :class:`ExportGraphSignature` from a flat graph and state maps.
``graph`` supplies the ordered placeholders and the output leaves;
the mapping arguments classify each one.
"""
buffer_mutations = buffer_mutations or {}
parameter_mutations = parameter_mutations or {}
user_input_mutations = user_input_mutations or {}
input_tokens = list(input_tokens)
output_tokens = list(output_tokens)
input_specs: list[InputSpec] = []
for node in graph.nodes:
if node.op != "placeholder":
continue
if node.name in input_tokens:
input_specs.append(InputSpec(InputKind.TOKEN, TokenArgument(node.name), None))
continue
if node.name in inputs_to_parameters:
input_specs.append(
InputSpec(InputKind.PARAMETER, TensorArgument(node.name), inputs_to_parameters[node.name])
)
elif node.name in inputs_to_buffers:
target = inputs_to_buffers[node.name]
input_specs.append(
InputSpec(
InputKind.BUFFER,
TensorArgument(node.name),
target,
persistent=target not in set(non_persistent_buffers),
)
)
else:
input_specs.append(
InputSpec(InputKind.USER_INPUT, _make_argument_spec(node, input_tokens), None)
)
output_node = graph.output_node
leaves = []
stack = [output_node.args[0]]
while stack:
item = stack.pop(0)
if isinstance(item, Node):
leaves.append(item)
elif isinstance(item, (tuple, list)):
stack.extend(item)
elif isinstance(item, dict):
stack.extend(item.values())
else:
leaves.append(item)
output_specs: list[OutputSpec] = []
mutation_slots = len(buffer_mutations) + len(parameter_mutations) + len(user_input_mutations) + len(output_tokens)
for index, value in enumerate(leaves):
spec = _make_argument_spec(value, output_tokens)
if isinstance(spec, TokenArgument):
output_specs.append(OutputSpec(OutputKind.TOKEN, spec, None))
continue
if index < mutation_slots and isinstance(spec, TensorArgument):
if spec.name in buffer_mutations:
output_specs.append(OutputSpec(OutputKind.BUFFER_MUTATION, spec, buffer_mutations[spec.name]))
continue
if spec.name in parameter_mutations:
output_specs.append(OutputSpec(OutputKind.PARAMETER_MUTATION, spec, parameter_mutations[spec.name]))
continue
if spec.name in user_input_mutations:
output_specs.append(OutputSpec(OutputKind.USER_INPUT_MUTATION, spec, user_input_mutations[spec.name]))
continue
output_specs.append(OutputSpec(OutputKind.USER_OUTPUT, spec, None))
if user_outputs is not None:
named = set(user_outputs)
for spec in output_specs:
if spec.kind is OutputKind.USER_OUTPUT and spec.arg.name in named:
spec.kind = OutputKind.USER_OUTPUT
return ExportGraphSignature(input_specs=input_specs, output_specs=output_specs)Help improve this page
Found an error, an unclear step, or a missing example?
Was this page helpful?

