Files
Artificial-Sweetener-Simple…/tests/segmentation/detection/test_sam_loader.py
T

345 lines
11 KiB
Python

# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Tests for SAM loader runtime service."""
from __future__ import annotations
import sys
from dataclasses import dataclass, field
from pathlib import Path
from types import ModuleType
import pytest
from support.helpers import FakeFolderPaths
from simple_syrup.runtime.loaded_models import LoadedSAMModel
from simple_syrup.runtime.model_downloads import DownloadRequest, DownloadResult
from simple_syrup.runtime.sam_loader import SAMLoaderService, SAMModelCacheKey
@dataclass
class _RecordingPhaseProgress:
"""Record semantic SAM load phases emitted by the runtime service."""
phases: list[str] = field(default_factory=list)
def advance(self, phase: str) -> None:
"""Record one phase transition."""
self.phases.append(phase)
class RecordingDownloader:
"""Downloader double that writes requested artifacts."""
def __init__(self) -> None:
"""Create an empty recording downloader."""
self.requests: list[DownloadRequest] = []
def download(
self,
request: DownloadRequest,
progress: object | None = None,
) -> DownloadResult:
"""Record and satisfy a download request."""
self.requests.append(request)
request.destination_path.parent.mkdir(parents=True, exist_ok=True)
request.destination_path.write_bytes(b"model")
return DownloadResult(request.destination_path, 5, False)
def test_sam_loader_downloads_missing_known_artifact(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""SAM loader downloads known missing models when enabled."""
downloader = RecordingDownloader()
_install_fake_segment_anything(monkeypatch)
loaded = SAMLoaderService(
downloader=downloader, # type: ignore[arg-type]
folder_paths_module=FakeFolderPaths(tmp_path),
).load_model("sam_vit_b (375MB)", True)
assert isinstance(loaded, LoadedSAMModel)
assert loaded.managed_model is not None
assert downloader.requests
def test_sam_loader_uses_process_cache_for_identical_resolved_model(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Identical SAM loads reuse the same loaded container and SAM model."""
state = _install_fake_segment_anything(monkeypatch)
_create_sam_file(tmp_path, "sam_vit_b_01ec64.pth")
cache: dict[SAMModelCacheKey, LoadedSAMModel] = {}
service = SAMLoaderService(
folder_paths_module=FakeFolderPaths(tmp_path),
cache=cache,
)
first = service.load_model("sam_vit_b (375MB)", auto_download=True)
second = service.load_model("sam_vit_b (375MB)", auto_download=True)
assert second is first
assert state.checkpoints == [str(tmp_path / "sams" / "sam_vit_b_01ec64.pth")]
assert len(cache) == 1
def test_sam_loader_reports_cache_miss_and_hit_phases(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""SAM loads should expose checkpoint construction and cache reuse."""
_install_fake_segment_anything(monkeypatch)
_create_sam_file(tmp_path, "sam_vit_b_01ec64.pth")
service = SAMLoaderService(
folder_paths_module=FakeFolderPaths(tmp_path),
cache={},
)
cache_miss_progress = _RecordingPhaseProgress()
cache_hit_progress = _RecordingPhaseProgress()
service.load_model(
"sam_vit_b (375MB)",
auto_download=True,
phase_progress=cache_miss_progress,
)
service.load_model(
"sam_vit_b (375MB)",
auto_download=True,
phase_progress=cache_hit_progress,
)
assert cache_miss_progress.phases == [
"resolving_artifacts",
"loading_checkpoint",
"registering_device_management",
"completed",
]
assert cache_hit_progress.phases == [
"resolving_artifacts",
"cache_hit",
"completed",
]
def test_sam_loader_cache_separates_model_selections(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Different SAM catalog selections produce separate loaded containers."""
state = _install_fake_segment_anything(monkeypatch)
_create_sam_file(tmp_path, "sam_vit_b_01ec64.pth")
_create_sam_file(tmp_path, "sam_vit_l_0b3195.pth")
cache: dict[SAMModelCacheKey, LoadedSAMModel] = {}
service = SAMLoaderService(
folder_paths_module=FakeFolderPaths(tmp_path),
cache=cache,
)
first = service.load_model("sam_vit_b (375MB)", auto_download=True)
second = service.load_model("sam_vit_l (1.25GB)", auto_download=True)
assert second is not first
assert state.checkpoints == [
str(tmp_path / "sams" / "sam_vit_b_01ec64.pth"),
str(tmp_path / "sams" / "sam_vit_l_0b3195.pth"),
]
assert len(cache) == 2
def test_sam_loader_does_not_cache_failed_registry_load(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A failed SAM registry construction leaves the cache empty for retry."""
state = _install_fake_segment_anything(monkeypatch, fail_once=True)
_create_sam_file(tmp_path, "sam_vit_b_01ec64.pth")
cache: dict[SAMModelCacheKey, LoadedSAMModel] = {}
service = SAMLoaderService(
folder_paths_module=FakeFolderPaths(tmp_path),
cache=cache,
)
phase_progress = _RecordingPhaseProgress()
with pytest.raises(RuntimeError, match="SAM failed"):
service.load_model(
"sam_vit_b (375MB)",
auto_download=True,
phase_progress=phase_progress,
)
assert phase_progress.phases == [
"resolving_artifacts",
"loading_checkpoint",
"failed",
]
loaded = service.load_model("sam_vit_b (375MB)", auto_download=True)
assert isinstance(loaded, LoadedSAMModel)
assert len(state.checkpoints) == 2
assert len(cache) == 1
def test_sam_loader_loads_sam_hq_from_owned_runtime(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""SAM-HQ catalog entries load from SimpleSyrup's vendored runtime."""
state = _install_fake_sam_hq_runtime(monkeypatch)
_create_sam_file(tmp_path, "sam_hq_vit_b.pth")
loaded = SAMLoaderService(
folder_paths_module=FakeFolderPaths(tmp_path),
).load_model("sam_hq_vit_b (379MB)", auto_download=True)
assert isinstance(loaded, LoadedSAMModel)
assert loaded.model_id == "sam_hq_vit_b"
assert loaded.managed_model is not None
assert state.checkpoints == [str(tmp_path / "sams" / "sam_hq_vit_b.pth")]
def test_sam_loader_downloads_and_loads_fast_sam_s(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""FastSAM-s uses the existing SAM download, cache, and model wrapper path."""
downloader = RecordingDownloader()
constructed: list[str] = []
class FakeFastSAM:
"""Minimal FastSAM-compatible model fake."""
def __init__(self, checkpoint: str) -> None:
"""Record the checkpoint passed by the shared SAM loader."""
constructed.append(checkpoint)
def to(self, device: object) -> None:
"""Accept device management."""
def eval(self) -> None:
"""Accept evaluation mode."""
ultralytics = ModuleType("ultralytics")
ultralytics.FastSAM = FakeFastSAM # type: ignore[attr-defined]
monkeypatch.setitem(sys.modules, "ultralytics", ultralytics)
loaded = SAMLoaderService(
downloader=downloader, # type: ignore[arg-type]
folder_paths_module=FakeFolderPaths(tmp_path),
).load_model("FastSAM-s (23MB)", auto_download=True)
expected = str(tmp_path / "sams" / "FastSAM-s.pt")
assert downloader.requests[0].destination_path == Path(expected)
assert constructed == [expected]
assert loaded.model_id == "fast_sam_s"
assert loaded.managed_model is not None
def test_sam_loader_errors_when_missing_and_download_disabled(tmp_path: Path) -> None:
"""SAM loader fails clearly when downloads are disabled."""
with pytest.raises(FileNotFoundError, match="auto_download is disabled"):
SAMLoaderService(folder_paths_module=FakeFolderPaths(tmp_path)).load_model(
"sam_vit_b (375MB)",
False,
)
@dataclass
class _FakeSegmentAnythingState:
"""Record fake Segment Anything model construction."""
checkpoints: list[str] = field(default_factory=list)
fail_once: bool = False
def _install_fake_segment_anything(
monkeypatch: pytest.MonkeyPatch,
fail_once: bool = False,
) -> _FakeSegmentAnythingState:
"""Install a fake segment_anything registry for loader tests."""
state = _FakeSegmentAnythingState(fail_once=fail_once)
class FakeModel:
"""Minimal PyTorch-like model fake."""
def __init__(self) -> None:
"""Create a model that records device movement."""
self.to_calls = 0
def to(self, device: object) -> None:
"""Record forbidden loader-time device movement."""
self.to_calls += 1
def eval(self) -> None:
"""Accept eval mode."""
def build_model(checkpoint: str) -> FakeModel:
"""Record checkpoint construction and optionally fail once."""
state.checkpoints.append(checkpoint)
if state.fail_once:
state.fail_once = False
raise RuntimeError("SAM failed")
return FakeModel()
module = ModuleType("segment_anything")
module.sam_model_registry = {"vit_b": build_model, "vit_l": build_model} # type: ignore[attr-defined]
monkeypatch.setitem(sys.modules, "segment_anything", module)
return state
def _install_fake_sam_hq_runtime(
monkeypatch: pytest.MonkeyPatch,
) -> _FakeSegmentAnythingState:
"""Install a fake SimpleSyrup SAM-HQ registry for loader tests."""
state = _FakeSegmentAnythingState()
class FakeModel:
"""Minimal SAM-HQ model fake."""
def to(self, device: object) -> None:
"""Accept device movement."""
def eval(self) -> None:
"""Accept eval mode."""
def build_model(checkpoint: str) -> FakeModel:
"""Record SAM-HQ checkpoint construction."""
state.checkpoints.append(checkpoint)
return FakeModel()
module = ModuleType("simple_syrup.third_party.sam_hq_runtime.build_sam_hq")
module.sam_model_registry = {"sam_hq_vit_b": build_model} # type: ignore[attr-defined]
monkeypatch.setitem(sys.modules, module.__name__, module)
return state
def _create_sam_file(tmp_path: Path, filename: str) -> None:
"""Create one local SAM checkpoint file."""
model_dir = tmp_path / "sams"
model_dir.mkdir(parents=True, exist_ok=True)
(model_dir / filename).write_bytes(b"model")