feat(models): add curated ultralytics downloads
This commit is contained in:
@@ -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),)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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, ...],
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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",)
|
||||
|
||||
@@ -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()
|
||||
)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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()
|
||||
)
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
Vendored
+13
-1
@@ -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);
|
||||
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user