File size: 11,371 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
# SPDX-FileCopyrightText: 2026 Tenstorrent AI ULC
# SPDX-License-Identifier: Apache-2.0

"""Focused contracts for the family-neutral model composition root."""

import ast
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock

import pytest

import ttnn
from models.common.llm_runtime.config import PagedKVCacheConfig, TraceConfig, WarmupConfig
from models.common.models import executor as executor_module
from models.common.models.executor import ModelExecutor, ModelExecutorConfig

_EXECUTOR_PATH = Path(__file__).parents[2] / "models" / "executor.py"
_MODELS_ROOT = _EXECUTOR_PATH.parent


def _config(*, device_sampling_enabled: bool = False) -> ModelExecutorConfig:
    return ModelExecutorConfig(
        trace=TraceConfig(mode="none"),
        warmup=WarmupConfig(),
        paged_kv_cache=PagedKVCacheConfig(
            block_size=32,
            max_num_blocks=128,
            num_blocks=128,
            dtype=ttnn.bfloat8_b,
        ),
        device_sampling_enabled=device_sampling_enabled,
    )


def test_common_executor_has_no_concrete_model_dependencies_or_dispatch() -> None:
    tree = ast.parse(_EXECUTOR_PATH.read_text())
    imports = {node.module for node in ast.walk(tree) if isinstance(node, ast.ImportFrom) and node.module is not None}
    imported_names = {alias.name for node in ast.walk(tree) if isinstance(node, ast.Import) for alias in node.names}
    assert not any(module.startswith("models.common.models.") for module in imports)
    assert not any(name.startswith("models.common.models.") for name in imported_names)

    control_flow = [
        node.test for node in ast.walk(tree) if isinstance(node, (ast.If, ast.IfExp, ast.While, ast.Assert))
    ]
    dispatch_names = {"model_name", "model_id", "provider_id", "checkpoint_path", "model_version"}
    assert not any(
        isinstance(node, ast.Name) and node.id in dispatch_names
        for expression in control_flow
        for node in ast.walk(expression)
    )

    torch_calls = [
        node
        for node in ast.walk(tree)
        if isinstance(node, ast.Call)
        and isinstance(node.func, ast.Attribute)
        and isinstance(node.func.value, ast.Name)
        and node.func.value.id == "torch"
    ]
    assert torch_calls == []


def test_model_layer_has_only_the_approved_family_modules_and_readmes() -> None:
    assert (_MODELS_ROOT / "llama3_executor.py").is_file()
    assert (_MODELS_ROOT / "qwen2_executor.py").is_file()
    assert not (_MODELS_ROOT / "qwen3_executor.py").exists()

    model_directories = sorted(
        path for path in _MODELS_ROOT.iterdir() if path.is_dir() and (path / "model.py").is_file()
    )
    assert len(model_directories) == 12
    assert all((path / "README.md").is_file() for path in model_directories)


@pytest.mark.parametrize(
    "relative_path",
    (
        "deepseek_r1_distill_qwen_14b/executor.py",
        "mistral_7b/executor.py",
        "phi4/executor.py",
    ),
)
def test_direct_composition_examples_do_not_depend_on_shared_family_executors(relative_path: str) -> None:
    tree = ast.parse((_MODELS_ROOT / relative_path).read_text())
    imports = {node.module for node in ast.walk(tree) if isinstance(node, ast.ImportFrom)}
    assert "models.common.models.executor" not in imports
    assert "models.common.models.llama3_executor" not in imports
    assert "models.common.models.qwen2_executor" not in imports


def test_common_config_is_frozen_and_rejects_non_exact_nested_config_types(expect_error) -> None:
    config = _config()
    with expect_error(AttributeError, "cannot assign to field"):
        config.device_sampling_enabled = True

    class TraceConfigSubclass(TraceConfig):
        pass

    with expect_error(TypeError, "trace must be exactly TraceConfig"):
        ModelExecutorConfig(
            trace=TraceConfigSubclass(mode="none"),
            warmup=config.warmup,
            paged_kv_cache=config.paged_kv_cache,
            device_sampling_enabled=False,
        )


def test_sampling_state_inputs_are_an_optional_owned_pair(expect_error) -> None:
    with expect_error(ValueError, "must be supplied together"):
        ModelExecutor(
            None,
            None,
            _config(device_sampling_enabled=True),
            sampling_state_controller=object(),
        )
    with expect_error(ValueError, "requires device sampling"):
        ModelExecutor(
            None,
            None,
            _config(),
            sampling_state_controller=object(),
            sampling_state=object(),
        )


@pytest.mark.parametrize(
    ("enable_trace", "expected"),
    [(True, ["prime", "coordinator"]), (False, ["coordinator", "prime"])],
)
def test_prefill_warmup_policy_controls_order_through_one_continuation(enable_trace, expected) -> None:
    events = []

    def policy(executor, default_warmup, *, kv_cache, can_sample_on_device, enable_trace):
        assert executor is target
        assert kv_cache is cache
        assert can_sample_on_device
        if enable_trace:
            events.append("prime")
        default_warmup()
        if not enable_trace:
            events.append("prime")

    cache = object()
    target = object.__new__(ModelExecutor)
    target._terminal = False
    target.prefill_runtime = SimpleNamespace(transient_orphan_count=0)
    target.decode_runtime = SimpleNamespace(transient_orphan_count=0)
    target._prefill_warmup = policy
    target.warmup = SimpleNamespace(
        warmup_prefill=lambda **kwargs: events.append("coordinator"),
    )

    target.warmup_model_prefill(
        kv_cache=cache,
        can_sample_on_device=True,
        enable_trace=enable_trace,
    )

    assert events == expected


def test_request_state_and_execution_target_are_forwarded_by_identity() -> None:
    target = object.__new__(ModelExecutor)
    target._terminal = False
    target.prefill_runtime = SimpleNamespace(transient_orphan_count=0)
    target.decode_runtime = SimpleNamespace(transient_orphan_count=0)
    target._validate_bound_cache = MagicMock()
    target._ensure_sampling_for = MagicMock()
    target._prefill_execution = MagicMock()
    target._request_state_fields = ("prompt_tokens", "output_tokens", "slot_remap")

    values = {name: object() for name in ("tokens", "page_table", "prompt_tokens", "output_tokens", "slot_remap")}
    target.compile_prefill(
        tokens=values["tokens"],
        page_table=values["page_table"],
        prompt_tokens=values["prompt_tokens"],
        output_tokens=values["output_tokens"],
        slot_remap=values["slot_remap"],
    )

    forwarded = target._prefill_execution.compile_prefill.call_args.kwargs
    for name, value in values.items():
        assert forwarded[name] is value


def test_layout_refresh_preserves_owner_and_sampling_state_identity(monkeypatch) -> None:
    state = object()
    layout = object()
    prefill_config = SimpleNamespace(
        max_batch_size=4,
        max_prefill_chunk_size=2048,
        device_sampling_enabled=True,
        can_enable_trace=lambda *_: True,
        supports_batched_prefill=True,
        disable_batched_prefill=True,
        max_prefill_batch_size=4,
        batched_prefill_batched_extract=True,
        trace_capture_prime_sequence_lengths=(128,),
        sampling_state_controller=object(),
        sampling_state=state,
    )
    decode_config = SimpleNamespace(
        lane_capacity=4,
        device_sampling_enabled=True,
        force_greedy_top_k=True,
        sampling_state_controller=prefill_config.sampling_state_controller,
        sampling_state=state,
    )
    warmup_config = SimpleNamespace(warmup=object(), prefill_sequence_lengths=(128,))
    resolved_prefill = SimpleNamespace(page_table_layout=layout, sampling_state=state)
    resolved_decode = SimpleNamespace(page_table_layout=layout, sampling_state=state)
    resolved_warmup = SimpleNamespace(page_table_layout=layout)
    prefill_resolve = MagicMock(return_value=resolved_prefill)
    decode_resolve = MagicMock(return_value=resolved_decode)
    warmup_resolve = MagicMock(return_value=resolved_warmup)
    monkeypatch.setattr(executor_module.PrefillRuntimeConfig, "resolve", prefill_resolve)
    monkeypatch.setattr(executor_module.DecodeRuntimeConfig, "resolve", decode_resolve)
    monkeypatch.setattr(executor_module.WarmupCoordinatorConfig, "resolve", warmup_resolve)

    target = object.__new__(ModelExecutor)
    target.model = object()
    target.output_reader = object()
    target.config = SimpleNamespace(trace=object())
    target.prefill_runtime = SimpleNamespace(config=prefill_config)
    target.decode_runtime = SimpleNamespace(config=decode_config)
    target.warmup = SimpleNamespace(config=warmup_config)
    target._resolve_page_table_layout = lambda: layout
    owners = (target.prefill_runtime, target.decode_runtime, target.warmup)

    target._refresh_page_table_layout()

    assert (target.prefill_runtime, target.decode_runtime, target.warmup) == owners
    assert target.page_table_layout is layout
    assert all(owner.config.page_table_layout is layout for owner in owners)
    assert target.prefill_runtime.config.sampling_state is state
    assert target.decode_runtime.config.sampling_state is state
    assert prefill_resolve.call_args.kwargs["trace_capture_prime_sequence_lengths"] == (128,)
    assert prefill_resolve.call_args.kwargs["sampling_state_controller"] is prefill_config.sampling_state_controller
    assert prefill_resolve.call_args.kwargs["sampling_state"] is state
    assert decode_resolve.call_args.kwargs["sampling_state_controller"] is prefill_config.sampling_state_controller
    assert decode_resolve.call_args.kwargs["sampling_state"] is state


def test_cleanup_is_ordered_retryable_idempotent_and_terminal(expect_error) -> None:
    events = []
    failing = {"reader", "trace"}

    class _Owner:
        def __init__(self, name):
            self.name = name

        def action(self, *args):
            events.append(self.name)
            if self.name in failing:
                raise RuntimeError(self.name)

        cleanup = action
        drain = action
        drain_external_outputs = action
        cleanup_transients = action
        release = action

    target = object.__new__(ModelExecutor)
    target._terminal = False
    target._cleaned_up = False
    target._owner_name = "TestExecutor"
    target.decode_runtime = _Owner("decode")
    target.output_reader = _Owner("reader")
    target.prefill_runtime = _Owner("prefill")
    target.trace_compiler = _Owner("trace")
    target.program_compiler = _Owner("program")
    target.config = SimpleNamespace(device_sampling_enabled=True)
    target.sampling_state_controller = _Owner("sampling-state")
    target.sampling_state = object()
    target.model = SimpleNamespace(sampling=_Owner("sampling"))
    target.kv_cache_manager = _Owner("kv")

    expected = ["decode", "reader", "prefill", "decode", "trace", "program", "sampling-state", "sampling", "kv"]
    with expect_error(RuntimeError, "reader") as raised:
        target.cleanup()
    assert events == expected
    assert [str(error) for error in raised.value.cleanup_failures] == ["trace"]
    assert target.terminal
    assert not target._cleaned_up

    failing.clear()
    target.cleanup()
    target.cleanup()
    assert events == expected * 2
    assert target._cleaned_up