Download code/models/common/llm_runtime/trace_compiler.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 20.9 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/trace_compiler.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/llm_runtime/trace_compiler.py
-
curl -L -o trace_compiler.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/trace_compiler.py
20.9 kB
| # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC | |
| # SPDX-License-Identifier: Apache-2.0 | |
| """Trace capture, replay, and persistent-resource ownership.""" | |
| from __future__ import annotations | |
| import ctypes | |
| from collections.abc import Callable | |
| from dataclasses import dataclass, field | |
| from typing import Any | |
| from loguru import logger | |
| from ttnn.tools import trace_allocation_tracker | |
| import ttnn | |
| from models.common.llm_runtime.program_compiler import ( | |
| ProgramCompiler, | |
| ProgramKey, | |
| signature_digest, | |
| validate_sha256_digest, | |
| ) | |
| from models.common.llm_runtime.tensor_resources import ( | |
| TensorResourceOrphan, | |
| attach_cleanup_failures, | |
| best_effort_deallocate_owned_tensors, | |
| raise_cleanup_failures, | |
| release_orphans, | |
| ) | |
| _TRACE_KEY_DOMAIN = "tttv2.llm-runtime.trace" | |
| _TRACE_KEY_SCHEMA_VERSION = 1 | |
| class TraceKey: | |
| """Full content digest for one operation-produced trace signature.""" | |
| digest: str | |
| def __post_init__(self) -> None: | |
| validate_sha256_digest(self.digest, "trace") | |
| def from_signature(cls, signature: Any) -> "TraceKey": | |
| return cls(signature_digest(_TRACE_KEY_DOMAIN, _TRACE_KEY_SCHEMA_VERSION, signature)) | |
| class PersistentInputs: | |
| """Trace-owned persistent replay inputs opaque to public runtime APIs.""" | |
| values: Any | |
| class InputRefreshPolicy: | |
| every_replay: tuple[str, ...] = () | |
| full_on_batch_reset: bool = True | |
| full_on_graph_switch: bool = True | |
| full_without_device_feedback: bool = True | |
| refresh_page_table_on_change: bool = True | |
| class RefreshDecision: | |
| full: bool | |
| page_table: bool | |
| fields: tuple[str, ...] | |
| class TraceArtifact: | |
| trace_id: int | |
| persistent_inputs: PersistentInputs | |
| outputs: Any | |
| refresh_policy: InputRefreshPolicy | |
| trace_released: bool = False | |
| deallocated_tensor_ids: set[int] = field(default_factory=set, repr=False) | |
| class TraceCapturePlan: | |
| """Operation-produced specification for one trace-capable compiled program.""" | |
| program_key: ProgramKey | |
| trace_signature: Any | |
| operation: str | |
| prepare_inputs: Callable[[], PersistentInputs | Any] | |
| capture: Callable[[PersistentInputs], Any] | |
| refresh_policy: InputRefreshPolicy = InputRefreshPolicy() | |
| schema_fingerprint: Any = None | |
| prepare_workspace: Callable[[], Any] | None = None | |
| workspace_fingerprint: Any = None | |
| prime: Callable[[PersistentInputs], Any] | None = None | |
| release_prime_output: Callable[[Any], list[BaseException]] | None = None | |
| def __post_init__(self) -> None: | |
| if self.operation not in ("prefill", "decode"): | |
| raise ValueError(f"Unsupported trace operation: {self.operation!r}") | |
| if (self.prime is None) is not (self.release_prime_output is None): | |
| raise ValueError("trace capture prime and output releaser must be configured together") | |
| class TraceRecord: | |
| signature: Any | |
| operation: str | |
| artifact: TraceArtifact | None = None | |
| class TraceAliasRecord: | |
| """Program-local postprocess state kept outside the shared hidden trace.""" | |
| trace_key: TraceKey | |
| workspace_fingerprint: Any | |
| prepare_workspace: Callable[[], Any] | None = field(default=None, repr=False) | |
| workspace: Any = None | |
| deallocated_tensor_ids: set[int] = field(default_factory=set, repr=False) | |
| class TraceCompiler: | |
| """Register, capture, replay, and release traces for compiled programs. | |
| ``TracedExecutor.compile_*`` first compiles an eager program and calls | |
| `register_capture_plan`. `WarmupCoordinator` calls | |
| `capture_all` only after the complete configured program set exists. | |
| Forward execution then calls `replay` with operation-owned refresh | |
| logic. The compiler owns trace artifacts and persistent inputs, while the | |
| composed ``ProgramCompiler`` remains the sole program registry. | |
| """ | |
| def __init__(self, program_compiler: ProgramCompiler): | |
| if not isinstance(program_compiler, ProgramCompiler): | |
| raise TypeError("program_compiler must be a ProgramCompiler") | |
| self.program_compiler = program_compiler | |
| self.mesh_device = program_compiler.mesh_device | |
| self._traces: dict[TraceKey, TraceRecord] = {} | |
| self._plans: dict[TraceKey, TraceCapturePlan] = {} | |
| self._program_to_trace: dict[ProgramKey, TraceKey] = {} | |
| self._aliases: dict[ProgramKey, TraceAliasRecord] = {} | |
| self._rollback_orphans: list[TensorResourceOrphan] = [] | |
| self._capture_in_progress = False | |
| self._activated = False | |
| self._released = False | |
| self._previous_replay_key: TraceKey | None = None | |
| self._replay_count = 0 | |
| self._replay_counts = {"prefill": 0, "decode": 0} | |
| # Public API | |
| def trace_active(self) -> bool: | |
| return self._activated | |
| def replay_count(self) -> int: | |
| """Return successfully submitted trace replays across all operations.""" | |
| return self._replay_count | |
| def replay_counts(self) -> dict[str, int]: | |
| """Return a snapshot of successfully submitted replays by operation.""" | |
| return dict(self._replay_counts) | |
| def trace_count(self) -> int: | |
| """Return the number of semantic hidden traces in the registry.""" | |
| return len(self._traces) | |
| def trace_association_count(self) -> int: | |
| """Return the number of compiled-program aliases associated to traces.""" | |
| return len(self._program_to_trace) | |
| def registered_coverage(self, operation: str) -> tuple[tuple[TraceKey, Any], ...]: | |
| """Return registered trace keys/signatures for one operation.""" | |
| if operation not in ("prefill", "decode"): | |
| raise ValueError(f"Unsupported trace operation: {operation!r}") | |
| return tuple( | |
| (trace_key, record.signature) for trace_key, record in self._traces.items() if record.operation == operation | |
| ) | |
| def get(self, key: TraceKey) -> TraceRecord | None: | |
| """Return the record needed to finish an operation-specific replay.""" | |
| return self._traces.get(key) | |
| def trace_key_for_program(self, program_key: ProgramKey) -> TraceKey | None: | |
| """Return the registered trace association for one compiled program.""" | |
| return self._program_to_trace.get(program_key) | |
| def workspace_for_program(self, program_key: ProgramKey) -> Any: | |
| """Return one alias's postprocess workspace after capture allocation.""" | |
| alias = self._aliases.get(program_key) | |
| if alias is None: | |
| raise RuntimeError(f"Program key {program_key.digest} has no trace alias workspace") | |
| return alias.workspace | |
| def register_capture_plan(self, plan: TraceCapturePlan) -> TraceKey: | |
| """Validate a compiled source and register one explicit trace association.""" | |
| self._ensure_live() | |
| if self._capture_in_progress or self._activated: | |
| raise RuntimeError("Cannot register trace capture plans during capture or after trace activation") | |
| self.program_compiler.require_compiled(plan.program_key) | |
| trace_key = TraceKey.from_signature(plan.trace_signature) | |
| existing_association = self._program_to_trace.get(plan.program_key) | |
| if existing_association is not None and existing_association != trace_key: | |
| raise ValueError(f"Program key {plan.program_key.digest} already has a different trace association") | |
| existing_alias = self._aliases.get(plan.program_key) | |
| if existing_alias is not None and existing_alias.workspace_fingerprint != plan.workspace_fingerprint: | |
| raise ValueError( | |
| f"Program key {plan.program_key.digest} was registered with a different workspace fingerprint" | |
| ) | |
| record = self._traces.get(trace_key) | |
| if record is None: | |
| record = TraceRecord( | |
| signature=plan.trace_signature, | |
| operation=plan.operation, | |
| ) | |
| self._traces[trace_key] = record | |
| self._plans[trace_key] = plan | |
| else: | |
| if record.signature != plan.trace_signature: | |
| raise RuntimeError(f"Trace key collision for digest {trace_key.digest}: retained signature differs") | |
| if record.operation != plan.operation: | |
| raise RuntimeError(f"Trace key collision for digest {trace_key.digest}: operation differs") | |
| if self._plans[trace_key].refresh_policy != plan.refresh_policy: | |
| raise ValueError(f"Trace key {trace_key.digest} was registered with a different refresh policy") | |
| if self._plans[trace_key].schema_fingerprint != plan.schema_fingerprint: | |
| raise ValueError(f"Trace key {trace_key.digest} was registered with a different schema fingerprint") | |
| self._program_to_trace[plan.program_key] = trace_key | |
| if existing_alias is None: | |
| self._aliases[plan.program_key] = TraceAliasRecord( | |
| trace_key=trace_key, | |
| workspace_fingerprint=plan.workspace_fingerprint, | |
| prepare_workspace=plan.prepare_workspace, | |
| ) | |
| return trace_key | |
| def capture_all(self) -> None: | |
| """Allocate every persistent input before beginning the first capture.""" | |
| self._ensure_live() | |
| if self._activated: | |
| return | |
| if self._capture_in_progress: | |
| raise RuntimeError("Trace capture is already in progress") | |
| if not self._plans: | |
| return | |
| if self.program_compiler.compile_orphan_count: | |
| raise RuntimeError("Cannot capture while unreleased compile outputs remain") | |
| prepared: dict[TraceKey, tuple[PersistentInputs, TraceCapturePlan]] = {} | |
| captured_keys: set[TraceKey] = set() | |
| self._capture_in_progress = True | |
| try: | |
| for trace_key, plan in self._plans.items(): | |
| self.program_compiler.require_compiled(plan.program_key) | |
| values = plan.prepare_inputs() | |
| persistent = values if isinstance(values, PersistentInputs) else PersistentInputs(values) | |
| prepared[trace_key] = (persistent, plan) | |
| for program_key, alias in self._aliases.items(): | |
| if alias.prepare_workspace is not None: | |
| alias.workspace = alias.prepare_workspace() | |
| capture_order = sorted( | |
| prepared, | |
| key=lambda trace_key: self._traces[trace_key].operation == "prefill", | |
| ) | |
| for trace_key in capture_order: | |
| persistent, plan = prepared[trace_key] | |
| record = self._traces[trace_key] | |
| # Program signatures intentionally describe padded trace | |
| # identity, not active-row cardinality. Selected operation | |
| # plans therefore prime their exact persistent-input body | |
| # immediately before capturing that same body. No unrelated | |
| # trace can perturb allocator/program state between the prime | |
| # and ``begin_trace_capture``. | |
| if plan.prime is not None: | |
| prime_output = None | |
| try: | |
| prime_output = plan.prime(persistent) | |
| ttnn.synchronize_device(self.mesh_device) | |
| except BaseException as primary: | |
| cleanup_failures = plan.release_prime_output(prime_output) | |
| try: | |
| ttnn.synchronize_device(self.mesh_device) | |
| except BaseException as error: | |
| cleanup_failures.append(error) | |
| attach_cleanup_failures(primary, cleanup_failures) | |
| raise | |
| release_failures = plan.release_prime_output(prime_output) | |
| try: | |
| ttnn.synchronize_device(self.mesh_device) | |
| except BaseException as error: | |
| release_failures.append(error) | |
| if release_failures: | |
| raise_cleanup_failures(release_failures) | |
| logger.info(f"Primed {plan.operation} trace capture body: signature={plan.trace_signature!r}") | |
| self.program_compiler.set_trace_capture_in_progress(True) | |
| # Whatever the model allocates between begin/end_trace_capture belongs to the trace | |
| # being recorded and must stay allocated for replay; recording N traces means capture | |
| # N runs while 1..N-1 are live, which ordering cannot avoid. Acknowledge the window | |
| # (no-op unless TT_METAL_TRACE_ALLOC_TRACKING=1), as tt_transformers' generator does. | |
| with trace_allocation_tracker.corruptible_allocation_scope(self.mesh_device): | |
| trace_id = ttnn.begin_trace_capture(self.mesh_device, cq_id=0) | |
| outputs = None | |
| capture_ended = False | |
| try: | |
| outputs = plan.capture(persistent) | |
| ttnn.end_trace_capture(self.mesh_device, trace_id, cq_id=0) | |
| capture_ended = True | |
| ttnn.synchronize_device(self.mesh_device) | |
| except BaseException as primary: | |
| cleanup_failures = [] | |
| if not capture_ended: | |
| try: | |
| ttnn.end_trace_capture(self.mesh_device, trace_id, cq_id=0) | |
| except BaseException as error: | |
| cleanup_failures.append(error) | |
| record.artifact = TraceArtifact( | |
| trace_id=trace_id, | |
| persistent_inputs=persistent, | |
| outputs=outputs, | |
| refresh_policy=plan.refresh_policy, | |
| ) | |
| captured_keys.add(trace_key) | |
| cleanup_failures.extend(self._release_trace(record)) | |
| attach_cleanup_failures(primary, cleanup_failures) | |
| raise | |
| record.artifact = TraceArtifact( | |
| trace_id=trace_id, | |
| persistent_inputs=persistent, | |
| outputs=outputs, | |
| refresh_policy=plan.refresh_policy, | |
| ) | |
| logger.info(f"Captured {plan.operation} trace: signature={plan.trace_signature!r}") | |
| captured_keys.add(trace_key) | |
| self.program_compiler.set_trace_capture_in_progress(False) | |
| ttnn.synchronize_device(self.mesh_device) | |
| self._capture_in_progress = False | |
| self.program_compiler.set_trace_capture_in_progress(False) | |
| self._activated = True | |
| self.program_compiler.set_trace_active(True) | |
| _trim_host_allocator() | |
| except BaseException as primary: | |
| cleanup_failures = self._release_trace_resources() | |
| cleanup_failures.extend(self._release_alias_workspaces()) | |
| for trace_key, (persistent, _) in prepared.items(): | |
| if trace_key in captured_keys: | |
| continue | |
| orphan = TensorResourceOrphan(persistent.values) | |
| orphan_failures = best_effort_deallocate_owned_tensors( | |
| orphan.values, | |
| orphan.deallocated_tensor_ids, | |
| ) | |
| cleanup_failures.extend(orphan_failures) | |
| if orphan_failures: | |
| self._rollback_orphans.append(orphan) | |
| self._activated = ( | |
| bool(self._rollback_orphans) | |
| or any(record.artifact is not None for record in self._traces.values()) | |
| or any(alias.workspace is not None for alias in self._aliases.values()) | |
| ) | |
| self._capture_in_progress = False | |
| self.program_compiler.set_trace_capture_in_progress(False) | |
| self.program_compiler.set_trace_active(self._activated) | |
| attach_cleanup_failures(primary, cleanup_failures) | |
| raise | |
| def replay( | |
| self, | |
| program_key: ProgramKey, | |
| refresh_inputs: Callable[[TraceArtifact, RefreshDecision], None], | |
| *, | |
| reset_batch: bool = False, | |
| device_feedback_enabled: bool = False, | |
| feedback_compatible: bool = False, | |
| page_table_changed: bool = False, | |
| ) -> Any: | |
| """Refresh persistent inputs and enqueue one non-blocking trace replay.""" | |
| self._ensure_live() | |
| self.program_compiler.require_compiled(program_key) | |
| trace_key = self._program_to_trace.get(program_key) | |
| if trace_key is None: | |
| raise RuntimeError(f"Program key {program_key.digest} has no trace association") | |
| record = self._traces[trace_key] | |
| artifact = record.artifact | |
| if artifact is None: | |
| raise RuntimeError(f"Trace key {trace_key.digest} has not been captured") | |
| policy = artifact.refresh_policy | |
| switched = self._previous_replay_key != trace_key | |
| full = ( | |
| (policy.full_on_batch_reset and reset_batch) | |
| or (policy.full_on_graph_switch and switched) | |
| or (policy.full_without_device_feedback and not (device_feedback_enabled and feedback_compatible)) | |
| ) | |
| decision = RefreshDecision( | |
| full=full, | |
| page_table=policy.refresh_page_table_on_change and page_table_changed, | |
| fields=policy.every_replay, | |
| ) | |
| refresh_inputs(artifact, decision) | |
| ttnn.execute_trace(self.mesh_device, artifact.trace_id, cq_id=0, blocking=False) | |
| self._replay_count += 1 | |
| self._replay_counts[record.operation] += 1 | |
| self._previous_replay_key = trace_key | |
| return artifact.outputs | |
| def cleanup(self) -> None: | |
| """Release traces and persistent inputs, then reopen the program gate.""" | |
| if self._released: | |
| return | |
| failures = self._release_trace_resources() | |
| failures.extend(self._release_alias_workspaces()) | |
| failures.extend(release_orphans(self._rollback_orphans)) | |
| if failures: | |
| self._activated = True | |
| self.program_compiler.set_trace_active(True) | |
| error = RuntimeError(f"Failed to release {len(failures)} trace resource(s)") | |
| attach_cleanup_failures(error, failures) | |
| raise error from failures[0] | |
| self._capture_in_progress = False | |
| self._activated = False | |
| self._previous_replay_key = None | |
| self.program_compiler.set_trace_capture_in_progress(False) | |
| self.program_compiler.set_trace_active(False) | |
| self._released = True | |
| # Private implementation | |
| def _release_trace_resources(self) -> list[BaseException]: | |
| failures: list[BaseException] = [] | |
| for record in self._traces.values(): | |
| failures.extend(self._release_trace(record)) | |
| return failures | |
| def _release_alias_workspaces(self) -> list[BaseException]: | |
| failures: list[BaseException] = [] | |
| for alias in self._aliases.values(): | |
| if alias.workspace is None: | |
| continue | |
| alias_failures = best_effort_deallocate_owned_tensors( | |
| alias.workspace, | |
| alias.deallocated_tensor_ids, | |
| ) | |
| failures.extend(alias_failures) | |
| if not alias_failures: | |
| alias.workspace = None | |
| return failures | |
| def _release_trace(self, record: TraceRecord) -> list[BaseException]: | |
| artifact = record.artifact | |
| if artifact is None: | |
| return [] | |
| if not artifact.trace_released: | |
| try: | |
| ttnn.release_trace(self.mesh_device, artifact.trace_id) | |
| except BaseException as error: | |
| return [error] | |
| artifact.trace_released = True | |
| failures = best_effort_deallocate_owned_tensors( | |
| (artifact.persistent_inputs.values, artifact.outputs), | |
| artifact.deallocated_tensor_ids, | |
| ) | |
| if failures: | |
| return failures | |
| record.artifact = None | |
| return [] | |
| def _ensure_live(self) -> None: | |
| if self._released: | |
| raise RuntimeError("TraceCompiler has been released") | |
| def _trim_host_allocator() -> None: | |
| """Return released trace-capture staging arenas to the OS when supported.""" | |
| try: | |
| malloc_trim = ctypes.CDLL(None).malloc_trim | |
| except (AttributeError, OSError): | |
| return | |
| malloc_trim.argtypes = (ctypes.c_size_t,) | |
| malloc_trim.restype = ctypes.c_int | |
| malloc_trim(0) | |