diff --git a/simple_syrup/nodes/grounded_sam_model_info.py b/simple_syrup/nodes/grounded_sam_model_info.py index de6c874..7dc7e42 100644 --- a/simple_syrup/nodes/grounded_sam_model_info.py +++ b/simple_syrup/nodes/grounded_sam_model_info.py @@ -8,7 +8,7 @@ from __future__ import annotations from typing import Any -from ..runtime.model_catalog import grounding_dino_choices, sam_choices +from ..runtime.model_choices import ModelChoiceService, default_choice from ..runtime.model_metadata import GroundedSAMModelMetadata from . import tooltips @@ -17,6 +17,7 @@ class GroundedSAMModelInfo: """Expose selected grounded SAM source and local path metadata.""" _metadata = GroundedSAMModelMetadata() + _choices = ModelChoiceService() RETURN_TYPES = ("STRING",) RETURN_NAMES = ("model_info",) @@ -31,19 +32,27 @@ class GroundedSAMModelInfo: def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: """Declare deterministic model metadata inputs.""" + sam_model_choices = cls._choices.sam_choices() + grounding_dino_model_choices = cls._choices.grounding_dino_choices() return { "required": { "sam_model": ( - sam_choices(), + sam_model_choices, { - "default": "sam_hq_vit_b (379MB)", + "default": default_choice( + sam_model_choices, + "sam_hq_vit_b (379MB)", + ), "tooltip": tooltips.SAM_MODEL_INPUT, }, ), "grounding_dino_model": ( - grounding_dino_choices(), + grounding_dino_model_choices, { - "default": "GroundingDINO_SwinT_OGC (694MB)", + "default": default_choice( + grounding_dino_model_choices, + "GroundingDINO_SwinT_OGC (694MB)", + ), "tooltip": tooltips.GROUNDING_DINO_MODEL_INPUT, }, ), @@ -53,4 +62,6 @@ class GroundedSAMModelInfo: def describe(self, sam_model: str, grounding_dino_model: str) -> tuple[str]: """Return JSON metadata for selected model entries.""" + self._choices.reject_sentinel(sam_model) + self._choices.reject_sentinel(grounding_dino_model) return (self._metadata.describe_selection(sam_model, grounding_dino_model),) diff --git a/simple_syrup/nodes/load_ultralytics_model.py b/simple_syrup/nodes/load_ultralytics_model.py index 368b1cb..7810738 100644 --- a/simple_syrup/nodes/load_ultralytics_model.py +++ b/simple_syrup/nodes/load_ultralytics_model.py @@ -8,6 +8,7 @@ from __future__ import annotations from typing import Any, ClassVar +from ..runtime.model_downloads import ComfyProgressReporter from ..runtime.ultralytics_loader import UltralyticsLoaderService @@ -40,7 +41,8 @@ class LoadUltralyticsModel: { "default": choices[0], "tooltip": ( - "Ultralytics model file in the ComfyUI models folder." + "A local Ultralytics model or a curated model that " + "downloads to ComfyUI's Impact Pack-compatible folders." ), }, ) @@ -50,5 +52,8 @@ class LoadUltralyticsModel: def load(self, model_name: str) -> tuple[object, object, object]: """Load the selected detector and paired compatibility facades.""" - loaded = self.service_class().load(model_name) + loaded = self.service_class().load( + model_name, + progress=ComfyProgressReporter(), + ) return loaded.detector_model, loaded.bbox_detector, loaded.segm_detector diff --git a/simple_syrup/runtime/model_catalog.py b/simple_syrup/runtime/model_catalog.py index dfa5abf..9d6a316 100644 --- a/simple_syrup/runtime/model_catalog.py +++ b/simple_syrup/runtime/model_catalog.py @@ -2,7 +2,7 @@ # Copyright (C) 2026 Artificial Sweetener and contributors # SPDX-License-Identifier: AGPL-3.0-or-later -"""Known model metadata for grounded SAM masking.""" +"""Known model metadata for downloadable SimpleSyrup model loaders.""" from __future__ import annotations @@ -11,13 +11,14 @@ from enum import StrEnum class ModelFamily(StrEnum): - """Catalog families used by grounded SAM model selection.""" + """Catalog families used by SimpleSyrup model selection.""" SAM = "sam" GROUNDING_DINO = "grounding_dino" TEXT_ENCODER = "text_encoder" VITMATTE = "vitmatte" WD14_TAGGER = "wd14_tagger" + ULTRALYTICS = "ultralytics" @dataclass(frozen=True) @@ -29,6 +30,7 @@ class ModelArtifact: folder_name: str source_url: str description: str + sha256: str | None = None @dataclass(frozen=True) @@ -386,6 +388,298 @@ WD14_TAGGER_ENTRIES: tuple[ModelEntry, ...] = ( ) +_ANZHCS_YOLOS_REVISION = "f5a2306d7fed4f3cfc26c25ff1ab2e3f3cfce855" +_ANZHCS_YOLOS_REPOSITORY = "Anzhc/Anzhcs_YOLOs" + + +def _huggingface_yolo_entry( + *, + entry_id: str, + display_name: str, + filename: str, + folder_name: str, + model_type: str, + source_repo: str, + revision: str, + license_note: str, + description: str, + sha256: str, +) -> ModelEntry: + """Build one revision-pinned Hugging Face Ultralytics catalog entry.""" + + encoded_filename = filename.replace(" ", "%20") + return ModelEntry( + entry_id=entry_id, + display_name=display_name, + family=ModelFamily.ULTRALYTICS, + model_type=model_type, + source_repo=source_repo, + license_note=license_note, + artifacts=( + ModelArtifact( + artifact_id=f"{entry_id}_checkpoint", + filename=filename, + folder_name=folder_name, + source_url=( + f"https://huggingface.co/{source_repo}/resolve/{revision}/" + f"{encoded_filename}" + ), + description=description, + sha256=sha256, + ), + ), + ) + + +def _anzhc_yolo_entry( + *, + entry_id: str, + display_name: str, + filename: str, + folder_name: str, + model_type: str, + description: str, + sha256: str, +) -> ModelEntry: + """Build one revision-pinned Anzhc Ultralytics catalog entry.""" + + return _huggingface_yolo_entry( + entry_id=entry_id, + display_name=display_name, + filename=filename, + folder_name=folder_name, + model_type=model_type, + source_repo=_ANZHCS_YOLOS_REPOSITORY, + revision=_ANZHCS_YOLOS_REVISION, + license_note="AGPL-3.0", + description=description, + sha256=sha256, + ) + + +ULTRALYTICS_ENTRIES: tuple[ModelEntry, ...] = ( + _anzhc_yolo_entry( + entry_id="anzhc_face_seg", + display_name="Anzhc Face -seg (6.52MB)", + filename="Anzhc Face -seg.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc face segmentation model", + sha256="dbf083201298a495e332113de0612d1be1ae8307628628eb7972a31979cdbbb3", + ), + _anzhc_yolo_entry( + entry_id="anzhc_face_seg_640_v2_y8n", + display_name="Anzhc Face seg 640 v2 y8n (6.56MB)", + filename="Anzhc Face seg 640 v2 y8n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc face segmentation model", + sha256="d473e8bccc4c833d8eb36c95e566ce6460ffdc8b2899c859910e380c85def276", + ), + _anzhc_yolo_entry( + entry_id="anzhc_face_seg_768_v2_y8n", + display_name="Anzhc Face seg 768 v2 y8n (6.58MB)", + filename="Anzhc Face seg 768 v2 y8n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc face segmentation model", + sha256="9a1e5b154c1d190812447431bda6b8f260f132877812b4a2f163981f54558355", + ), + _anzhc_yolo_entry( + entry_id="anzhc_face_seg_768ms_v2_y8n", + display_name="Anzhc Face seg 768MS v2 y8n (6.60MB)", + filename="Anzhc Face seg 768MS v2 y8n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc multi-scale face segmentation model", + sha256="429e88d9aecb9fa4167ffd41a6ebc42c97b7fa785aa5468a7eb302ceb9837aae", + ), + _anzhc_yolo_entry( + entry_id="anzhc_face_seg_1024_v2_y8n", + display_name="Anzhc Face seg 1024 v2 y8n (6.63MB)", + filename="Anzhc Face seg 1024 v2 y8n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc face segmentation model", + sha256="1bbcfd7a9f407c6f6e4389a371dbcc392f9444421cf7f824152e92bf563dc6a3", + ), + _anzhc_yolo_entry( + entry_id="anzhc_face_seg_640_v3_y11n", + display_name="Anzhc Face seg 640 v3 y11n (5.80MB)", + filename="Anzhc Face seg 640 v3 y11n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc YOLO11 face segmentation model", + sha256="96437afc773bacd118e275e6cddc1fb7263c78dc11299989c7a00a26506c45bf", + ), + _anzhc_yolo_entry( + entry_id="anzhc_face_seg_640_v4_y11n", + display_name="Anzhc Face seg 640 v4 y11n (5.74MB)", + filename="Anzhc Face seg 640 v4 y11n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc YOLO11 face segmentation model", + sha256="1e77ad7bd349babd8a4a90478bfc965348642b63a8d95d3b43ee13db42fd0a64", + ), + _anzhc_yolo_entry( + entry_id="anzhcs_manface_v02_1024_y8n", + display_name="Anzhcs ManFace v02 1024 y8n (6.06MB)", + filename="Anzhcs ManFace v02 1024 y8n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc male face segmentation model", + sha256="184b9a680afb3c4a559e46e2fe692338fe7bdd6267979fa4ef10526fa96c1b31", + ), + _anzhc_yolo_entry( + entry_id="anzhcs_womanface_v05_1024_y8n", + display_name="Anzhcs WomanFace v05 1024 y8n (6.07MB)", + filename="Anzhcs WomanFace v05 1024 y8n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc female face segmentation model", + sha256="84db37616e1ca975c4e23fa5a300acf0edd9144ec287bbbdbd1ad0f4a3afa9c1", + ), + _anzhc_yolo_entry( + entry_id="anzhc_eyes_seg_hd", + display_name="Anzhc Eyes -seg-hd (6.59MB)", + filename="Anzhc Eyes -seg-hd.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc eye segmentation model", + sha256="6be1c13ca7a51c2425e278e07e7ae3d4c94ee125b874a0104a142f4f5a35a308", + ), + _anzhc_yolo_entry( + entry_id="anzhc_headhair_seg_y8n", + display_name="Anzhc HeadHair seg y8n (6.50MB)", + filename="Anzhc HeadHair seg y8n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc head and hair segmentation model", + sha256="a6e99b1305f600c35e7f6400741c2322b198ae03755f91dc1c59d7a78d77f13c", + ), + _anzhc_yolo_entry( + entry_id="anzhc_headhair_seg_y8m", + display_name="Anzhc HeadHair seg y8m (52.34MB)", + filename="Anzhc HeadHair seg y8m.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc head and hair segmentation model", + sha256="f63aa1cdb63a26c0025a4a984588248241a5838aff4edfeea93d9c155efe0b5e", + ), + _anzhc_yolo_entry( + entry_id="anzhc_breasts_seg_v1_1024n", + display_name="Anzhc Breasts Seg v1 1024n (6.58MB)", + filename="Anzhc Breasts Seg v1 1024n.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc breast segmentation model", + sha256="d469bd7abdcbe32a946e0e342bc1fe96aa021987787d51245f97a29e114cb31b", + ), + _anzhc_yolo_entry( + entry_id="anzhc_breasts_seg_v1_1024s", + display_name="Anzhc Breasts Seg v1 1024s (22.86MB)", + filename="Anzhc Breasts Seg v1 1024s.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc breast segmentation model", + sha256="413a9b948a40f96a83769a882816ef0dd2b91b49673c91bff75463660077b395", + ), + _anzhc_yolo_entry( + entry_id="anzhc_breasts_seg_v1_1024m", + display_name="Anzhc Breasts Seg v1 1024m (52.39MB)", + filename="Anzhc Breasts Seg v1 1024m.pt", + folder_name="ultralytics_segm", + model_type="segment", + description="Anzhc breast segmentation model", + sha256="53d15e82a8308f8056f4929838e00e42c8da576b661e0c2b4fef5837d8b5b2b4", + ), + _huggingface_yolo_entry( + entry_id="bingsu_face_yolov8n_v2", + display_name="Bingsu Face YOLOv8n v2 (6.23MB)", + filename="face_yolov8n_v2.pt", + folder_name="ultralytics_bbox", + model_type="detect", + source_repo="Bingsu/adetailer", + revision="53cc19de382014514d9d4038601d261a7faa9b7b", + license_note="Apache-2.0", + description="Bingsu ADetailer face detection model", + sha256="8f5f2110f83c4e00712993fab48c771d26036e2e80ec62bd5b9cb37c29e36b36", + ), + _huggingface_yolo_entry( + entry_id="bingsu_face_yolov8s", + display_name="Bingsu Face YOLOv8s (22.5MB)", + filename="face_yolov8s.pt", + folder_name="ultralytics_bbox", + model_type="detect", + source_repo="Bingsu/adetailer", + revision="53cc19de382014514d9d4038601d261a7faa9b7b", + license_note="Apache-2.0", + description="Bingsu ADetailer face detection model", + sha256="c7237eff25787377de196961140ceaed324d859ee8de5a775d93d33a0e3fab78", + ), + _huggingface_yolo_entry( + entry_id="bingsu_hand_yolov8n", + display_name="Bingsu Hand YOLOv8n (6.23MB)", + filename="hand_yolov8n.pt", + folder_name="ultralytics_bbox", + model_type="detect", + source_repo="Bingsu/adetailer", + revision="53cc19de382014514d9d4038601d261a7faa9b7b", + license_note="Apache-2.0", + description="Bingsu ADetailer hand detection model", + sha256="3991202eb69e9ddcb3b9ba80cdeb41e734ffaf844403d6c9f47d515cd88c6f29", + ), + _huggingface_yolo_entry( + entry_id="bingsu_hand_yolov8s", + display_name="Bingsu Hand YOLOv8s (22.5MB)", + filename="hand_yolov8s.pt", + folder_name="ultralytics_bbox", + model_type="detect", + source_repo="Bingsu/adetailer", + revision="53cc19de382014514d9d4038601d261a7faa9b7b", + license_note="Apache-2.0", + description="Bingsu ADetailer hand detection model", + sha256="70b540063fbc385736d8258970744a4afbc4cbf7932134bae3b24cdadeadec06", + ), + _huggingface_yolo_entry( + entry_id="bingsu_person_yolov8n_seg", + display_name="Bingsu Person YOLOv8n-seg (6.78MB)", + filename="person_yolov8n-seg.pt", + folder_name="ultralytics_segm", + model_type="segment", + source_repo="Bingsu/adetailer", + revision="53cc19de382014514d9d4038601d261a7faa9b7b", + license_note="Apache-2.0", + description="Bingsu ADetailer person segmentation model", + sha256="38fc8aaae97cb6e70be4ec44770005b26ed473471362afcda62a0037d7ccf432", + ), + _huggingface_yolo_entry( + entry_id="bingsu_person_yolov8s_seg", + display_name="Bingsu Person YOLOv8s-seg (23.9MB)", + filename="person_yolov8s-seg.pt", + folder_name="ultralytics_segm", + model_type="segment", + source_repo="Bingsu/adetailer", + revision="53cc19de382014514d9d4038601d261a7faa9b7b", + license_note="Apache-2.0", + description="Bingsu ADetailer person segmentation model", + sha256="53c54aec2239355faffc6c5b70d0f3d05042f386f956cbec39cec46ad456f050", + ), + _huggingface_yolo_entry( + entry_id="fuyucchi_yolov8x6_animeface", + display_name="Fuyucchi YOLOv8x6 Anime Face (195MB)", + filename="yolov8x6_animeface.pt", + folder_name="ultralytics_bbox", + model_type="detect", + source_repo="Fuyucchi/yolov8_animeface", + revision="b0841ce930453c0f23ceb8086d6554c17de5fe4a", + license_note="AGPL-3.0", + description="Fuyucchi high-resolution anime face detection model", + sha256="f3cdc1a6266347322439fd9b3c8f5a1222668eb10c8adf00e17b28c48b95213c", + ), +) + + def sam_choices() -> list[str]: """Return deterministic SAM dropdown choices.""" @@ -410,6 +704,12 @@ def wd14_tagger_choices() -> list[str]: return [entry.display_name for entry in WD14_TAGGER_ENTRIES] +def ultralytics_choices() -> list[str]: + """Return deterministic Ultralytics dropdown choices.""" + + return [entry.display_name for entry in ULTRALYTICS_ENTRIES] + + def get_sam_entry(selection: str) -> ModelEntry: """Return the SAM catalog entry matching an id or display name.""" @@ -434,6 +734,12 @@ def get_wd14_tagger_entry(selection: str) -> ModelEntry: return _get_entry(selection, WD14_TAGGER_ENTRIES, "WD14 tagger") +def get_ultralytics_entry(selection: str) -> ModelEntry: + """Return the Ultralytics catalog entry matching an id or display name.""" + + return _get_entry(selection, ULTRALYTICS_ENTRIES, "Ultralytics") + + def _get_entry( selection: str, entries: tuple[ModelEntry, ...], diff --git a/simple_syrup/runtime/model_choices.py b/simple_syrup/runtime/model_choices.py index 0a451a4..834b8d1 100644 --- a/simple_syrup/runtime/model_choices.py +++ b/simple_syrup/runtime/model_choices.py @@ -12,11 +12,13 @@ from typing import Protocol from .model_catalog import ( GROUNDING_DINO_ENTRIES, SAM_ENTRIES, + ULTRALYTICS_ENTRIES, VITMATTE_ENTRIES, WD14_TAGGER_ENTRIES, ModelEntry, grounding_dino_choices, sam_choices, + ultralytics_choices, vitmatte_choices, wd14_tagger_choices, ) @@ -108,6 +110,18 @@ class ModelChoiceService: ] return choices or [NO_LOCAL_WD14_TAGGER_MODELS] + def ultralytics_choices(self) -> list[str]: + """Return settings-aware curated Ultralytics dropdown choices.""" + + if self._show_downloadable_models(): + return ultralytics_choices() + + return [ + entry.display_name + for entry in ULTRALYTICS_ENTRIES + if self._entry_artifacts_are_local(entry) + ] + def reject_sentinel(self, selection: str) -> None: """Reject placeholder dropdown selections before loader work begins.""" diff --git a/simple_syrup/runtime/model_metadata.py b/simple_syrup/runtime/model_metadata.py index 68173b6..9c578a8 100644 --- a/simple_syrup/runtime/model_metadata.py +++ b/simple_syrup/runtime/model_metadata.py @@ -78,6 +78,7 @@ class GroundedSAMModelMetadata: "artifact_id": artifact.artifact_id, "filename": artifact.filename, "source_url": artifact.source_url, + "sha256": artifact.sha256, "expected_path": str(expected), "local_path": str(local_path) if local_path else None, "installed": local_path is not None, diff --git a/simple_syrup/runtime/ultralytics_loader.py b/simple_syrup/runtime/ultralytics_loader.py index 3d8f3e2..629ee8f 100644 --- a/simple_syrup/runtime/ultralytics_loader.py +++ b/simple_syrup/runtime/ultralytics_loader.py @@ -14,7 +14,14 @@ from types import ModuleType from typing import Any, TypeAlias, cast from ..shared.logging import get_logger -from .model_folders import SUPPORTED_MODEL_EXTENSIONS +from .model_catalog import ULTRALYTICS_ENTRIES, ModelEntry +from .model_choices import ModelChoiceService +from .model_downloads import DownloadRequest, ModelDownloader, ProgressReporter +from .model_folders import ( + SUPPORTED_MODEL_EXTENSIONS, + expected_model_file, + resolve_model_file, +) from .model_instance_cache import ModelInstanceCache LOGGER = get_logger(__name__) @@ -52,7 +59,6 @@ class LoadedUltralyticsDetector: class UltralyticsModelCacheKey: """Identify a loaded Ultralytics detector for process-level reuse.""" - model_name: str model_path: Path @@ -68,6 +74,8 @@ class UltralyticsLoaderService: self, folder_paths_module: ModuleType | None = None, ultralytics_module: ModuleType | None = None, + downloader: ModelDownloader | None = None, + choice_service: ModelChoiceService | None = None, cache: ( MutableMapping[UltralyticsModelCacheKey, LoadedUltralyticsDetector] | None ) = None, @@ -76,6 +84,10 @@ class UltralyticsLoaderService: self._folder_paths_module = folder_paths_module self._ultralytics_module = ultralytics_module + self._downloader = downloader or ModelDownloader() + self._choice_service = choice_service or ModelChoiceService( + folder_paths_module=folder_paths_module + ) self._cache: ModelInstanceCache[ UltralyticsModelCacheKey, LoadedUltralyticsDetector ] = ModelInstanceCache( @@ -83,9 +95,18 @@ class UltralyticsLoaderService: ) def model_choices(self) -> list[str]: - """Return local Ultralytics model choices for ComfyUI dropdowns.""" + """Return curated and local Ultralytics model choices for ComfyUI dropdowns.""" - choices = self.available_models() + self._register_model_folders() + curated_choices = self._choice_service.ultralytics_choices() + curated_local_paths = { + _catalog_selection(entry) for entry in ULTRALYTICS_ENTRIES + } + choices = curated_choices + [ + choice + for choice in self.available_models() + if choice not in curated_local_paths + ] return choices or [NO_LOCAL_ULTRALYTICS_MODELS] def available_models(self) -> list[str]: @@ -133,16 +154,23 @@ class UltralyticsLoaderService: return sorted(choices) - def load(self, model_name: str) -> LoadedUltralyticsDetector: + def load( + self, + model_name: str, + progress: ProgressReporter | None = None, + ) -> LoadedUltralyticsDetector: """Load one Ultralytics model and create compatibility facades.""" self.reject_sentinel(model_name) - model_path = self.resolve_model_path(model_name) - normalized_name = _normalized_model_name(model_name) - key = UltralyticsModelCacheKey( - model_name=normalized_name, - model_path=model_path.resolve(), - ) + entry = _catalog_entry_or_none(model_name) + if entry is None: + model_path = self.resolve_model_path(model_name) + normalized_name = _normalized_model_name(model_name) + else: + model_path = self._resolve_catalog_entry(entry, progress) + normalized_name = _catalog_selection(entry) + + key = UltralyticsModelCacheKey(model_path=model_path.resolve()) already_loaded = key in self._cache.entries loaded = self._cache.get_or_load( key, @@ -160,6 +188,40 @@ class UltralyticsLoaderService: ) return loaded + def _resolve_catalog_entry( + self, + entry: ModelEntry, + progress: ProgressReporter | None, + ) -> Path: + """Resolve or securely download one curated Ultralytics checkpoint.""" + + if len(entry.artifacts) != 1: + raise RuntimeError( + f"Ultralytics catalog entry '{entry.entry_id}' must have one artifact." + ) + + self._register_model_folders() + artifact = entry.artifacts[0] + existing = resolve_model_file( + artifact.folder_name, + artifact.filename, + self._folder_paths_module, + ) + destination = existing or expected_model_file( + artifact.folder_name, artifact.filename, self._folder_paths_module + ) + result = self._downloader.download( + DownloadRequest( + source_url=artifact.source_url, + destination_path=destination, + expected_folder=destination.parent, + description=artifact.description, + expected_sha256=artifact.sha256, + ), + progress, + ) + return result.path + def _load_uncached_detector( self, model_name: str, @@ -227,9 +289,10 @@ class UltralyticsLoaderService: if model_name == NO_LOCAL_ULTRALYTICS_MODELS: raise ValueError( - "No local Ultralytics models are available. Install a model in " - "models\\ultralytics, models\\ultralytics\\bbox, or " - "models\\ultralytics\\segm." + "No local Ultralytics models are available. Enable 'Show " + "downloadable models in loader dropdowns' in SimpleSyrup settings " + "or install a model in models\\ultralytics, " + "models\\ultralytics\\bbox, or models\\ultralytics\\segm." ) def resolve_model_path(self, model_name: str) -> Path: @@ -387,6 +450,42 @@ def _normalized_model_name(model_name: str) -> str: return model_name.replace("\\", "/") +def _catalog_entry_or_none(selection: str) -> ModelEntry | None: + """Return a curated Ultralytics entry when a dropdown label matches it.""" + + return next( + ( + entry + for entry in ULTRALYTICS_ENTRIES + if selection in (entry.entry_id, entry.display_name) + ), + None, + ) + + +def _catalog_selection(entry: ModelEntry) -> str: + """Return the local conventional selection path for one catalog entry.""" + + if len(entry.artifacts) != 1: + raise ValueError( + f"Ultralytics catalog entry '{entry.entry_id}' must have one artifact." + ) + + artifact = entry.artifacts[0] + prefix_by_folder = { + ULTRALYTICS_BBOX_FOLDER: "bbox", + ULTRALYTICS_SEGM_FOLDER: "segm", + } + try: + prefix = prefix_by_folder[artifact.folder_name] + except KeyError as error: + raise ValueError( + f"Ultralytics catalog entry '{entry.entry_id}' has unsupported folder " + f"'{artifact.folder_name}'." + ) from error + return f"{prefix}/{artifact.filename}" + + def _model_task(model_name: str, raw_model: object) -> str: """Infer detector task from choice prefix or model metadata.""" diff --git a/tests/test_grounded_sam_model_info_node.py b/tests/test_grounded_sam_model_info_node.py index 7a74052..0377911 100644 --- a/tests/test_grounded_sam_model_info_node.py +++ b/tests/test_grounded_sam_model_info_node.py @@ -6,7 +6,9 @@ from __future__ import annotations -from typing import Any +from typing import Any, cast + +import pytest from simple_syrup.nodes.grounded_sam_model_info import GroundedSAMModelInfo @@ -20,9 +22,25 @@ def test_model_info_node_contract_constants() -> None: assert GroundedSAMModelInfo.CATEGORY == "SimpleSyrup/Masking" -def test_model_info_node_declares_expected_inputs() -> None: +def test_model_info_node_declares_expected_inputs( + monkeypatch: pytest.MonkeyPatch, +) -> None: """Model info node exposes model selectors.""" + class FakeChoices: + """Return the known downloadable selections for declaration tests.""" + + def sam_choices(self) -> list[str]: + """Return the expected SAM choice.""" + + return ["sam_hq_vit_b (379MB)"] + + def grounding_dino_choices(self) -> list[str]: + """Return the expected GroundingDINO choice.""" + + return ["GroundingDINO_SwinT_OGC (694MB)"] + + monkeypatch.setattr(GroundedSAMModelInfo, "_choices", cast(Any, FakeChoices())) input_types: dict[str, dict[str, tuple[Any, ...]]] = ( GroundedSAMModelInfo.INPUT_TYPES() ) @@ -33,6 +51,38 @@ def test_model_info_node_declares_expected_inputs() -> None: assert "GroundingDINO_SwinT_OGC (694MB)" in required["grounding_dino_model"][0] +def test_model_info_node_uses_settings_aware_choices() -> None: + """Model metadata selectors follow the downloadable-models preference.""" + + class FakeChoices: + """Return the local-only choices supplied by settings policy.""" + + def sam_choices(self) -> list[str]: + """Return the available SAM choices.""" + + return ["local-sam"] + + def grounding_dino_choices(self) -> list[str]: + """Return the available GroundingDINO choices.""" + + return ["local-dino"] + + def reject_sentinel(self, selection: str) -> None: + """Accept the deterministic test selections.""" + + del selection + + original = GroundedSAMModelInfo._choices + GroundedSAMModelInfo._choices = cast(Any, FakeChoices()) + try: + required = GroundedSAMModelInfo.INPUT_TYPES()["required"] + finally: + GroundedSAMModelInfo._choices = original + + assert required["sam_model"][0] == ["local-sam"] + assert required["grounding_dino_model"][0] == ["local-dino"] + + def test_model_info_node_delegates_to_metadata_provider() -> None: """Node execution delegates metadata creation to its metadata provider.""" @@ -44,12 +94,23 @@ def test_model_info_node_delegates_to_metadata_provider() -> None: return f"{sam_model}|{grounding_dino_model}" + class FakeChoices: + """Accept all model selections while exercising metadata delegation.""" + + def reject_sentinel(self, selection: str) -> None: + """Accept the deterministic test selections.""" + + del selection + node = GroundedSAMModelInfo() original = GroundedSAMModelInfo._metadata + original_choices = GroundedSAMModelInfo._choices GroundedSAMModelInfo._metadata = FakeMetadata() # type: ignore[assignment] + GroundedSAMModelInfo._choices = cast(Any, FakeChoices()) try: result = node.describe("sam", "dino") finally: GroundedSAMModelInfo._metadata = original + GroundedSAMModelInfo._choices = original_choices assert result == ("sam|dino",) diff --git a/tests/test_grounding_dino_model_loader_node.py b/tests/test_grounding_dino_model_loader_node.py index acf05a6..9029b93 100644 --- a/tests/test_grounding_dino_model_loader_node.py +++ b/tests/test_grounding_dino_model_loader_node.py @@ -28,9 +28,21 @@ def test_grounding_dino_model_loader_contract() -> None: assert GroundingDINOModelLoader.CATEGORY == "SimpleSyrup/Masking" -def test_grounding_dino_model_loader_declares_expected_inputs() -> None: +def test_grounding_dino_model_loader_declares_expected_inputs( + monkeypatch: pytest.MonkeyPatch, +) -> None: """GroundingDINO loader makes text encoder selection explicit.""" + def catalog_choices() -> list[str]: + """Return the catalog choice expected by this declaration test.""" + + return ["GroundingDINO_SwinT_OGC (694MB)"] + + monkeypatch.setattr( + GroundingDINOModelLoader._choices, + "grounding_dino_choices", + catalog_choices, + ) input_types: dict[str, dict[str, tuple[Any, ...]]] = ( GroundingDINOModelLoader.INPUT_TYPES() ) diff --git a/tests/test_load_ultralytics_model_node.py b/tests/test_load_ultralytics_model_node.py index 264956c..56a3a3c 100644 --- a/tests/test_load_ultralytics_model_node.py +++ b/tests/test_load_ultralytics_model_node.py @@ -11,6 +11,7 @@ from typing import Any, cast import pytest from simple_syrup.nodes.load_ultralytics_model import LoadUltralyticsModel +from simple_syrup.runtime.model_downloads import ProgressReporter from simple_syrup.runtime.ultralytics_loader import LoadedUltralyticsDetector @@ -50,8 +51,13 @@ class _FakeLoaderService: return ["model.pt"] - def load(self, model_name: str) -> LoadedUltralyticsDetector: + def load( + self, + model_name: str, + progress: ProgressReporter | None = None, + ) -> LoadedUltralyticsDetector: """Return deterministic loaded outputs.""" + del progress assert model_name == "model.pt" return LoadedUltralyticsDetector(cast(Any, "native"), "bbox", "segm") diff --git a/tests/test_model_catalog.py b/tests/test_model_catalog.py index 7ffb9a6..7d336e7 100644 --- a/tests/test_model_catalog.py +++ b/tests/test_model_catalog.py @@ -16,10 +16,13 @@ from simple_syrup.runtime.model_catalog import ( BERT_ENTRY, GROUNDING_DINO_ENTRIES, SAM_ENTRIES, + ULTRALYTICS_ENTRIES, get_grounding_dino_entry, get_sam_entry, + get_ultralytics_entry, grounding_dino_choices, sam_choices, + ultralytics_choices, ) @@ -54,6 +57,77 @@ def test_catalog_choices_are_deterministic() -> None: assert grounding_dino_choices() == [ entry.display_name for entry in GROUNDING_DINO_ENTRIES ] + assert ultralytics_choices() == [ + entry.display_name for entry in ULTRALYTICS_ENTRIES + ] + + +def test_ultralytics_catalog_has_pinned_verified_anzhc_checkpoints() -> None: + """Curated Anzhc models are revision-pinned, verified, and task-foldered.""" + + anzhc_entries = tuple( + entry + for entry in ULTRALYTICS_ENTRIES + if entry.source_repo == "Anzhc/Anzhcs_YOLOs" + ) + + assert len(anzhc_entries) == 15 + assert all(entry.source_repo == "Anzhc/Anzhcs_YOLOs" for entry in anzhc_entries) + assert all(len(entry.artifacts) == 1 for entry in anzhc_entries) + assert all( + artifact.source_url.startswith( + "https://huggingface.co/Anzhc/Anzhcs_YOLOs/resolve/" + "f5a2306d7fed4f3cfc26c25ff1ab2e3f3cfce855/" + ) + for entry in anzhc_entries + for artifact in entry.artifacts + ) + assert all( + artifact.folder_name == "ultralytics_segm" + and artifact.sha256 is not None + and len(artifact.sha256) == 64 + for entry in anzhc_entries + for artifact in entry.artifacts + ) + assert all( + "Drones" not in artifact.filename + and "Score" not in artifact.filename + and "Breast size" not in artifact.filename + for entry in anzhc_entries + for artifact in entry.artifacts + ) + + +def test_ultralytics_catalog_has_verified_adetailer_and_anime_models() -> None: + """ADetailer and anime face checkpoints have compatible curated metadata.""" + + assert len(ULTRALYTICS_ENTRIES) == 22 + + face = get_ultralytics_entry("bingsu_face_yolov8n_v2") + hand = get_ultralytics_entry("bingsu_hand_yolov8s") + person = get_ultralytics_entry("bingsu_person_yolov8s_seg") + anime_face = get_ultralytics_entry("fuyucchi_yolov8x6_animeface") + + assert face.artifacts[0].folder_name == "ultralytics_bbox" + assert hand.artifacts[0].folder_name == "ultralytics_bbox" + assert person.artifacts[0].folder_name == "ultralytics_segm" + assert anime_face.artifacts[0].folder_name == "ultralytics_bbox" + assert face.source_repo == "Bingsu/adetailer" + assert face.license_note == "Apache-2.0" + assert ( + "/resolve/53cc19de382014514d9d4038601d261a7faa9b7b/" + in face.artifacts[0].source_url + ) + assert anime_face.source_repo == "Fuyucchi/yolov8_animeface" + assert anime_face.license_note == "AGPL-3.0" + assert "/resolve/b0841ce930453c0f23ceb8086d6554c17de5fe4a/" in ( + anime_face.artifacts[0].source_url + ) + assert all( + artifact.sha256 is not None and len(artifact.sha256) == 64 + for entry in (face, hand, person, anime_face) + for artifact in entry.artifacts + ) def test_catalog_lookup_rejects_unknown_selection() -> None: diff --git a/tests/test_model_choices.py b/tests/test_model_choices.py index 0570e74..61210ca 100644 --- a/tests/test_model_choices.py +++ b/tests/test_model_choices.py @@ -47,6 +47,7 @@ def test_downloadable_mode_includes_catalog_entries(tmp_path: Path) -> None: assert "GroundingDINO_SwinT_OGC (694MB)" in service.grounding_dino_choices() assert "vitmatte-small-composition-1k" in service.vitmatte_choices() assert "wd-eva02-large-tagger-v3" in service.wd14_tagger_choices() + assert "Bingsu Hand YOLOv8n (6.23MB)" in service.ultralytics_choices() def test_local_only_mode_returns_sentinels_when_no_models_exist( @@ -60,6 +61,7 @@ def test_local_only_mode_returns_sentinels_when_no_models_exist( assert service.grounding_dino_choices() == [NO_LOCAL_GROUNDING_DINO_MODELS] assert service.vitmatte_choices() == [NO_LOCAL_VITMATTE_MODELS] assert service.wd14_tagger_choices() == [NO_LOCAL_WD14_TAGGER_MODELS] + assert service.ultralytics_choices() == [] def test_sam_local_only_lists_installed_catalog_artifacts(tmp_path: Path) -> None: diff --git a/tests/test_sam_model_loader_node.py b/tests/test_sam_model_loader_node.py index 6ac3030..7a84df1 100644 --- a/tests/test_sam_model_loader_node.py +++ b/tests/test_sam_model_loader_node.py @@ -23,9 +23,17 @@ def test_sam_model_loader_contract() -> None: assert SAMModelLoader.CATEGORY == "SimpleSyrup/Masking" -def test_sam_model_loader_declares_expected_inputs() -> None: +def test_sam_model_loader_declares_expected_inputs( + monkeypatch: pytest.MonkeyPatch, +) -> None: """SAM loader inputs are deterministic and loader-owned.""" + def catalog_choices() -> list[str]: + """Return the catalog choices expected by this declaration test.""" + + return ["sam_vit_b (375MB)", "FastSAM-s (23MB)"] + + monkeypatch.setattr(SAMModelLoader._choices, "sam_choices", catalog_choices) input_types: dict[str, dict[str, tuple[Any, ...]]] = SAMModelLoader.INPUT_TYPES() required = input_types["required"] diff --git a/tests/test_ultralytics_loader.py b/tests/test_ultralytics_loader.py index dd14b4d..21f5ec7 100644 --- a/tests/test_ultralytics_loader.py +++ b/tests/test_ultralytics_loader.py @@ -13,6 +13,15 @@ from typing import Any, cast import pytest +from simple_syrup.runtime.model_catalog import get_ultralytics_entry +from simple_syrup.runtime.model_choices import ModelChoiceService +from simple_syrup.runtime.model_downloads import ( + DownloadRequest, + DownloadResult, + ModelDownloader, + ProgressReporter, +) +from simple_syrup.runtime.settings import SimpleSyrupSettings from simple_syrup.runtime.ultralytics_loader import ( NO_LOCAL_ULTRALYTICS_MODELS, LoadedUltralyticsDetector, @@ -32,7 +41,11 @@ def test_model_choices_list_conventional_folders(tmp_path: Path) -> None: (models_dir / "ultralytics" / "bbox" / "face.pt").write_bytes(b"") (models_dir / "ultralytics" / "segm" / "person.pt").write_bytes(b"") - service = UltralyticsLoaderService(folder_paths_module=_folder_paths(models_dir)) + folder_paths = _folder_paths(models_dir) + service = UltralyticsLoaderService( + folder_paths_module=folder_paths, + choice_service=_choice_service(folder_paths, show_downloadable_models=False), + ) assert service.model_choices() == ["bbox/face.pt", "root.pt", "segm/person.pt"] @@ -43,11 +56,53 @@ def test_model_choices_returns_sentinel_when_no_models(tmp_path: Path) -> None: models_dir = tmp_path / "models" models_dir.mkdir() - service = UltralyticsLoaderService(folder_paths_module=_folder_paths(models_dir)) + folder_paths = _folder_paths(models_dir) + service = UltralyticsLoaderService( + folder_paths_module=folder_paths, + choice_service=_choice_service(folder_paths, show_downloadable_models=False), + ) assert service.model_choices() == [NO_LOCAL_ULTRALYTICS_MODELS] +def test_model_choices_include_curated_downloadable_models(tmp_path: Path) -> None: + """Downloadable mode exposes the complete curated Anzhc model selection.""" + + folder_paths = _folder_paths(tmp_path / "models") + service = UltralyticsLoaderService( + folder_paths_module=folder_paths, + choice_service=_choice_service(folder_paths, show_downloadable_models=True), + ) + + choices = service.model_choices() + + assert len(choices) == 22 + assert "Anzhc Face -seg (6.52MB)" in choices + assert "Bingsu Hand YOLOv8n (6.23MB)" in choices + assert "Fuyucchi YOLOv8x6 Anime Face (195MB)" in choices + assert "Anzhcs Breast size det cls v8 640 y11m (38.70MB)" not in choices + assert not any("Drone" in choice for choice in choices) + assert not any("Score" in choice for choice in choices) + + +def test_local_only_choices_use_curated_label_for_installed_model( + tmp_path: Path, +) -> None: + """Installed curated files keep their friendly dropdown label when hidden.""" + + models_dir = tmp_path / "models" + checkpoint = models_dir / "ultralytics" / "segm" / "Anzhc Face -seg.pt" + checkpoint.parent.mkdir(parents=True) + checkpoint.write_bytes(b"checkpoint") + folder_paths = _folder_paths(models_dir) + service = UltralyticsLoaderService( + folder_paths_module=folder_paths, + choice_service=_choice_service(folder_paths, show_downloadable_models=False), + ) + + assert service.model_choices() == ["Anzhc Face -seg (6.52MB)"] + + def test_missing_model_raises_value_error(tmp_path: Path) -> None: """Loading rejects unknown model choices before importing Ultralytics.""" @@ -102,6 +157,112 @@ def test_loader_returns_native_and_compatibility_outputs(tmp_path: Path) -> None assert loaded.bbox_detector is cast(Any, loaded.segm_detector).bbox_detector +def test_curated_model_downloads_to_impact_pack_compatible_folder( + tmp_path: Path, +) -> None: + """A curated selection downloads with checksum verification into segm.""" + + models_dir = tmp_path / "models" + folder_paths = _folder_paths(models_dir) + downloader = _RecordingDownloader() + ultralytics_module = ModuleType("ultralytics") + cast(Any, ultralytics_module).YOLO = _FakeYOLO + entry = get_ultralytics_entry("anzhc_face_seg") + service = UltralyticsLoaderService( + folder_paths_module=folder_paths, + ultralytics_module=ultralytics_module, + downloader=downloader, + choice_service=_choice_service(folder_paths, show_downloadable_models=True), + cache={}, + ) + + loaded = service.load(entry.display_name) + + expected_path = models_dir / "ultralytics" / "segm" / "Anzhc Face -seg.pt" + assert loaded.detector_model.model_path == expected_path + assert loaded.detector_model.model_name == "segm/Anzhc Face -seg.pt" + assert loaded.detector_model.supports_segmentation is True + assert downloader.requests[0].destination_path == expected_path + assert downloader.requests[0].expected_folder == expected_path.parent + assert downloader.requests[0].expected_sha256 == entry.artifacts[0].sha256 + + +def test_curated_bbox_model_downloads_to_impact_pack_compatible_folder( + tmp_path: Path, +) -> None: + """A curated bbox selection downloads into the conventional bbox folder.""" + + models_dir = tmp_path / "models" + folder_paths = _folder_paths(models_dir) + downloader = _RecordingDownloader() + ultralytics_module = ModuleType("ultralytics") + cast(Any, ultralytics_module).YOLO = _FakeYOLO + entry = get_ultralytics_entry("bingsu_hand_yolov8n") + service = UltralyticsLoaderService( + folder_paths_module=folder_paths, + ultralytics_module=ultralytics_module, + downloader=downloader, + choice_service=_choice_service(folder_paths, show_downloadable_models=True), + cache={}, + ) + + loaded = service.load(entry.display_name) + + expected_path = models_dir / "ultralytics" / "bbox" / "hand_yolov8n.pt" + assert loaded.detector_model.model_path == expected_path + assert loaded.detector_model.supports_segmentation is False + assert downloader.requests[0].destination_path == expected_path + assert downloader.requests[0].expected_sha256 == entry.artifacts[0].sha256 + + +def test_curated_existing_model_must_match_its_catalog_checksum( + tmp_path: Path, +) -> None: + """A pre-existing curated checkpoint cannot bypass checksum verification.""" + + models_dir = tmp_path / "models" + checkpoint = models_dir / "ultralytics" / "segm" / "Anzhc Face -seg.pt" + checkpoint.parent.mkdir(parents=True) + checkpoint.write_bytes(b"wrong checkpoint") + folder_paths = _folder_paths(models_dir) + service = UltralyticsLoaderService( + folder_paths_module=folder_paths, + choice_service=_choice_service(folder_paths, show_downloadable_models=True), + cache={}, + ) + + with pytest.raises(ValueError, match="checksum mismatch"): + service.load("Anzhc Face -seg (6.52MB)") + + +def test_catalog_and_local_selection_share_one_loaded_model(tmp_path: Path) -> None: + """Catalog and conventional-path selections share the loaded model instance.""" + + models_dir = tmp_path / "models" + folder_paths = _folder_paths(models_dir) + downloader = _RecordingDownloader() + ultralytics_module = ModuleType("ultralytics") + yolo_factory = _RecordingYOLOFactory() + cast(Any, ultralytics_module).YOLO = yolo_factory + entry = get_ultralytics_entry("anzhc_face_seg") + cache: dict[UltralyticsModelCacheKey, LoadedUltralyticsDetector] = {} + service = UltralyticsLoaderService( + folder_paths_module=folder_paths, + ultralytics_module=ultralytics_module, + downloader=downloader, + choice_service=_choice_service(folder_paths, show_downloadable_models=True), + cache=cache, + ) + + catalog_loaded = service.load(entry.display_name) + local_loaded = service.load("segm/Anzhc Face -seg.pt") + + assert local_loaded is catalog_loaded + assert len(downloader.requests) == 1 + assert len(yolo_factory.paths) == 1 + assert len(cache) == 1 + + def test_bbox_prefix_marks_model_as_bbox_only(tmp_path: Path) -> None: """BBox-prefixed models do not claim segmentation support.""" @@ -233,6 +394,61 @@ class _RecordingYOLOFactory: return _FakeYOLO(path) +class _RecordingDownloader(ModelDownloader): + """Download boundary double that records verified catalog requests.""" + + def __init__(self) -> None: + """Initialize the recorded request collection.""" + + self.requests: list[DownloadRequest] = [] + + def download( + self, + request: DownloadRequest, + progress: ProgressReporter | None = None, + ) -> DownloadResult: + """Materialize a placeholder checkpoint at the requested destination.""" + + del progress + self.requests.append(request) + request.destination_path.parent.mkdir(parents=True, exist_ok=True) + request.destination_path.write_bytes(b"checkpoint") + return DownloadResult( + path=request.destination_path, + bytes_downloaded=len(b"checkpoint"), + skipped_existing=False, + ) + + +class _FakeSettingsRepository: + """Settings boundary double for Ultralytics dropdown tests.""" + + def __init__(self, show_downloadable_models: bool) -> None: + """Store the configured dropdown visibility preference.""" + + self._settings = SimpleSyrupSettings( + show_downloadable_models=show_downloadable_models + ) + + def load(self) -> SimpleSyrupSettings: + """Return the configured settings value.""" + + return self._settings + + +def _choice_service( + folder_paths: ModuleType, + *, + show_downloadable_models: bool, +) -> ModelChoiceService: + """Build an Ultralytics choice service with deterministic settings.""" + + return ModelChoiceService( + _FakeSettingsRepository(show_downloadable_models), + folder_paths, + ) + + def _folder_paths(models_dir: Path) -> ModuleType: """Build a minimal fake ComfyUI folder_paths module.""" diff --git a/tests/test_vitmatte_model_loader_node.py b/tests/test_vitmatte_model_loader_node.py index 6b04187..ecf0b14 100644 --- a/tests/test_vitmatte_model_loader_node.py +++ b/tests/test_vitmatte_model_loader_node.py @@ -23,9 +23,24 @@ def test_vitmatte_model_loader_contract() -> None: assert ViTMatteModelLoader.CATEGORY == "SimpleSyrup/Masking" -def test_vitmatte_model_loader_declares_expected_inputs() -> None: +def test_vitmatte_model_loader_declares_expected_inputs( + monkeypatch: pytest.MonkeyPatch, +) -> None: """ViTMatte loader inputs are asset-only and deterministic.""" + def catalog_choices() -> list[str]: + """Return the catalog choices expected by this declaration test.""" + + return [ + "vitmatte-small-composition-1k", + "vitmatte-base-composition-1k", + ] + + monkeypatch.setattr( + ViTMatteModelLoader._choices, + "vitmatte_choices", + catalog_choices, + ) input_types: dict[str, dict[str, tuple[Any, ...]]] = ( ViTMatteModelLoader.INPUT_TYPES() ) diff --git a/tests/test_wd14_tagger_loader_node.py b/tests/test_wd14_tagger_loader_node.py index 76c4cb3..1c4d090 100644 --- a/tests/test_wd14_tagger_loader_node.py +++ b/tests/test_wd14_tagger_loader_node.py @@ -23,9 +23,21 @@ def test_wd14_tagger_loader_contract() -> None: assert WD14TaggerLoader.CATEGORY == "SimpleSyrup/Tagging" -def test_wd14_tagger_loader_declares_expected_inputs() -> None: +def test_wd14_tagger_loader_declares_expected_inputs( + monkeypatch: pytest.MonkeyPatch, +) -> None: """WD14 loader inputs are asset-only and deterministic.""" + def catalog_choices() -> list[str]: + """Return the catalog choice expected by this declaration test.""" + + return ["wd-eva02-large-tagger-v3"] + + monkeypatch.setattr( + WD14TaggerLoader._choices, + "wd14_tagger_choices", + catalog_choices, + ) input_types: dict[str, dict[str, tuple[Any, ...]]] = WD14TaggerLoader.INPUT_TYPES() required = input_types["required"] diff --git a/tests/test_wd14_tagger_loader_v3_node.py b/tests/test_wd14_tagger_loader_v3_node.py index 64d9702..2004794 100644 --- a/tests/test_wd14_tagger_loader_v3_node.py +++ b/tests/test_wd14_tagger_loader_v3_node.py @@ -8,13 +8,25 @@ from __future__ import annotations from typing import Any +import pytest + from simple_syrup.nodes.wd14_tagger_loader import WD14TaggerLoader +from simple_syrup.nodes_v3 import wd14_tagger_loader as wd14_tagger_loader_v3 from simple_syrup.nodes_v3.wd14_tagger_loader import WD14TaggerLoaderV3 -def test_wd14_tagger_loader_v3_schema() -> None: +def test_wd14_tagger_loader_v3_schema(monkeypatch: pytest.MonkeyPatch) -> None: """The v3 loader schema exposes the WD14 tagger loader contract.""" + class FakeChoices: + """Return the catalog choice expected by this schema test.""" + + def wd14_tagger_choices(self) -> list[str]: + """Return the expected WD14 tagger choice.""" + + return ["wd-eva02-large-tagger-v3"] + + monkeypatch.setattr(wd14_tagger_loader_v3, "ModelChoiceService", FakeChoices) schema = WD14TaggerLoaderV3.define_schema() assert schema.node_id == "SimpleSyrup.WD14TaggerLoader" diff --git a/web/dist/simple-syrup.js b/web/dist/simple-syrup.js index 932e7bb..5e20fa3 100644 --- a/web/dist/simple-syrup.js +++ b/web/dist/simple-syrup.js @@ -232,7 +232,7 @@ async function backendErrorMessage(response, fallback) { // web/src/downloadableModelsSetting.ts var SIMPLE_SYRUP_SETTING_ID = "SimpleSyrup.ShowDownloadableModels"; var SIMPLE_SYRUP_SETTING_LABEL = "SimpleSyrup: Show downloadable models in loader dropdowns"; -var SIMPLE_SYRUP_SETTING_DESCRIPTION = "Show known downloadable SAM, GroundingDINO, and ViTMatte models even when they are not installed locally."; +var SIMPLE_SYRUP_SETTING_DESCRIPTION = "Show known downloadable SAM, GroundingDINO, ViTMatte, WD14 tagger, and Ultralytics models even when they are not installed locally."; function registerDownloadableModelsSetting(app2, context, logger) { const setting = app2.ui.settings.addSetting({ id: SIMPLE_SYRUP_SETTING_ID, @@ -255,6 +255,15 @@ function registerDownloadableModelsSetting(app2, context, logger) { error ); setting.value = previous.show_downloadable_models; + return; + } + try { + await context.refreshModelChoices(); + } catch (error) { + logger.warn( + "Could not refresh Comfy loader model choices after saving SimpleSyrup settings.", + error + ); } } }); @@ -692,6 +701,9 @@ async function registerSimpleSyrupSettings(app2, api = defaultApi(), logger = co saveSettings: (settings) => api.saveSettings(settings), setSettings: (settings) => { savedSettings = settings; + }, + refreshModelChoices: async () => { + await app2.refreshComboInNodes?.(); } }; registerDownloadableModelsSetting(app2, settingsContext, logger); diff --git a/web/src/downloadableModelsSetting.ts b/web/src/downloadableModelsSetting.ts index de8b082..0f3ea9f 100644 --- a/web/src/downloadableModelsSetting.ts +++ b/web/src/downloadableModelsSetting.ts @@ -9,12 +9,13 @@ export const SIMPLE_SYRUP_SETTING_ID = "SimpleSyrup.ShowDownloadableModels"; export const SIMPLE_SYRUP_SETTING_LABEL = "SimpleSyrup: Show downloadable models in loader dropdowns"; export const SIMPLE_SYRUP_SETTING_DESCRIPTION = - "Show known downloadable SAM, GroundingDINO, and ViTMatte models even when they are not installed locally."; + "Show known downloadable SAM, GroundingDINO, ViTMatte, WD14 tagger, and Ultralytics models even when they are not installed locally."; export interface GeneralSettingsContext { getSettings(): SimpleSyrupSettings; saveSettings(settings: SimpleSyrupSettings): Promise; setSettings(settings: SimpleSyrupSettings): void; + refreshModelChoices(): Promise; } export function registerDownloadableModelsSetting( @@ -43,6 +44,16 @@ export function registerDownloadableModelsSetting( error ); setting.value = previous.show_downloadable_models; + return; + } + + try { + await context.refreshModelChoices(); + } catch (error) { + logger.warn( + "Could not refresh Comfy loader model choices after saving SimpleSyrup settings.", + error + ); } } }); diff --git a/web/src/settingsRegistration.ts b/web/src/settingsRegistration.ts index 93a07c3..8b5dc27 100644 --- a/web/src/settingsRegistration.ts +++ b/web/src/settingsRegistration.ts @@ -63,6 +63,9 @@ export async function registerSimpleSyrupSettings( saveSettings: (settings) => api.saveSettings(settings), setSettings: (settings) => { savedSettings = settings; + }, + refreshModelChoices: async () => { + await app.refreshComboInNodes?.(); } }; registerDownloadableModelsSetting(app, settingsContext, logger); diff --git a/web/tests/settings.test.ts b/web/tests/settings.test.ts index 2c3df08..5cb1684 100644 --- a/web/tests/settings.test.ts +++ b/web/tests/settings.test.ts @@ -6,6 +6,7 @@ import { describe, expect, it, vi } from "vitest"; import { SIMPLE_SYRUP_SETTING_ID, + SIMPLE_SYRUP_SETTING_DESCRIPTION, SIMPLE_SYRUP_SETTING_LABEL } from "../src/downloadableModelsSetting"; import { @@ -31,8 +32,11 @@ describe("Comfy settings registration", () => { id: SIMPLE_SYRUP_SETTING_ID, name: SIMPLE_SYRUP_SETTING_LABEL, type: "boolean", - defaultValue: false + defaultValue: false, + tooltip: SIMPLE_SYRUP_SETTING_DESCRIPTION }); + expect(SIMPLE_SYRUP_SETTING_DESCRIPTION).toContain("WD14 tagger"); + expect(SIMPLE_SYRUP_SETTING_DESCRIPTION).toContain("Ultralytics"); expect(app.ui.settings.settings[0]?.value).toBe(false); expect(app.ui.settings.definitions[1]).toMatchObject({ id: QUANT_CACHE_SETTING_ID, @@ -53,6 +57,8 @@ describe("Comfy settings registration", () => { it("saves setting changes to the backend", async () => { const app = createFakeComfyApp(); + const refreshComboInNodes = vi.fn().mockResolvedValue(undefined); + app.refreshComboInNodes = refreshComboInNodes; const saveSettings = vi .fn() .mockResolvedValue({ @@ -81,6 +87,7 @@ describe("Comfy settings registration", () => { quant_cache_limit_gib: 20 }); expect(app.ui.settings.settings[0]?.value).toBe(true); + expect(refreshComboInNodes).toHaveBeenCalledOnce(); }); it("falls back to the default and warns when backend load fails", async () => { @@ -144,6 +151,22 @@ describe("Comfy settings registration", () => { expect(app.ui.settings.settings[0]?.value).toBe(true); }); + it("keeps a saved setting when live model-choice refresh fails", async () => { + const app = createFakeComfyApp(); + const logger = { warn: vi.fn() }; + app.refreshComboInNodes = vi.fn().mockRejectedValue(new Error("offline")); + const api = fakeSettingsApi(false); + + await registerSimpleSyrupSettings(app, api, logger); + await app.ui.settings.definitions[0]?.onChange?.(true); + + expect(app.ui.settings.settings[0]?.value).toBe(true); + expect(logger.warn).toHaveBeenCalledWith( + expect.stringContaining("Could not refresh Comfy loader model choices"), + expect.any(Error) + ); + }); + it("shows global quant cache usage and saves its GiB limit", async () => { const app = createFakeComfyApp(); const api = fakeSettingsApi(true);