File size: 20,899 Bytes
2415c4c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 | # 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
@dataclass(frozen=True)
class TraceKey:
"""Full content digest for one operation-produced trace signature."""
digest: str
def __post_init__(self) -> None:
validate_sha256_digest(self.digest, "trace")
@classmethod
def from_signature(cls, signature: Any) -> "TraceKey":
return cls(signature_digest(_TRACE_KEY_DOMAIN, _TRACE_KEY_SCHEMA_VERSION, signature))
@dataclass
class PersistentInputs:
"""Trace-owned persistent replay inputs opaque to public runtime APIs."""
values: Any
@dataclass(frozen=True)
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
@dataclass(frozen=True)
class RefreshDecision:
full: bool
page_table: bool
fields: tuple[str, ...]
@dataclass
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)
@dataclass(frozen=True)
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")
@dataclass
class TraceRecord:
signature: Any
operation: str
artifact: TraceArtifact | None = None
@dataclass
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
@property
def trace_active(self) -> bool:
return self._activated
@property
def replay_count(self) -> int:
"""Return successfully submitted trace replays across all operations."""
return self._replay_count
@property
def replay_counts(self) -> dict[str, int]:
"""Return a snapshot of successfully submitted replays by operation."""
return dict(self._replay_counts)
@property
def trace_count(self) -> int:
"""Return the number of semantic hidden traces in the registry."""
return len(self._traces)
@property
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)
|