# SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC # SPDX-License-Identifier: Apache-2.0 from dataclasses import dataclass from itertools import permutations import pytest import torch import models.common.llm_runtime.trace_compiler as trace_compiler_module import ttnn from models.common.llm_runtime.decode import DecodeDeviceInputs, DecodePersistentInputs from models.common.llm_runtime.prefill.inputs import PrefillDeviceInputs, PrefillPositionInputs from models.common.llm_runtime.prefill.trace import PrefillHiddenPersistentInputs, PrefillReplayState from models.common.llm_runtime.program_compiler import ProgramCompiler from models.common.llm_runtime.trace_compiler import ( InputRefreshPolicy, PersistentInputs, TraceCapturePlan, TraceCompiler, ) @dataclass(frozen=True) class _Signature: kind: str variant: int @property def key_material(self): return (("kind", self.kind), ("variant", self.variant)) def _patch_backend(monkeypatch, events): next_trace_id = iter(range(100, 200)) monkeypatch.setattr(ttnn, "synchronize_device", lambda mesh: events.append(("sync", mesh))) monkeypatch.setattr( ttnn, "begin_trace_capture", lambda mesh, cq_id: events.append(("begin", mesh, cq_id)) or next(next_trace_id), ) monkeypatch.setattr( ttnn, "end_trace_capture", lambda mesh, trace_id, cq_id: events.append(("end", trace_id, cq_id)), ) monkeypatch.setattr( ttnn, "execute_trace", lambda mesh, trace_id, cq_id, blocking: events.append(("execute", trace_id, cq_id, blocking)), ) monkeypatch.setattr(ttnn, "release_trace", lambda mesh, trace_id: events.append(("release", trace_id))) monkeypatch.setattr(trace_compiler_module, "_trim_host_allocator", lambda: events.append(("trim",))) def _compiled_program(program_compiler, monkeypatch, variant): monkeypatch.setattr(ttnn, "synchronize_device", lambda mesh: None) return program_compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) def _plan(program, variant, events, *, operation="decode", policy=InputRefreshPolicy()): return TraceCapturePlan( program_key=program.key, trace_signature=_Signature("trace", variant), operation=operation, prepare_inputs=lambda: events.append(("prepare", variant)) or (), capture=lambda persistent: events.append(("capture", variant)) or torch.zeros(1), refresh_policy=policy, ) def test_trace_compiler_retains_exact_program_compiler_and_separate_registries(monkeypatch): compiler = ProgramCompiler("mesh", lambda: object()) program = _compiled_program(compiler, monkeypatch, 1) trace = TraceCompiler(compiler) trace_key = trace.register_capture_plan(_plan(program, 1, [])) assert trace.program_compiler is compiler assert trace.mesh_device is compiler.mesh_device assert trace.trace_key_for_program(program.key) == trace_key assert trace.get(trace_key) is not None assert compiler.require_compiled(program.key) is program assert not hasattr(program, "artifact") assert not hasattr(trace, "programs") def test_trace_aliases_share_one_artifact_without_copying_program_records(monkeypatch): events = [] _patch_backend(monkeypatch, events) compiler = ProgramCompiler("mesh", lambda: object()) first = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1)) second = compiler.compile(_Signature("program", 2), lambda context: torch.zeros(1)) trace = TraceCompiler(compiler) shared_signature = _Signature("trace", 1) first_key = trace.register_capture_plan( TraceCapturePlan(first.key, shared_signature, "decode", lambda: (), lambda persistent: torch.zeros(1)) ) second_key = trace.register_capture_plan( TraceCapturePlan(second.key, shared_signature, "decode", lambda: (), lambda persistent: torch.zeros(1)) ) trace.capture_all() assert first_key == second_key assert trace.get(first_key) is not None assert trace.trace_key_for_program(first.key) == first_key assert trace.trace_key_for_program(second.key) == first_key assert compiler.require_compiled(first.key) is first assert compiler.require_compiled(second.key) is second assert [event[0] for event in events].count("begin") == 1 assert trace.trace_count == 1 assert trace.trace_association_count == 2 assert trace.registered_coverage("decode") == ((first_key, shared_signature),) assert trace.registered_coverage("prefill") == () @pytest.mark.parametrize("order", tuple(permutations(("logits", "argmax", "topk")))) def test_hidden_trace_alias_workspaces_are_registration_order_independent(monkeypatch, order): events = [] _patch_backend(monkeypatch, events) compiler = ProgramCompiler("mesh", lambda: object()) programs = { path: compiler.compile(_Signature("program", index), lambda context: torch.zeros(1)) for index, path in enumerate(("logits", "argmax", "topk"), 1) } trace = TraceCompiler(compiler) signature = _Signature("shared-hidden", 1) for path in order: trace.register_capture_plan( TraceCapturePlan( programs[path].key, signature, "prefill", lambda: "hidden-inputs", lambda persistent: "hidden-output", schema_fingerprint=("hidden-v2",), prepare_workspace=lambda path=path: f"{path}-workspace", workspace_fingerprint=("postprocess-v1", path), ) ) trace.capture_all() assert [event[0] for event in events].count("begin") == 1 assert {path: trace.workspace_for_program(program.key) for path, program in programs.items()} == { "logits": "logits-workspace", "argmax": "argmax-workspace", "topk": "topk-workspace", } def test_trace_alias_rejects_a_mismatched_persistent_schema(monkeypatch, expect_error): compiler = ProgramCompiler("mesh", lambda: object()) first = _compiled_program(compiler, monkeypatch, 1) second = _compiled_program(compiler, monkeypatch, 2) trace = TraceCompiler(compiler) signature = _Signature("trace", 1) trace.register_capture_plan( TraceCapturePlan(first.key, signature, "prefill", lambda: (), lambda _: (), schema_fingerprint=("v1",)) ) with expect_error(ValueError, "different schema fingerprint"): trace.register_capture_plan( TraceCapturePlan(second.key, signature, "prefill", lambda: (), lambda _: (), schema_fingerprint=("v2",)) ) def test_capture_allocates_every_input_before_capture_and_coordinates_gates(monkeypatch, expect_error): events = [] _patch_backend(monkeypatch, events) compiler = ProgramCompiler("mesh", lambda: object()) programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)] trace = TraceCompiler(compiler) def capture_with_gate(persistent): events.append(("capture", 1)) with expect_error(RuntimeError, "capture is in progress"): compiler.compile(_Signature("program", 3), lambda context: torch.zeros(1)) return torch.zeros(1) trace.register_capture_plan( TraceCapturePlan( programs[0].key, _Signature("trace", 1), "decode", lambda: events.append(("prepare", 1)) or (), capture_with_gate, ) ) trace.register_capture_plan(_plan(programs[1], 2, events)) events.clear() trace.capture_all() first_begin = next(index for index, event in enumerate(events) if event[0] == "begin") assert events[:first_begin] == [("prepare", 1), ("prepare", 2)] assert trace.trace_active and compiler.trace_active assert not compiler.trace_capture_in_progress with expect_error(RuntimeError, "after trace activation"): compiler.compile(_Signature("program", 3), lambda context: torch.zeros(1)) def test_opt_in_capture_prime_runs_after_allocations_and_before_capture_gate(monkeypatch): events = [] _patch_backend(monkeypatch, events) compiler = ProgramCompiler("mesh", lambda: object()) program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1)) trace = TraceCompiler(compiler) def prime(persistent): assert not compiler.trace_capture_in_progress events.append(("prime", persistent)) return "prime-output" def release_prime_output(output): events.append(("release-prime", output)) return [] trace.register_capture_plan( TraceCapturePlan( program.key, _Signature("trace", 1), "prefill", lambda: events.append(("prepare", 1)) or "persistent", lambda persistent: events.append(("capture", persistent)) or torch.zeros(1), prepare_workspace=lambda: events.append(("workspace", 1)) or "workspace", workspace_fingerprint=("workspace",), prime=prime, release_prime_output=release_prime_output, ) ) events.clear() trace.capture_all() first_begin = next(index for index, event in enumerate(events) if event[0] == "begin") assert events[:first_begin] == [ ("prepare", 1), ("workspace", 1), ("prime", PersistentInputs("persistent")), ("sync", "mesh"), ("release-prime", "prime-output"), ("sync", "mesh"), ] assert ("capture", PersistentInputs("persistent")) in events[first_begin:] def test_capture_prime_failure_rolls_back_without_beginning_trace(monkeypatch, expect_error): events = [] _patch_backend(monkeypatch, events) class OwnedTensor: pass released = [] monkeypatch.setattr(ttnn, "Tensor", OwnedTensor) monkeypatch.setattr(ttnn, "deallocate", released.append) compiler = ProgramCompiler("mesh", lambda: object()) program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1)) trace = TraceCompiler(compiler) persistent = OwnedTensor() trace.register_capture_plan( TraceCapturePlan( program.key, _Signature("trace", 1), "prefill", lambda: persistent, lambda _: pytest.fail("capture must not begin after prime failure"), prime=lambda _: (_ for _ in ()).throw(RuntimeError("prime failed")), release_prime_output=lambda _: [], ) ) with expect_error(RuntimeError, "prime failed"): trace.capture_all() assert not any(event[0] == "begin" for event in events) assert released == [persistent] assert not trace.trace_active and not compiler.trace_active def test_capture_prime_release_failure_still_synchronizes_before_rollback(monkeypatch, expect_error): events = [] _patch_backend(monkeypatch, events) compiler = ProgramCompiler("mesh", lambda: object()) program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1)) trace = TraceCompiler(compiler) release_error = RuntimeError("prime release failed") trace.register_capture_plan( TraceCapturePlan( program.key, _Signature("trace", 1), "prefill", lambda: (), lambda _: pytest.fail("capture must not begin after prime release failure"), prime=lambda _: events.append(("prime",)) or "prime-output", release_prime_output=lambda output: events.append(("release-prime", output)) or [release_error], ) ) events.clear() with expect_error(RuntimeError, "prime release failed") as caught: trace.capture_all() assert caught.value is release_error assert events[:4] == [ ("prime",), ("sync", "mesh"), ("release-prime", "prime-output"), ("sync", "mesh"), ] assert not any(event[0] == "begin" for event in events) def test_capture_orders_decode_before_prefill_after_allocating_every_input(monkeypatch): events = [] _patch_backend(monkeypatch, events) compiler = ProgramCompiler("mesh", lambda: object()) programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)] trace = TraceCompiler(compiler) trace.register_capture_plan(_plan(programs[0], 1, events, operation="prefill")) trace.register_capture_plan(_plan(programs[1], 2, events, operation="decode")) events.clear() trace.capture_all() first_begin = next(index for index, event in enumerate(events) if event[0] == "begin") assert events[:first_begin] == [("prepare", 1), ("prepare", 2)] assert [event for event in events if event[0] == "capture"] == [("capture", 2), ("capture", 1)] def test_capture_failure_rolls_back_traces_and_uncaptured_inputs(monkeypatch, expect_error): events = [] _patch_backend(monkeypatch, events) class OwnedTensor: pass released = [] monkeypatch.setattr(ttnn, "Tensor", OwnedTensor) monkeypatch.setattr(ttnn, "deallocate", released.append) compiler = ProgramCompiler("mesh", lambda: object()) programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)] trace = TraceCompiler(compiler) first_input, first_output, second_input = OwnedTensor(), OwnedTensor(), OwnedTensor() trace.register_capture_plan( TraceCapturePlan( programs[0].key, _Signature("trace", 1), "decode", lambda: first_input, lambda persistent: first_output, ) ) primary = RuntimeError("second capture failed") trace.register_capture_plan( TraceCapturePlan( programs[1].key, _Signature("trace", 2), "decode", lambda: second_input, lambda persistent: (_ for _ in ()).throw(primary), ) ) with expect_error(RuntimeError, "second capture failed") as caught: trace.capture_all() assert caught.value is primary assert released.count(first_input) == 1 assert released.count(first_output) == 1 assert released.count(second_input) == 1 assert not trace.trace_active and not compiler.trace_active for program in programs: trace_key = trace.trace_key_for_program(program.key) assert trace_key is not None assert trace.get(trace_key).artifact is None def test_incomplete_capture_rollback_keeps_program_gate_closed_until_cleanup(monkeypatch, expect_error): events = [] _patch_backend(monkeypatch, events) class OwnedTensor: pass retry = OwnedTensor() attempts = [] def deallocate(value): attempts.append(value) if value is retry and attempts.count(retry) == 1: raise RuntimeError("release once") monkeypatch.setattr(ttnn, "Tensor", OwnedTensor) monkeypatch.setattr(ttnn, "deallocate", deallocate) compiler = ProgramCompiler("mesh", lambda: object()) programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)] trace = TraceCompiler(compiler) trace.register_capture_plan( TraceCapturePlan(programs[0].key, _Signature("trace", 1), "decode", lambda: retry, lambda _: torch.zeros(1)) ) trace.register_capture_plan( TraceCapturePlan( programs[1].key, _Signature("trace", 2), "decode", lambda: (_ for _ in ()).throw(RuntimeError("prepare failed")), lambda _: torch.zeros(1), ) ) with expect_error(RuntimeError, "prepare failed"): trace.capture_all() assert trace.trace_active and compiler.trace_active with expect_error(RuntimeError, "after trace activation"): compiler.compile(_Signature("program", 3), lambda context: torch.zeros(1)) trace.cleanup() assert attempts.count(retry) == 2 assert not compiler.trace_active def test_replay_refresh_decisions_cover_first_replay_page_change_feedback_and_switch(monkeypatch): events = [] _patch_backend(monkeypatch, events) compiler = ProgramCompiler("mesh", lambda: object()) programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)] policy = InputRefreshPolicy(every_replay=("position", "sampling")) trace = TraceCompiler(compiler) for variant, program in enumerate(programs, 1): trace.register_capture_plan(_plan(program, variant, events, policy=policy)) trace.capture_all() decisions = [] trace.replay( programs[0].key, lambda artifact, decision: decisions.append(decision), device_feedback_enabled=True, feedback_compatible=True, ) trace.replay( programs[0].key, lambda artifact, decision: decisions.append(decision), device_feedback_enabled=True, feedback_compatible=True, page_table_changed=True, ) trace.replay( programs[1].key, lambda artifact, decision: decisions.append(decision), device_feedback_enabled=True, feedback_compatible=True, ) trace.replay( programs[1].key, lambda artifact, decision: decisions.append(decision), device_feedback_enabled=False, ) assert [decision.full for decision in decisions] == [True, False, True, True] assert [decision.page_table for decision in decisions] == [False, True, False, False] assert all(decision.fields == ("position", "sampling") for decision in decisions) assert [event[0] for event in events].count("execute") == 4 assert trace.replay_count == 4 assert trace.replay_counts == {"prefill": 0, "decode": 4} def test_replay_counters_increment_only_after_successful_submission(monkeypatch, expect_error): events = [] _patch_backend(monkeypatch, events) compiler = ProgramCompiler("mesh", lambda: object()) program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1)) trace = TraceCompiler(compiler) trace.register_capture_plan(_plan(program, 1, events, operation="prefill")) trace.capture_all() monkeypatch.setattr(ttnn, "execute_trace", lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("submit"))) with expect_error(RuntimeError, "submit"): trace.replay(program.key, lambda artifact, decision: None) assert trace.replay_count == 0 assert trace.replay_counts == {"prefill": 0, "decode": 0} def test_cleanup_retries_trace_release_before_deallocating_and_does_not_own_programs(monkeypatch, expect_error): events = [] _patch_backend(monkeypatch, events) class OwnedTensor: pass persistent, output = OwnedTensor(), OwnedTensor() deallocated = [] release_attempts = [] def release(mesh, trace_id): release_attempts.append(trace_id) if len(release_attempts) == 1: raise RuntimeError("trace release once") monkeypatch.setattr(ttnn, "Tensor", OwnedTensor) monkeypatch.setattr(ttnn, "deallocate", deallocated.append) monkeypatch.setattr(ttnn, "release_trace", release) compiler = ProgramCompiler("mesh", lambda: object()) program = compiler.compile(_Signature("program", 1), lambda context: torch.zeros(1)) trace = TraceCompiler(compiler) trace.register_capture_plan( TraceCapturePlan( program.key, _Signature("trace", 1), "decode", lambda: persistent, lambda values: output, ) ) trace.capture_all() with expect_error(RuntimeError, "Failed to release"): trace.cleanup() assert deallocated == [] trace.cleanup() trace.cleanup() assert release_attempts == [100, 100] assert deallocated.count(persistent) == 1 assert deallocated.count(output) == 1 compiler.cleanup() def test_cleanup_releases_operation_owned_persistent_dataclasses_once(monkeypatch): events = [] _patch_backend(monkeypatch, events) class OwnedTensor: pass values = [OwnedTensor() for _ in range(20)] decode = DecodePersistentInputs( DecodeDeviceInputs(*values[:4]), tuple(values[4:7]), ) prefill = PrefillHiddenPersistentInputs(PrefillDeviceInputs(*values[7:14])) prefill_workspace = PrefillReplayState( PrefillPositionInputs(*values[14:17]), (values[17], values[17], values[17]), values[18], ) deallocated = [] monkeypatch.setattr(ttnn, "Tensor", OwnedTensor) monkeypatch.setattr(ttnn, "deallocate", deallocated.append) compiler = ProgramCompiler("mesh", lambda: object()) programs = [compiler.compile(_Signature("program", variant), lambda context: torch.zeros(1)) for variant in (1, 2)] trace = TraceCompiler(compiler) trace.register_capture_plan( TraceCapturePlan(programs[0].key, _Signature("trace", 1), "decode", lambda: decode, lambda _: values[0]) ) trace.register_capture_plan( TraceCapturePlan( programs[1].key, _Signature("trace", 2), "prefill", lambda: prefill, lambda _: values[19], prepare_workspace=lambda: prefill_workspace, workspace_fingerprint=("prefill-workspace",), ) ) trace.capture_all() trace.cleanup() assert len(deallocated) == len(values) assert {id(value) for value in deallocated} == {id(value) for value in values}