TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

Source code for tensorplay.graph._utils

from __future__ import annotations

import inspect
import keyword
import re
from contextlib import contextmanager
from contextvars import ContextVar
from typing import Any, Callable, Dict, Iterable, Optional


class _LazyString:
    """String-like value whose formatting work is postponed until display."""

    __slots__ = ("_factory",)

    def __init__(self, factory: Callable[[], str]) -> None:
        self._factory = factory

    def __str__(self) -> str:
        return self._factory()

    def __repr__(self) -> str:
        return str(self)


def lazy_format_graph_code(
    name: str, graph_module: Any, maybe_id: int | None = None, **kwargs: Any
) -> _LazyString:
    """Return a lazily rendered description of a graph module."""

    del kwargs
    label = f"{name} {maybe_id}" if maybe_id is not None else name

    def render() -> str:
        code = getattr(graph_module, "code", "")
        return _format_graph_code(
            f"===== {label} =====\n",
            getattr(getattr(graph_module, "forward", None), "__code__", None),
            code,
        )

    return _LazyString(render)


def _format_graph_code(name: str, filename: Any, graph_str: str) -> str:
    filename_str = getattr(filename, "co_filename", filename)
    return f"TRACED GRAPH\n {name} {filename_str} {graph_str}\n"


def first_call_function_nn_module_stack(graph: Any) -> dict[str, Any] | None:
    """Return module-stack metadata from the first function node that has it."""

    for node in graph.nodes:
        if node.op == "call_function" and "nn_module_stack" in node.meta:
            return node.meta["nn_module_stack"]
    return None


def get_node_context(node: Any, num_nodes: int = 2) -> str:
    """Return a short source-order context ending at ``node``."""

    nodes = list(node.graph.nodes)
    try:
        index = nodes.index(node)
    except ValueError:
        return str(node)
    start = max(0, index - max(1, num_nodes) + 1)
    return "\n".join(str(item) for item in nodes[start : index + 1])



[docs]
class GraphCaptureError(RuntimeError):
    """Raised when Python code cannot be represented by the current graph."""



_compiling: ContextVar[bool] = ContextVar("tensorplay_graph_compiling", default=False)

_capture_disabled: ContextVar[bool] = ContextVar(
    "tensorplay_graph_capture_disabled", default=False
)

#: Set while a tracer executes a recorded node to obtain its example value.
#: Sample execution must run eagerly; capture-aware factory functions check
#: this so they do not record a nested node while the tracer is only sampling.
_executing_sample: ContextVar[bool] = ContextVar(
    "tensorplay_graph_executing_sample", default=False
)

#: Allocating tensor factories that must be recorded as graph nodes even
#: with no proxy argument.  They return fresh storage; freezing their eager
#: result into a constant would detach in-place fills and runtime-dependent
#: factory calls from the compiled region.
_FACTORY_NAMES = frozenset(
    {
        "empty",
        "empty_like",
        "empty_strided",
        "empty_permuted",
        "empty_quantized",
        "full",
        "full_like",
        "zeros",
        "zeros_like",
        "ones",
        "ones_like",
        "arange",
        "linspace",
        "logspace",
        "eye",
        "new_empty",
        "new_empty_strided",
        "new_full",
        "new_zeros",
        "new_ones",
    }
)

# The active tracer is exposed through a context variable so small graph
# markers can participate in capture even when they have no tensor argument.
# Keeping this state thread-local is important for nested captures and for
# concurrent compiler workers.
_active_tracer: ContextVar[Any] = ContextVar(
    "tensorplay_graph_active_tracer", default=None
)


def get_active_tracer() -> Any:
    return _active_tracer.get()


# Set while an operator traces a piece of the program it holds -- a branch, a
# loop body -- in a trace of its own.  A value of the enclosing trace that the
# piece closes over is read there as the tensor it stood for; the operator then
# hands that tensor back to the enclosing trace as one more input, where it is
# recognised as the value it came from.
_reading_enclosing_values: ContextVar[bool] = ContextVar(
    "tensorplay_graph_reading_enclosing_values", default=False
)


@contextmanager
def reading_enclosing_values():
    token = _reading_enclosing_values.set(True)
    try:
        yield
    finally:
        _reading_enclosing_values.reset(token)


def _native_capture_state(
    entering: bool,
    *,
    compiling: bool = False,
    exporting: bool = False,
    disabled: bool = False,
) -> bool:
    """Update the thread-local state owned by the native graph runtime."""

    try:
        import tensorplay

        native = getattr(getattr(tensorplay, "_C", None), "_stax", None)
        operation = getattr(
            native,
            "capture_state_enter" if entering else "capture_state_exit",
            None,
        )
    except (AttributeError, ImportError):
        return False
    if operation is None:
        return False
    operation(compiling, exporting, disabled)
    return True


@contextmanager
def compiler_context(*, require_native: bool = False) -> Any:
    token = _compiling.set(True)
    native_entered = False
    try:
        native_entered = _native_capture_state(True, compiling=True)
        if require_native and not native_entered:
            raise GraphCaptureError(
                "TensorPlay native capture state is unavailable"
            )
        yield
    finally:
        if native_entered:
            _native_capture_state(False, compiling=True)
        _compiling.reset(token)


def _map_arg(value: Any, fn: Callable[[Any], Any]) -> Any:
    from .node import Node
    from .proxy import Proxy

    if isinstance(value, (Node, Proxy)):
        return fn(value)
    if isinstance(value, tuple):
        return tuple(_map_arg(item, fn) for item in value)
    if isinstance(value, list):
        return [_map_arg(item, fn) for item in value]
    if isinstance(value, dict):
        return {key: _map_arg(item, fn) for key, item in value.items()}
    if isinstance(value, slice):
        return slice(
            _map_arg(value.start, fn),
            _map_arg(value.stop, fn),
            _map_arg(value.step, fn),
        )
    return value


def _iter_proxies(value: Any) -> Iterable["Proxy"]:
    from .proxy import Proxy

    if isinstance(value, Proxy):
        yield value
        return
    if isinstance(value, (tuple, list)):
        for item in value:
            yield from _iter_proxies(item)
        return
    if isinstance(value, dict):
        for item in value.values():
            yield from _iter_proxies(item)
        return
    if isinstance(value, slice):
        yield from _iter_proxies(value.start)
        yield from _iter_proxies(value.stop)
        yield from _iter_proxies(value.step)


_TRACE_DEPTH = 0


def capturing() -> bool:
    return _TRACE_DEPTH > 0


def capture_call(
    target: Callable[..., Any], args: tuple[Any, ...], kwargs: dict[str, Any]
) -> Optional["Proxy"]:
    found = False
    for value in args:
        for _proxy in _iter_proxies(value):
            found = True
            break
        if found:
            break
    if not found and kwargs:
        for value in kwargs.values():
            for _proxy in _iter_proxies(value):
                found = True
                break
            if found:
                break
    # Factory operations with no proxy argument are recorded anyway.  A
    # stochastic factory (rand, randn, randint, ...) samples at call time, so
    # freezing its eager result into a graph constant would make every
    # compiled call reuse one random draw.  An allocating factory (empty,
    # zeros, arange, ...) creates fresh storage, so freezing its eager result
    # would pin the compiled call to one captured buffer -- which breaks
    # in-place fills (``empty().uniform_()`` must re-sample per call) and
    # leaves runtime shapes computed from metadata baked at capture time.
    # During compile capture both kinds become graph nodes and run per call.
    if not found:
        name = getattr(target, "__name__", "")
        stochastic = name.startswith("rand") or name in {
            "bernoulli",
            "multinomial",
            "normal",
            "poisson",
            "exponential",
            "geometric",
            "cauchy",
            "log_normal",
        }
        factory = name in _FACTORY_NAMES
        if not (
            (stochastic or factory)
            and _compiling.get()
            and not _executing_sample.get()
        ):
            return None
    proxies = list(_iter_proxies(args))
    proxies.extend(_iter_proxies(kwargs))
    if _capture_disabled.get():
        raise GraphCaptureError("graph capture is disabled for this operation")
    if not proxies:
        tracer = get_active_tracer()
        if tracer is None:
            return None
    else:
        tracer = proxies[0].tracer
        if any(proxy.tracer is not tracer for proxy in proxies[1:]):
            raise GraphCaptureError("cannot combine proxies from different traces")
    return tracer.create_proxy("call_function", target, args, kwargs)


def _iter_nodes(value: Any) -> Iterable["Node"]:
    from .node import Node

    if isinstance(value, Node):
        yield value
        return
    if isinstance(value, (tuple, list)):
        for item in value:
            yield from _iter_nodes(item)
        return
    if isinstance(value, dict):
        for item in value.values():
            yield from _iter_nodes(item)
        return
    if isinstance(value, slice):
        yield from _iter_nodes(value.start)
        yield from _iter_nodes(value.stop)
        yield from _iter_nodes(value.step)


def _snake_case(name: str) -> str:
    s1 = re.sub("(.)([A-Z][a-z]+)", r"\1_\2", name)
    return re.sub("([a-z0-9])([A-Z])", r"\1_\2", s1).lower()


def gate_outcome(kind: str, sample: Any) -> Any:
    if kind == "iter":
        return ("iter",) + tuple(sample)
    item = sample.item() if hasattr(sample, "item") else sample
    if kind == "bool":
        return bool(item)
    if kind in ("int", "index"):
        return int(item)
    if kind == "float":
        return float(item)
    raise GraphCaptureError(f"unknown control-flow gate kind {kind!r}")


_SANITIZED_NAMES: Dict[str, str] = {}


def _sanitize_name(name: str) -> str:
    cached = _SANITIZED_NAMES.get(name)
    if cached is not None:
        return cached
    sanitized = re.sub(r"[^0-9a-zA-Z_]", "_", name)
    if not sanitized or sanitized[0].isdigit() or keyword.iskeyword(sanitized):
        sanitized = f"_{sanitized}"
    if len(_SANITIZED_NAMES) < 8192:
        _SANITIZED_NAMES[name] = sanitized
    return sanitized


_TARGET_STRINGS: Dict[Any, str] = {}


def _target_to_str(target: Any) -> str:
    try:
        cached = _TARGET_STRINGS.get(target)
    except TypeError:
        cached = None
    if cached is not None:
        return cached
    if isinstance(target, str):
        result = _snake_case(target.split(".")[-1])
    elif callable(target):
        atom = getattr(target, "__name__", None) or type(target).__name__
        result = _snake_case(str(atom))
    else:
        result = type(target).__name__
    if len(_TARGET_STRINGS) < 8192:
        try:
            _TARGET_STRINGS[target] = result
        except TypeError:
            pass
    return result


def _format_target(target: Any) -> str:
    name = getattr(target, "__name__", None)
    if isinstance(target, str):
        return target
    if callable(target) and name:
        module = getattr(target, "__module__", "") or ""
        if module == "builtins":
            return str(name)
        # The private accelerator module is presented under its public
        # facade so user-facing renderings share one spelling.
        if module == "_operator":
            module = "operator"
        if module:
            return f"{module}.{name}"
        return str(name)
    if name:
        return str(name)
    return repr(target)
Ask DeepWiki