feat(models): add curated ultralytics downloads

This commit is contained in:
Artificial Sweetener
2026-09-19 00:25:58 -04:00
parent 22e4a5d202
commit 41e8a2b61c
20 changed files with 939 additions and 36 deletions
+16 -5
View File
@@ -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),)
+7 -2
View File
@@ -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
+308 -2
View File
@@ -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, ...],
+14
View File
@@ -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."""
+1
View File
@@ -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,
+113 -14
View File
@@ -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."""
+63 -2
View File
@@ -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",)
+13 -1
View File
@@ -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()
)
+7 -1
View File
@@ -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")
+74
View File
@@ -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:
+2
View File
@@ -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:
+9 -1
View File
@@ -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"]
+218 -2
View File
@@ -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."""
+16 -1
View File
@@ -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()
)
+13 -1
View File
@@ -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"]
+13 -1
View File
@@ -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"
+13 -1
View File
@@ -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);
+12 -1
View File
@@ -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<SimpleSyrupSettings>;
setSettings(settings: SimpleSyrupSettings): void;
refreshModelChoices(): Promise<void>;
}
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
);
}
}
});
+3
View File
@@ -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);
+24 -1
View File
@@ -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<SimpleSyrupSettingsApi["saveSettings"]>()
.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);