Download code/tt_diffusion_planner/ttaw/trace.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 70.1 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/trace.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/ttaw/trace.py
-
curl -L -o trace.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/trace.py
70.1 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """C02: ``TraceRunner`` -- every device stage of a port runs inside metal traces. | |
| What it enforces (TT_PLATFORM.md section 3, REFERENCE_PATTERNS.md section 1.4): | |
| 1. **Persistent I/O before any capture.** Inputs, RT-dev parameters, state buffers and (optionally) output buffers | |
| are allocated (and given defined initial values: never uninitialised index buffers) when they are added, i.e. | |
| before the first capture. A variant returns either the tensors its last ops produce (allocated during capture, | |
| valid while the trace lives) or persistent outputs written with ``ctx.write_output`` (stable addresses shared by | |
| several variants, never overwritten by another variant's replay). | |
| 2. **Warm-up, then capture.** Every variant runs eagerly ``warmup_runs`` times (kernel JIT, program cache, prepared | |
| conv weights) before *any* trace is captured. Adding a variant after a capture releases all traces, warms the new | |
| one and recaptures everything, because a warm-up after a capture could allocate into a trace's freed | |
| intermediates. | |
| 3. **Strict capture.** ``device.set_program_cache_misses_allowed(False)`` during capture (a miss raises with the op | |
| name instead of aborting with "Writes are not supported during trace capture"); ``end_trace_capture`` runs in a | |
| ``finally`` block and a failed capture is released (an open capture left the process spinning in close_device: | |
| RP section 1.4); the program-cache entry count must not change during capture. | |
| 4. **Variants.** Several traces keyed by name (shape buckets, modes, segments), chosen by the caller per run. | |
| 5. **RT-dev parameters.** Per-frame values that must not change the program (thresholds, timesteps, poses) live | |
| in persistent device tensors refreshed before ``execute_trace``; unchanged values are not re-uploaded. | |
| 6. **1CQ / 2CQ.** With 2 CQs inputs are uploaded on CQ1 and ordered with events, exactly like | |
| ``common/tools/check_dispatch.py`` (CQ1 waits for the last trace, uploads, records; CQ0 waits, replays, records). | |
| ``stage_inputs=True`` uses the ``tt_cnn`` executor pattern instead: CQ1 writes a DRAM staging copy while the | |
| previous trace still runs, and an eager copy on CQ0 moves it into the trace input before the replay. Every CQ0 | |
| write into a trace input outside the protocol (:meth:`TraceRunner.write_input`) re-records the event CQ1 waits for. | |
| 7. **State inside the trace.** ``ctx.write_state(name, value)`` is ``ttnn.copy(value, buffer)`` into a persistent | |
| buffer (FLOAT32 / UINT32 / BFLOAT16 / ... all supported), or no op at all when ``value`` was computed straight | |
| into ``ctx.write_target(name)`` (``output_tensor=``). ``pingpong=True`` states use two buffers and two traces | |
| per variant (phase 0 reads A writes B, phase 1 reads B writes A); the phase flips after every run of a variant | |
| that writes the ping-pong states. Per-stream state (PLAN.md D16): ``add_state(..., banks=n)`` + | |
| :meth:`TraceRunner.save_state` / :meth:`TraceRunner.load_state`, with the policy in :class:`StreamBanks`. | |
| 8. **Readback.** One packed output (:func:`pack_outputs`) gives one D2H; reads go into preallocated host tensors, | |
| on CQ0 or (segmented pipelines, 2CQ) on CQ1 after a host-side event wait. | |
| 9. **Alloc tracking.** With ``TT_METAL_TRACE_ALLOC_TRACKING=1`` set before ``import ttnn``, ``ttnn.execute_trace`` | |
| refuses to replay over live unsafe buffers; the runner acknowledges the outputs of traces captured after the | |
| first one (they may be overwritten by an earlier trace's replay: read outputs before running another variant). | |
| Example:: | |
| runner = TraceRunner(device, num_command_queues=2) | |
| runner.add_input("x", shape=(1, 1, 64, 64), dtype="bfloat16", layout=ttnn.TILE_LAYOUT) | |
| runner.add_param("scale", 1.0) # fp32 [1,1,1,1], TILE | |
| runner.add_variant("default", lambda ctx: ttnn.relu(ttnn.multiply(ctx["x"], ctx["scale"]))) | |
| runner.capture() # warm-up, then capture | |
| y = runner("default", inputs={"x": x_np}, params={"scale": 0.5}) # upload, replay, read -> numpy | |
| """ | |
| from __future__ import annotations | |
| import contextlib | |
| import math | |
| import time | |
| from dataclasses import dataclass, field | |
| from typing import Any, Callable, Dict, Iterator, List, Mapping, Optional, Sequence, Tuple | |
| import numpy as np | |
| from .io import InputError | |
| from .tensors import TILE, dtype_name, round_up, to_host_tensor, to_numpy, ttnn_dtype | |
| __all__ = [ | |
| "CQ_COMPUTE", | |
| "CQ_INPUT", | |
| "PackEntry", | |
| "PackLayout", | |
| "Packed", | |
| "pack_outputs", | |
| "SINGLE_ROW_MAX_ELEMS", | |
| "PACK_ROW_ELEMS", | |
| "TraceContext", | |
| "TraceRunner", | |
| "StreamBanks", | |
| "alloc_tracking_enabled", | |
| ] | |
| CQ_COMPUTE = 0 # programs, traces and (by default) readback | |
| CQ_INPUT = 1 # host -> device uploads when the device has 2 command queues | |
| def alloc_tracking_enabled() -> bool: | |
| """True when ``TT_METAL_TRACE_ALLOC_TRACKING=1`` was set before ttnn was imported (read from tt-metal).""" | |
| try: | |
| from ttnn.tools import trace_allocation_tracker as tracker | |
| except ImportError: | |
| return False | |
| return bool(getattr(tracker, "TRACE_ALLOC_TRACKING", False)) | |
| # ----------------------------------------------------------------------------------------------- packing | |
| # Packed layouts (PORT_LOG Q9 of the YOLOX port). A ROW_MAJOR tensor is stored page by page, one page per row, and | |
| # the RM reshape / concat programs stage whole pages in L1 (reshape_rm_program_factory.cpp: 2 x the destination page | |
| # per kernel copy when the pages are not 16-byte aligned), so a single-row pack of more than ~0.6 MB fails with | |
| # "RM reshape dest staging does not fit in L1". Above SINGLE_ROW_MAX_ELEMS the pack is a [1, 1, rows, R] tensor of | |
| # R = PACK_ROW_ELEMS elements per row: every RM page the packing programs touch stays below ROW_PAGE_MAX_BYTES, | |
| # whatever the size of the outputs (tens of MB), and the readback is still ONE device-to-host copy. | |
| SINGLE_ROW_MAX_ELEMS = 131072 # one [1, 1, 1, total] row up to 512 KiB of float32 (the 0.1.0 - 0.14.0 layout) | |
| PACK_ROW_ELEMS = 1024 # elements per row of the multi-row layout (4 KiB float32 pages) | |
| FLAT_MAX_BYTES = 64 << 10 # multi-row: tensors up to this size are flattened to one row, then cut into rows | |
| ROW_PAGE_MAX_BYTES = 128 << 10 # multi-row: largest RM page (last dim x element size) a packed tensor may have | |
| _ELEM_BYTES = {"float32": 4, "uint32": 4, "int32": 4, "bfloat16": 2, "uint16": 2, "uint8": 1} | |
| class PackEntry: | |
| """One tensor inside a packed readback: elements ``[offset, offset + numel)`` reshaped to ``shape``. | |
| ``pitch > 0`` (multi-row layout only): the tensor's rows (its last dim, ``shape[-1]`` elements) | |
| are stored ``pitch`` elements apart, zero-padded, i.e. elements ``[offset, offset + numel // shape[-1] * pitch)`` | |
| viewed as ``(rows, pitch)`` hold it in their first ``shape[-1]`` columns. ``0``: contiguous.""" | |
| name: str | |
| offset: int | |
| numel: int | |
| shape: Tuple[int, ...] | |
| pitch: int = 0 | |
| def view(self, flat: np.ndarray) -> np.ndarray: | |
| """This tensor inside the flat readback ``flat`` (a view, strided when ``pitch`` is set).""" | |
| if not self.pitch: | |
| return flat[self.offset:self.offset + self.numel].reshape(self.shape) | |
| cols = int(self.shape[-1]) if self.shape else 1 | |
| rows = self.numel // cols | |
| return flat[self.offset:self.offset + rows * self.pitch].reshape(rows, self.pitch)[:, :cols].reshape( | |
| self.shape) | |
| class PackLayout: | |
| """Host-side description of a packed output (built at capture time, applied at every read). | |
| ``entries`` index the packed tensor's elements in row-major order, so :meth:`unpack` does not depend on the | |
| device shape ``(1, 1, rows, row_elems)``: one row of ``total`` elements (``rows == 1``) or rows of | |
| ``row_elems`` elements, every tensor starting on a row boundary (multi-row layout; an entry with a ``pitch`` | |
| stores its rows zero-padded to that many elements, see :class:`PackEntry`).""" | |
| entries: Tuple[PackEntry, ...] | |
| total: int | |
| rows: int = 1 | |
| row_elems: int = 0 # 0: one row of ``total`` elements | |
| def shape(self) -> Tuple[int, int, int, int]: | |
| """Shape of the packed device tensor.""" | |
| return (1, 1, self.rows, self.row_elems or self.total) | |
| def unpack(self, flat: Any) -> Dict[str, np.ndarray]: | |
| """Flat readback (any shape with ``total`` elements) -> ``{name: array}`` (views, no copies; the view of an | |
| entry with a ``pitch`` is strided, not C-contiguous).""" | |
| a = np.asarray(flat).reshape(-1) | |
| if a.size != self.total: | |
| raise ValueError(f"packed readback has {a.size} elements, layout expects {self.total}") | |
| return {e.name: e.view(a) for e in self.entries} | |
| class Packed: | |
| """A packed device tensor plus its layout; return it from a variant function to get ``{name: array}`` back.""" | |
| tensor: Any | |
| layout: PackLayout | |
| def _elem_bytes(dtype: Any) -> int: | |
| return _ELEM_BYTES.get(dtype_name(dtype), 4) | |
| def _row_major(ttnn: Any, t: Any, dtype: Any) -> Any: | |
| if t.dtype != dtype: | |
| t = ttnn.typecast(t, dtype) | |
| if t.layout != ttnn.ROW_MAJOR_LAYOUT: | |
| t = ttnn.to_layout(t, ttnn.ROW_MAJOR_LAYOUT) | |
| return t | |
| def _flat_rows(ttnn: Any, t: Any, shape: Tuple[int, ...], numel: int, row_elems: int) -> Tuple[Any, int]: | |
| """A small RM tensor -> one row -> zero-padded to whole rows -> ``[1, 1, k, row_elems]``.""" | |
| padded = round_up(numel, row_elems) | |
| if shape != (1, 1, 1, numel): | |
| t = ttnn.reshape(t, (1, 1, 1, numel)) | |
| if padded != numel: | |
| t = ttnn.pad(t, [(0, 0), (0, 0), (0, 0), (0, padded - numel)], 0.0) | |
| if padded != row_elems: | |
| t = ttnn.reshape(t, (1, 1, padded // row_elems, row_elems)) | |
| return t, padded | |
| def _fill_rows(rows: int, cols: int, row_elems: int) -> int: | |
| """Rows of ``cols`` elements a ``[rows, cols]`` tensor is zero-padded to so that it fills whole packed rows.""" | |
| return round_up(rows, row_elems // math.gcd(cols, row_elems)) | |
| def _row_pitches(cols: int, row_elems: int, elem: int) -> List[int]: | |
| """Row pitches > ``cols`` worth trying: ``cols`` rounded up to every power of two dividing ``row_elems`` (rows of | |
| that pitch fill a packed row every ``row_elems / gcd`` rows) and to ``row_elems`` (every row fills whole packed | |
| rows), within the page budget.""" | |
| steps, m = {row_elems}, 2 | |
| while row_elems % m == 0: | |
| steps.add(m) | |
| m *= 2 | |
| return sorted({p for p in (round_up(cols, m) for m in steps) if p != cols and p * elem <= ROW_PAGE_MAX_BYTES}) | |
| def _pack_rows_segment(ttnn: Any, name: str, t: Any, shape: Tuple[int, ...], numel: int, row_elems: int, | |
| elem: int) -> Tuple[Any, int, int]: | |
| """One ROW_MAJOR tensor -> ``[1, 1, k, row_elems]`` holding its elements in order, zero-padded to whole rows. | |
| Returns ``(segment, k * row_elems, pitch)`` (``pitch``: see :class:`PackEntry`). No RM page above | |
| ``ROW_PAGE_MAX_BYTES`` is created or reshaped.""" | |
| cols = int(shape[-1]) if shape else 1 | |
| rows = numel // cols | |
| if cols * elem <= ROW_PAGE_MAX_BYTES: | |
| # [rows, cols] is a view of the RM tensor; rows * cols fills whole packed rows when rows % unit == 0 | |
| rows_p = _fill_rows(rows, cols, row_elems) | |
| if rows_p != rows and numel * elem <= FLAT_MAX_BYTES: | |
| return (*_flat_rows(ttnn, t, shape, numel, row_elems), 0) # small: pad < row_elems elements, not rows | |
| if rows_p != rows: | |
| # Appending zero rows costs up to row_elems / gcd(cols, row_elems) - 1 rows: 80 MB for a [1, 1, 2, 20001] | |
| # fp32 tensor. Take one flat row (when it fits the page budget) or rows zero-padded to a pitch instead | |
| # when that is smaller by more than 1/8 of the tensor (typical shapes keep the contiguous layout). | |
| options = [(p * _fill_rows(rows, p, row_elems), p) for p in _row_pitches(cols, row_elems, elem)] | |
| if round_up(numel, row_elems) * elem <= ROW_PAGE_MAX_BYTES: | |
| options.append((round_up(numel, row_elems), 0)) | |
| if options and rows_p * cols - min(options)[0] > numel // 8: | |
| size, pitch = min(options) # ties: the flat row, then the narrower pitch | |
| if not pitch: | |
| return (*_flat_rows(ttnn, t, shape, numel, row_elems), 0) | |
| prows = size // pitch | |
| if shape != (1, 1, rows, cols): | |
| t = ttnn.reshape(t, (1, 1, rows, cols)) | |
| t = ttnn.pad(t, [(0, 0), (0, 0), (0, prows - rows), (0, pitch - cols)], 0.0) | |
| if pitch != row_elems: | |
| t = ttnn.reshape(t, (1, 1, size // row_elems, row_elems)) | |
| return t, size, pitch | |
| # zero rows appended, then one RM reshape into rows of row_elems (source pages of cols elements, destination | |
| # pages of row_elems elements: both small) | |
| if shape != (1, 1, rows, cols): | |
| t = ttnn.reshape(t, (1, 1, rows, cols)) | |
| if rows_p != rows: | |
| t = ttnn.pad(t, [(0, 0), (0, 0), (0, rows_p - rows), (0, 0)], 0.0) | |
| if cols != row_elems: | |
| t = ttnn.reshape(t, (1, 1, rows_p * cols // row_elems, row_elems)) | |
| return t, rows_p * cols, 0 | |
| if rows != 1: | |
| raise ValueError(f"pack_outputs: {name!r} {shape} has rows of {cols * elem} B; the multi-row layout reads " | |
| f"RM rows of at most {ROW_PAGE_MAX_BYTES} B: give it a narrower last dim (e.g. " | |
| f"[1, 1, -1, {row_elems}]) before packing") | |
| # one wide row (a flat vector): cut it into chunks of whole packed rows, each well inside the L1 budget | |
| if shape != (1, 1, 1, numel): | |
| t = ttnn.reshape(t, (1, 1, 1, numel)) | |
| width = max(row_elems, (ROW_PAGE_MAX_BYTES // elem) // row_elems * row_elems) | |
| parts, total = [], 0 | |
| for start in range(0, numel, width): | |
| stop = min(start + width, numel) | |
| chunk = ttnn.slice(t, [0, 0, 0, start], [1, 1, 1, stop]) | |
| chunk, padded = _flat_rows(ttnn, chunk, (1, 1, 1, stop - start), stop - start, row_elems) | |
| parts.append(chunk) | |
| total += padded | |
| return (parts[0] if len(parts) == 1 else ttnn.concat(parts, dim=2)), total, 0 | |
| def pack_outputs(tensors: Mapping[str, Any], *, dtype: Any = "float32", align: int = 32, | |
| row_elems: Optional[int] = None) -> Packed: | |
| """Pack several device tensors into ONE ROW_MAJOR tensor so the host reads them with one device-to-host copy. | |
| Each tensor is typecast to ``dtype`` (float32 default: exact for bf16 and for integers below 2**24) and converted | |
| to ROW_MAJOR; call this inside the variant function (the ops become part of the trace) and return the result. | |
| :meth:`PackLayout.unpack` (done by ``TraceRunner.read``) gives ``{name: array}`` back in the original shapes. | |
| Device layout: while the packed total is at most :data:`SINGLE_ROW_MAX_ELEMS`, one ``[1, 1, 1, total]`` row (each | |
| tensor flattened and zero-padded to ``align`` elements, then concatenated): a ``[1, 1, N, 1]`` RM readback is N | |
| pages and costs milliseconds, one row is one page (RP section 1.4). Above it (or with ``row_elems=R``) a | |
| ``[1, 1, rows, R]`` tensor (default R = :data:`PACK_ROW_ELEMS`): each tensor is zero-padded to whole rows of R | |
| elements and the segments are concatenated on the row dim, so no RM page exceeds :data:`ROW_PAGE_MAX_BYTES` | |
| and outputs of tens of MB pack (a single row of more than ~0.6 MB does not fit the RM reshape's L1 staging: | |
| YOLOX PORT_LOG Q9). A tensor is padded with zero rows (contiguous), or, when that wastes more than 1/8 of it | |
| (an awkward last dim with few rows: e.g. ``[1, 1, 2, 20001]`` would take 1024 rows), flattened (up to the page | |
| budget) or stored with its rows zero-padded to a wider pitch (:class:`PackEntry` ``pitch``). A last dim wider | |
| than ``ROW_PAGE_MAX_BYTES`` is accepted for a flat vector (``[1, 1, 1, N]``, cut into chunks with | |
| ``ttnn.slice``); any other tensor needs a narrower last dim (``ValueError``).""" | |
| import ttnn | |
| if not tensors: | |
| raise ValueError("pack_outputs needs at least one tensor") | |
| if align <= 0: | |
| raise ValueError("align must be positive") | |
| if row_elems is not None and (int(row_elems) <= 0 or int(row_elems) % TILE): | |
| raise ValueError(f"row_elems={row_elems}: expected a positive multiple of {TILE}") | |
| dtype = ttnn_dtype(dtype) | |
| elem = _elem_bytes(dtype) | |
| items = [(str(name), t, tuple(int(s) for s in t.shape)) for name, t in tensors.items()] | |
| empty = [name for name, _, shape in items if math.prod(shape) == 0] | |
| if empty: | |
| raise ValueError(f"pack_outputs: {empty} have no elements") | |
| single_total = sum(round_up(math.prod(shape), align) for _, _, shape in items) | |
| if row_elems is None and single_total <= SINGLE_ROW_MAX_ELEMS: | |
| segments, entries, offset = [], [], 0 | |
| for name, t, shape in items: # the single-row layout of ttaw 0.1.0 - 0.14.0 | |
| numel = math.prod(shape) | |
| t = _row_major(ttnn, t, dtype) | |
| if shape != (1, 1, 1, numel): | |
| t = ttnn.reshape(t, (1, 1, 1, numel)) | |
| padded = round_up(numel, align) | |
| if padded != numel: | |
| t = ttnn.pad(t, [(0, 0), (0, 0), (0, 0), (0, padded - numel)], 0.0) | |
| segments.append(t) | |
| entries.append(PackEntry(name, offset, numel, shape)) | |
| offset += padded | |
| packed = segments[0] if len(segments) == 1 else ttnn.concat(segments, dim=-1) | |
| return Packed(packed, PackLayout(tuple(entries), offset)) | |
| r = int(row_elems or PACK_ROW_ELEMS) | |
| segments, entries, offset = [], [], 0 | |
| for name, t, shape in items: | |
| numel = math.prod(shape) | |
| t, padded, pitch = _pack_rows_segment(ttnn, name, _row_major(ttnn, t, dtype), shape, numel, r, elem) | |
| segments.append(t) | |
| entries.append(PackEntry(name, offset, numel, shape, pitch)) | |
| offset += padded | |
| packed = segments[0] if len(segments) == 1 else ttnn.concat(segments, dim=2) | |
| return Packed(packed, PackLayout(tuple(entries), offset, rows=offset // r, row_elems=r)) | |
| # --------------------------------------------------------------------------------------------- internals | |
| class _Slot: | |
| """A persistent device tensor (input, parameter or state) allocated before any capture.""" | |
| name: str | |
| kind: str # "input" | "param" | "state" | "output" | |
| shape: Tuple[int, ...] | |
| dtype: Any | |
| layout: Any | |
| memory_config: Any | |
| buffers: List[Any] # 1 buffer, or 2 for a ping-pong state | |
| init: Any # host tensor holding the initial value (states, params) / warm-up value (inputs) | |
| staging: Any = None # DRAM staging copy (2CQ stage_inputs mode) | |
| stage_fn: Optional[Callable[[Any, Any], Any]] = None | |
| last_value: Optional[bytes] = None # params: bytes of the last uploaded value (skip unchanged uploads) | |
| banks: List[Any] = field(default_factory=list) # states: save / load buffers (D16 per-stream state banks) | |
| def pingpong(self) -> bool: | |
| return len(self.buffers) == 2 | |
| def tensors(self) -> List[Any]: | |
| """Every device tensor the slot owns.""" | |
| return self.buffers + ([self.staging] if self.staging is not None else []) + self.banks | |
| class _Variant: | |
| name: str | |
| fn: Callable[["TraceContext"], Any] | |
| warmup_runs: int | |
| class _Trace: | |
| variant: str | |
| phase: int | |
| trace_id: Any | |
| outputs: Any | |
| leaves: List[Tuple[Tuple, Any, Optional[PackLayout]]] # (path, device tensor, pack layout or None) | |
| steps: bool # writes the ping-pong states (flips the phase) | |
| capture_ms: float | |
| host_buffers: Optional[List[Any]] = None | |
| done_event: Any = None | |
| def _flatten(obj: Any, path: Tuple = ()) -> Iterator[Tuple[Tuple, Any]]: | |
| """Leaves of a tensor / Packed / list / tuple / dict structure with their paths.""" | |
| if isinstance(obj, Mapping): | |
| for k, v in obj.items(): | |
| yield from _flatten(v, path + (k,)) | |
| elif isinstance(obj, (list, tuple)): | |
| for i, v in enumerate(obj): | |
| yield from _flatten(v, path + (i,)) | |
| else: | |
| yield path, obj | |
| def _rebuild(obj: Any, values: Dict[Tuple, Any], path: Tuple = ()) -> Any: | |
| if isinstance(obj, Mapping): | |
| return {k: _rebuild(v, values, path + (k,)) for k, v in obj.items()} | |
| if isinstance(obj, (list, tuple)): | |
| return type(obj)(_rebuild(v, values, path + (i,)) for i, v in enumerate(obj)) | |
| return values[path] | |
| def _buffer_address(t: Any) -> Optional[int]: | |
| try: | |
| return int(t.buffer_address()) | |
| except Exception: # noqa: BLE001 -- host tensor, deallocated tensor or a fake | |
| return None | |
| def _same_buffer(a: Any, b: Any) -> bool: | |
| """``a`` is ``b`` or a tensor over the same device buffer.""" | |
| if a is b: | |
| return True | |
| addr = _buffer_address(a) | |
| return addr is not None and addr == _buffer_address(b) | |
| def _check_copy(what: str, value: Any, slot: "_Slot") -> None: | |
| """The preconditions of ``ttnn.copy(value, <slot buffer>)``, checked in Python so a mistake is a clear error | |
| instead of a TT_FATAL inside an open capture: same logical shape and layout; a dtype change only in TILE.""" | |
| import ttnn | |
| shape = tuple(int(s) for s in value.shape) | |
| if shape != slot.shape: | |
| raise ValueError(f"{what}: value shape {shape} != {slot.shape}") | |
| if value.layout != slot.layout: | |
| raise ValueError(f"{what}: value layout {value.layout} != {slot.layout} (ttnn.copy keeps the layout)") | |
| if value.dtype != slot.dtype and slot.layout != ttnn.TILE_LAYOUT: | |
| raise ValueError(f"{what}: dtype {value.dtype} -> {slot.dtype} needs TILE layout (ttnn.copy)") | |
| class TraceContext(Mapping): | |
| """What a variant function receives: persistent tensors by name (``ctx["x"]``), and state writes. | |
| ``ctx[name]`` returns an input, a parameter, a persistent output buffer, or the buffer a state is *read* from in | |
| this phase. ``ctx.write_state(name, value)`` copies ``value`` into the buffer the state is *written* to (in-place | |
| states: the same buffer, so read everything you need from it before writing); ``ctx.write_output(name, value)`` | |
| copies into a persistent output and returns it.""" | |
| def __init__(self, runner: "TraceRunner", variant: str, phase: int, capturing: bool): | |
| self._runner = runner | |
| self.variant = variant | |
| self.phase = phase | |
| self.capturing = capturing | |
| self.writes: set = set() | |
| def device(self): | |
| return self._runner.device | |
| def __getitem__(self, name: str): | |
| slot = self._runner._slot(name) | |
| if slot.kind == "state": | |
| return self.state(name) | |
| return slot.buffers[0] | |
| def __iter__(self) -> Iterator[str]: | |
| return iter(self._runner._slots) | |
| def __len__(self) -> int: | |
| return len(self._runner._slots) | |
| def state(self, name: str): | |
| """The buffer state ``name`` is read from in this phase.""" | |
| slot = self._runner._slot(name, "state") | |
| return slot.buffers[self.phase] if slot.pingpong else slot.buffers[0] | |
| def write_target(self, name: str): | |
| """The buffer state ``name`` is written to in this phase (ping-pong: the other buffer; in place: the same | |
| one). Pass it as ``output_tensor=`` of the op that produces the new state and then call | |
| ``write_state(name, it)``: no copy program is traced (probe P13: saves one program per state per frame).""" | |
| slot = self._runner._slot(name, "state") | |
| return slot.buffers[1 - self.phase] if slot.pingpong else slot.buffers[0] | |
| def write_state(self, name: str, value) -> None: | |
| """``ttnn.copy(value, <write buffer>)``: same logical shape and layout as the state (dtype may differ in | |
| TILE layout). The copy is part of the trace; it is skipped when ``value`` already is the write buffer | |
| (:meth:`write_target`).""" | |
| import ttnn | |
| slot = self._runner._slot(name, "state") | |
| target = self.write_target(name) | |
| _check_copy(f"state {name!r}", value, slot) | |
| if not _same_buffer(value, target): | |
| ttnn.copy(value, target) | |
| self.writes.add(name) | |
| def write_output(self, name: str, value): | |
| """``ttnn.copy(value, <persistent output>)`` (part of the trace); returns the output buffer, which the | |
| variant can return as (part of) its outputs. Skipped when ``value`` already is that buffer.""" | |
| import ttnn | |
| slot = self._runner._slot(name, "output") | |
| _check_copy(f"output {name!r}", value, slot) | |
| if not _same_buffer(value, slot.buffers[0]): | |
| ttnn.copy(value, slot.buffers[0]) | |
| return slot.buffers[0] | |
| class TraceRunner: | |
| """Persistent device I/O + warm-up + capture + replay of one model's traced stages (see the module docstring). | |
| Args: | |
| device: an open ttnn device. | |
| num_command_queues: 1 or 2. ``None`` uses what :func:`.device.open_device` recorded (1 if unknown). | |
| warmup_runs: eager runs of each variant (and phase) before any capture. | |
| stage_inputs: 2CQ only: upload into DRAM staging buffers and copy them into the trace inputs with an eager | |
| op on CQ0 (the upload of frame k+1 then overlaps the replay of frame k). | |
| forbid_cache_misses: call ``device.set_program_cache_misses_allowed(False)`` during capture. | |
| alloc_tracking: ``True`` raises unless the process tracks trace allocations | |
| (``TT_METAL_TRACE_ALLOC_TRACKING=1`` set before ``import ttnn``). Whenever tracking is active, the outputs | |
| of traces captured after the first are acknowledged as corruptible (see the module docstring). | |
| name: label used in messages and ``describe()``. | |
| """ | |
| def __init__(self, device, *, num_command_queues: Optional[int] = None, warmup_runs: int = 1, | |
| stage_inputs: bool = False, forbid_cache_misses: bool = True, alloc_tracking: Optional[bool] = None, | |
| name: str = "model"): | |
| from .device import open_info | |
| opened = open_info(device).get("num_command_queues") | |
| if num_command_queues is None: | |
| num_command_queues = opened or 1 | |
| if num_command_queues not in (1, 2): | |
| raise ValueError(f"num_command_queues={num_command_queues}: expected 1 or 2") | |
| if opened is not None and num_command_queues > opened: | |
| raise ValueError(f"the device was opened with {opened} command queue(s); cannot run {num_command_queues}") | |
| if stage_inputs and num_command_queues != 2: | |
| raise ValueError("stage_inputs=True needs num_command_queues=2") | |
| if warmup_runs < 1: | |
| raise ValueError("warmup_runs must be >= 1 (capture needs a warm program cache)") | |
| tracking = alloc_tracking_enabled() | |
| if alloc_tracking and not tracking: | |
| raise RuntimeError("alloc_tracking=True but trace allocation tracking is off: export " | |
| "TT_METAL_TRACE_ALLOC_TRACKING=1 (optionally TT_METAL_TRACE_ALLOC_TRACEBACKS=1) " | |
| "before Python imports ttnn") | |
| self.device = device | |
| self.name = name | |
| self.num_command_queues = int(num_command_queues) | |
| self.warmup_runs = int(warmup_runs) | |
| self.stage_inputs = bool(stage_inputs) | |
| self.forbid_cache_misses = bool(forbid_cache_misses) | |
| self.alloc_tracking = bool(tracking) | |
| self._slots: Dict[str, _Slot] = {} | |
| self._variants: Dict[str, _Variant] = {} | |
| self._traces: Dict[Tuple[str, int], _Trace] = {} | |
| self._phase = 0 | |
| self._last_phase: Dict[str, int] = {} | |
| self._last_variant: Optional[str] = None | |
| self._pending_params: Dict[str, Any] = {} | |
| self._op_event = None | |
| self._stage_free_event = None | |
| self._captures = 0 | |
| self._capturing = False | |
| self._eager_warmups: List[Callable[[], Any]] = [] | |
| self._eager_warmed = False | |
| self._closed = False | |
| self.timings_ms: Dict[str, Dict[str, float]] = {} | |
| if hasattr(device, "enable_program_cache"): | |
| device.enable_program_cache() | |
| # ------------------------------------------------------------------------------------- registration | |
| def _check_registration(self, what: str) -> None: | |
| self._check_open() | |
| if self._capturing: | |
| raise RuntimeError(f"cannot add {what} while a variant is being warmed up or captured: register " | |
| "inputs, params, states, outputs and variants before capture()") | |
| def _new_slot(self, name: str, kind: str, init: Any, shape: Optional[Sequence[int]], dtype: Any, layout: Any, | |
| memory_config: Any, n_buffers: int, stage_fn: Optional[Callable], n_banks: int = 0) -> _Slot: | |
| import ttnn | |
| self._check_registration(f"{kind} {name!r}") | |
| if name in self._slots: | |
| raise ValueError(f"{name!r} is already registered (as {self._slots[name].kind})") | |
| if self._traces: | |
| raise RuntimeError(f"cannot add {kind} {name!r} after capture: persistent tensors must exist before the " | |
| "first capture (release() and rebuild)") | |
| layout = ttnn.ROW_MAJOR_LAYOUT if layout is None else layout | |
| memory_config = ttnn.DRAM_MEMORY_CONFIG if memory_config is None else memory_config | |
| dtype = ttnn_dtype(dtype) | |
| if init is None: | |
| if shape is None: | |
| raise ValueError(f"{kind} {name!r}: give init= or shape=") | |
| init = np.zeros(tuple(shape), np.float32 if dtype_name(dtype) in ("float32", "bfloat16", "bfloat8_b", | |
| "bfloat4_b") else np.int64) | |
| if isinstance(init, ttnn.Tensor): | |
| host = to_host_tensor(init, dtype, layout, shape=shape) | |
| else: | |
| arr = to_numpy(init) if hasattr(init, "detach") else np.asarray(init) | |
| if shape is not None: | |
| arr = np.broadcast_to(arr, tuple(shape)) | |
| host = to_host_tensor(arr, dtype, layout) | |
| shape_t = tuple(int(s) for s in host.shape) | |
| def allocate(config: Any): | |
| buf = ttnn.allocate_tensor_on_device(ttnn.Shape(list(shape_t)), dtype, layout, self.device, config) | |
| ttnn.copy_host_to_device_tensor(host, buf, cq_id=CQ_COMPUTE) # defined contents, never garbage | |
| return buf | |
| buffers = [allocate(memory_config) for _ in range(n_buffers)] | |
| staging = allocate(ttnn.DRAM_MEMORY_CONFIG) if self.stage_inputs and kind in ("input", "param") else None | |
| banks = [allocate(ttnn.DRAM_MEMORY_CONFIG) for _ in range(n_banks)] | |
| slot = _Slot(name, kind, shape_t, dtype, layout, memory_config, buffers, host, staging, | |
| stage_fn or (lambda src, dst: ttnn.copy(src, dst)), banks=banks) | |
| self._slots[name] = slot | |
| return slot | |
| def add_input(self, name: str, init: Any = None, *, shape: Optional[Sequence[int]] = None, | |
| dtype: Any = "bfloat16", layout: Any = None, memory_config: Any = None, | |
| stage_fn: Optional[Callable[[Any, Any], Any]] = None): | |
| """A persistent trace input (default ROW_MAJOR in DRAM). ``init`` (array / torch / ttnn host tensor) is the | |
| initial content and the warm-up input; zeros of ``shape`` otherwise -- give real data when the graph | |
| gathers with these values (garbage indices can hang the chip). ``stage_fn(staging, persistent)`` replaces | |
| the eager ``ttnn.copy`` in ``stage_inputs`` mode (e.g. a reshard into a sharded L1 input). Returns the | |
| device tensor.""" | |
| return self._new_slot(name, "input", init, shape, dtype, layout, memory_config, 1, stage_fn).buffers[0] | |
| def add_param(self, name: str, value: Any = 0.0, *, shape: Sequence[int] = (1, 1, 1, 1), dtype: Any = "float32", | |
| layout: Any = None, memory_config: Any = None): | |
| """An RT-dev parameter: a persistent device tensor (default fp32 ``[1, 1, 1, 1]`` TILE, which broadcasts in | |
| ttnn binary ops) refreshed before a replay when ``run(params=...)`` / :meth:`set_params` changes it.""" | |
| import ttnn | |
| layout = ttnn.TILE_LAYOUT if layout is None else layout | |
| slot = self._new_slot(name, "param", np.broadcast_to(np.asarray(value), tuple(shape)), tuple(shape), dtype, | |
| layout, memory_config, 1, None) | |
| slot.last_value = self._param_bytes(slot, value) | |
| return slot.buffers[0] | |
| def add_state(self, name: str, init: Any = None, *, shape: Optional[Sequence[int]] = None, dtype: Any = "float32", | |
| layout: Any = None, memory_config: Any = None, pingpong: bool = False, banks: int = 0): | |
| """Temporal state kept on the device across replays (memory queues, previous BEV, ring buffers). | |
| In-place (default): one buffer, read via ``ctx[name]`` and written via ``ctx.write_state``. ``pingpong``: | |
| two buffers and two traces per variant. ``init`` is restored by :meth:`reset_state`. ``banks``: extra DRAM | |
| buffers of the same spec for :meth:`save_state` / :meth:`load_state` (one per stream id of a | |
| :class:`StreamBanks`, PLAN.md D16), allocated now because nothing may be allocated after a capture. | |
| Returns the buffer(s).""" | |
| import ttnn | |
| if banks < 0: | |
| raise ValueError("banks must be >= 0") | |
| layout = ttnn.TILE_LAYOUT if layout is None else layout | |
| slot = self._new_slot(name, "state", init, shape, dtype, layout, memory_config, 2 if pingpong else 1, None, | |
| n_banks=int(banks)) | |
| return tuple(slot.buffers) if pingpong else slot.buffers[0] | |
| def add_output(self, name: str, *, shape: Sequence[int], dtype: Any = "float32", layout: Any = None, | |
| memory_config: Any = None): | |
| """A persistent output buffer (default TILE in DRAM, zeros) allocated before any capture. Variants write it | |
| with ``ctx.write_output(name, value)`` (one traced ``ttnn.copy``) and return it; its address survives | |
| recaptures and is shared by every variant (e.g. shape buckets with one readback). Returns the buffer.""" | |
| import ttnn | |
| layout = ttnn.TILE_LAYOUT if layout is None else layout | |
| return self._new_slot(name, "output", None, shape, dtype, layout, memory_config, 1, None).buffers[0] | |
| def add_variant(self, name: str, fn: Callable[[TraceContext], Any], *, warmup_runs: Optional[int] = None) -> None: | |
| """Register a traced function. ``fn(ctx)`` runs ttnn ops on ``ctx[...]`` tensors and returns a device tensor, | |
| a :class:`Packed`, or a list / tuple / dict of them (the persistent trace outputs). No host I/O, no | |
| ``synchronize``, no torch ops inside ``fn``: it is called for warm-up and then recorded.""" | |
| self._check_registration(f"variant {name!r}") | |
| if name in self._variants: | |
| raise ValueError(f"variant {name!r} already exists") | |
| self._variants[name] = _Variant(name, fn, self.warmup_runs if warmup_runs is None else int(warmup_runs)) | |
| def add_eager_warmup(self, fn: Callable[[], Any]) -> None: | |
| """Register eager device work that the model runs *between* replays (host-fallback glue, an eager layout | |
| change of a read-back tensor, ...): ``fn()`` runs once in the first :meth:`capture`, before any trace is | |
| captured, so its programs are compiled -- and their kernel binaries allocated in DRAM -- before the first | |
| capture. A program compiled after a capture shares the address space of the traces' freed intermediates and | |
| a replay can overwrite its binaries (tt-metal ``tech_reports/.../TraceCorrectness.md``: corruption or a | |
| hang); ``TT_METAL_TRACE_ALLOC_TRACKING=1`` reports it.""" | |
| self._check_registration("an eager warm-up") | |
| if self._traces or self._eager_warmed: | |
| raise RuntimeError("add eager warm-ups before the first capture (release() and rebuild)") | |
| self._eager_warmups.append(fn) | |
| # ------------------------------------------------------------------------------------------ capture | |
| def phases(self) -> int: | |
| """2 when a ping-pong state exists (two traces per variant), else 1.""" | |
| return 2 if any(s.pingpong for s in self._slots.values()) else 1 | |
| def captured(self) -> bool: | |
| return bool(self._variants) and all((v, p) in self._traces for v in self._variants for p in range(self.phases)) | |
| def capture(self) -> None: | |
| """Warm up and capture every registered variant that has no trace yet (idempotent). If traces already exist | |
| and a new variant was added, all traces are released, the new variants warmed up, and all recaptured.""" | |
| import ttnn | |
| self._check_open() | |
| if not self._variants: | |
| raise RuntimeError("no variant registered (add_variant)") | |
| pending = [v for v in self._variants if any((v, p) not in self._traces for p in range(self.phases))] | |
| if not pending: | |
| return | |
| if self._traces: | |
| ttnn.synchronize_device(self.device) # no replay of a trace being released is still in flight | |
| self._release_traces() | |
| warm, pending = pending, list(self._variants) | |
| else: | |
| warm = pending | |
| self._capturing = True | |
| try: | |
| if not self._eager_warmed: | |
| self._warm_eager_programs() | |
| for vname in warm: | |
| t0 = time.perf_counter() | |
| variant = self._variants[vname] | |
| for phase in range(self.phases): | |
| for _ in range(variant.warmup_runs): | |
| ctx = TraceContext(self, vname, phase, capturing=False) | |
| outputs = variant.fn(ctx) | |
| ttnn.synchronize_device(self.device) | |
| self._free_transient(outputs) | |
| self.timings_ms.setdefault(vname, {})["warmup"] = (time.perf_counter() - t0) * 1e3 | |
| self._phase = 0 | |
| self.reset_state() | |
| for vname in pending: | |
| self.timings_ms.setdefault(vname, {})["capture"] = 0.0 | |
| for phase in range(self.phases): | |
| self._capture_one(vname, phase) | |
| finally: | |
| self._capturing = False | |
| ttnn.synchronize_device(self.device) | |
| if self.num_command_queues == 2: | |
| self._op_event = ttnn.record_event(self.device, CQ_COMPUTE) | |
| self._stage_free_event = self._op_event | |
| def _warm_eager_programs(self) -> None: | |
| """Compile the runner's own eager programs before the first capture: the staging copies (``stage_inputs``), | |
| the state <-> bank copies (``banks``) and the registered eager warm-ups. Contents are kept: a staging buffer | |
| mirrors its input (:meth:`write_input` writes both), and a state and its first bank hold the same value | |
| before the states are reset ahead of the captures.""" | |
| import ttnn | |
| for slot in self._slots.values(): | |
| if slot.staging is not None: | |
| slot.stage_fn(slot.staging, slot.buffers[0]) | |
| if slot.banks: | |
| ttnn.copy(slot.buffers[0], slot.banks[0]) | |
| ttnn.copy(slot.banks[0], slot.buffers[0]) | |
| for fn in self._eager_warmups: | |
| fn() | |
| ttnn.synchronize_device(self.device) | |
| self._eager_warmed = True | |
| def _eager(self, what: str) -> Iterator[None]: | |
| """Eager runner work after a capture must not compile a program (see :meth:`add_eager_warmup`).""" | |
| dev = self.device | |
| count = getattr(dev, "num_program_cache_entries", None) | |
| before = count() if (self._traces and count is not None) else None | |
| yield | |
| if before is not None and count() != before: | |
| raise RuntimeError(f"{self.name}: {what} compiled a new program after capture; its kernel binaries share " | |
| "DRAM with the traces' freed intermediates, so a replay can overwrite them. Run it " | |
| "before the first capture (add_eager_warmup) or keep the tensor specs of the warmed " | |
| "program") | |
| def _capture_one(self, vname: str, phase: int) -> None: | |
| import ttnn | |
| dev = self.device | |
| variant = self._variants[vname] | |
| ctx = TraceContext(self, vname, phase, capturing=True) | |
| entries_before = dev.num_program_cache_entries() if hasattr(dev, "num_program_cache_entries") else None | |
| forbid = self.forbid_cache_misses and hasattr(dev, "set_program_cache_misses_allowed") | |
| t0 = time.perf_counter() | |
| if forbid: | |
| dev.set_program_cache_misses_allowed(False) | |
| trace_id = ttnn.begin_trace_capture(dev, cq_id=CQ_COMPUTE) | |
| ok = ended = False | |
| try: | |
| outputs = variant.fn(ctx) | |
| ok = True | |
| finally: | |
| try: | |
| ttnn.end_trace_capture(dev, trace_id, cq_id=CQ_COMPUTE) | |
| ended = True | |
| finally: | |
| if forbid: | |
| dev.set_program_cache_misses_allowed(True) | |
| if not (ok and ended): | |
| self._safe_release(trace_id) | |
| try: | |
| entries_after = dev.num_program_cache_entries() if entries_before is not None else None | |
| if entries_before is not None and entries_after != entries_before: | |
| raise RuntimeError(f"{self.name}/{vname}: the program cache grew during capture ({entries_before} -> " | |
| f"{entries_after}); warm-up does not cover the traced graph") | |
| leaves = [] | |
| for path, leaf in _flatten(outputs): | |
| if isinstance(leaf, Packed): | |
| leaves.append((path, leaf.tensor, leaf.layout)) | |
| elif isinstance(leaf, ttnn.Tensor): | |
| leaves.append((path, leaf, None)) | |
| else: | |
| raise TypeError(f"{self.name}/{vname}: output {path} is {type(leaf).__name__}, expected a device " | |
| "tensor or Packed") | |
| if not leaves: | |
| raise ValueError(f"{self.name}/{vname}: the variant returned no output tensor") | |
| pingpong = {s.name for s in self._slots.values() if s.pingpong} | |
| written = ctx.writes & pingpong | |
| if written and written != pingpong: | |
| raise RuntimeError(f"{self.name}/{vname}: writes ping-pong states {sorted(written)} but not " | |
| f"{sorted(pingpong - written)}; a stepping variant must write all of them") | |
| except BaseException: | |
| self._safe_release(trace_id) | |
| raise | |
| if self.alloc_tracking and self._captures > 0: | |
| from ttnn.tools import trace_allocation_tracker as tracker | |
| keep = self._persistent_addresses() | |
| for _, tensor, _ in leaves: | |
| if _buffer_address(tensor) not in keep: | |
| tracker.acknowledge_corruptible(tensor) | |
| self._captures += 1 | |
| ms = (time.perf_counter() - t0) * 1e3 | |
| self._traces[(vname, phase)] = _Trace(vname, phase, trace_id, outputs, leaves, bool(written), ms) | |
| self.timings_ms[vname]["capture"] += ms | |
| def _persistent_addresses(self) -> set: | |
| addrs = set() | |
| for slot in self._slots.values(): | |
| for t in slot.tensors(): | |
| a = _buffer_address(t) | |
| if a is not None: | |
| addrs.add(a) | |
| return addrs | |
| def _deallocate_outputs(self, tensors: Iterator[Any]) -> None: | |
| """Deallocate op-produced output tensors, never a persistent buffer or a view of one. ``force=False``: a | |
| tensor sharing its device memory with another owner (a view of a model weight returned as an output) is | |
| left to its owners -- ``ttnn.deallocate`` forces by default and would free the weight.""" | |
| import ttnn | |
| keep = self._persistent_addresses() | |
| seen = set() | |
| for tensor in tensors: | |
| if not isinstance(tensor, ttnn.Tensor) or id(tensor) in seen: | |
| continue | |
| seen.add(id(tensor)) | |
| if tensor.is_allocated() and _buffer_address(tensor) not in keep: | |
| ttnn.deallocate(tensor, False) | |
| def _free_transient(self, outputs: Any) -> None: | |
| """Deallocate warm-up / eager outputs (see :meth:`_deallocate_outputs`).""" | |
| self._deallocate_outputs(leaf.tensor if isinstance(leaf, Packed) else leaf for _, leaf in _flatten(outputs)) | |
| def _safe_release(self, trace_id) -> None: | |
| import ttnn | |
| try: | |
| ttnn.release_trace(self.device, trace_id) | |
| except Exception: # noqa: BLE001 -- best effort on an error path; the original error is re-raised | |
| pass | |
| # -------------------------------------------------------------------------------------------- inputs | |
| def _slot(self, name: str, kind: Optional[str] = None) -> _Slot: | |
| slot = self._slots.get(name) | |
| if slot is None: | |
| raise KeyError(f"{self.name}: no input / param / state named {name!r}") | |
| if kind is not None and slot.kind != kind: | |
| raise KeyError(f"{self.name}: {name!r} is a {slot.kind}, not a {kind}") | |
| return slot | |
| def _param_array(slot: _Slot, value: Any) -> np.ndarray: | |
| dt = np.float32 if dtype_name(slot.dtype) in ("float32", "bfloat16", "bfloat8_b", "bfloat4_b") else np.int64 | |
| return np.ascontiguousarray(np.broadcast_to(np.asarray(value, dtype=dt), slot.shape)) | |
| def _param_bytes(self, slot: _Slot, value: Any) -> bytes: | |
| return self._param_array(slot, value).tobytes() | |
| def set_params(self, **values: Any) -> None: | |
| """Queue RT-dev parameter values for the next run (uploaded only if they changed).""" | |
| for name in values: | |
| self._slot(name, "param") | |
| self._pending_params.update(values) | |
| def _collect_uploads(self, inputs: Optional[Mapping[str, Any]], | |
| params: Optional[Mapping[str, Any]]) -> List[Tuple[_Slot, Any, Optional[bytes]]]: | |
| uploads = [] | |
| for name, value in (inputs or {}).items(): | |
| slot = self._slot(name, "input") | |
| uploads.append((slot, to_host_tensor(value, slot.dtype, slot.layout, shape=slot.shape), None)) | |
| merged = dict(self._pending_params) | |
| merged.update(params or {}) | |
| for name, value in merged.items(): | |
| slot = self._slot(name, "param") | |
| arr = self._param_array(slot, value) | |
| key = arr.tobytes() | |
| if key != slot.last_value: | |
| uploads.append((slot, to_host_tensor(arr, slot.dtype, slot.layout), key)) | |
| return uploads | |
| def _ensure_events(self) -> None: | |
| """2CQ: the CQ0 events CQ1 waits for exist (``capture()`` records them; a partially failed capture, which | |
| leaves earlier variants runnable, does not get that far).""" | |
| import ttnn | |
| if self._op_event is None: | |
| self._op_event = ttnn.record_event(self.device, CQ_COMPUTE) | |
| if self._stage_free_event is None: | |
| self._stage_free_event = self._op_event | |
| def _enqueue_uploads(self, uploads: List[Tuple[_Slot, Any, Optional[bytes]]]) -> None: | |
| import ttnn | |
| if uploads: | |
| if self.num_command_queues == 1: | |
| for slot, host, _ in uploads: | |
| ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_COMPUTE) | |
| elif self.stage_inputs: | |
| self._ensure_events() | |
| ttnn.wait_for_event(CQ_INPUT, self._stage_free_event) | |
| for slot, host, _ in uploads: | |
| ttnn.copy_host_to_device_tensor(host, slot.staging, cq_id=CQ_INPUT) | |
| written = ttnn.record_event(self.device, CQ_INPUT) | |
| ttnn.wait_for_event(CQ_COMPUTE, written) | |
| with self._eager("a stage_fn copy"): | |
| for slot, _, _ in uploads: | |
| slot.stage_fn(slot.staging, slot.buffers[0]) | |
| self._stage_free_event = ttnn.record_event(self.device, CQ_COMPUTE) | |
| else: | |
| self._ensure_events() | |
| ttnn.wait_for_event(CQ_INPUT, self._op_event) | |
| for slot, host, _ in uploads: | |
| ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_INPUT) | |
| written = ttnn.record_event(self.device, CQ_INPUT) | |
| ttnn.wait_for_event(CQ_COMPUTE, written) | |
| for slot, _, key in uploads: | |
| if key is not None: | |
| slot.last_value = key | |
| self._pending_params.clear() | |
| def write_input(self, name: str, value: Any) -> None: | |
| """Upload ``value`` into input ``name`` now (CQ0, outside any trace): e.g. a realistic warm-up sample. | |
| With 2 CQs the event that later CQ1 uploads wait for is re-recorded after this write, so an upload of the | |
| next frame can never land before it (both write the same buffer from different queues).""" | |
| import ttnn | |
| self._check_open() | |
| slot = self._slot(name, "input") | |
| host = to_host_tensor(value, slot.dtype, slot.layout, shape=slot.shape) | |
| ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_COMPUTE) | |
| if slot.staging is not None: # the staging buffer mirrors the input (the stage-copy warm-up keeps it) | |
| ttnn.copy_host_to_device_tensor(host, slot.staging, cq_id=CQ_COMPUTE) | |
| if self.num_command_queues == 2 and self._op_event is not None: | |
| self._op_event = ttnn.record_event(self.device, CQ_COMPUTE) | |
| if self.stage_inputs: | |
| self._stage_free_event = self._op_event | |
| # --------------------------------------------------------------------------------------------- run | |
| def _trace_for(self, variant: Optional[str], phase: Optional[int] = None) -> _Trace: | |
| self._check_open() | |
| name = variant or self._last_variant or (next(iter(self._variants)) if len(self._variants) == 1 else None) | |
| if name is None: | |
| raise ValueError("variant name required") | |
| if name not in self._variants: | |
| raise KeyError(f"{self.name}: unknown variant {name!r}; have {sorted(self._variants)}") | |
| p = self._phase if phase is None else phase | |
| trace = self._traces.get((name, p)) | |
| if trace is None: | |
| raise RuntimeError(f"{self.name}: variant {name!r} is not captured; call capture() first") | |
| return trace | |
| def _execute(self, trace: _Trace) -> None: | |
| import ttnn | |
| ttnn.execute_trace(self.device, trace.trace_id, cq_id=CQ_COMPUTE, blocking=False) | |
| if self.num_command_queues == 2: | |
| self._op_event = ttnn.record_event(self.device, CQ_COMPUTE) | |
| trace.done_event = self._op_event | |
| self._last_variant = trace.variant | |
| self._last_phase[trace.variant] = trace.phase | |
| if trace.steps: | |
| self._phase ^= 1 | |
| def upload(self, inputs: Optional[Mapping[str, Any]] = None, params: Optional[Mapping[str, Any]] = None) -> int: | |
| """Enqueue the uploads of ``inputs`` and changed ``params`` (CQ0, or CQ1 + events with 2 CQs) without | |
| replaying anything; the next :meth:`run` / :meth:`replay` consumes them. Returns the number of tensors | |
| uploaded.""" | |
| self._check_open() | |
| if not self.captured: | |
| raise RuntimeError(f"{self.name}: capture() before uploading") | |
| uploads = self._collect_uploads(inputs, params) | |
| self._enqueue_uploads(uploads) | |
| return len(uploads) | |
| def run(self, variant: Optional[str] = None, inputs: Optional[Mapping[str, Any]] = None, | |
| params: Optional[Mapping[str, Any]] = None): | |
| """Upload ``inputs`` (name -> numpy / torch / ttnn host tensor) and changed ``params``, then replay the | |
| variant's trace without blocking. Returns its device outputs (valid once the replay finished).""" | |
| trace = self._trace_for(variant) | |
| self._enqueue_uploads(self._collect_uploads(inputs, params)) | |
| self._execute(trace) | |
| return trace.outputs | |
| def replay(self, variant: Optional[str] = None, n: int = 1) -> None: | |
| """Replay ``n`` times with no uploads (back-to-back device timing); honours ping-pong phases.""" | |
| for _ in range(n): | |
| self._execute(self._trace_for(variant)) | |
| def outputs(self, variant: Optional[str] = None, phase: Optional[int] = None): | |
| """Device outputs of a variant's trace (default: the phase that ran last for it, else phase 0).""" | |
| name = variant or self._last_variant | |
| if phase is None and name is not None: | |
| phase = self._last_phase.get(name, 0) | |
| return self._trace_for(name, phase if phase is not None else 0).outputs | |
| def read(self, variant: Optional[str] = None, *, cq_id: int = CQ_COMPUTE, as_torch: bool = False): | |
| """Read the outputs of the last run of ``variant`` into preallocated host tensors (blocking). | |
| Returns the output structure with numpy arrays (``as_torch=True``: torch tensors); a :class:`Packed` leaf | |
| becomes ``{name: array}``. ``cq_id=1`` (2CQ) waits on the host for the trace's completion event and reads on | |
| CQ1, so CQ0 can already replay the next segment (segmented D2H).""" | |
| import ttnn | |
| name = variant or self._last_variant | |
| if name is None or name not in self._last_phase: | |
| raise RuntimeError(f"{self.name}: variant {name!r} has not run yet") | |
| trace = self._trace_for(name, self._last_phase[name]) | |
| if cq_id not in (CQ_COMPUTE, CQ_INPUT): | |
| raise ValueError(f"cq_id={cq_id}: expected 0 or 1") | |
| if cq_id == CQ_INPUT: | |
| if self.num_command_queues != 2: | |
| raise ValueError("cq_id=1 needs a device opened with 2 command queues") | |
| ttnn.event_synchronize(trace.done_event) | |
| if trace.host_buffers is None: | |
| trace.host_buffers = [ttnn.allocate_tensor_on_host(t.spec, self.device) for _, t, _ in trace.leaves] | |
| values: Dict[Tuple, Any] = {} | |
| for (path, tensor, layout), host in zip(trace.leaves, trace.host_buffers): | |
| ttnn.copy_device_to_host_tensor(tensor, host, blocking=True, cq_id=cq_id) | |
| arr = to_numpy(host) | |
| value: Any = layout.unpack(arr) if layout is not None else arr | |
| if as_torch: | |
| import torch | |
| value = ({k: torch.from_numpy(np.ascontiguousarray(v)) for k, v in value.items()} | |
| if isinstance(value, dict) else torch.from_numpy(np.ascontiguousarray(value))) | |
| values[path] = value | |
| structure = trace.outputs | |
| if isinstance(structure, Packed) or not isinstance(structure, (Mapping, list, tuple)): | |
| return values[()] | |
| return _rebuild(structure, values) | |
| def __call__(self, variant: Optional[str] = None, inputs: Optional[Mapping[str, Any]] = None, | |
| params: Optional[Mapping[str, Any]] = None, *, as_torch: bool = False): | |
| """``run`` + ``read`` (the common synchronous path).""" | |
| self.run(variant, inputs, params) | |
| return self.read(variant, as_torch=as_torch) | |
| def run_eager(self, variant: Optional[str] = None, inputs: Optional[Mapping[str, Any]] = None, | |
| params: Optional[Mapping[str, Any]] = None): | |
| """Upload, run the variant's function *eagerly* (no trace) and return its outputs read to numpy, freeing | |
| every eager buffer before returning (safe next to captured traces). State writes happen as in a replay. | |
| Use it for replay-vs-eager bit checks.""" | |
| import ttnn | |
| self._check_open() | |
| name = variant or self._last_variant or (next(iter(self._variants)) if self._variants else None) | |
| if name not in self._variants: | |
| raise KeyError(f"{self.name}: unknown variant {name!r}; have {sorted(self._variants)}") | |
| variant_obj = self._variants[name] | |
| uploads = self._collect_uploads(inputs, params) | |
| for slot, host, key in uploads: | |
| ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_COMPUTE) | |
| if key is not None: | |
| slot.last_value = key | |
| self._pending_params.clear() | |
| ctx = TraceContext(self, name, self._phase, capturing=False) | |
| with self._eager(f"run_eager({name!r})"): | |
| outputs = variant_obj.fn(ctx) | |
| values = {} | |
| for path, leaf in _flatten(outputs): | |
| if isinstance(leaf, Packed): | |
| values[path] = leaf.layout.unpack(to_numpy(leaf.tensor)) | |
| else: | |
| values[path] = to_numpy(leaf) | |
| self._free_transient(outputs) | |
| pingpong = {s.name for s in self._slots.values() if s.pingpong} | |
| if pingpong and (ctx.writes & pingpong) == pingpong: | |
| self._phase ^= 1 | |
| if isinstance(outputs, Packed) or not isinstance(outputs, (Mapping, list, tuple)): | |
| return values[()] | |
| return _rebuild(outputs, values) | |
| # ------------------------------------------------------------------------------------------- state | |
| def phase(self) -> int: | |
| """Which buffer of each ping-pong state the next run reads (0 = the first buffer).""" | |
| return self._phase | |
| def state_buffer(self, name: str): | |
| """The device buffer the next run reads for state ``name`` (ping-pong: the buffer of the current | |
| :attr:`phase`).""" | |
| slot = self._slot(name, "state") | |
| return slot.buffers[self._phase] if slot.pingpong else slot.buffers[0] | |
| def reset_state(self, name: Optional[str] = None, value: Any = None) -> None: | |
| """Write the initial value (or ``value``) of state ``name`` (all states when ``None``) into the buffer the | |
| next run reads. ``value``: numpy / scalar / torch / ttnn host tensor (uploaded) or a ttnn *device* tensor of | |
| the same shape (``ttnn.copy`` on the device). Enqueued on CQ0, so it lands after any replay already | |
| enqueued.""" | |
| import ttnn | |
| self._check_open() | |
| names = [name] if name is not None else [s.name for s in self._slots.values() if s.kind == "state"] | |
| for n in names: | |
| slot = self._slot(n, "state") | |
| target = self.state_buffer(n) | |
| if isinstance(value, ttnn.Tensor) and value.storage_type() == ttnn.StorageType.DEVICE: | |
| _check_copy(f"reset_state({n!r})", value, slot) | |
| with self._eager(f"reset_state({n!r}) from a device tensor"): | |
| ttnn.copy(value, target) | |
| continue | |
| if value is None: | |
| host = slot.init | |
| elif isinstance(value, ttnn.Tensor): | |
| host = to_host_tensor(value, slot.dtype, slot.layout, shape=slot.shape) | |
| else: | |
| arr = to_numpy(value) if hasattr(value, "detach") else np.asarray(value) | |
| host = to_host_tensor(np.broadcast_to(arr, slot.shape), slot.dtype, slot.layout) | |
| ttnn.copy_host_to_device_tensor(host, target, cq_id=CQ_COMPUTE) | |
| def read_state(self, name: str) -> np.ndarray: | |
| """The value the next run will read for state ``name`` (blocking read on CQ0).""" | |
| return to_numpy(self.state_buffer(name)) | |
| def _bank(self, name: str, bank: int): | |
| slot = self._slot(name, "state") | |
| if not 0 <= bank < len(slot.banks): | |
| raise IndexError(f"{self.name}: state {name!r} has {len(slot.banks)} bank(s) (add_state(banks=...)), " | |
| f"not bank {bank}") | |
| return slot.banks[bank] | |
| def save_state(self, name: str, bank: int) -> None: | |
| """Copy the current value of state ``name`` into its ``bank`` (``ttnn.copy`` on CQ0, eager, after any | |
| replay already enqueued): the first half of a stream switch (PLAN.md D16; :class:`StreamBanks`).""" | |
| import ttnn | |
| self._check_open() | |
| with self._eager(f"save_state({name!r})"): | |
| ttnn.copy(self.state_buffer(name), self._bank(name, bank)) | |
| def load_state(self, name: str, bank: int) -> None: | |
| """Copy ``bank`` back into the buffer the next run reads for state ``name`` (``ttnn.copy`` on CQ0).""" | |
| import ttnn | |
| self._check_open() | |
| with self._eager(f"load_state({name!r})"): | |
| ttnn.copy(self._bank(name, bank), self.state_buffer(name)) | |
| # ------------------------------------------------------------------------------------------- misc | |
| def trace_ids(self) -> Dict[Tuple[str, int], Any]: | |
| """``{(variant, phase): trace_id}`` for profiling scripts.""" | |
| return {k: t.trace_id for k, t in self._traces.items()} | |
| def describe(self) -> Dict[str, Any]: | |
| """A JSON-able summary for ``model.info`` / OPT_BASELINE.""" | |
| def spec(s: _Slot) -> Dict[str, Any]: | |
| return {"shape": list(s.shape), "dtype": dtype_name(s.dtype), "layout": str(s.layout).rsplit(".", 1)[-1], | |
| **({"pingpong": s.pingpong, "banks": len(s.banks)} if s.kind == "state" else {})} | |
| entries = None | |
| if hasattr(self.device, "num_program_cache_entries"): | |
| entries = int(self.device.num_program_cache_entries()) | |
| return { | |
| "name": self.name, "num_command_queues": self.num_command_queues, "stage_inputs": self.stage_inputs, | |
| "warmup_runs": self.warmup_runs, "phases": self.phases, "alloc_tracking": self.alloc_tracking, | |
| "variants": sorted(self._variants), "traces": len(self._traces), | |
| "inputs": {s.name: spec(s) for s in self._slots.values() if s.kind == "input"}, | |
| "params": {s.name: spec(s) for s in self._slots.values() if s.kind == "param"}, | |
| "states": {s.name: spec(s) for s in self._slots.values() if s.kind == "state"}, | |
| "outputs": {s.name: spec(s) for s in self._slots.values() if s.kind == "output"}, | |
| "timings_ms": {k: {kk: round(vv, 3) for kk, vv in v.items()} for k, v in self.timings_ms.items()}, | |
| "program_cache_entries": entries, | |
| } | |
| def _release_traces(self) -> None: | |
| traces, self._traces = self._traces, {} | |
| for trace in traces.values(): | |
| self._safe_release(trace.trace_id) | |
| self._last_phase.clear() | |
| self._last_variant = None | |
| self._captures = 0 | |
| self._deallocate_outputs(tensor for trace in traces.values() for _, tensor, _ in trace.leaves) | |
| def release(self) -> None: | |
| """Release every trace and deallocate the persistent tensors. Idempotent; the persistent tensors are freed | |
| even when the device sync or a trace release raises (the error propagates afterwards).""" | |
| import ttnn | |
| if self._closed: | |
| return | |
| self._closed = True | |
| try: | |
| try: | |
| ttnn.synchronize_device(self.device) | |
| finally: | |
| self._release_traces() | |
| finally: | |
| slots, self._slots = list(self._slots.values()), {} | |
| for slot in slots: | |
| for t in slot.tensors(): | |
| if t.is_allocated(): | |
| ttnn.deallocate(t) | |
| def _check_open(self) -> None: | |
| if self._closed: | |
| raise RuntimeError(f"{self.name}: the TraceRunner was released") | |
| def __enter__(self) -> "TraceRunner": | |
| return self | |
| def __exit__(self, *exc) -> None: | |
| self.release() | |
| class StreamBanks: | |
| """Per-stream device state of a temporal model (PLAN.md D16) on top of the states of a :class:`TraceRunner`. | |
| The traces read and write one set of state buffers: the *active* stream's. With ``max_streams > 1`` every known | |
| stream id owns one bank of each state (``add_state(..., banks=max_streams)``) and a switch saves the active | |
| stream into its bank and loads the selected one (``ttnn.copy`` on CQ0, after the replays already enqueued; | |
| probe P13 measured ~80 us eager for a StreamPETR-size state). A new id beyond ``max_streams`` is refused | |
| (``on_full="reject"``: :class:`~.io.InputError`, HTTP 400) or takes over the least recently used stream | |
| (``"evict"``; with ``max_streams=1`` that simply restarts the one state). A stream starts fresh -- its states | |
| reset to their ``init`` values -- when it is new, when ``reset=True``, or when its timestamp goes backwards or | |
| jumps by more than ``max_gap_s``; :meth:`select` returns True then, so the model can also set its first-frame | |
| RT-dev params (BEVFormer ``use_prev_bev=0``, BEVDet ``flag``, ...). Create it after the last ``capture()`` (a | |
| capture resets every state) and call :meth:`select` under the model lock, before the run of each frame:: | |
| S_MAX = 1 # first publish (D16) | |
| runner.add_state("prev_bev", shape=(1, 1, 22500, 256), dtype="bfloat16", pingpong=True, | |
| banks=S_MAX if S_MAX > 1 else 0) | |
| streams = StreamBanks(runner, ["prev_bev"], max_streams=S_MAX, on_full="evict", max_gap_s=2.0) | |
| ... | |
| fresh = streams.select(stream.get("id", "default"), reset=stream.get("reset", False), | |
| timestamp_s=stream.get("timestamp_s")) | |
| out = runner("frame", inputs=..., params={"use_prev_bev": 0.0 if fresh else 1.0}) | |
| """ | |
| def __init__(self, runner: TraceRunner, states: Sequence[str], *, max_streams: int = 1, on_full: str = "reject", | |
| max_gap_s: Optional[float] = None): | |
| if int(max_streams) < 1: | |
| raise ValueError("max_streams must be >= 1") | |
| if on_full not in ("reject", "evict"): | |
| raise ValueError(f"on_full={on_full!r}: expected 'reject' or 'evict'") | |
| self.runner = runner | |
| self.states = tuple(states) | |
| if not self.states: | |
| raise ValueError("StreamBanks needs at least one state") | |
| self.max_streams = int(max_streams) | |
| for name in self.states: | |
| banks = len(runner._slot(name, "state").banks) | |
| if self.max_streams > 1 and banks < self.max_streams: | |
| raise ValueError(f"state {name!r} has {banks} bank(s); StreamBanks(max_streams={self.max_streams}) " | |
| f"needs add_state(..., banks={self.max_streams})") | |
| self.on_full = on_full | |
| self.max_gap_s = None if max_gap_s is None else float(max_gap_s) | |
| self.active: Optional[str] = None | |
| self._bank: Dict[str, int] = {} # stream id -> bank index (max_streams > 1) | |
| self._last_t: Dict[str, float] = {} # stream id -> timestamp of its last frame | |
| self._used: Dict[str, int] = {} # stream id -> last use (LRU clock) | |
| self._clock = 0 | |
| def streams(self) -> List[str]: | |
| """Known stream ids, most recently used first.""" | |
| return sorted(self._used, key=lambda s: -self._used[s]) | |
| def _save_active(self) -> None: | |
| if self.active is not None and self.active in self._bank: | |
| for name in self.states: | |
| self.runner.save_state(name, self._bank[self.active]) | |
| def forget(self, stream_id: str) -> None: | |
| """Drop a stream (its bank becomes free; if it was active, the next :meth:`select` starts fresh).""" | |
| sid = str(stream_id) | |
| self._used.pop(sid, None) | |
| self._bank.pop(sid, None) | |
| self._last_t.pop(sid, None) | |
| if self.active == sid: | |
| self.active = None | |
| def select(self, stream_id: Any = "default", *, reset: bool = False, timestamp_s: Optional[float] = None) -> bool: | |
| """Make ``stream_id`` the active stream for the next run; returns True when its state starts fresh.""" | |
| sid = str(stream_id) | |
| fresh = bool(reset) | |
| if sid != self.active: | |
| if sid in self._used: # a known stream parked in its bank | |
| self._save_active() | |
| if not fresh: | |
| for name in self.states: | |
| self.runner.load_state(name, self._bank[sid]) | |
| else: # a new stream | |
| if len(self._used) >= self.max_streams: | |
| if self.on_full == "reject": | |
| raise InputError(f"stream {sid!r}: this model keeps device state for {self.max_streams} " | |
| f"stream(s), in use by {self.streams}; reuse an id") | |
| self.forget(min(self._used, key=self._used.__getitem__)) | |
| self._save_active() | |
| if self.max_streams > 1: | |
| self._bank[sid] = min(set(range(self.max_streams)) - set(self._bank.values())) | |
| fresh = True | |
| self.active = sid | |
| if not fresh and timestamp_s is not None and self.max_gap_s is not None and sid in self._last_t: | |
| dt = float(timestamp_s) - self._last_t[sid] | |
| fresh = dt < 0 or dt > self.max_gap_s | |
| if fresh: | |
| for name in self.states: | |
| self.runner.reset_state(name) | |
| if timestamp_s is not None: | |
| self._last_t[sid] = float(timestamp_s) | |
| elif fresh: | |
| self._last_t.pop(sid, None) | |
| self._clock += 1 | |
| self._used[sid] = self._clock | |
| return fresh | |
| def describe(self) -> Dict[str, Any]: | |
| """JSON-able summary for ``model.info``.""" | |
| return {"max_streams": self.max_streams, "on_full": self.on_full, "max_gap_s": self.max_gap_s, | |
| "states": list(self.states), "active": self.active, "streams": self.streams} | |