stisiTT's picture
Add files using upload-large-folder tool
9aa90e0 verified
Raw History Blame Contribute Delete
33.1 kB
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import functools
import inspect
import weakref
from types import NoneType
from typing import TYPE_CHECKING, Any, overload
import torch
from loguru import logger
import ttnn
from models.tt_dit.utils import tensor
if TYPE_CHECKING:
from collections.abc import Callable, Sequence
from typing import ClassVar
_OMITTED = object()
class StateTensor:
"""A persistent device buffer refreshed in place under tracing.
A ttnn trace bakes its inputs' absolute addresses, so a traced input must keep the same buffer
and only update its contents; untraced, ``update`` rebinds to the new value.
"""
def __init__(self) -> None:
self._data: ttnn.Tensor | None = None
@property
def data(self) -> ttnn.Tensor | None:
return self._data
@property
def value(self) -> ttnn.Tensor | None:
return self._data
def update(
self,
value: torch.Tensor | ttnn.Tensor,
traced: bool,
dtype: ttnn.DataType = ttnn.bfloat16,
mesh_axes: list[int] | None = None,
device: ttnn.Device | None = None,
) -> None:
"""Update the state tensor with a new value.
If the tensor is used during tracing, the value is copied to (or used to initalize) the underlying tensor.
Args:
value: The new value to update the state tensor with.
traced: Whether the tensor is used during tracing.
dtype: The data type of the value.
mesh_axes: The mesh axes of the value.
device: The device of the value.
"""
if torch.is_tensor(value):
assert device is not None, "device must be provided if using torch tensor"
value = tensor.from_torch(value, device=device, mesh_axes=mesh_axes, dtype=dtype)
if self._data is None or not traced:
self._data = value
else:
ttnn.copy(value, self._data)
class Tracer:
"""Wrapper for capturing and executing a trace of a given function.
All inputs and outputs of the traced function must be ``ttnn.Tensor`` instances or plain
Python scalars (``int``, ``float``, ``str``, ``bool``, ``None``), optionally nested in
tuples, lists, or dicts.
Important caveats:
1. Tensors allocated after trace capture may be overwritten during trace execution.
Host tensors are not affected. Input tensors are copied before trace execution, so
they can safely be allocated on device if their content is not needed after execution.
2. The tracer returns the same output tensor objects every time; a subsequent call
overwrites previous results in place.
"""
_traces_live: ClassVar[dict[int, int]] = {}
@overload
def __init__(
self,
function: Callable[..., Any],
/,
*,
device: ttnn.MeshDevice,
prep_run: bool = True,
clone_prep_inputs: bool = True,
) -> None:
...
@overload
def __init__(
self,
function: Callable[..., Any],
/,
*,
devices: Sequence[ttnn.MeshDevice],
prep_run: bool = True,
clone_prep_inputs: bool = True,
) -> None:
...
def __init__(
self,
function: Callable[..., Any],
/,
*,
device: ttnn.MeshDevice | None = None,
devices: Sequence[ttnn.MeshDevice] | None = None,
prep_run: bool = True,
clone_prep_inputs: bool = True,
) -> None:
"""Initialize the tracer.
If the function modifies its input tensors in place, set ``clone_prep_inputs`` to ``True``
so that preparation runs operate on cloned inputs, leaving the originals intact for trace
capture.
Exactly one of ``device`` or ``devices`` must be provided. When ``devices`` is given, a
trace is captured and executed on each device simultaneously.
Args:
function: Function to be traced.
device: Single device on which to capture and execute the trace.
devices: Multiple devices on which to capture and execute the trace simultaneously.
prep_run: Whether to run the function once before capturing the trace.
clone_prep_inputs: Whether to clone tensor inputs for the preparation run.
"""
if (device is None) == (devices is None):
msg = "exactly one of 'device' or 'devices' must be provided"
raise ValueError(msg)
if devices is None:
assert device is not None
devices = (device,)
for d in devices:
if d.id() not in Tracer._traces_live:
Tracer._traces_live[d.id()] = 0
self._function = function
self._devices: tuple[ttnn.MeshDevice, ...] = tuple(devices)
self._prep_run = prep_run
self._clone_prep_inputs = clone_prep_inputs
self._args: tuple[Any, ...] = ()
self._kwargs: dict[str, Any] = {}
self._outputs: Any = None
self._trace_ids: tuple[ttnn.MeshTraceId, ...] | None = None
def __call__(
self,
*args: Any,
traced: bool = True,
tracer_cq_id: int = 0,
tracer_blocking_execution: bool = True,
tracer_execute_on_capture: bool = True,
**kwargs: Any,
) -> Any:
"""Capture or execute trace.
In traced mode, the first call captures the trace and subsequent calls execute it. The
first call's inputs initialize the stored input state; subsequent calls must pass the
same number of positional arguments and the same set of keyword argument names. Tensor
inputs are validated for shape/dtype/layout match; non-tensor inputs (scalars, ``None``)
must compare equal to the captured values. Host tensor inputs are automatically moved
to the tracer device.
Args:
traced: Whether to capture/execute the trace on this call. If ``False``, the wrapped
function is called directly without tracing.
tracer_cq_id: Command queue id.
tracer_blocking_execution: Whether ``ttnn.execute_trace`` should block.
tracer_execute_on_capture: Whether to execute the trace immediately after capturing it
on the first call. If ``False``, only the trace is captured and outputs are not
computed.
*args: Positional inputs to pass to the wrapped function.
**kwargs: Named inputs to pass to the wrapped function.
Returns:
The outputs of the wrapped function.
Raises:
TypeError: If outputs have unsupported types, or if subsequent traced calls do not
match the captured arity/keyword set.
Any exception raised by the wrapped function during first invocation.
"""
if not traced:
if self._function is None:
msg = "untraced execution is not possible after the captured function has been released"
raise RuntimeError(msg)
if self._trace_ids is not None:
msg = "untraced execution is not allowed after the trace has been captured"
raise RuntimeError(msg)
self._args = args
self._kwargs = kwargs
return self._function(*args, **kwargs)
if tracer_blocking_execution and len(self._devices) != 1:
tracer_blocking_execution = False
logger.warning("blocking execution is not supported with multiple devices")
if self._trace_ids is None:
self._capture(
args,
kwargs,
cq_id=tracer_cq_id,
blocking_execution=tracer_blocking_execution,
execute=tracer_execute_on_capture,
)
else:
self._execute(args, kwargs, cq_id=tracer_cq_id, blocking=tracer_blocking_execution)
return self._outputs
def _capture(
self,
args: tuple[Any, ...],
kwargs: dict[str, Any],
*,
cq_id: int,
blocking_execution: bool,
execute: bool,
) -> None:
if self._function is None:
msg = "tracer cannot be reused after the trace was released"
raise RuntimeError(msg)
args = _tree_map(_verify_value, args, path_label="args")
kwargs = _tree_map(_verify_value, kwargs, path_label="kwargs")
self._args = _tree_map(self._tensor_to_device, args, path_label="args")
self._kwargs = _tree_map(self._tensor_to_device, kwargs, path_label="kwargs")
if self._prep_run:
if self._clone_prep_inputs:
prep_args = _tree_map(_clone_tensor, self._args, path_label="args")
prep_kwargs = _tree_map(_clone_tensor, self._kwargs, path_label="kwargs")
else:
prep_args = self._args
prep_kwargs = self._kwargs
self._function(*prep_args, **prep_kwargs)
del prep_args, prep_kwargs
# capture trace
logger.debug("capturing trace...")
trace_ids = tuple(ttnn.begin_trace_capture(d, cq_id=cq_id) for d in self._devices)
try:
try:
outputs = self._function(*self._args, **self._kwargs)
finally:
for d, trace_id in zip(self._devices, trace_ids, strict=True):
ttnn.end_trace_capture(d, trace_id, cq_id=cq_id)
outputs = _tree_map(_verify_value, outputs, path_label="outputs")
except Exception:
for d, trace_id in zip(self._devices, trace_ids, strict=True):
ttnn.release_trace(d, trace_id)
raise
for d in self._devices:
Tracer._traces_live[d.id()] += 1
self._trace_ids = trace_ids
self._outputs = outputs
if execute:
# Trace capture records commands but does not execute them. Execute the trace to
# actually compute outputs.
for d, trace_id in zip(self._devices, trace_ids, strict=True):
ttnn.execute_trace(d, trace_id, cq_id=cq_id, blocking=blocking_execution)
def _execute(
self,
args: tuple[Any, ...],
kwargs: dict[str, Any],
*,
cq_id: int,
blocking: bool,
) -> None:
trace_ids = self._trace_ids
assert trace_ids is not None
if len(args) != len(self._args):
msg = f"expected {len(self._args)} positional args, got {len(args)}"
raise TypeError(msg)
if kwargs.keys() != self._kwargs.keys():
msg = f"expected kwargs {sorted(self._kwargs)}, got {sorted(kwargs)}"
raise TypeError(msg)
_tree_map(self._update_input, self._args, args, path_label="args")
for name, new in kwargs.items():
_tree_map(self._update_input, self._kwargs[name], new, path_label=f'kwargs["{name}"]')
for d, trace_id in zip(self._devices, trace_ids, strict=True):
ttnn.execute_trace(d, trace_id, cq_id=cq_id, blocking=blocking)
@property
def trace_captured(self) -> bool:
"""Whether a trace has been captured and is ready for execution."""
return self._trace_ids is not None
@property
def inputs(self) -> dict[int | str, Any]:
"""Stored inputs from the most recent call (empty before the first call).
Keyed by parameter name (``str``) for keyword inputs, or by index (``int``) for positional
inputs. After the trace has been captured, tensor entries are the trace's input buffers, so
the reference is a safe handle to the latest value of that input, which is useful when the
caller needs to consume an input after trace execution. In untraced mode, tensor entries are
simply the most recently passed references.
"""
return {**dict(enumerate(self._args)), **self._kwargs}
def release_trace(self) -> None:
"""Release the captured trace and clear inputs and outputs."""
trace_ids = self._trace_ids
if trace_ids is None:
return
self._trace_ids = None
self._args = ()
self._kwargs = {}
self._outputs = None
for d, trace_id in zip(self._devices, trace_ids, strict=True):
Tracer._traces_live[d.id()] -= 1
ttnn.release_trace(d, trace_id)
def release_function(self) -> None:
"""Drop the reference to the wrapped function.
This allows resources held by the function to be garbage collected.
"""
self._function = None
def _tensor_to_device(self, value: Any, *, path_label: str) -> Any:
if not isinstance(value, ttnn.Tensor):
return value
if value.device() is None:
if len(self._devices) != 1:
msg = (
f"input '{path_label}' is a host tensor; host tensors are not supported "
"during capture when the tracer manages multiple devices"
)
raise ValueError(msg)
return value.to(self._devices[0])
if any(value.device() == d for d in self._devices):
return value
msg = f"input '{path_label}' device {value.device()} does not match any tracer device"
raise ValueError(msg)
def _update_input(self, prev: Any, new: Any, *, path_label: str) -> None:
"""Copy ``new`` into the slot ``prev`` in place."""
if type(new) is not type(prev):
msg = f"input '{path_label}' type {type(new)} does not match the initial type {type(prev)}"
raise TypeError(msg)
if isinstance(new, ttnn.Tensor):
if new.shape != prev.shape or new.dtype != prev.dtype or new.layout != prev.layout:
msg = f"input '{path_label}' tensor properties do not match the initial value"
raise ValueError(msg)
if new.device() is None:
ttnn.copy_host_to_device_tensor(new, prev)
else:
if new.device() != prev.device():
msg = f"input '{path_label}' tensor device does not match the initial device"
raise ValueError(msg)
if new.buffer_address() != prev.buffer_address():
ttnn.copy(new, prev)
elif new != prev:
msg = f"input '{path_label}' does not match the initial value"
raise ValueError(msg)
_TRACER_VALID_INPUT_TYPES = (ttnn.Tensor, int, float, str, bool, NoneType)
_MESH_DEVICE_PARAM_NAMES: tuple[str, ...] = ("mesh_device", "device")
"""Parameter names that ``traced_function`` recognises as the Tracer's mesh device, in priority order."""
def _verify_value(value: Any, *, path_label: str) -> Any:
if not isinstance(value, _TRACER_VALID_INPUT_TYPES):
msg = f"value '{path_label}' has unsupported type {type(value)}"
raise TypeError(msg)
return value
def _is_tracer_valid_value(value: Any) -> bool:
"""Return True if ``value`` can be passed through to a ``Tracer`` as an input.
Matches ``_TRACER_VALID_INPUT_TYPES`` plus any nesting of those inside ``tuple``,
``list``, or ``dict`` (the containers supported by ``Tracer._tree_map``). Used at
call time by ``traced_function`` to classify each argument as either a tracer input
or a bindable config value.
"""
if isinstance(value, _TRACER_VALID_INPUT_TYPES):
return True
if isinstance(value, (tuple, list)):
return all(_is_tracer_valid_value(v) for v in value)
if isinstance(value, dict):
return all(isinstance(k, _TRACER_VALID_INPUT_TYPES) and _is_tracer_valid_value(v) for k, v in value.items())
return False
def _clone_tensor(value: Any, *, path_label: str) -> Any:
"""Clone a tensor, passing through non-tensor values unchanged."""
del path_label
return ttnn.clone(value) if isinstance(value, ttnn.Tensor) else value
def _tree_map(f: Callable[..., Any], x: Any, /, *xs: Any, path_label: str) -> Any:
"""Apply a function to leaves of nested data structures.
Recursively traverses nested structures (tuples, lists, dicts) and applies
the given function to corresponding leaf elements across all input structures.
Args:
f: A callable that takes N arguments, where N is the number of input
structures (1 + len(xs)). Applied to leaf elements.
x: The first nested data structure to traverse.
*xs: Additional nested data structures with the same shape as x.
path_label: String representing the current traversal path (used for error messages).
Returns:
A new nested structure with the same shape as the inputs, where each leaf has been
transformed by applying f to the corresponding leaves from all input structures.
Raises:
TypeError: If the input structures don't have matching types at corresponding positions.
ValueError: If tuples/lists have different lengths or dicts have different keys.
"""
if not isinstance(x, (tuple, list, dict)):
return f(x, *xs, path_label=path_label)
for y in xs:
if not isinstance(y, type(x)):
msg = f"types of '{path_label}' should be the same: {type(x)} != {type(y)}"
raise TypeError(msg)
if isinstance(x, tuple) and all(isinstance(y, tuple) for y in xs):
for y in xs:
if len(x) != len(y):
msg = f"tuple lengths of '{path_label}' should be the same: {len(x)} != {len(y)}"
raise ValueError(msg)
return tuple(
_tree_map(f, *elts, path_label=f"{path_label}[{i}]") for i, elts in enumerate(zip(x, *xs, strict=True))
)
if isinstance(x, list) and all(isinstance(y, list) for y in xs):
for y in xs:
if len(x) != len(y):
msg = f"list lengths of '{path_label}' should be the same: {len(x)} != {len(y)}"
raise ValueError(msg)
return [_tree_map(f, *elts, path_label=f"{path_label}[{i}]") for i, elts in enumerate(zip(x, *xs, strict=True))]
if isinstance(x, dict) and all(isinstance(y, dict) for y in xs):
for y in xs:
if x.keys() != y.keys():
msg = f"dict keys of '{path_label}' should be the same: {x.keys()} != {y.keys()}"
raise ValueError(msg)
return {key: _tree_map(f, *(d[key] for d in (x, *xs)), path_label=f'{path_label}["{key}"]') for key in x}
raise AssertionError # unreachable
_TRACER_CALL_KWARGS = frozenset(
name
for name, param in inspect.signature(Tracer.__call__).parameters.items()
if param.kind == inspect.Parameter.KEYWORD_ONLY
)
def traced_function(
_fn: Callable[..., Any] | None = None,
*,
device: ttnn.MeshDevice | Callable[..., ttnn.MeshDevice] | None = None,
inject_mesh_device: bool = False,
prep_run: bool = True,
clone_prep_inputs: bool = True,
) -> Any:
"""Decorator that adds optional tracing to any function or method via the ``Tracer`` class.
Can be applied with or without arguments::
# Method — device resolved lazily via callable from the bound context (self):
@traced_function(device=lambda self: self.mesh_device, clone_prep_inputs=False)
def my_method(self, x: ttnn.Tensor, scale: float) -> ttnn.Tensor: ...
# Standalone function — device and other non-tracer-valid args are classified
# dynamically on the first traced call. The parameter must be named `mesh_device`:
@traced_function(clone_prep_inputs=False)
def my_function(x, mesh_device, ccl_manager=None, pre_transfer_fn=None): ...
# Standalone function — `mesh_device` is auto-injected as a wrapper-only kwarg
# (like `traced`) and never forwarded to the wrapped function:
@traced_function(inject_mesh_device=True, clone_prep_inputs=False)
def my_pure_function(x: ttnn.Tensor, scale: float) -> ttnn.Tensor: ...
# Call without tracing (original function, no overhead):
result = my_function(x, mesh_device=md)
result = my_pure_function(x, scale=1.0)
result = model.my_method(x, scale=1.0)
# Call with tracing (captures on first call, replays on subsequent calls):
result = my_function(x, mesh_device=md, traced=True)
result = my_pure_function(x, scale=1.0, mesh_device=md, traced=True)
result = model.my_method(x, scale=1.0, traced=True)
# With optional Tracer call-time kwargs:
result = my_function(x, mesh_device=md, traced=True, tracer_cq_id=1, tracer_blocking_execution=False)
The decorated callable gains a ``traced`` keyword argument at call time. When
``traced=False`` (the default) the original function is called directly. When
``traced=True`` a ``Tracer`` is lazily created on the first call and subsequent
calls execute the captured trace.
Three modes for supplying the Tracer's device:
1. **Context-bound** — ``device=<callable>`` (or literal). The first
positional argument is treated as a bindable context (typical for
methods: ``self``), bound away via ``functools.partial``, and — when
``device`` is callable — passed through it to resolve the device. One
``Tracer`` per unique context in a ``WeakKeyDictionary``.
2. **Auto-discovered** — ``device=None`` (omitted). The function must
declare a parameter literally named ``mesh_device`` or ``device``
(first match in that priority order). On the first traced call,
each argument is classified at runtime: any value that isn't a
``Tracer``-valid input (``ttnn.Tensor``, scalar, ``None``, or a
nested ``tuple``/``list``/``dict`` of those) is bound into a
``functools.partial``. The bind set is frozen after the first call
and reused. The matched mesh parameter drives the ``Tracer``'s
device. One ``Tracer`` is cached per unique tuple of bound-value
identities.
3. **Injected** — ``inject_mesh_device=True``. The wrapper accepts
either ``mesh_device=`` or ``device=`` as an auto-added kwarg
(analogous to ``traced=``), consumes it, and never forwards it to
the wrapped function. Exactly one of the two may be supplied per
call. The function itself must *not* declare a ``mesh_device`` or
``device`` parameter. Other non-tracer-valid args from the wrapped
function's signature are still classified and bound on the first
traced call, exactly as in mode (2).
Decoration-time validation:
- ``device=`` and ``inject_mesh_device=True`` are mutually exclusive.
- ``inject_mesh_device=True`` requires the wrapped function *not* to
declare a ``mesh_device`` or ``device`` parameter (to avoid ambiguity).
- Omitting ``device=`` with ``inject_mesh_device=False`` on a function
that declares neither ``mesh_device`` nor ``device`` raises ``ValueError``.
Tracer call-time kwargs (``tracer_cq_id``, ``tracer_blocking_execution``,
``tracer_execute_on_capture``) are forwarded to the ``Tracer`` when tracing
and stripped before calling the original function in the untraced path.
A wrapper-only ``tracer_trace_key`` kwarg selects which trace to capture/replay for a
given context, so one instance can hold several traces (e.g. the same step method captured
at two input shapes). It defaults to ``None`` and is never forwarded to the function.
Args:
_fn: The function to wrap when used without parentheses (``@traced_function``).
device: Device for tracing. For methods, pass a callable
(e.g. ``lambda self: self.mesh_device``) so it can be resolved lazily
from the bound context. Omit entirely for standalone functions that
declare a ``mesh_device`` parameter or set ``inject_mesh_device=True``.
inject_mesh_device: If ``True``, the wrapper accepts ``mesh_device=`` as an
auto-added kwarg and consumes it without forwarding to the wrapped
function. Mutually exclusive with ``device=``.
prep_run: Forwarded to ``Tracer.__init__``.
clone_prep_inputs: Forwarded to ``Tracer.__init__``.
"""
def _resolve_device(context: Any) -> ttnn.MeshDevice:
if device is None:
msg = (
"device= must be provided to @traced_function. "
"Pass a ttnn.MeshDevice directly, or a callable that accepts the bound context "
"and returns one (e.g. device=lambda self: self.mesh_device)."
)
raise ValueError(msg)
return device(context) if callable(device) else device
if inject_mesh_device and device is not None:
msg = "@traced_function: inject_mesh_device=True is mutually exclusive with device="
raise ValueError(msg)
def decorator(fn: Callable[..., Any]) -> Callable[..., Any]:
sig = inspect.signature(fn)
# First matching parameter name (in priority order) is what we'll look up at call time.
sig_mesh_param = next((n for n in _MESH_DEVICE_PARAM_NAMES if n in sig.parameters), None)
if inject_mesh_device and sig_mesh_param is not None:
msg = (
f"@traced_function: {fn.__qualname__} declares a {sig_mesh_param!r} parameter but "
"inject_mesh_device=True would also add one. Rename the parameter, or drop "
"inject_mesh_device=True to use signature-based auto-discovery instead."
)
raise ValueError(msg)
if device is None and not inject_mesh_device and sig_mesh_param is None:
msg = (
f"@traced_function: {fn.__qualname__} was decorated without device= and has no "
f"parameter named any of {_MESH_DEVICE_PARAM_NAMES}. Provide "
"device=<callable returning a ttnn.MeshDevice> (e.g. lambda self: self.mesh_device) "
"for methods, declare a 'mesh_device' (or 'device') parameter on the function for "
"auto-discovery, or pass inject_mesh_device=True to have the wrapper accept one as "
"an auto-added kwarg."
)
raise ValueError(msg)
_tracers: weakref.WeakKeyDictionary[Any, Tracer] = weakref.WeakKeyDictionary()
# Set only when ``tracer_trace_key`` is used: one context holding several traces keyed by it
# — e.g. one method captured at multiple input shapes. Keeps ``_tracers`` (the default,
# one-per-context store) unchanged for callers that don't key.
_tracers_keyed: weakref.WeakKeyDictionary[Any, dict[Any, Tracer]] = weakref.WeakKeyDictionary()
_tracers_auto: dict[tuple[int, ...], Tracer] = {}
_auto_bind_names: tuple[str, ...] | None = None # frozen on the first traced call
def _needs_new_tracer(tracer: Tracer | None) -> bool:
return tracer is None or (not tracer.trace_captured and tracer._function is None)
@functools.wraps(fn)
def wrapper(*args: Any, traced: bool = False, **kwargs: Any) -> Any:
nonlocal _auto_bind_names
# Wrapper-only: selects the trace for this context; never forwarded to fn or the Tracer.
trace_key = kwargs.pop("tracer_trace_key", None)
# When injecting, accept either name but only one at a time. The kwarg is consumed
# by the wrapper in every path — it is never forwarded to the wrapped function.
injected_mesh_device: Any = _OMITTED
if inject_mesh_device:
supplied = [(n, kwargs.pop(n)) for n in _MESH_DEVICE_PARAM_NAMES if n in kwargs]
if len(supplied) > 1:
msg = (
f"@traced_function: pass only one of {_MESH_DEVICE_PARAM_NAMES} as the "
f"injected mesh-device kwarg; got "
f"{ {n: v for n, v in supplied} !r}"
)
raise TypeError(msg)
if supplied:
injected_mesh_device = supplied[0][1]
if not traced:
for k in _TRACER_CALL_KWARGS:
kwargs.pop(k, None)
return fn(*args, **kwargs)
# Auto-discovery path: classify args at runtime on the first traced call.
# Everything that isn't a Tracer-valid value is bound into the partial.
if device is None:
tracer_call_kwargs = {k: kwargs.pop(k) for k in list(kwargs) if k in _TRACER_CALL_KWARGS}
bound = sig.bind(*args, **kwargs)
bound.apply_defaults()
if _auto_bind_names is None:
_auto_bind_names = tuple(
name
for name in sig.parameters
if name in bound.arguments and not _is_tracer_valid_value(bound.arguments[name])
)
if not inject_mesh_device and sig_mesh_param not in _auto_bind_names:
md = bound.arguments.get(sig_mesh_param, None) if sig_mesh_param else None
msg = (
f"@traced_function: {fn.__qualname__} expected {sig_mesh_param!r} to "
f"be a non-tracer-valid value (e.g. ttnn.MeshDevice), got "
f"{type(md).__name__}={md!r}"
)
raise TypeError(msg)
bind_values = {n: bound.arguments.pop(n) for n in _auto_bind_names}
if inject_mesh_device:
if injected_mesh_device is _OMITTED:
msg = (
f"@traced_function: {fn.__qualname__} was called with traced=True "
f"but none of {_MESH_DEVICE_PARAM_NAMES} were supplied as kwargs "
"(required when inject_mesh_device=True)."
)
raise TypeError(msg)
mesh_device_val = injected_mesh_device
key = (id(mesh_device_val), *(id(v) for v in bind_values.values()))
else:
mesh_device_val = bind_values[sig_mesh_param]
key = tuple(id(v) for v in bind_values.values())
t = _tracers_auto.get(key)
if _needs_new_tracer(t):
_tracers_auto[key] = Tracer(
functools.partial(fn, **bind_values),
device=mesh_device_val,
prep_run=prep_run,
clone_prep_inputs=clone_prep_inputs,
)
return _tracers_auto[key](*bound.args, **bound.kwargs, **tracer_call_kwargs)
# Context-bound path: first arg is a non-tracer-valid context (e.g. self); bind it away.
# Default (no trace_key) is one Tracer per instance in _tracers; a trace_key selects
# among several per instance in _tracers_keyed.
if args and not isinstance(args[0], _TRACER_VALID_INPUT_TYPES):
context, rest = args[0], args[1:]
if trace_key is None:
store, store_key = _tracers, context
else:
store, store_key = _tracers_keyed.setdefault(context, {}), trace_key
if _needs_new_tracer(store.get(store_key)):
store[store_key] = Tracer(
functools.partial(fn, context),
device=_resolve_device(context),
prep_run=prep_run,
clone_prep_inputs=clone_prep_inputs,
)
return store[store_key](*rest, **kwargs)
# device= was supplied but the first positional argument is a valid Tracer input
# (or absent), so there's no context to bind. Standalone functions should omit
# device= and declare a 'mesh_device' parameter for auto-discovery instead.
msg = (
f"@traced_function: {fn.__qualname__} was called with traced=True, but device= "
"was provided at decoration time and the first positional argument is not a "
"bindable context (it is a tracer-valid type). For standalone functions, omit "
"device= and declare a 'mesh_device' parameter to use auto-discovery; for "
"methods, ensure self is passed positionally."
)
raise TypeError(msg)
wrapper._tracers = _tracers # type: ignore[attr-defined]
wrapper._tracers_keyed = _tracers_keyed # type: ignore[attr-defined]
wrapper._tracers_auto = _tracers_auto # type: ignore[attr-defined]
return wrapper
if _fn is not None:
return decorator(_fn)
return decorator