File size: 4,629 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 | # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
# SPDX-License-Identifier: Apache-2.0
from dataclasses import FrozenInstanceError, fields
import pytest
import ttnn
from models.common.llm_runtime import config as runtime_config
from models.common.llm_runtime.config import PagedKVCacheConfig, PageTableLayout, TraceConfig, WarmupConfig
from models.common.models.llama3_8b.executor import Llama3ExecutorConfig
class _TraceConfigSubclass(TraceConfig):
pass
def _paged_config(**overrides):
kwargs = {
"block_size": 32,
"max_num_blocks": 1024,
"dtype": ttnn.bfloat8_b,
}
kwargs.update(overrides)
return PagedKVCacheConfig(**kwargs)
def test_executor_config_has_exact_static_policy_owners_and_is_frozen(expect_error):
config = Llama3ExecutorConfig(
trace=TraceConfig(mode="all"),
warmup=WarmupConfig(),
paged_kv_cache=_paged_config(),
device_sampling_enabled=True,
)
assert [field.name for field in fields(config)] == [
"trace",
"warmup",
"paged_kv_cache",
"device_sampling_enabled",
"allow_batched_prefill_with_device_sampling_for_diagnostics",
]
assert not config.allow_batched_prefill_with_device_sampling_for_diagnostics
forbidden = {
"model",
"mesh_device",
"hf_model",
"tokenizer",
"dtype",
"n_layers",
"sampling_config",
"sampling_output_dtype",
}
assert forbidden.isdisjoint(field.name for field in fields(config))
assert not hasattr(runtime_config, "LLMGraphCompilerConfig")
assert not hasattr(runtime_config, "LLMExecutorConfig")
assert not hasattr(runtime_config, "Sampling1DConfig")
with expect_error(FrozenInstanceError, ""):
config.device_sampling_enabled = False
@pytest.mark.parametrize(
("field_name", "invalid_value"),
[
("trace", WarmupConfig()),
("trace", _TraceConfigSubclass()),
("warmup", TraceConfig()),
("paged_kv_cache", WarmupConfig()),
],
)
def test_executor_config_rejects_non_exact_nested_config_types(field_name, invalid_value, expect_error):
values = {
"trace": TraceConfig(),
"warmup": WarmupConfig(),
"paged_kv_cache": _paged_config(),
"device_sampling_enabled": False,
}
values[field_name] = invalid_value
with expect_error(TypeError, rf"{field_name} must be exactly"):
Llama3ExecutorConfig(**values)
@pytest.mark.parametrize(
("mode", "prefill", "decode"),
[("none", False, False), ("decode_only", False, True), ("all", True, True)],
)
def test_trace_config_selects_static_coverage(mode, prefill, decode, expect_error):
config = TraceConfig(mode=mode)
assert config.prefill_enabled is prefill
assert config.decode_enabled is decode
with expect_error(FrozenInstanceError, ""):
config.mode = "none"
def test_trace_config_rejects_unknown_mode(expect_error):
with expect_error(ValueError, "Unsupported trace mode"):
TraceConfig(mode="prefill_only")
def test_warmup_config_keeps_model_derived_defaults_and_is_deeply_immutable(expect_error):
config = WarmupConfig()
assert config.prefill_seq_lens is None
assert config.prefill_batch_sizes == (1, 2, 4, 8, 16, 32)
assert config.include_decode_top_k is False
with expect_error(TypeError, "must be a tuple"):
WarmupConfig(prefill_batch_sizes=[1, 2])
def test_paged_kv_config_has_plan_fields_and_resolved_capacity(expect_error):
unresolved = _paged_config()
resolved = _paged_config(num_blocks=512)
assert [field.name for field in fields(unresolved)] == [
"block_size",
"max_num_blocks",
"dtype",
"memory_config",
"num_blocks",
]
assert unresolved.memory_config == ttnn.DRAM_MEMORY_CONFIG
assert not unresolved.is_resolved()
assert resolved.is_resolved()
with expect_error(FrozenInstanceError, ""):
resolved.num_blocks = 256
def test_paged_kv_config_rejects_invalid_capacity(expect_error):
with expect_error(ValueError, "exceeds max_num_blocks"):
_paged_config(num_blocks=1025)
with expect_error(ValueError, "block_size"):
_paged_config(block_size=0)
def test_page_table_layout_is_resolved_without_warmup_policy():
layout = PageTableLayout.resolve(
block_size=32,
model_max_sequence_length=4096,
physical_num_blocks=100,
max_prefill_chunk_size=2048,
)
assert layout.raw_capacity_width == 100
assert layout.decode_width == 104
assert layout.prefill_width == 168
|