Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 14 additions & 12 deletions src/tserve/runtime/bootstrap.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,11 +89,10 @@ def loaded_models(self) -> ModelsResult:
def bootstrap(model: list[str | Path | tuple[str, Any]]) -> Runtime:
"""Resolve, construct, load, warmup, and register each selected model.

For every item: ``resolve_model`` → ``create_executor(info.executor)``
→ ``load`` → ``warmup`` → ``stats.register``. Tuple items pass the
second element (craft spec or object) to ``load``; strings and
paths pass the item itself. Duplicate ``ModelInfo.id`` values raise
before a second load.
Resolve and check all items before creating any executor. Then, for
every item: ``create_executor(info.executor)`` → ``load`` → ``warmup``
→ ``stats.register``. Tuple items pass the second element (craft spec
or object) to ``load``; strings and paths pass the item itself.

Parameters
----------
Expand Down Expand Up @@ -132,6 +131,15 @@ def bootstrap(model: list[str | Path | tuple[str, Any]]) -> Runtime:
tserve.scheduling.scheduler.Scheduler
Forecast dispatch is not implemented in this package.
"""
resolved: list[tuple[ModelInfo, Any]] = []
seen: set[str] = set()
for item in model:
info = resolve_model(item)
if info.id in seen:
raise ValueError(f"duplicate model id {info.id!r} in model")
seen.add(info.id)
resolved.append((info, item[1] if isinstance(item, tuple) else item))

stats = Stats()
executors: dict[str, Executor] = {}
models: dict[str, ModelInfo] = {}
Expand All @@ -145,13 +153,7 @@ def bootstrap(model: list[str | Path | tuple[str, Any]]) -> Runtime:
f"{paint('is always loaded as a baseline', '2')}"
)

for position, item in enumerate(model, start=1):
info = resolve_model(item)
item = item[1] if isinstance(item, tuple) else item

if info.id in models:
raise ValueError(f"duplicate model id {info.id!r} in model")

for position, (info, item) in enumerate(resolved, start=1):
label = f"{info.id} via {info.executor}"
dots = paint("." * max(3, 40 - len(label)), "2")
prefix = (
Expand Down
23 changes: 22 additions & 1 deletion src/tserve/runtime/tests/test_bootstrap.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,10 +77,31 @@ def test_bootstrap_loads_craft():
def test_bootstrap_rejects_duplicate():
with (
patch("tserve.runtime.bootstrap.resolve_model", return_value=_info()),
patch("tserve.runtime.bootstrap.create_executor", return_value=_executor()),
patch("tserve.runtime.bootstrap.create_executor") as create_executor,
pytest.raises(ValueError, match="duplicate model id 'naive' in model"),
):
bootstrap(["naive", "naive"])
create_executor.assert_not_called()


def test_bootstrap_validates_before_loading():
executor = _executor()
with (
patch(
"tserve.runtime.bootstrap.resolve_model",
side_effect=[_info(), ValueError("unknown model 'missing'")],
) as resolve_model,
patch(
"tserve.runtime.bootstrap.create_executor", return_value=executor
) as create_executor,
pytest.raises(ValueError, match="unknown model 'missing'"),
):
bootstrap(["naive", "missing"])

assert resolve_model.call_count == 2
create_executor.assert_not_called()
executor.load.assert_not_called()
executor.warmup.assert_not_called()


def test_loaded_models():
Expand Down
Loading