Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0583ba2675 | ||
|
|
dcc37d7ab6 | ||
|
|
583b22a4bf | ||
|
|
ebe01efc49 | ||
|
|
41e8a2b61c | ||
|
|
22e4a5d202 | ||
|
|
017e3fc7fe |
@@ -1,3 +1,24 @@
|
||||
# [1.8.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.7.1...v1.8.0) (2026-09-19)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **downloads:** keep unknown sizes indeterminate ([a31467c](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/a31467cd3a4299e9e3929d44281018dba0322b3e))
|
||||
* **models:** hide installed catalog choices ([d887e87](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/d887e87e03eb6d0d9d5325fe43b67fbd9e47cf71))
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **models:** add curated ultralytics downloads ([b907fa2](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/b907fa2a17afa30170bc22cf1134a750241e55c2))
|
||||
* **models:** prioritize installed ultralytics choices ([d30e04f](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/d30e04f229366d3e1d2388bd4d713b2b47a12a1f))
|
||||
|
||||
## [1.7.1](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.7.0...v1.7.1) (2026-09-11)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **regional:** preserve shared model patch ancestry ([6059a3f](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/6059a3f913a9502671666e83faeb8686a7a8da27))
|
||||
|
||||
# [1.7.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.6.0...v1.7.0) (2026-09-05)
|
||||
|
||||
|
||||
|
||||
Generated
+2
-2
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"name": "simple-syrup-comfyui",
|
||||
"version": "1.7.0",
|
||||
"version": "1.8.0",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "simple-syrup-comfyui",
|
||||
"version": "1.7.0",
|
||||
"version": "1.8.0",
|
||||
"license": "AGPL-3.0-or-later",
|
||||
"devDependencies": {
|
||||
"@eslint/js": "^9.39.1",
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "simple-syrup-comfyui",
|
||||
"version": "1.7.0",
|
||||
"version": "1.8.0",
|
||||
"private": true,
|
||||
"license": "AGPL-3.0-or-later",
|
||||
"type": "module",
|
||||
|
||||
+1
-1
@@ -5,7 +5,7 @@ build-backend = "setuptools.build_meta"
|
||||
[project]
|
||||
name = "SimpleSyrup"
|
||||
description = "Workflow-focused ComfyUI extensions for image generation."
|
||||
version = "1.7.0"
|
||||
version = "1.8.0"
|
||||
license = "AGPL-3.0-or-later"
|
||||
license-files = ["LICENSE"]
|
||||
requires-python = ">=3.11"
|
||||
|
||||
@@ -6,6 +6,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
__version__ = "1.7.0"
|
||||
__version__ = "1.8.0"
|
||||
|
||||
__all__: list[str] = ["__version__"]
|
||||
|
||||
@@ -30,6 +30,7 @@ class ProcessedRegionalAttentionEntry:
|
||||
schedule: ConditioningScheduleRange
|
||||
cross_attention: torch.Tensor
|
||||
strength: float
|
||||
cross_attention_value_multiplier: torch.Tensor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate entry order, model context, and finite scalar strength."""
|
||||
@@ -62,6 +63,21 @@ class ProcessedRegionalAttentionEntry:
|
||||
if not math.isfinite(float(self.strength)):
|
||||
raise ValueError("Processed conditioning strength must be finite.")
|
||||
object.__setattr__(self, "strength", float(self.strength))
|
||||
multiplier = self.cross_attention_value_multiplier
|
||||
if multiplier is None:
|
||||
return
|
||||
if (
|
||||
not isinstance(multiplier, torch.Tensor)
|
||||
or multiplier.shape != (*self.cross_attention.shape[:2], 1)
|
||||
or not multiplier.is_floating_point()
|
||||
or multiplier.device != self.cross_attention.device
|
||||
or multiplier.dtype != self.cross_attention.dtype
|
||||
or not bool(torch.isfinite(multiplier).all().item())
|
||||
):
|
||||
raise ValueError(
|
||||
"Processed attention value multiplier must be a finite floating "
|
||||
"BxSx1 tensor aligned with cross_attention."
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
||||
@@ -48,6 +48,7 @@ class BatchedRegionalAttentionEntry:
|
||||
entry_index: int
|
||||
context: torch.Tensor
|
||||
strengths: tuple[float, ...]
|
||||
cross_attention_value_multiplier: torch.Tensor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate entry order, aligned context, and finite sample strengths."""
|
||||
@@ -70,6 +71,11 @@ class BatchedRegionalAttentionEntry:
|
||||
)
|
||||
if not math.isfinite(float(strength)):
|
||||
raise ValueError("Regional attention entry strength must be finite.")
|
||||
_validate_value_multiplier(
|
||||
self.cross_attention_value_multiplier,
|
||||
self.context,
|
||||
name="entry",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
@@ -109,6 +115,7 @@ class BatchedRegionalAttentionContexts:
|
||||
chunks: tuple[RegionalAttentionChunkBatch, ...]
|
||||
base_context: torch.Tensor
|
||||
regions: tuple[BatchedRegionalAttentionRegion, ...]
|
||||
base_value_multiplier: torch.Tensor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate complete chunk and tensor alignment."""
|
||||
@@ -141,6 +148,11 @@ class BatchedRegionalAttentionContexts:
|
||||
expected_batch=expected_start,
|
||||
name="base",
|
||||
)
|
||||
_validate_value_multiplier(
|
||||
self.base_value_multiplier,
|
||||
self.base_context,
|
||||
name="base",
|
||||
)
|
||||
if not isinstance(self.regions, tuple):
|
||||
raise TypeError("Regional attention regions must be a tuple.")
|
||||
if tuple(region.region_index for region in self.regions) != tuple(
|
||||
@@ -189,3 +201,27 @@ def _validate_aligned_context(
|
||||
raise ValueError(
|
||||
f"Regional attention {name} context must contain finite floating values."
|
||||
)
|
||||
|
||||
|
||||
def _validate_value_multiplier(
|
||||
multiplier: object,
|
||||
context: torch.Tensor,
|
||||
*,
|
||||
name: str,
|
||||
) -> None:
|
||||
"""Validate one optional value multiplier against its aligned context."""
|
||||
|
||||
if multiplier is None:
|
||||
return
|
||||
if (
|
||||
not isinstance(multiplier, torch.Tensor)
|
||||
or multiplier.shape != (*context.shape[:2], 1)
|
||||
or not multiplier.is_floating_point()
|
||||
or multiplier.device != context.device
|
||||
or multiplier.dtype != context.dtype
|
||||
or not bool(torch.isfinite(multiplier).all().item())
|
||||
):
|
||||
raise ValueError(
|
||||
f"Regional attention {name} value multiplier must be a finite "
|
||||
"floating BxSx1 tensor aligned with its context."
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -10,6 +10,7 @@ from dataclasses import dataclass
|
||||
|
||||
from ..model_attention_patch_mutations import ModelAttn2PatchesMutation
|
||||
from ..patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation
|
||||
from ..ppm_negpip_interop import PpmNegpipInterop
|
||||
from ..regional_lora.standard_unet_native_admission import (
|
||||
StandardUnetNativeLoraAdmission,
|
||||
)
|
||||
@@ -43,6 +44,7 @@ class StandardUnetAttentionBackend:
|
||||
model: object,
|
||||
state: StandardUnetAttentionState,
|
||||
admission: StandardUnetNativeLoraAdmission,
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> StandardUnetAttentionModel:
|
||||
"""Return a direct MODEL child containing only the paired UNet patches."""
|
||||
|
||||
@@ -55,6 +57,8 @@ class StandardUnetAttentionBackend:
|
||||
"Standard UNet admission and processed conditioning must share "
|
||||
"the same regional LoRA plan."
|
||||
)
|
||||
if negpip is not None and not isinstance(negpip, PpmNegpipInterop):
|
||||
raise TypeError("Standard UNet backend NegPiP state has an invalid type.")
|
||||
attention_phase = StandardUnetAttentionPhaseSession()
|
||||
template = (
|
||||
STANDARD_UNET_VARIANT_TEMPLATE_CACHE.resolve(model, admission)
|
||||
@@ -68,6 +72,7 @@ class StandardUnetAttentionBackend:
|
||||
admission,
|
||||
attention_phase,
|
||||
template,
|
||||
negpip,
|
||||
),
|
||||
)
|
||||
if template is not None
|
||||
@@ -85,6 +90,7 @@ class StandardUnetAttentionBackend:
|
||||
ModelAttn2PatchesMutation(
|
||||
patches.input_patch,
|
||||
patches.output_patch,
|
||||
(() if negpip is None else (negpip.attention_patch,)),
|
||||
),
|
||||
)
|
||||
derived = PATCHER_LIFECYCLE.derive_model(
|
||||
|
||||
@@ -27,6 +27,7 @@ from ..services.attention_coupling_preparation_service import (
|
||||
AttentionCouplingPreparation,
|
||||
)
|
||||
from .attention_coupling.context_validation import RegionalContextValidator
|
||||
from .ppm_negpip_interop import PpmNegpipInterop
|
||||
|
||||
|
||||
class ComfyRegionalConditioningProcessor:
|
||||
@@ -40,6 +41,7 @@ class ComfyRegionalConditioningProcessor:
|
||||
noise: torch.Tensor,
|
||||
device: torch.device,
|
||||
context_validator: RegionalContextValidator,
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> ProcessedRegionalAttentionPlan:
|
||||
"""Return model-ready positive and negative context banks."""
|
||||
|
||||
@@ -57,6 +59,8 @@ class ComfyRegionalConditioningProcessor:
|
||||
raise TypeError("Regional context processing device must be torch.device.")
|
||||
if not isinstance(context_validator, RegionalContextValidator):
|
||||
raise TypeError("Regional context validator has an invalid type.")
|
||||
if negpip is not None and not isinstance(negpip, PpmNegpipInterop):
|
||||
raise TypeError("Regional conditioning NegPiP state has an invalid type.")
|
||||
base_model = getattr(model, "model", None)
|
||||
extra_conds = getattr(base_model, "extra_conds", None)
|
||||
if not callable(extra_conds):
|
||||
@@ -75,6 +79,7 @@ class ComfyRegionalConditioningProcessor:
|
||||
noise=noise,
|
||||
device=device,
|
||||
context_validator=context_validator,
|
||||
negpip=negpip,
|
||||
)
|
||||
negative = self._process_branch(
|
||||
preparation.plan.negative,
|
||||
@@ -84,6 +89,7 @@ class ComfyRegionalConditioningProcessor:
|
||||
noise=noise,
|
||||
device=device,
|
||||
context_validator=context_validator,
|
||||
negpip=negpip,
|
||||
)
|
||||
return ProcessedRegionalAttentionPlan(
|
||||
positive=positive,
|
||||
@@ -102,6 +108,7 @@ class ComfyRegionalConditioningProcessor:
|
||||
noise: torch.Tensor,
|
||||
device: torch.device,
|
||||
context_validator: RegionalContextValidator,
|
||||
negpip: PpmNegpipInterop | None,
|
||||
) -> ProcessedRegionalAttentionBranch:
|
||||
"""Process one base plus its ordered regional context bank."""
|
||||
|
||||
@@ -115,6 +122,7 @@ class ComfyRegionalConditioningProcessor:
|
||||
noise=noise,
|
||||
device=device,
|
||||
context_validator=context_validator,
|
||||
negpip=negpip,
|
||||
)
|
||||
regional = tuple(
|
||||
self._process_context(
|
||||
@@ -127,6 +135,7 @@ class ComfyRegionalConditioningProcessor:
|
||||
noise=noise,
|
||||
device=device,
|
||||
context_validator=context_validator,
|
||||
negpip=negpip,
|
||||
)
|
||||
for context in branch.regional_contexts
|
||||
)
|
||||
@@ -144,6 +153,7 @@ class ComfyRegionalConditioningProcessor:
|
||||
noise: torch.Tensor,
|
||||
device: torch.device,
|
||||
context_validator: RegionalContextValidator,
|
||||
negpip: PpmNegpipInterop | None,
|
||||
) -> ProcessedRegionalAttentionContext:
|
||||
"""Convert and extract one exact post-adapter Anima context tensor."""
|
||||
|
||||
@@ -168,6 +178,7 @@ class ComfyRegionalConditioningProcessor:
|
||||
conditioning_index=conditioning_index,
|
||||
prompt_type=prompt_type,
|
||||
context_validator=context_validator,
|
||||
negpip=negpip,
|
||||
)
|
||||
for entry_index, encoded_item in enumerate(encoded)
|
||||
)
|
||||
@@ -185,6 +196,7 @@ class ComfyRegionalConditioningProcessor:
|
||||
conditioning_index: int,
|
||||
prompt_type: str,
|
||||
context_validator: RegionalContextValidator,
|
||||
negpip: PpmNegpipInterop | None,
|
||||
) -> ProcessedRegionalAttentionEntry:
|
||||
"""Extract one exact post-adapter Anima context and Comfy strength."""
|
||||
|
||||
@@ -233,6 +245,11 @@ class ComfyRegionalConditioningProcessor:
|
||||
),
|
||||
cross_attention=context,
|
||||
strength=float(strength),
|
||||
cross_attention_value_multiplier=(
|
||||
None
|
||||
if negpip is None
|
||||
else negpip.extract_value_multiplier(model_conds, context)
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -22,6 +22,7 @@ def _apply_paired_attention_patches(
|
||||
attention_name: str,
|
||||
input_patch: Callable[..., object],
|
||||
output_patch: Callable[..., object],
|
||||
trailing_input_patches: tuple[Callable[..., object], ...],
|
||||
) -> None:
|
||||
"""Validate and atomically install one paired attention callback surface."""
|
||||
|
||||
@@ -29,6 +30,12 @@ def _apply_paired_attention_patches(
|
||||
raise TypeError(f"MODEL {attention_name} input patch must be callable.")
|
||||
if not callable(output_patch):
|
||||
raise TypeError(f"MODEL {attention_name} output patch must be callable.")
|
||||
if not isinstance(trailing_input_patches, tuple) or any(
|
||||
not callable(patch) for patch in trailing_input_patches
|
||||
):
|
||||
raise TypeError(
|
||||
f"MODEL {attention_name} preserved input patches must be callables."
|
||||
)
|
||||
input_setter = _require_bound_method(
|
||||
model,
|
||||
f"set_model_{attention_name}_patch",
|
||||
@@ -54,11 +61,21 @@ def _apply_paired_attention_patches(
|
||||
output_name = f"{attention_name}_output_patch"
|
||||
input_exists = _require_callable_patch_list(patches, input_name)
|
||||
output_exists = _require_callable_patch_list(patches, output_name)
|
||||
if input_exists:
|
||||
existing_input = patches.get(input_name, [])
|
||||
if input_exists and (
|
||||
not trailing_input_patches or existing_input != list(trailing_input_patches)
|
||||
):
|
||||
raise ValueError(f"MODEL {attention_name} input patch is already installed.")
|
||||
if not input_exists and trailing_input_patches:
|
||||
raise ValueError(
|
||||
f"MODEL {attention_name} preserved input patch is not installed."
|
||||
)
|
||||
if output_exists:
|
||||
raise ValueError(f"MODEL {attention_name} output patch is already installed.")
|
||||
input_setter(input_patch)
|
||||
if trailing_input_patches:
|
||||
patches[input_name] = [input_patch, *trailing_input_patches]
|
||||
else:
|
||||
input_setter(input_patch)
|
||||
output_setter(output_patch)
|
||||
|
||||
|
||||
@@ -68,6 +85,7 @@ class ModelAttn2PatchesMutation:
|
||||
|
||||
input_patch: Callable[..., object]
|
||||
output_patch: Callable[..., object]
|
||||
trailing_input_patches: tuple[Callable[..., object], ...] = ()
|
||||
|
||||
def apply(self, model: object) -> None:
|
||||
"""Validate both attn2 surfaces before either mutation."""
|
||||
@@ -77,4 +95,5 @@ class ModelAttn2PatchesMutation:
|
||||
attention_name="attn2",
|
||||
input_patch=self.input_patch,
|
||||
output_patch=self.output_patch,
|
||||
trailing_input_patches=self.trailing_input_patches,
|
||||
)
|
||||
|
||||
@@ -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, ...],
|
||||
|
||||
@@ -6,24 +6,17 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import ModuleType
|
||||
from typing import Protocol
|
||||
|
||||
from .model_catalog import (
|
||||
GROUNDING_DINO_ENTRIES,
|
||||
SAM_ENTRIES,
|
||||
VITMATTE_ENTRIES,
|
||||
WD14_TAGGER_ENTRIES,
|
||||
ModelEntry,
|
||||
grounding_dino_choices,
|
||||
sam_choices,
|
||||
ultralytics_choices,
|
||||
vitmatte_choices,
|
||||
wd14_tagger_choices,
|
||||
)
|
||||
from .model_folders import resolve_model_file
|
||||
from .settings import SimpleSyrupSettings
|
||||
from .settings_repository import SimpleSyrupSettingsRepository
|
||||
from .vitmatte_loader import ViTMatteLoaderService
|
||||
|
||||
NO_LOCAL_SAM_MODELS = "No local SAM models found"
|
||||
NO_LOCAL_GROUNDING_DINO_MODELS = "No local GroundingDINO models found"
|
||||
@@ -44,69 +37,52 @@ class ModelChoiceService:
|
||||
def __init__(
|
||||
self,
|
||||
settings_repository: SettingsProvider | None = None,
|
||||
folder_paths_module: ModuleType | None = None,
|
||||
) -> None:
|
||||
"""Create the choice service with injectable external boundaries."""
|
||||
"""Create the choice service with an injectable settings boundary."""
|
||||
|
||||
self._settings_repository = (
|
||||
settings_repository or SimpleSyrupSettingsRepository()
|
||||
)
|
||||
self._folder_paths_module = folder_paths_module
|
||||
self._vitmatte_loader = ViTMatteLoaderService(
|
||||
folder_paths_module=folder_paths_module
|
||||
)
|
||||
|
||||
def sam_choices(self) -> list[str]:
|
||||
"""Return settings-aware SAM dropdown choices."""
|
||||
"""Return curated SAM choices when catalog models are visible."""
|
||||
|
||||
if self._show_downloadable_models():
|
||||
return sam_choices()
|
||||
|
||||
choices = [
|
||||
entry.display_name
|
||||
for entry in SAM_ENTRIES
|
||||
if self._entry_artifacts_are_local(entry)
|
||||
]
|
||||
return choices or [NO_LOCAL_SAM_MODELS]
|
||||
return [NO_LOCAL_SAM_MODELS]
|
||||
|
||||
def grounding_dino_choices(self) -> list[str]:
|
||||
"""Return settings-aware GroundingDINO dropdown choices."""
|
||||
"""Return curated GroundingDINO choices when catalog models are visible."""
|
||||
|
||||
if self._show_downloadable_models():
|
||||
return grounding_dino_choices()
|
||||
|
||||
choices = [
|
||||
entry.display_name
|
||||
for entry in GROUNDING_DINO_ENTRIES
|
||||
if self._entry_artifacts_are_local(entry)
|
||||
]
|
||||
return choices or [NO_LOCAL_GROUNDING_DINO_MODELS]
|
||||
return [NO_LOCAL_GROUNDING_DINO_MODELS]
|
||||
|
||||
def vitmatte_choices(self) -> list[str]:
|
||||
"""Return settings-aware ViTMatte dropdown choices."""
|
||||
"""Return curated ViTMatte choices when catalog models are visible."""
|
||||
|
||||
if self._show_downloadable_models():
|
||||
return vitmatte_choices()
|
||||
|
||||
choices = [
|
||||
entry.display_name
|
||||
for entry in VITMATTE_ENTRIES
|
||||
if self._vitmatte_entry_is_local(entry)
|
||||
]
|
||||
return choices or [NO_LOCAL_VITMATTE_MODELS]
|
||||
return [NO_LOCAL_VITMATTE_MODELS]
|
||||
|
||||
def wd14_tagger_choices(self) -> list[str]:
|
||||
"""Return settings-aware WD14 tagger dropdown choices."""
|
||||
"""Return curated WD14 tagger choices when catalog models are visible."""
|
||||
|
||||
if self._show_downloadable_models():
|
||||
return wd14_tagger_choices()
|
||||
|
||||
choices = [
|
||||
entry.display_name
|
||||
for entry in WD14_TAGGER_ENTRIES
|
||||
if self._entry_artifacts_are_local(entry)
|
||||
]
|
||||
return choices or [NO_LOCAL_WD14_TAGGER_MODELS]
|
||||
return [NO_LOCAL_WD14_TAGGER_MODELS]
|
||||
|
||||
def ultralytics_choices(self) -> list[str]:
|
||||
"""Return curated Ultralytics choices when catalog models are visible."""
|
||||
|
||||
if self._show_downloadable_models():
|
||||
return ultralytics_choices()
|
||||
|
||||
return []
|
||||
|
||||
def reject_sentinel(self, selection: str) -> None:
|
||||
"""Reject placeholder dropdown selections before loader work begins."""
|
||||
@@ -143,28 +119,6 @@ class ModelChoiceService:
|
||||
|
||||
return self._settings_repository.load().show_downloadable_models
|
||||
|
||||
def _entry_artifacts_are_local(self, entry: ModelEntry) -> bool:
|
||||
"""Return whether every catalog artifact exists locally."""
|
||||
|
||||
return all(
|
||||
resolve_model_file(
|
||||
artifact.folder_name,
|
||||
artifact.filename,
|
||||
self._folder_paths_module,
|
||||
)
|
||||
is not None
|
||||
for artifact in entry.artifacts
|
||||
)
|
||||
|
||||
def _vitmatte_entry_is_local(self, entry: ModelEntry) -> bool:
|
||||
"""Return whether a valid ViTMatte directory exists locally."""
|
||||
|
||||
try:
|
||||
self._vitmatte_loader.resolve_model_directory(entry, auto_download=False)
|
||||
except FileNotFoundError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def default_choice(choices: list[str], preferred: str) -> str:
|
||||
"""Return the preferred default when visible, otherwise the first choice."""
|
||||
|
||||
@@ -59,30 +59,30 @@ class ComfyProgressReporter:
|
||||
"""Initialize an empty ComfyUI progress reporter."""
|
||||
|
||||
self._progress_bar: object | None = None
|
||||
self._total = 1
|
||||
self._total: int | None = None
|
||||
|
||||
def start(self, label: str, total: int | None) -> None:
|
||||
"""Create a ComfyUI progress bar for one artifact."""
|
||||
|
||||
comfy_utils = importlib.import_module("comfy.utils")
|
||||
progress_bar_class = comfy_utils.ProgressBar
|
||||
self._total = total if total and total > 0 else 1
|
||||
self._progress_bar = progress_bar_class(self._total)
|
||||
self._total = total if total and total > 0 else None
|
||||
progress_total = self._total or 1
|
||||
self._progress_bar = progress_bar_class(progress_total)
|
||||
self.advance(0, total)
|
||||
LOGGER.info("download progress started", extra={"label": label, "total": total})
|
||||
|
||||
def advance(self, current: int, total: int | None) -> None:
|
||||
"""Update the ComfyUI progress bar."""
|
||||
"""Update known-size downloads without falsely completing unknown ones."""
|
||||
|
||||
if self._progress_bar is None:
|
||||
return
|
||||
if total and total > 0 and total != self._total:
|
||||
if total is None or total <= 0:
|
||||
return
|
||||
if total != self._total:
|
||||
self._total = total
|
||||
value = (
|
||||
current if total and total > 0 else min(current // CHUNK_SIZE, self._total)
|
||||
)
|
||||
progress_bar = cast(_ComfyProgressBar, self._progress_bar)
|
||||
progress_bar.update_absolute(value, self._total)
|
||||
progress_bar.update_absolute(current, self._total)
|
||||
|
||||
def finish(self) -> None:
|
||||
"""Mark the current ComfyUI progress bar complete."""
|
||||
@@ -90,7 +90,8 @@ class ComfyProgressReporter:
|
||||
if self._progress_bar is None:
|
||||
return
|
||||
progress_bar = cast(_ComfyProgressBar, self._progress_bar)
|
||||
progress_bar.update_absolute(self._total, self._total)
|
||||
total = self._total or 1
|
||||
progress_bar.update_absolute(total, total)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -7,7 +7,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from typing import Protocol, TypeVar, cast
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol, TypeVar, cast
|
||||
|
||||
from .clip_patcher_model_alignment import align_clip_text_encoder_with_patcher
|
||||
|
||||
@@ -27,6 +28,14 @@ class ClipMutation(Protocol):
|
||||
|
||||
|
||||
PatcherValue = TypeVar("PatcherValue")
|
||||
_MODEL_FALLBACK_BOUNDARY_ATTACHMENT = "simple_syrup.model_fallback_boundary"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ModelFallbackBoundary:
|
||||
"""Retain the durable Comfy patcher beneath consecutive Syrup derivations."""
|
||||
|
||||
patcher: object
|
||||
|
||||
|
||||
class ComfyPatcherLifecycle:
|
||||
@@ -48,6 +57,7 @@ class ComfyPatcherLifecycle:
|
||||
disable_dynamic=disable_dynamic,
|
||||
)
|
||||
self._require_direct_parent(source, derived, operation=operation)
|
||||
self._stabilize_model_fallback(source, derived)
|
||||
for mutation in mutations:
|
||||
mutation.apply(derived)
|
||||
return derived
|
||||
@@ -66,13 +76,19 @@ class ComfyPatcherLifecycle:
|
||||
getter = getattr(model_override_source, "get_clone_model_override", None)
|
||||
if not callable(getter):
|
||||
raise TypeError(f"{operation} requires a model-override source.")
|
||||
model_override = getter()
|
||||
derived = self._clone(
|
||||
source,
|
||||
operation=operation,
|
||||
disable_dynamic=disable_dynamic,
|
||||
model_override=getter(),
|
||||
model_override=model_override,
|
||||
)
|
||||
self._require_direct_parent(source, derived, operation=operation)
|
||||
self._stabilize_model_fallback(
|
||||
source,
|
||||
derived,
|
||||
explicit_boundary=model_override_source,
|
||||
)
|
||||
for mutation in mutations:
|
||||
mutation.apply(derived)
|
||||
return derived
|
||||
@@ -107,6 +123,7 @@ class ComfyPatcherLifecycle:
|
||||
derived_patcher,
|
||||
operation=operation,
|
||||
)
|
||||
self._stabilize_model_fallback(source_patcher, derived_patcher)
|
||||
align_clip_text_encoder_with_patcher(derived)
|
||||
for mutation in mutations:
|
||||
mutation.apply(derived)
|
||||
@@ -181,6 +198,41 @@ class ComfyPatcherLifecycle:
|
||||
f"{operation} produced a derived patcher without its source as parent."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _stabilize_model_fallback(
|
||||
source: object,
|
||||
derived: object,
|
||||
*,
|
||||
explicit_boundary: object | None = None,
|
||||
) -> None:
|
||||
"""Collapse Syrup-only lineage onto the durable same-model boundary."""
|
||||
|
||||
derived_attachments = getattr(derived, "attachments", None)
|
||||
if not isinstance(derived_attachments, dict):
|
||||
return
|
||||
boundary = explicit_boundary
|
||||
if boundary is None:
|
||||
source_attachments = getattr(source, "attachments", None)
|
||||
inherited = (
|
||||
source_attachments.get(_MODEL_FALLBACK_BOUNDARY_ATTACHMENT)
|
||||
if isinstance(source_attachments, dict)
|
||||
else None
|
||||
)
|
||||
if inherited is not None and not isinstance(
|
||||
inherited,
|
||||
_ModelFallbackBoundary,
|
||||
):
|
||||
raise TypeError("SimpleSyrup MODEL fallback boundary is invalid.")
|
||||
boundary = inherited.patcher if inherited is not None else source
|
||||
|
||||
if getattr(boundary, "model", None) is not getattr(derived, "model", None):
|
||||
derived_attachments.pop(_MODEL_FALLBACK_BOUNDARY_ATTACHMENT, None)
|
||||
return
|
||||
cast(Any, derived).parent = boundary
|
||||
derived_attachments[_MODEL_FALLBACK_BOUNDARY_ATTACHMENT] = (
|
||||
_ModelFallbackBoundary(boundary)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _required_clip_patcher(value: object, *, value_name: str) -> object:
|
||||
"""Return the CLIP patcher required for lineage validation."""
|
||||
|
||||
@@ -0,0 +1,233 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Adapt the complete installed PPM NegPiP patch family to regional execution."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
from typing import cast
|
||||
|
||||
import torch
|
||||
from comfy.patcher_extension import WrappersMP
|
||||
|
||||
from ..domain.regional_model_capabilities import RegionalModelFamily
|
||||
|
||||
_MODEL_MARKER = "ppm_negpip"
|
||||
_ANIMA_WRAPPER_KEY = "ppm_negpip_anima"
|
||||
_ANIMA_CONDITION_KEY = "c_ppm_negpip_mask"
|
||||
_ANIMA_TRANSFORMER_KEY = "ppm_negpip_mask"
|
||||
_EXTRA_CONDS_PATH = "extra_conds"
|
||||
_ATTN2_PATCH_NAME = "attn2_patch"
|
||||
_UNET_CALLBACK = (
|
||||
"src.negpip.unet_negpip",
|
||||
"sdxl_attn2_negpip",
|
||||
)
|
||||
_ANIMA_CALLBACK = (
|
||||
"src.negpip.anima_negpip",
|
||||
"cosmos_attn2_negpip",
|
||||
)
|
||||
_ANIMA_WRAPPER = (
|
||||
"src.negpip.anima_negpip",
|
||||
"cosmos_diffusion_negpip_wrapper",
|
||||
)
|
||||
_ANIMA_EXTRA_CONDS = (
|
||||
"src.negpip.anima_negpip",
|
||||
"anima_extra_conds_negpip_wrapper.<locals>._anima_extra_conds_negpip_wrapper",
|
||||
)
|
||||
|
||||
|
||||
class PpmNegpipSemantics(StrEnum):
|
||||
"""Identify the family-specific NegPiP conditioning representation."""
|
||||
|
||||
STANDARD_UNET_SPLIT_KEY_VALUE = "standard-unet-split-key-value"
|
||||
ANIMA_VALUE_MASK = "anima-value-mask"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PpmNegpipInterop:
|
||||
"""Retain identity-validated PPM objects needed by regional execution."""
|
||||
|
||||
semantics: PpmNegpipSemantics
|
||||
attention_patch: Callable[..., object]
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require one typed semantic mode and callable preserved callback."""
|
||||
|
||||
if not isinstance(self.semantics, PpmNegpipSemantics):
|
||||
raise TypeError("NegPiP semantics have an invalid type.")
|
||||
if not callable(self.attention_patch):
|
||||
raise TypeError("NegPiP attention patch must be callable.")
|
||||
|
||||
def extract_value_multiplier(
|
||||
self,
|
||||
model_conditions: dict[object, object],
|
||||
context: torch.Tensor,
|
||||
) -> torch.Tensor | None:
|
||||
"""Return one validated Anima value multiplier or no UNet multiplier."""
|
||||
|
||||
if self.semantics is PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE:
|
||||
return None
|
||||
condition = model_conditions.get(_ANIMA_CONDITION_KEY)
|
||||
if condition is None:
|
||||
return context.new_ones((*context.shape[:2], 1))
|
||||
multiplier = getattr(condition, "cond", None)
|
||||
if not isinstance(multiplier, torch.Tensor):
|
||||
raise TypeError("Anima NegPiP value mask condition must contain a tensor.")
|
||||
if (
|
||||
multiplier.ndim != 3
|
||||
or int(multiplier.shape[0]) != int(context.shape[0])
|
||||
or int(multiplier.shape[1]) != int(context.shape[1])
|
||||
or int(multiplier.shape[2]) != 1
|
||||
):
|
||||
raise ValueError(
|
||||
"Anima NegPiP value mask must match the conditioning batch and "
|
||||
"sequence with one multiplier channel."
|
||||
)
|
||||
multiplier = multiplier.to(device=context.device, dtype=context.dtype)
|
||||
if not bool(((multiplier == 1) | (multiplier == -1)).all().item()):
|
||||
raise ValueError("Anima NegPiP value mask must contain only -1 and 1.")
|
||||
return multiplier
|
||||
|
||||
def prepare_anima_transformer_options(
|
||||
self,
|
||||
source: dict[str, object],
|
||||
packed_multiplier: torch.Tensor,
|
||||
) -> dict[str, object]:
|
||||
"""Publish a packed mask on an isolated Anima cross-attention call."""
|
||||
|
||||
if self.semantics is not PpmNegpipSemantics.ANIMA_VALUE_MASK:
|
||||
raise ValueError("Only Anima NegPiP semantics can publish a value mask.")
|
||||
if not isinstance(packed_multiplier, torch.Tensor):
|
||||
raise TypeError("Packed Anima NegPiP multiplier must be a tensor.")
|
||||
prepared = source.copy()
|
||||
prepared[_ANIMA_TRANSFORMER_KEY] = packed_multiplier
|
||||
return prepared
|
||||
|
||||
|
||||
class PpmNegpipInteropValidator:
|
||||
"""Admit only one complete identity-validated PPM NegPiP family."""
|
||||
|
||||
def validate(
|
||||
self,
|
||||
family: RegionalModelFamily,
|
||||
*,
|
||||
model_options: dict[object, object],
|
||||
wrappers: dict[str, dict[object, list[object]]],
|
||||
object_patches: dict[object, object],
|
||||
transformer_patches: dict[str, list[object]],
|
||||
) -> PpmNegpipInterop | None:
|
||||
"""Return preserved NegPiP state or reject every partial/conflicting form."""
|
||||
|
||||
if not isinstance(family, RegionalModelFamily):
|
||||
raise TypeError("NegPiP interop requires a model family.")
|
||||
marker = model_options.get(_MODEL_MARKER, False)
|
||||
if not isinstance(marker, bool):
|
||||
raise TypeError("MODEL ppm_negpip marker must be boolean.")
|
||||
attention = transformer_patches.get(_ATTN2_PATCH_NAME, [])
|
||||
anima_wrappers = wrappers.get(WrappersMP.DIFFUSION_MODEL, {}).get(
|
||||
_ANIMA_WRAPPER_KEY,
|
||||
[],
|
||||
)
|
||||
extra_conds = object_patches.get(_EXTRA_CONDS_PATH)
|
||||
recognized_surface = any(
|
||||
(
|
||||
any(_is_identity(item, *_UNET_CALLBACK) for item in attention),
|
||||
any(_is_identity(item, *_ANIMA_CALLBACK) for item in attention),
|
||||
bool(anima_wrappers),
|
||||
_is_identity(extra_conds, *_ANIMA_EXTRA_CONDS),
|
||||
)
|
||||
)
|
||||
if not marker:
|
||||
if recognized_surface:
|
||||
raise ValueError(
|
||||
"MODEL contains an incomplete NegPiP patch family without its "
|
||||
"marker. Reapply CLIP NegPip to a clean MODEL."
|
||||
)
|
||||
return None
|
||||
if family is RegionalModelFamily.STANDARD_UNET:
|
||||
return self._validate_standard_unet(
|
||||
attention,
|
||||
anima_wrappers=anima_wrappers,
|
||||
extra_conds=extra_conds,
|
||||
)
|
||||
if family is RegionalModelFamily.ANIMA:
|
||||
return self._validate_anima(
|
||||
attention,
|
||||
anima_wrappers=anima_wrappers,
|
||||
extra_conds=extra_conds,
|
||||
)
|
||||
raise ValueError(f"NegPiP does not support model family {family.value!r}.")
|
||||
|
||||
@staticmethod
|
||||
def _validate_standard_unet(
|
||||
attention: list[object],
|
||||
*,
|
||||
anima_wrappers: list[object],
|
||||
extra_conds: object,
|
||||
) -> PpmNegpipInterop:
|
||||
"""Require exactly PPM's single UNet split-K/V callback surface."""
|
||||
|
||||
if (
|
||||
len(attention) != 1
|
||||
or not _is_identity(attention[0], *_UNET_CALLBACK)
|
||||
or anima_wrappers
|
||||
or _is_identity(extra_conds, *_ANIMA_EXTRA_CONDS)
|
||||
):
|
||||
raise ValueError(
|
||||
"Standard UNet NegPiP requires exactly its PPM split-K/V attention "
|
||||
"patch and no Anima NegPiP surfaces."
|
||||
)
|
||||
return PpmNegpipInterop(
|
||||
PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE,
|
||||
cast(Callable[..., object], attention[0]),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_anima(
|
||||
attention: list[object],
|
||||
*,
|
||||
anima_wrappers: list[object],
|
||||
extra_conds: object,
|
||||
) -> PpmNegpipInterop:
|
||||
"""Require PPM's exact callback, wrapper, and extra-condition surfaces."""
|
||||
|
||||
if (
|
||||
len(attention) != 1
|
||||
or not _is_identity(attention[0], *_ANIMA_CALLBACK)
|
||||
or len(anima_wrappers) != 1
|
||||
or not _is_identity(anima_wrappers[0], *_ANIMA_WRAPPER)
|
||||
or not _is_identity(extra_conds, *_ANIMA_EXTRA_CONDS)
|
||||
):
|
||||
raise ValueError(
|
||||
"Anima NegPiP requires exactly its PPM attention patch, keyed "
|
||||
"diffusion wrapper, and extra_conds object patch."
|
||||
)
|
||||
return PpmNegpipInterop(
|
||||
PpmNegpipSemantics.ANIMA_VALUE_MASK,
|
||||
cast(Callable[..., object], attention[0]),
|
||||
)
|
||||
|
||||
|
||||
def _is_identity(
|
||||
value: object,
|
||||
module_suffix: str,
|
||||
qualified_name: str,
|
||||
) -> bool:
|
||||
"""Match one callable by its stable defining module suffix and qualified name."""
|
||||
|
||||
if not callable(value):
|
||||
return False
|
||||
module = getattr(value, "__module__", None)
|
||||
qualname = getattr(value, "__qualname__", None)
|
||||
return (
|
||||
isinstance(module, str)
|
||||
and (module == module_suffix or module.endswith(f".{module_suffix}"))
|
||||
and qualname == qualified_name
|
||||
)
|
||||
|
||||
|
||||
PPM_NEGPIP_INTEROP_VALIDATOR = PpmNegpipInteropValidator()
|
||||
@@ -7,6 +7,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as functional
|
||||
from comfy.utils import repeat_to_batch_size
|
||||
|
||||
from ..domain.processed_regional_attention import (
|
||||
@@ -143,6 +144,13 @@ class RegionalAttentionBatchingService:
|
||||
)
|
||||
for region_index in range(plan.mask_bank.region_count)
|
||||
),
|
||||
base_value_multiplier=self._align_value_multiplier(
|
||||
tuple(chunk.base_entry for chunk in selected),
|
||||
latent_batch_size=latent_batch_size,
|
||||
device=aligned_base_context.device,
|
||||
dtype=aligned_base_context.dtype,
|
||||
target_sequence_length=target_sequence_length,
|
||||
),
|
||||
)
|
||||
|
||||
def _align_region(
|
||||
@@ -163,6 +171,7 @@ class RegionalAttentionBatchingService:
|
||||
for entry_index in range(entry_count):
|
||||
context_parts: list[torch.Tensor] = []
|
||||
strengths: list[float] = []
|
||||
multiplier_entries: list[ProcessedRegionalAttentionEntry] = []
|
||||
for chunk, source in zip(chunks, sources, strict=True):
|
||||
if source is None:
|
||||
entry = chunk.base_entry
|
||||
@@ -186,12 +195,20 @@ class RegionalAttentionBatchingService:
|
||||
target_length=target_sequence_length,
|
||||
)
|
||||
context_parts.append(repeated)
|
||||
multiplier_entries.append(entry)
|
||||
strengths.extend((strength,) * latent_batch_size)
|
||||
entries.append(
|
||||
BatchedRegionalAttentionEntry(
|
||||
entry_index,
|
||||
torch.cat(context_parts, dim=0),
|
||||
tuple(strengths),
|
||||
self._align_value_multiplier(
|
||||
tuple(multiplier_entries),
|
||||
latent_batch_size=latent_batch_size,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
target_sequence_length=target_sequence_length,
|
||||
),
|
||||
)
|
||||
)
|
||||
return BatchedRegionalAttentionRegion(region_index, tuple(entries))
|
||||
@@ -201,6 +218,7 @@ class RegionalAttentionBatchingService:
|
||||
|
||||
entries = self._entries(plan)
|
||||
authority = entries[0].cross_attention
|
||||
has_value_multiplier = entries[0].cross_attention_value_multiplier is not None
|
||||
for entry in entries[1:]:
|
||||
tensor = entry.cross_attention
|
||||
shape_mismatch = (
|
||||
@@ -216,6 +234,48 @@ class RegionalAttentionBatchingService:
|
||||
raise ValueError("Regional attention context devices must match.")
|
||||
if tensor.dtype != authority.dtype:
|
||||
raise ValueError("Regional attention context dtypes must match.")
|
||||
if (
|
||||
entry.cross_attention_value_multiplier is not None
|
||||
) is not has_value_multiplier:
|
||||
raise ValueError(
|
||||
"Regional attention value multiplier presence must be uniform."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _align_value_multiplier(
|
||||
entries: tuple[ProcessedRegionalAttentionEntry, ...],
|
||||
*,
|
||||
latent_batch_size: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
target_sequence_length: int,
|
||||
) -> torch.Tensor | None:
|
||||
"""Repeat and sequence-align one chunk-major value-multiplier bank."""
|
||||
|
||||
if not entries or entries[0].cross_attention_value_multiplier is None:
|
||||
return None
|
||||
parts: list[torch.Tensor] = []
|
||||
for entry in entries:
|
||||
multiplier = entry.cross_attention_value_multiplier
|
||||
if multiplier is None:
|
||||
raise ValueError(
|
||||
"Regional attention value multiplier presence must be uniform."
|
||||
)
|
||||
repeated = repeat_to_batch_size(multiplier, latent_batch_size).to(
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
sequence_length = int(repeated.shape[1])
|
||||
if sequence_length < target_sequence_length:
|
||||
repeated = functional.pad(
|
||||
repeated,
|
||||
(0, 0, 0, target_sequence_length - sequence_length),
|
||||
value=1.0,
|
||||
)
|
||||
elif sequence_length > target_sequence_length:
|
||||
repeated = repeated[:, :target_sequence_length]
|
||||
parts.append(repeated)
|
||||
return torch.cat(parts, dim=0)
|
||||
|
||||
@staticmethod
|
||||
def _entries(
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from ..patcher_lifecycle import ModelMutation
|
||||
from ..ppm_negpip_interop import PpmNegpipInterop
|
||||
from .anima_activation_context import (
|
||||
ANIMA_ACTIVATION_CONTEXT,
|
||||
AnimaActivationContext,
|
||||
@@ -75,6 +76,7 @@ def anima_attention_coupling_mutations(
|
||||
),
|
||||
phase_context: AnimaCompositionPhaseContext = (ANIMA_COMPOSITION_PHASE_CONTEXT),
|
||||
query_mask_context: AnimaQueryMaskContext = ANIMA_QUERY_MASK_CONTEXT,
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> tuple[ModelMutation, ...]:
|
||||
"""Return attention-only or complete regional-LoRA mutation composition."""
|
||||
|
||||
@@ -106,6 +108,7 @@ def anima_attention_coupling_mutations(
|
||||
invocation_context=cross_attention_context,
|
||||
phase_context=phase_context,
|
||||
query_activity=query_activity,
|
||||
negpip=negpip,
|
||||
)
|
||||
phase_wrapper = anima_composition_phase_wrapper_mutation(
|
||||
surface,
|
||||
|
||||
@@ -18,6 +18,7 @@ from ...domain.regional_conditioning_output import (
|
||||
RegionalConditioningOutputCombiner,
|
||||
)
|
||||
from ..model_patcher_mutations import ModelExactObjectPatchMutation
|
||||
from ..ppm_negpip_interop import PpmNegpipInterop
|
||||
from .anima_activation_context import (
|
||||
ANIMA_ACTIVATION_CONTEXT,
|
||||
AnimaActivationContext,
|
||||
@@ -26,6 +27,7 @@ from .anima_activation_context import (
|
||||
from .anima_attention_execution import AnimaRegionalAttentionExecution
|
||||
from .anima_branch_batch import (
|
||||
ANIMA_BASE_BRANCH_KEY,
|
||||
AnimaRegionalBranchBatch,
|
||||
AnimaRegionalBranchKey,
|
||||
)
|
||||
from .anima_composition_phase_context import AnimaCompositionPhaseContext
|
||||
@@ -71,6 +73,7 @@ class AnimaRegionalCrossAttentionPatch(nn.Module):
|
||||
),
|
||||
weighting: RegionalAttentionWeightingPolicy | None = None,
|
||||
entry_combiner: RegionalConditioningOutputCombiner | None = None,
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> None:
|
||||
"""Retain the exact original attention owner and focused collaborators."""
|
||||
|
||||
@@ -92,6 +95,9 @@ class AnimaRegionalCrossAttentionPatch(nn.Module):
|
||||
self._query_activity = query_activity
|
||||
self._weighting = weighting or ANIMA_CROSS_ATTENTION_WEIGHTING_POLICY
|
||||
self._entry_combiner = entry_combiner or REGIONAL_CONDITIONING_OUTPUT_COMBINER
|
||||
if negpip is not None and not isinstance(negpip, PpmNegpipInterop):
|
||||
raise TypeError("Anima cross-attention NegPiP state has an invalid type.")
|
||||
self._negpip = negpip
|
||||
|
||||
def __setattr__(self, name: str, value: Any) -> None:
|
||||
"""Keep later exact child patches synchronized with installed attention."""
|
||||
@@ -143,14 +149,18 @@ class AnimaRegionalCrossAttentionPatch(nn.Module):
|
||||
branch_batch = activity.attention_branches
|
||||
branch_x = branch_batch.pack_source(x)
|
||||
branch_context = branch_batch.pack_branch_values(branch_values)
|
||||
original_options = {} if transformer_options is None else transformer_options
|
||||
forwarded_options = self._prepare_transformer_options(
|
||||
original_options,
|
||||
branch_batch,
|
||||
execution_contexts,
|
||||
)
|
||||
with self._invocation_context.activate(branch_batch.invocation):
|
||||
branch_output = self._backing.module(
|
||||
branch_x,
|
||||
branch_context,
|
||||
rope_emb=rope_emb,
|
||||
transformer_options=(
|
||||
{} if transformer_options is None else transformer_options
|
||||
),
|
||||
transformer_options=forwarded_options,
|
||||
)
|
||||
if not isinstance(branch_output, torch.Tensor):
|
||||
raise TypeError("Original Anima cross-attention must return a tensor.")
|
||||
@@ -186,6 +196,41 @@ class AnimaRegionalCrossAttentionPatch(nn.Module):
|
||||
regional_outputs=torch.stack(regional_outputs),
|
||||
)
|
||||
|
||||
def _prepare_transformer_options(
|
||||
self,
|
||||
source: dict[str, Any],
|
||||
branch_batch: AnimaRegionalBranchBatch,
|
||||
contexts: BatchedRegionalAttentionContexts,
|
||||
) -> dict[str, object]:
|
||||
"""Pack NegPiP multipliers or preserve ordinary option identity."""
|
||||
|
||||
if self._negpip is None:
|
||||
if contexts.base_value_multiplier is not None:
|
||||
raise ValueError(
|
||||
"Anima value multipliers require admitted NegPiP semantics."
|
||||
)
|
||||
return source
|
||||
if not isinstance(branch_batch, AnimaRegionalBranchBatch):
|
||||
raise TypeError("Anima NegPiP branch batch has an invalid type.")
|
||||
base_multiplier = contexts.base_value_multiplier
|
||||
if base_multiplier is None:
|
||||
raise ValueError("Anima NegPiP requires an aligned base value multiplier.")
|
||||
values = {ANIMA_BASE_BRANCH_KEY: base_multiplier}
|
||||
for region in contexts.regions:
|
||||
for entry in region.entries:
|
||||
multiplier = entry.cross_attention_value_multiplier
|
||||
if multiplier is None:
|
||||
raise ValueError(
|
||||
"Anima NegPiP requires every regional value multiplier."
|
||||
)
|
||||
values[
|
||||
AnimaRegionalBranchKey(region.region_index, entry.entry_index)
|
||||
] = multiplier
|
||||
return self._negpip.prepare_anima_transformer_options(
|
||||
source,
|
||||
branch_batch.pack_branch_values(values),
|
||||
)
|
||||
|
||||
def _validate_inputs(
|
||||
self,
|
||||
*,
|
||||
@@ -246,6 +291,7 @@ def anima_cross_attention_mutations(
|
||||
query_activity: AnimaRegionalQueryActivityContext = (
|
||||
ANIMA_REGIONAL_QUERY_ACTIVITY_CONTEXT
|
||||
),
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> tuple[ModelExactObjectPatchMutation, ...]:
|
||||
"""Build one exact clone-local cross-attention replacement per Anima block."""
|
||||
|
||||
@@ -262,6 +308,7 @@ def anima_cross_attention_mutations(
|
||||
invocation_context=invocation_context,
|
||||
phase_context=phase_context,
|
||||
query_activity=query_activity,
|
||||
negpip=negpip,
|
||||
),
|
||||
)
|
||||
for block in surface.blocks
|
||||
|
||||
@@ -10,6 +10,7 @@ from dataclasses import dataclass
|
||||
|
||||
from ...domain.processed_regional_attention import ProcessedRegionalAttentionPlan
|
||||
from ..patcher_lifecycle import PATCHER_LIFECYCLE
|
||||
from ..ppm_negpip_interop import PpmNegpipInterop
|
||||
from ..regional_attention_template import build_regional_attention_template
|
||||
from ..regional_lora_plan_adapter import RegionalLoraPlanAdaptation
|
||||
from .anima_attention_context_wrapper import (
|
||||
@@ -44,6 +45,7 @@ class FullContextAnimaAttentionBackend:
|
||||
adaptation: RegionalLoraPlanAdaptation,
|
||||
region_strengths: tuple[float, ...],
|
||||
latent_batch_size: int,
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> FullContextAnimaAttentionModel:
|
||||
"""Return one collision-safe clone prepared for dynamic sampler calls."""
|
||||
|
||||
@@ -51,6 +53,8 @@ class FullContextAnimaAttentionBackend:
|
||||
raise TypeError("Anima backend requires a processed attention plan.")
|
||||
if not isinstance(adaptation, RegionalLoraPlanAdaptation):
|
||||
raise TypeError("Anima backend requires a regional LoRA adaptation.")
|
||||
if negpip is not None and not isinstance(negpip, PpmNegpipInterop):
|
||||
raise TypeError("Anima backend NegPiP state has an invalid type.")
|
||||
admitted = ANIMA_REGIONAL_LORA_PLAN_ADMISSION_SERVICE.admit(adaptation)
|
||||
ANIMA_GLOBAL_REGIONAL_LORA_OVERLAP_VALIDATOR.validate(model, admitted)
|
||||
template = build_regional_attention_template(
|
||||
@@ -89,6 +93,7 @@ class FullContextAnimaAttentionBackend:
|
||||
surface,
|
||||
attention,
|
||||
composition=composition,
|
||||
negpip=negpip,
|
||||
),
|
||||
)
|
||||
derived = PATCHER_LIFECYCLE.derive_model(
|
||||
|
||||
@@ -5,22 +5,62 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import ModuleType
|
||||
from typing import Protocol, cast, runtime_checkable
|
||||
|
||||
import torch
|
||||
|
||||
from .fused_active_accumulation_kernel import (
|
||||
MAX_FUSED_ADAPTERS_PER_LAUNCH,
|
||||
REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL,
|
||||
from .fused_active_accumulation_contract import MAX_FUSED_ADAPTERS_PER_LAUNCH
|
||||
from .triton_runtime import TRITON_RUNTIME_RESOLVER, TritonRuntimeResolver
|
||||
|
||||
_TRITON_BACKEND_MODULE = (
|
||||
"simple_syrup.runtime.regional_lora.fused_active_accumulation_kernel"
|
||||
)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _FusedAccumulationKernel(Protocol):
|
||||
"""Describe the lazily resolved fused CUDA launch surface."""
|
||||
|
||||
def launch(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
*,
|
||||
rank_values: torch.Tensor,
|
||||
up: torch.Tensor,
|
||||
multipliers: tuple[torch.Tensor, ...],
|
||||
indices: torch.Tensor | None,
|
||||
adapter_start: int,
|
||||
target_indices: torch.Tensor | None = None,
|
||||
) -> None:
|
||||
"""Launch one validated fused accumulation chunk."""
|
||||
|
||||
...
|
||||
|
||||
|
||||
class _FusedAccumulationBackend(Protocol):
|
||||
"""Describe the exported lazy backend module surface."""
|
||||
|
||||
REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL: _FusedAccumulationKernel
|
||||
|
||||
|
||||
class RegionalLoraFusedActiveAccumulator:
|
||||
"""Own verified CUDA fusion for B projection and ordered output updates."""
|
||||
|
||||
@staticmethod
|
||||
def supports(output: torch.Tensor, up: torch.Tensor) -> bool:
|
||||
def __init__(
|
||||
self,
|
||||
resolver: TritonRuntimeResolver = TRITON_RUNTIME_RESOLVER,
|
||||
) -> None:
|
||||
"""Retain the process-level optional acceleration authority."""
|
||||
|
||||
if not isinstance(resolver, TritonRuntimeResolver):
|
||||
raise TypeError("Fused accumulation requires a Triton resolver.")
|
||||
self._resolver = resolver
|
||||
|
||||
def supports(self, output: torch.Tensor, up: torch.Tensor) -> bool:
|
||||
"""Admit only verified contiguous CUDA bf16/fp16 projection shapes."""
|
||||
|
||||
return (
|
||||
eligible = (
|
||||
isinstance(output, torch.Tensor)
|
||||
and isinstance(up, torch.Tensor)
|
||||
and output.device.type == "cuda"
|
||||
@@ -32,6 +72,7 @@ class RegionalLoraFusedActiveAccumulator:
|
||||
and output.is_contiguous()
|
||||
and up.is_contiguous()
|
||||
)
|
||||
return eligible and self._backend() is not None
|
||||
|
||||
def add(
|
||||
self,
|
||||
@@ -59,7 +100,7 @@ class RegionalLoraFusedActiveAccumulator:
|
||||
return output
|
||||
for start in range(0, adapter_count, MAX_FUSED_ADAPTERS_PER_LAUNCH):
|
||||
stop = min(start + MAX_FUSED_ADAPTERS_PER_LAUNCH, adapter_count)
|
||||
REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL.launch(
|
||||
self._require_backend().launch(
|
||||
output,
|
||||
rank_values=rank_values,
|
||||
up=up,
|
||||
@@ -97,7 +138,7 @@ class RegionalLoraFusedActiveAccumulator:
|
||||
)
|
||||
if active_row_count == 0:
|
||||
return output
|
||||
REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL.launch(
|
||||
self._require_backend().launch(
|
||||
output,
|
||||
rank_values=rank_values,
|
||||
up=up,
|
||||
@@ -137,7 +178,7 @@ class RegionalLoraFusedActiveAccumulator:
|
||||
return output
|
||||
for start in range(0, len(multipliers), MAX_FUSED_ADAPTERS_PER_LAUNCH):
|
||||
stop = min(start + MAX_FUSED_ADAPTERS_PER_LAUNCH, len(multipliers))
|
||||
REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL.launch(
|
||||
self._require_backend().launch(
|
||||
output,
|
||||
rank_values=rank_values,
|
||||
up=up,
|
||||
@@ -148,9 +189,28 @@ class RegionalLoraFusedActiveAccumulator:
|
||||
)
|
||||
return output
|
||||
|
||||
@classmethod
|
||||
def _backend(self) -> _FusedAccumulationKernel | None:
|
||||
"""Return the cached optional fused kernel without hiding failures."""
|
||||
|
||||
backend = self._resolver.resolve(_TRITON_BACKEND_MODULE)
|
||||
if backend is None:
|
||||
return None
|
||||
module = cast(_FusedAccumulationBackend, cast(ModuleType, backend))
|
||||
kernel = module.REGIONAL_LORA_FUSED_ACCUMULATION_KERNEL
|
||||
if not isinstance(kernel, _FusedAccumulationKernel):
|
||||
raise TypeError("Triton fused backend has an invalid kernel surface.")
|
||||
return kernel
|
||||
|
||||
def _require_backend(self) -> _FusedAccumulationKernel:
|
||||
"""Return the admitted kernel or reject an invalid direct fused call."""
|
||||
|
||||
backend = self._backend()
|
||||
if backend is None:
|
||||
raise RuntimeError("Triton fused accumulation is unavailable.")
|
||||
return backend
|
||||
|
||||
def _validate_mapped_projection(
|
||||
cls,
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
rank_values: torch.Tensor,
|
||||
up: torch.Tensor,
|
||||
@@ -159,7 +219,7 @@ class RegionalLoraFusedActiveAccumulator:
|
||||
) -> None:
|
||||
"""Require one unique-target batch and valid declared group mapping."""
|
||||
|
||||
if not cls.supports(output, up):
|
||||
if not self.supports(output, up):
|
||||
raise ValueError("Mapped fused accumulation received an unsupported path.")
|
||||
if (
|
||||
rank_values.device != output.device
|
||||
@@ -191,9 +251,8 @@ class RegionalLoraFusedActiveAccumulator:
|
||||
):
|
||||
raise ValueError("Mapped fused output indices are misaligned.")
|
||||
|
||||
@classmethod
|
||||
def _validate_projection(
|
||||
cls,
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
rank_values: torch.Tensor,
|
||||
up: torch.Tensor,
|
||||
@@ -201,7 +260,7 @@ class RegionalLoraFusedActiveAccumulator:
|
||||
) -> None:
|
||||
"""Require one complete aligned CUDA projection contract."""
|
||||
|
||||
if not cls.supports(output, up):
|
||||
if not self.supports(output, up):
|
||||
raise ValueError("Fused active accumulation received an unsupported path.")
|
||||
if (
|
||||
rank_values.device != output.device
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Define the shared launch-width contract for fused LoRA accumulation."""
|
||||
|
||||
MAX_FUSED_ADAPTERS_PER_LAUNCH = 8
|
||||
@@ -14,7 +14,8 @@ import torch
|
||||
import triton # type: ignore[import-untyped]
|
||||
import triton.language as tl # type: ignore[import-untyped]
|
||||
|
||||
MAX_FUSED_ADAPTERS_PER_LAUNCH = 8
|
||||
from .fused_active_accumulation_contract import MAX_FUSED_ADAPTERS_PER_LAUNCH
|
||||
|
||||
_BLOCK_ROWS = 16
|
||||
_BLOCK_OUTPUT_FEATURES = 64
|
||||
|
||||
|
||||
@@ -1,67 +1,48 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
# mypy: disable-error-code="no-untyped-def"
|
||||
# ruff: noqa: ANN001, ANN202
|
||||
|
||||
"""Accumulate ordered adapter outputs with exact execution-dtype rounding."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, cast
|
||||
from types import ModuleType
|
||||
from typing import Protocol, cast
|
||||
|
||||
import torch
|
||||
import triton # type: ignore[import-untyped]
|
||||
import triton.language as tl # type: ignore[import-untyped]
|
||||
|
||||
_MAX_DELTAS_PER_LAUNCH = 8
|
||||
_BLOCK_SIZE = 256
|
||||
from .triton_runtime import TRITON_RUNTIME_RESOLVER, TritonRuntimeResolver
|
||||
|
||||
_TRITON_BACKEND_MODULE = (
|
||||
"simple_syrup.runtime.regional_lora.ordered_accumulation_triton"
|
||||
)
|
||||
|
||||
|
||||
@triton.jit # type: ignore[untyped-decorator]
|
||||
def _ordered_accumulation_kernel(
|
||||
output,
|
||||
base,
|
||||
delta_0,
|
||||
delta_1,
|
||||
delta_2,
|
||||
delta_3,
|
||||
delta_4,
|
||||
delta_5,
|
||||
delta_6,
|
||||
delta_7,
|
||||
element_count,
|
||||
delta_count: tl.constexpr,
|
||||
execution_dtype: tl.constexpr,
|
||||
block_size: tl.constexpr,
|
||||
):
|
||||
"""Add one ordered chunk and round after every declared adapter."""
|
||||
class _OrderedAccumulationBackend(Protocol):
|
||||
"""Describe the lazy CUDA backend surface consumed by this owner."""
|
||||
|
||||
offsets = tl.program_id(0) * block_size + tl.arange(0, block_size)
|
||||
active = offsets < element_count
|
||||
value = tl.load(base + offsets, mask=active)
|
||||
if delta_count > 0:
|
||||
value = (value + tl.load(delta_0 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 1:
|
||||
value = (value + tl.load(delta_1 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 2:
|
||||
value = (value + tl.load(delta_2 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 3:
|
||||
value = (value + tl.load(delta_3 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 4:
|
||||
value = (value + tl.load(delta_4 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 5:
|
||||
value = (value + tl.load(delta_5 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 6:
|
||||
value = (value + tl.load(delta_6 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 7:
|
||||
value = (value + tl.load(delta_7 + offsets, mask=active)).to(execution_dtype)
|
||||
tl.store(output + offsets, value, mask=active)
|
||||
def accumulate(
|
||||
self,
|
||||
base: torch.Tensor,
|
||||
deltas: tuple[torch.Tensor, ...],
|
||||
) -> torch.Tensor:
|
||||
"""Accumulate validated CUDA tensors in declared order."""
|
||||
|
||||
...
|
||||
|
||||
|
||||
class OrderedTensorAccumulator:
|
||||
"""Own exact ordered accumulation and its CUDA launch policy."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
resolver: TritonRuntimeResolver = TRITON_RUNTIME_RESOLVER,
|
||||
) -> None:
|
||||
"""Retain the process-level optional acceleration authority."""
|
||||
|
||||
if not isinstance(resolver, TritonRuntimeResolver):
|
||||
raise TypeError("Ordered accumulation requires a Triton resolver.")
|
||||
self._resolver = resolver
|
||||
|
||||
def accumulate(
|
||||
self,
|
||||
base: torch.Tensor,
|
||||
@@ -74,25 +55,13 @@ class OrderedTensorAccumulator:
|
||||
return base
|
||||
if base.device.type != "cuda":
|
||||
return self._torch_accumulate(base, deltas)
|
||||
execution_dtype = self._triton_dtype(base.dtype)
|
||||
remaining = deltas
|
||||
result = base
|
||||
while remaining:
|
||||
chunk = remaining[:_MAX_DELTAS_PER_LAUNCH]
|
||||
remaining = remaining[_MAX_DELTAS_PER_LAUNCH:]
|
||||
padded = (*chunk, *((result,) * (_MAX_DELTAS_PER_LAUNCH - len(chunk))))
|
||||
grid = (triton.cdiv(result.numel(), _BLOCK_SIZE),)
|
||||
kernel = cast(Any, _ordered_accumulation_kernel)
|
||||
kernel[grid](
|
||||
result,
|
||||
result,
|
||||
*padded,
|
||||
result.numel(),
|
||||
delta_count=len(chunk),
|
||||
execution_dtype=execution_dtype,
|
||||
block_size=_BLOCK_SIZE,
|
||||
)
|
||||
return result
|
||||
backend = self._resolver.resolve(_TRITON_BACKEND_MODULE)
|
||||
if backend is None:
|
||||
return self._torch_accumulate(base, deltas)
|
||||
return cast(_OrderedAccumulationBackend, cast(ModuleType, backend)).accumulate(
|
||||
base,
|
||||
deltas,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _torch_accumulate(
|
||||
@@ -106,18 +75,6 @@ class OrderedTensorAccumulator:
|
||||
result.add_(delta)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _triton_dtype(dtype: torch.dtype) -> Any:
|
||||
"""Map the admitted floating execution dtype to a Triton scalar dtype."""
|
||||
|
||||
if dtype is torch.bfloat16:
|
||||
return tl.bfloat16
|
||||
if dtype is torch.float16:
|
||||
return tl.float16
|
||||
if dtype is torch.float32:
|
||||
return tl.float32
|
||||
raise TypeError(f"Ordered CUDA accumulation does not support {dtype}.")
|
||||
|
||||
@staticmethod
|
||||
def _validate(
|
||||
base: torch.Tensor,
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
# mypy: disable-error-code="no-untyped-def"
|
||||
# ruff: noqa: ANN001, ANN202
|
||||
|
||||
"""Provide the lazily imported Triton ordered-accumulation backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
import torch
|
||||
import triton # type: ignore[import-untyped]
|
||||
import triton.language as tl # type: ignore[import-untyped]
|
||||
|
||||
_MAX_DELTAS_PER_LAUNCH = 8
|
||||
_BLOCK_SIZE = 256
|
||||
|
||||
|
||||
@triton.jit # type: ignore[untyped-decorator]
|
||||
def _ordered_accumulation_kernel(
|
||||
output,
|
||||
base,
|
||||
delta_0,
|
||||
delta_1,
|
||||
delta_2,
|
||||
delta_3,
|
||||
delta_4,
|
||||
delta_5,
|
||||
delta_6,
|
||||
delta_7,
|
||||
element_count,
|
||||
delta_count: tl.constexpr,
|
||||
execution_dtype: tl.constexpr,
|
||||
block_size: tl.constexpr,
|
||||
):
|
||||
"""Add one ordered chunk and round after every declared adapter."""
|
||||
|
||||
offsets = tl.program_id(0) * block_size + tl.arange(0, block_size)
|
||||
active = offsets < element_count
|
||||
value = tl.load(base + offsets, mask=active)
|
||||
if delta_count > 0:
|
||||
value = (value + tl.load(delta_0 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 1:
|
||||
value = (value + tl.load(delta_1 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 2:
|
||||
value = (value + tl.load(delta_2 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 3:
|
||||
value = (value + tl.load(delta_3 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 4:
|
||||
value = (value + tl.load(delta_4 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 5:
|
||||
value = (value + tl.load(delta_5 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 6:
|
||||
value = (value + tl.load(delta_6 + offsets, mask=active)).to(execution_dtype)
|
||||
if delta_count > 7:
|
||||
value = (value + tl.load(delta_7 + offsets, mask=active)).to(execution_dtype)
|
||||
tl.store(output + offsets, value, mask=active)
|
||||
|
||||
|
||||
def accumulate(base: torch.Tensor, deltas: tuple[torch.Tensor, ...]) -> torch.Tensor:
|
||||
"""Accumulate CUDA deltas with the established launch and rounding policy."""
|
||||
|
||||
execution_dtype = _triton_dtype(base.dtype)
|
||||
remaining = deltas
|
||||
result = base
|
||||
while remaining:
|
||||
chunk = remaining[:_MAX_DELTAS_PER_LAUNCH]
|
||||
remaining = remaining[_MAX_DELTAS_PER_LAUNCH:]
|
||||
padded = (*chunk, *((result,) * (_MAX_DELTAS_PER_LAUNCH - len(chunk))))
|
||||
grid = (triton.cdiv(result.numel(), _BLOCK_SIZE),)
|
||||
kernel = cast(Any, _ordered_accumulation_kernel)
|
||||
kernel[grid](
|
||||
result,
|
||||
result,
|
||||
*padded,
|
||||
result.numel(),
|
||||
delta_count=len(chunk),
|
||||
execution_dtype=execution_dtype,
|
||||
block_size=_BLOCK_SIZE,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def _triton_dtype(dtype: torch.dtype) -> Any:
|
||||
"""Map the admitted floating execution dtype to a Triton scalar dtype."""
|
||||
|
||||
if dtype is torch.bfloat16:
|
||||
return tl.bfloat16
|
||||
if dtype is torch.float16:
|
||||
return tl.float16
|
||||
if dtype is torch.float32:
|
||||
return tl.float32
|
||||
raise TypeError(f"Ordered CUDA accumulation does not support {dtype}.")
|
||||
@@ -10,6 +10,7 @@ from ..attention_coupling.unet_attn2_execution_resolver import (
|
||||
UnetAttn2ExecutionResolver,
|
||||
)
|
||||
from ..attention_coupling.unet_attn2_patch import UnetAttn2PatchPair
|
||||
from ..ppm_negpip_interop import PpmNegpipInterop
|
||||
|
||||
_INPUT_PATCH_KEY = "attn2_patch"
|
||||
_OUTPUT_PATCH_KEY = "attn2_output_patch"
|
||||
@@ -18,12 +19,20 @@ _OUTPUT_PATCH_KEY = "attn2_output_patch"
|
||||
class StandardUnetVariantBaseAttention:
|
||||
"""Install Attention Couple only on graph-local unpatched base execution."""
|
||||
|
||||
def __init__(self, resolver: UnetAttn2ExecutionResolver) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
resolver: UnetAttn2ExecutionResolver,
|
||||
*,
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> None:
|
||||
"""Retain one paired callback authority for the active request state."""
|
||||
|
||||
if not isinstance(resolver, UnetAttn2ExecutionResolver):
|
||||
raise TypeError("Standard UNet base attention requires a resolver.")
|
||||
self._patches = UnetAttn2PatchPair(resolver)
|
||||
if negpip is not None and not isinstance(negpip, PpmNegpipInterop):
|
||||
raise TypeError("Standard UNet base attention NegPiP state is invalid.")
|
||||
self._negpip = negpip
|
||||
|
||||
def prepare(self, source: dict[str, object]) -> dict[str, object]:
|
||||
"""Return isolated transformer options with one collision-free pair."""
|
||||
@@ -40,10 +49,22 @@ class StandardUnetVariantBaseAttention:
|
||||
patches = source_patches.copy()
|
||||
else:
|
||||
raise TypeError("Standard UNet transformer patches must be a dictionary.")
|
||||
for key in (_INPUT_PATCH_KEY, _OUTPUT_PATCH_KEY):
|
||||
if key in patches:
|
||||
raise ValueError(f"Standard UNet base graph already contains {key!r}.")
|
||||
patches[_INPUT_PATCH_KEY] = [self._patches.input_patch]
|
||||
expected_input = [] if self._negpip is None else [self._negpip.attention_patch]
|
||||
if (
|
||||
self._negpip is None
|
||||
and _INPUT_PATCH_KEY in patches
|
||||
or self._negpip is not None
|
||||
and patches.get(_INPUT_PATCH_KEY) != expected_input
|
||||
):
|
||||
raise ValueError(
|
||||
"Standard UNet base graph already contains an unadmitted attn2 "
|
||||
"input patch."
|
||||
)
|
||||
if _OUTPUT_PATCH_KEY in patches:
|
||||
raise ValueError(
|
||||
"Standard UNet base graph already contains 'attn2_output_patch'."
|
||||
)
|
||||
patches[_INPUT_PATCH_KEY] = [self._patches.input_patch, *expected_input]
|
||||
patches[_OUTPUT_PATCH_KEY] = [self._patches.output_patch]
|
||||
prepared["patches"] = patches
|
||||
return prepared
|
||||
|
||||
@@ -22,6 +22,7 @@ from ..model_patcher_mutations import (
|
||||
ModelKeyedCallbackMutation,
|
||||
ModelKeyedWrapperMutation,
|
||||
)
|
||||
from ..ppm_negpip_interop import PpmNegpipInterop
|
||||
from .standard_unet_cold_sampling import (
|
||||
StandardUnetColdSamplingDiagnosticsMutation,
|
||||
)
|
||||
@@ -48,6 +49,7 @@ class StandardUnetVariantRuntimeMutation:
|
||||
admission: StandardUnetNativeLoraAdmission
|
||||
attention_phase: StandardUnetAttentionPhaseSession
|
||||
template: StandardUnetVariantTemplate
|
||||
negpip: PpmNegpipInterop | None = None
|
||||
|
||||
def apply(self, model: object) -> None:
|
||||
"""Build persistent variants before installing the private root clone."""
|
||||
@@ -69,7 +71,8 @@ class StandardUnetVariantRuntimeMutation:
|
||||
),
|
||||
attention_phase=self.attention_phase,
|
||||
base_attention=StandardUnetVariantBaseAttention(
|
||||
StandardUnetAttn2ExecutionResolver(self.state)
|
||||
StandardUnetAttn2ExecutionResolver(self.state),
|
||||
negpip=self.negpip,
|
||||
),
|
||||
)
|
||||
execution.prime()
|
||||
|
||||
@@ -148,8 +148,10 @@ class StandardUnetVariantTemplate:
|
||||
)
|
||||
if not isinstance(request, ModelPatcher) or request.is_dynamic():
|
||||
raise TypeError("Comfy did not bind a static standard-UNet request.")
|
||||
if request.parent is not source:
|
||||
raise RuntimeError("Static standard-UNet request lost source lineage.")
|
||||
if request.parent is not self.model:
|
||||
raise RuntimeError(
|
||||
"Static standard-UNet request lost its template fallback boundary."
|
||||
)
|
||||
if "diffusion_model" in request.object_patches_backup:
|
||||
ModelSharedObjectPatchMutation(
|
||||
"diffusion_model",
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Resolve optional Triton backends without importing them on Torch paths."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from importlib import import_module
|
||||
from threading import Lock
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TritonRuntimeResolver:
|
||||
"""Own thread-safe lazy backend imports and optional-package fallback."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
import_module: Callable[[str], object] = import_module,
|
||||
) -> None:
|
||||
"""Retain an injectable importer and empty process-lifetime cache."""
|
||||
|
||||
if not callable(import_module):
|
||||
raise TypeError("Triton runtime importer must be callable.")
|
||||
self._import_module = import_module
|
||||
self._lock = Lock()
|
||||
self._backends: dict[str, object | None] = {}
|
||||
self._missing_warning_emitted = False
|
||||
|
||||
def resolve(self, backend_module: str) -> object | None:
|
||||
"""Return one cached backend or None only when Triton is absent."""
|
||||
|
||||
if not isinstance(backend_module, str) or not backend_module:
|
||||
raise ValueError("Triton backend module must be a non-empty string.")
|
||||
with self._lock:
|
||||
if backend_module in self._backends:
|
||||
return self._backends[backend_module]
|
||||
try:
|
||||
backend = self._import_module(backend_module)
|
||||
except ModuleNotFoundError as error:
|
||||
if error.name != "triton" and not (
|
||||
isinstance(error.name, str) and error.name.startswith("triton.")
|
||||
):
|
||||
raise RuntimeError(
|
||||
f"Triton backend {backend_module!r} failed to initialize."
|
||||
) from error
|
||||
backend = None
|
||||
if not self._missing_warning_emitted:
|
||||
LOGGER.warning(
|
||||
"Triton acceleration is unavailable; using the Torch "
|
||||
"execution path",
|
||||
extra={
|
||||
"backend_module": backend_module,
|
||||
"missing_dependency": error.name,
|
||||
},
|
||||
)
|
||||
self._missing_warning_emitted = True
|
||||
except Exception as error:
|
||||
raise RuntimeError(
|
||||
f"Triton backend {backend_module!r} failed to initialize."
|
||||
) from error
|
||||
self._backends[backend_module] = backend
|
||||
return backend
|
||||
|
||||
|
||||
TRITON_RUNTIME_RESOLVER = TritonRuntimeResolver()
|
||||
@@ -20,11 +20,13 @@ from ..domain.regional_model_capabilities import (
|
||||
RegionalModelFamily,
|
||||
RegionalPatchConflict,
|
||||
)
|
||||
from .ppm_negpip_interop import (
|
||||
PPM_NEGPIP_INTEROP_VALIDATOR,
|
||||
PpmNegpipInterop,
|
||||
)
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
_NEGPIP_MODEL_OPTION = "ppm_negpip"
|
||||
_NEGPIP_ANIMA_WRAPPER_KEY = "ppm_negpip_anima"
|
||||
_EASYCACHE_OPTION = "easycache"
|
||||
_ATTN2_PATCH_CONFLICTS = {
|
||||
RegionalPatchConflict.ATTN2_INPUT_PATCH: "attn2_patch",
|
||||
@@ -49,6 +51,7 @@ class RegionalModelPatchInteropReport:
|
||||
|
||||
model_family: RegionalModelFamily
|
||||
preserved_modifiers: tuple[RegionalPreservedModelModifier, ...]
|
||||
negpip: PpmNegpipInterop | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require a typed family and canonical unique modifier order."""
|
||||
@@ -62,6 +65,8 @@ class RegionalModelPatchInteropReport:
|
||||
raise TypeError("Regional interop report modifiers have invalid types.")
|
||||
if len(set(self.preserved_modifiers)) != len(self.preserved_modifiers):
|
||||
raise ValueError("Regional interop report modifiers must be unique.")
|
||||
if self.negpip is not None and not isinstance(self.negpip, PpmNegpipInterop):
|
||||
raise TypeError("Regional interop report NegPiP state has an invalid type.")
|
||||
|
||||
@property
|
||||
def cache_modifier(self) -> RegionalPreservedModelModifier | None:
|
||||
@@ -99,12 +104,22 @@ class RegionalModelPatchInteropValidator:
|
||||
model_weight_patches = _require_dictionary_attribute(model, "patches")
|
||||
patches = _require_optional_patch_state(transformer_options)
|
||||
|
||||
self._reject_negpip(model_options, wrappers)
|
||||
negpip = PPM_NEGPIP_INTEROP_VALIDATOR.validate(
|
||||
capabilities.model_family,
|
||||
model_options=model_options,
|
||||
wrappers=wrappers,
|
||||
object_patches=object_patches,
|
||||
transformer_patches=patches,
|
||||
)
|
||||
cache_modifier = self._validate_cache_state(
|
||||
transformer_options,
|
||||
wrappers,
|
||||
)
|
||||
self._reject_attention_collisions(patches, capabilities)
|
||||
self._reject_attention_collisions(
|
||||
patches,
|
||||
capabilities,
|
||||
admitted_negpip=negpip,
|
||||
)
|
||||
|
||||
modifiers: list[RegionalPreservedModelModifier] = []
|
||||
model_wrapper = model_options.get("model_function_wrapper")
|
||||
@@ -135,6 +150,7 @@ class RegionalModelPatchInteropValidator:
|
||||
report = RegionalModelPatchInteropReport(
|
||||
capabilities.model_family,
|
||||
tuple(modifiers),
|
||||
negpip,
|
||||
)
|
||||
LOGGER.info(
|
||||
"Regional MODEL patch interoperability admitted",
|
||||
@@ -188,30 +204,6 @@ class RegionalModelPatchInteropValidator:
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _reject_negpip(
|
||||
model_options: dict[object, object],
|
||||
wrappers: dict[str, dict[object, list[object]]],
|
||||
) -> None:
|
||||
"""Reject installed NegPiP before its mask can enter branch packing."""
|
||||
|
||||
marker = model_options.get(_NEGPIP_MODEL_OPTION, False)
|
||||
if not isinstance(marker, bool):
|
||||
raise TypeError("MODEL ppm_negpip marker must be boolean.")
|
||||
negpip_wrapper = bool(
|
||||
wrappers.get(WrappersMP.DIFFUSION_MODEL, {}).get(
|
||||
_NEGPIP_ANIMA_WRAPPER_KEY,
|
||||
(),
|
||||
)
|
||||
)
|
||||
if marker or negpip_wrapper:
|
||||
raise ValueError(
|
||||
"Attention Coupling does not support NegPiP because its attention "
|
||||
"mask is aligned to the ordinary conditioning batch rather than "
|
||||
"SimpleSyrup's regional branch batch. Remove CLIP NegPip before "
|
||||
"the Attention Coupling sampler."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_cache_state(
|
||||
transformer_options: dict[object, object],
|
||||
@@ -260,6 +252,8 @@ class RegionalModelPatchInteropValidator:
|
||||
def _reject_attention_collisions(
|
||||
patches: dict[str, list[object]],
|
||||
capabilities: RegionalModelCapabilities,
|
||||
*,
|
||||
admitted_negpip: PpmNegpipInterop | None,
|
||||
) -> None:
|
||||
"""Reject every populated attention surface owned by the backend."""
|
||||
|
||||
@@ -268,6 +262,11 @@ class RegionalModelPatchInteropValidator:
|
||||
for conflict, patch_name in _ATTN2_PATCH_CONFLICTS.items()
|
||||
if conflict in capabilities.known_patch_conflicts
|
||||
and patches.get(patch_name)
|
||||
and not (
|
||||
patch_name == "attn2_patch"
|
||||
and admitted_negpip is not None
|
||||
and patches[patch_name] == [admitted_negpip.attention_patch]
|
||||
)
|
||||
)
|
||||
if conflicts:
|
||||
raise ValueError(
|
||||
|
||||
@@ -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,8 @@ 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()
|
||||
self._cache: ModelInstanceCache[
|
||||
UltralyticsModelCacheKey, LoadedUltralyticsDetector
|
||||
] = ModelInstanceCache(
|
||||
@@ -83,9 +93,39 @@ class UltralyticsLoaderService:
|
||||
)
|
||||
|
||||
def model_choices(self) -> list[str]:
|
||||
"""Return local Ultralytics model choices for ComfyUI dropdowns."""
|
||||
"""Return installed choices first, followed by downloadable catalog choices."""
|
||||
|
||||
choices = self.available_models()
|
||||
self._register_model_folders()
|
||||
curated_choices = self._choice_service.ultralytics_choices()
|
||||
catalog_choice_labels = {
|
||||
_catalog_selection(entry): entry.display_name
|
||||
for entry in ULTRALYTICS_ENTRIES
|
||||
}
|
||||
available_choices = self.available_models()
|
||||
visible_catalog_choices = set(curated_choices)
|
||||
installed_catalog_choices = [
|
||||
entry.display_name
|
||||
for entry in ULTRALYTICS_ENTRIES
|
||||
if (
|
||||
entry.display_name in visible_catalog_choices
|
||||
and _catalog_selection(entry) in available_choices
|
||||
)
|
||||
]
|
||||
installed_non_catalog_choices = [
|
||||
choice
|
||||
for choice in available_choices
|
||||
if choice not in catalog_choice_labels
|
||||
]
|
||||
downloadable_choices = [
|
||||
choice
|
||||
for choice in curated_choices
|
||||
if choice not in installed_catalog_choices
|
||||
]
|
||||
choices = (
|
||||
installed_non_catalog_choices
|
||||
+ installed_catalog_choices
|
||||
+ downloadable_choices
|
||||
)
|
||||
return choices or [NO_LOCAL_ULTRALYTICS_MODELS]
|
||||
|
||||
def available_models(self) -> list[str]:
|
||||
@@ -133,16 +173,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 +207,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 +308,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 +469,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."""
|
||||
|
||||
|
||||
@@ -107,5 +107,6 @@ class AnimaAttentionCouplingModelFamily:
|
||||
adaptation=admission.adaptation,
|
||||
region_strengths=region_strengths,
|
||||
latent_batch_size=latent_batch_size,
|
||||
negpip=interop_report.negpip,
|
||||
)
|
||||
return built.model
|
||||
|
||||
@@ -218,6 +218,7 @@ class AttentionCouplingModelPreparationService:
|
||||
noise=samples.to(device),
|
||||
device=device,
|
||||
context_validator=model_family.context_validator,
|
||||
negpip=interop_report.negpip,
|
||||
)
|
||||
interop_validator.validate_execution(
|
||||
interop_report,
|
||||
|
||||
@@ -142,6 +142,7 @@ class StandardUnetAttentionCouplingModelFamily:
|
||||
model=model,
|
||||
state=state,
|
||||
admission=admission,
|
||||
negpip=interop_report.negpip,
|
||||
)
|
||||
.model
|
||||
)
|
||||
|
||||
@@ -76,6 +76,7 @@ def test_anima_family_retains_single_frame_context_and_backend_policy() -> None:
|
||||
"adaptation": adaptation,
|
||||
"region_strengths": (0.75,),
|
||||
"latent_batch_size": 2,
|
||||
"negpip": None,
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@@ -31,6 +31,10 @@ from simple_syrup.domain.spatial_views import (
|
||||
SpatialViewKind,
|
||||
)
|
||||
from simple_syrup.runtime.patcher_lifecycle import PATCHER_LIFECYCLE
|
||||
from simple_syrup.runtime.ppm_negpip_interop import (
|
||||
PpmNegpipInterop,
|
||||
PpmNegpipSemantics,
|
||||
)
|
||||
from simple_syrup.runtime.regional_lora.anima_activation_context import (
|
||||
AnimaActivationContext,
|
||||
AnimaActivationGeometry,
|
||||
@@ -258,6 +262,79 @@ def test_patch_batches_complete_branches_and_blends_expected_outputs(
|
||||
assert invocation_context.current_or_none() is None
|
||||
|
||||
|
||||
def test_patch_packs_anima_negpip_masks_with_the_same_branch_segments() -> None:
|
||||
"""Align each compact regional context with its own NegPiP value mask."""
|
||||
|
||||
activation_context = AnimaActivationContext()
|
||||
invocation_context = AnimaCrossAttentionInvocationContext()
|
||||
original = _DeterministicCrossAttention(invocation_context)
|
||||
base = torch.zeros((1, 2, 1))
|
||||
region = torch.ones_like(base)
|
||||
base_multiplier = torch.tensor([[[1.0], [-1.0]]])
|
||||
region_multiplier = torch.tensor([[[-1.0], [1.0]]])
|
||||
contexts = BatchedRegionalAttentionContexts(
|
||||
latent_batch_size=1,
|
||||
chunks=(
|
||||
RegionalAttentionChunkBatch(
|
||||
0,
|
||||
RegionalAttentionBranch.POSITIVE,
|
||||
0,
|
||||
1,
|
||||
),
|
||||
),
|
||||
base_context=base,
|
||||
regions=(
|
||||
BatchedRegionalAttentionRegion(
|
||||
0,
|
||||
(
|
||||
BatchedRegionalAttentionEntry(
|
||||
0,
|
||||
region,
|
||||
(1.0,),
|
||||
region_multiplier,
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
base_value_multiplier=base_multiplier,
|
||||
)
|
||||
execution = AnimaRegionalAttentionExecution(
|
||||
contexts,
|
||||
_bank(torch.full((1, 1, 1), 0.5)),
|
||||
(1.0,),
|
||||
)
|
||||
negpip = PpmNegpipInterop(
|
||||
PpmNegpipSemantics.ANIMA_VALUE_MASK,
|
||||
lambda *args, **kwargs: (args, kwargs),
|
||||
)
|
||||
patch = AnimaRegionalCrossAttentionPatch(
|
||||
original,
|
||||
execution,
|
||||
activation_context=activation_context,
|
||||
invocation_context=invocation_context,
|
||||
phase_context=_FullRegionalPhaseContext(),
|
||||
negpip=negpip,
|
||||
)
|
||||
old_mask = torch.ones_like(base_multiplier)
|
||||
options: dict[str, object] = {"ppm_negpip_mask": old_mask}
|
||||
|
||||
with activation_context.activate(_geometry(batch=1, height=1, width=1)):
|
||||
patch(
|
||||
torch.zeros((1, 1, 1)),
|
||||
base,
|
||||
transformer_options=options,
|
||||
)
|
||||
|
||||
observed_options = original.calls[0][3]
|
||||
assert isinstance(observed_options, dict)
|
||||
assert observed_options is not options
|
||||
assert torch.equal(
|
||||
observed_options["ppm_negpip_mask"],
|
||||
torch.cat((base_multiplier, region_multiplier)),
|
||||
)
|
||||
assert options["ppm_negpip_mask"] is old_mask
|
||||
|
||||
|
||||
def test_cross_attention_backing_module_does_not_leak_a_host_weight_namespace() -> None:
|
||||
"""Keep the retained installed attention outside PyTorch child discovery."""
|
||||
|
||||
|
||||
@@ -25,10 +25,14 @@ from simple_syrup.domain.regional_attention_execution import (
|
||||
RegionalAttentionExecutionMode,
|
||||
)
|
||||
from simple_syrup.domain.regional_lora_plan import EMPTY_REGIONAL_LORA_PLAN
|
||||
from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily
|
||||
from simple_syrup.runtime.attention_coupling.family_admission import (
|
||||
AttentionCouplingFamilyAdmission,
|
||||
)
|
||||
from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation
|
||||
from simple_syrup.runtime.regional_model_patch_interop import (
|
||||
RegionalModelPatchInteropReport,
|
||||
)
|
||||
from simple_syrup.services.attention_coupling_model_family import (
|
||||
AttentionCouplingPreparedModelReuse,
|
||||
AttentionCouplingSamplerConditioning,
|
||||
@@ -59,10 +63,16 @@ class _CapabilityService:
|
||||
class _InteropValidator:
|
||||
"""Record centralized modifier admission without requiring a real patcher."""
|
||||
|
||||
report: ClassVar[object] = object()
|
||||
report: ClassVar[RegionalModelPatchInteropReport] = RegionalModelPatchInteropReport(
|
||||
RegionalModelFamily.ANIMA, ()
|
||||
)
|
||||
calls: ClassVar[list[tuple[object, ...]]] = []
|
||||
|
||||
def validate(self, model: object, capabilities: object) -> object:
|
||||
def validate(
|
||||
self,
|
||||
model: object,
|
||||
capabilities: object,
|
||||
) -> RegionalModelPatchInteropReport:
|
||||
"""Record exact orchestration inputs without changing them."""
|
||||
|
||||
type(self).calls.append((model, capabilities))
|
||||
|
||||
@@ -14,7 +14,10 @@ from typing import Any, cast
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.runtime.patcher_lifecycle import ComfyPatcherLifecycle
|
||||
from simple_syrup.runtime.patcher_lifecycle import (
|
||||
PATCHER_LIFECYCLE,
|
||||
ComfyPatcherLifecycle,
|
||||
)
|
||||
from simple_syrup.runtime.regional_lora.execution_cache import ModelCloneLineage
|
||||
|
||||
|
||||
@@ -64,6 +67,46 @@ def test_real_comfy_anima_lifecycle_regression(
|
||||
_assert_supported_model_mutations_share_one_clone()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("derivation_count", (2, 3, 5))
|
||||
def test_stacked_model_derivations_survive_simultaneous_cyclic_release(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
derivation_count: int,
|
||||
) -> None:
|
||||
"""Keep Comfy on a foreign boundary when any Syrup stack dies together."""
|
||||
|
||||
import comfy.model_management
|
||||
from comfy.model_management import LoadedModel
|
||||
|
||||
encoder = AnimaTEModel_()
|
||||
loader = _patcher(encoder)
|
||||
foreign_boundary = loader.clone()
|
||||
derived_models: list[object] = []
|
||||
current = foreign_boundary
|
||||
for stage in range(derivation_count):
|
||||
current = PATCHER_LIFECYCLE.derive_model(
|
||||
current,
|
||||
(),
|
||||
operation=f"stacked lifecycle regression stage {stage}",
|
||||
)
|
||||
derived_models.append(current)
|
||||
loaded = LoadedModel(current)
|
||||
loaded.real_model = weakref.ref(encoder)
|
||||
monkeypatch.setattr(comfy.model_management, "current_loaded_models", [loaded])
|
||||
|
||||
execution_cycle: list[object] = [*derived_models]
|
||||
execution_cycle.append(execution_cycle)
|
||||
del current, derived_models, execution_cycle
|
||||
gc.collect()
|
||||
with caplog.at_level(logging.INFO):
|
||||
comfy.model_management.cleanup_models_gc()
|
||||
|
||||
assert loaded.model is foreign_boundary
|
||||
assert loaded.is_dead() is False
|
||||
assert "Potential memory leak detected" not in caplog.text
|
||||
assert "WARNING, memory leak" not in caplog.text
|
||||
|
||||
|
||||
def test_clip_alignment_precedes_mutations_after_dynamic_to_static_clone() -> None:
|
||||
"""Mutate the same independently reloaded encoder the returned CLIP executes."""
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
from uuid import UUID
|
||||
|
||||
import comfy.conds
|
||||
@@ -31,6 +31,10 @@ from simple_syrup.runtime.attention_coupling.unet_context import (
|
||||
from simple_syrup.runtime.comfy_conditioning_processing import (
|
||||
COMFY_REGIONAL_CONDITIONING_PROCESSOR,
|
||||
)
|
||||
from simple_syrup.runtime.ppm_negpip_interop import (
|
||||
PpmNegpipInterop,
|
||||
PpmNegpipSemantics,
|
||||
)
|
||||
from simple_syrup.services.attention_coupling_preparation_service import (
|
||||
ATTENTION_COUPLING_PREPARATION_SERVICE,
|
||||
AttentionCouplingPreparation,
|
||||
@@ -89,6 +93,26 @@ class _LinearModelSampling:
|
||||
return 100.0 * (1.0 - float(percent))
|
||||
|
||||
|
||||
class _NegpipRecordingAnimaModel(_RecordingAnimaModel):
|
||||
"""Return a distinct PPM-style value mask for every processed context."""
|
||||
|
||||
def extra_conds(self, **kwargs: Any) -> dict[str, object]:
|
||||
"""Add a binary value mask beside the ordinary cross-attention output."""
|
||||
|
||||
result = super().extra_conds(**kwargs)
|
||||
output = self.outputs[-1]
|
||||
value = int(kwargs["negpip_value"])
|
||||
result["c_ppm_negpip_mask"] = comfy.conds.CONDRegular(
|
||||
torch.full(
|
||||
(*output.shape[:2], 1),
|
||||
value,
|
||||
dtype=torch.int32,
|
||||
device=output.device,
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def test_processor_uses_comfy_conversion_and_model_post_adapter_contexts() -> None:
|
||||
"""Retain weighted padded model outputs in exact positive/negative order."""
|
||||
|
||||
@@ -146,6 +170,47 @@ def test_processor_uses_comfy_conversion_and_model_post_adapter_contexts() -> No
|
||||
assert source_positive_context.shape == (1, 3, 1024)
|
||||
|
||||
|
||||
def test_processor_retains_each_anima_negpip_mask_with_its_scheduled_entry() -> None:
|
||||
"""Keep value semantics attached through Comfy conversion and UUID ownership."""
|
||||
|
||||
model = _NegpipRecordingAnimaModel()
|
||||
preparation = _preparation(
|
||||
positive=(
|
||||
_conditioning(1.0, negpip_value=1),
|
||||
_conditioning(2.0, negpip_value=-1),
|
||||
),
|
||||
negative=(
|
||||
_conditioning(-1.0, negpip_value=-1),
|
||||
_conditioning(-2.0, negpip_value=1),
|
||||
),
|
||||
)
|
||||
negpip = PpmNegpipInterop(
|
||||
PpmNegpipSemantics.ANIMA_VALUE_MASK,
|
||||
lambda *args, **kwargs: (args, kwargs),
|
||||
)
|
||||
|
||||
processed = COMFY_REGIONAL_CONDITIONING_PROCESSOR.process(
|
||||
preparation,
|
||||
model=SimpleNamespace(model=model),
|
||||
noise=torch.zeros((1, 16, 8, 8)),
|
||||
device=torch.device("cpu"),
|
||||
context_validator=ANIMA_REGIONAL_CONTEXT_VALIDATOR,
|
||||
negpip=negpip,
|
||||
)
|
||||
|
||||
entries = (
|
||||
processed.positive.base_context.entries[0],
|
||||
processed.positive.regional_contexts[0].entries[0],
|
||||
processed.negative.base_context.entries[0],
|
||||
processed.negative.regional_contexts[0].entries[0],
|
||||
)
|
||||
assert all(entry.cross_attention_value_multiplier is not None for entry in entries)
|
||||
assert [
|
||||
int(cast(torch.Tensor, entry.cross_attention_value_multiplier)[0, 0, 0].item())
|
||||
for entry in entries
|
||||
] == [1, -1, -1, 1]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("sequence_length", "feature_width", "message"),
|
||||
[
|
||||
@@ -347,6 +412,7 @@ def _conditioning(
|
||||
strength: float | None = None,
|
||||
start_percent: float | None = None,
|
||||
end_percent: float | None = None,
|
||||
negpip_value: int | None = None,
|
||||
) -> list[list[object]]:
|
||||
"""Build one small standard conditioning with model-consumed metadata."""
|
||||
|
||||
@@ -357,6 +423,8 @@ def _conditioning(
|
||||
metadata["start_percent"] = start_percent
|
||||
if end_percent is not None:
|
||||
metadata["end_percent"] = end_percent
|
||||
if negpip_value is not None:
|
||||
metadata["negpip_value"] = negpip_value
|
||||
return [
|
||||
[
|
||||
torch.full((1, 3, ANIMA_CONTEXT_FEATURE_WIDTH), value),
|
||||
|
||||
@@ -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:
|
||||
|
||||
+23
-37
@@ -7,7 +7,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -40,13 +39,13 @@ def test_downloadable_mode_includes_catalog_entries(tmp_path: Path) -> None:
|
||||
|
||||
service = ModelChoiceService(
|
||||
FakeSettingsRepository(show_downloadable_models=True),
|
||||
fake_folder_paths(tmp_path),
|
||||
)
|
||||
|
||||
assert "sam_vit_b (375MB)" in service.sam_choices()
|
||||
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,23 +59,22 @@ 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:
|
||||
"""SAM local-only mode lists installed known checkpoint files."""
|
||||
def test_hidden_catalog_mode_excludes_installed_sam_artifacts(tmp_path: Path) -> None:
|
||||
"""Catalog mode hides installed SAM entries when disabled."""
|
||||
|
||||
(tmp_path / "models" / "sams").mkdir(parents=True)
|
||||
(tmp_path / "models" / "sams" / "sam_vit_b_01ec64.pth").write_bytes(b"sam")
|
||||
|
||||
choices = local_only_service(tmp_path).sam_choices()
|
||||
|
||||
assert choices == ["sam_vit_b (375MB)"]
|
||||
assert local_only_service(tmp_path).sam_choices() == [NO_LOCAL_SAM_MODELS]
|
||||
|
||||
|
||||
def test_grounding_dino_local_only_requires_complete_artifacts(
|
||||
def test_hidden_catalog_mode_excludes_installed_grounding_dino_artifacts(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""GroundingDINO local-only mode excludes partial config/checkpoint pairs."""
|
||||
"""Catalog mode hides installed GroundingDINO entries when disabled."""
|
||||
|
||||
model_dir = tmp_path / "models" / "grounding-dino"
|
||||
model_dir.mkdir(parents=True)
|
||||
@@ -84,37 +82,35 @@ def test_grounding_dino_local_only_requires_complete_artifacts(
|
||||
(model_dir / "groundingdino_swint_ogc.pth").write_bytes(b"dino")
|
||||
(model_dir / "GroundingDINO_SwinB.cfg.py").write_text("", encoding="utf-8")
|
||||
|
||||
choices = local_only_service(tmp_path).grounding_dino_choices()
|
||||
|
||||
assert choices == ["GroundingDINO_SwinT_OGC (694MB)"]
|
||||
assert local_only_service(tmp_path).grounding_dino_choices() == [
|
||||
NO_LOCAL_GROUNDING_DINO_MODELS
|
||||
]
|
||||
|
||||
|
||||
def test_vitmatte_local_only_lists_valid_canonical_directory(
|
||||
def test_hidden_catalog_mode_excludes_installed_vitmatte_directory(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""ViTMatte local-only mode accepts canonical SimpleSyrup directories."""
|
||||
"""Catalog mode hides installed ViTMatte entries when disabled."""
|
||||
|
||||
create_vitmatte_snapshot(
|
||||
tmp_path / "models" / "vitmatte" / "vitmatte-small-composition-1k"
|
||||
)
|
||||
|
||||
choices = local_only_service(tmp_path).vitmatte_choices()
|
||||
|
||||
assert choices == ["vitmatte-small-composition-1k"]
|
||||
assert local_only_service(tmp_path).vitmatte_choices() == [NO_LOCAL_VITMATTE_MODELS]
|
||||
|
||||
|
||||
def test_vitmatte_local_only_lists_layerstyle_directory(tmp_path: Path) -> None:
|
||||
"""ViTMatte local-only mode accepts LayerStyle-compatible directories."""
|
||||
def test_hidden_catalog_mode_excludes_layerstyle_vitmatte_directory(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Catalog mode hides LayerStyle-compatible entries when disabled."""
|
||||
|
||||
create_vitmatte_snapshot(tmp_path / "models" / "vitmatte-base-composition-1k")
|
||||
|
||||
choices = local_only_service(tmp_path).vitmatte_choices()
|
||||
|
||||
assert choices == ["vitmatte-base-composition-1k"]
|
||||
assert local_only_service(tmp_path).vitmatte_choices() == [NO_LOCAL_VITMATTE_MODELS]
|
||||
|
||||
|
||||
def test_wd14_local_only_requires_complete_artifacts(tmp_path: Path) -> None:
|
||||
"""WD14 local-only mode excludes partial ONNX/CSV pairs."""
|
||||
def test_hidden_catalog_mode_excludes_installed_wd14_artifacts(tmp_path: Path) -> None:
|
||||
"""Catalog mode hides installed WD14 entries when disabled."""
|
||||
|
||||
model_dir = tmp_path / "models" / "wd14_tagger"
|
||||
model_dir.mkdir(parents=True)
|
||||
@@ -125,9 +121,9 @@ def test_wd14_local_only_requires_complete_artifacts(tmp_path: Path) -> None:
|
||||
)
|
||||
(model_dir / "wd-vit-tagger-v3.onnx").write_bytes(b"onnx")
|
||||
|
||||
choices = local_only_service(tmp_path).wd14_tagger_choices()
|
||||
|
||||
assert choices == ["wd-eva02-large-tagger-v3"]
|
||||
assert local_only_service(tmp_path).wd14_tagger_choices() == [
|
||||
NO_LOCAL_WD14_TAGGER_MODELS
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@@ -157,19 +153,9 @@ def local_only_service(tmp_path: Path) -> ModelChoiceService:
|
||||
|
||||
return ModelChoiceService(
|
||||
FakeSettingsRepository(show_downloadable_models=False),
|
||||
fake_folder_paths(tmp_path),
|
||||
)
|
||||
|
||||
|
||||
def fake_folder_paths(tmp_path: Path) -> ModuleType:
|
||||
"""Create a minimal fake Comfy folder_paths module."""
|
||||
|
||||
module = ModuleType("folder_paths")
|
||||
module.models_dir = str(tmp_path / "models") # type: ignore[attr-defined]
|
||||
module.folder_names_and_paths = {} # type: ignore[attr-defined]
|
||||
return module
|
||||
|
||||
|
||||
def create_vitmatte_snapshot(path: Path) -> None:
|
||||
"""Create the minimal file set required for a valid ViTMatte directory."""
|
||||
|
||||
|
||||
@@ -113,6 +113,41 @@ def test_comfy_progress_reporter_updates_the_active_node_progress(
|
||||
assert updates == [(-1, 6), (0, 6), (3, 6), (6, 6)]
|
||||
|
||||
|
||||
def test_comfy_progress_reporter_does_not_falsely_complete_unknown_downloads(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Unknown response lengths emit no misleading progress until completion."""
|
||||
|
||||
updates: list[tuple[int, int | None]] = []
|
||||
|
||||
class FakeProgressBar:
|
||||
"""Record ComfyUI absolute progress updates."""
|
||||
|
||||
def __init__(self, total: int) -> None:
|
||||
"""Record the total selected for the progress bar."""
|
||||
|
||||
updates.append((-1, total))
|
||||
|
||||
def update_absolute(self, value: int, total: int | None = None) -> None:
|
||||
"""Record one absolute progress update."""
|
||||
|
||||
updates.append((value, total))
|
||||
|
||||
comfy_module = ModuleType("comfy")
|
||||
comfy_utils = ModuleType("comfy.utils")
|
||||
comfy_utils.ProgressBar = FakeProgressBar # type: ignore[attr-defined]
|
||||
comfy_module.utils = comfy_utils # type: ignore[attr-defined]
|
||||
monkeypatch.setitem(sys.modules, "comfy", comfy_module)
|
||||
monkeypatch.setitem(sys.modules, "comfy.utils", comfy_utils)
|
||||
|
||||
reporter = ComfyProgressReporter()
|
||||
reporter.start("Downloading unknown-size model", None)
|
||||
reporter.advance(1024 * 1024, None)
|
||||
reporter.finish()
|
||||
|
||||
assert updates == [(-1, 1), (1, 1)]
|
||||
|
||||
|
||||
def test_downloader_streams_file_and_reports_progress(
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
|
||||
@@ -235,6 +235,35 @@ def test_collision_safe_mutations_integrate_through_one_real_comfy_clone() -> No
|
||||
assert ModelPatcher.set_model_attn2_patch is original_attn2_setter
|
||||
|
||||
|
||||
def test_attn2_mutation_prepends_before_exact_preserved_input_patch() -> None:
|
||||
"""Compose regional packing before an identity-admitted input transformer."""
|
||||
|
||||
model = _patcher(torch.nn.Linear(1, 1))
|
||||
|
||||
def preserved(*args: object) -> tuple[object, ...]:
|
||||
"""Return preserved callback arguments."""
|
||||
|
||||
return args
|
||||
|
||||
def regional(*args: object) -> tuple[object, ...]:
|
||||
"""Return regional callback arguments."""
|
||||
|
||||
return args
|
||||
|
||||
def output(*args: object) -> tuple[object, ...]:
|
||||
"""Return output callback arguments."""
|
||||
|
||||
return args
|
||||
|
||||
model.set_model_attn2_patch(preserved)
|
||||
|
||||
ModelAttn2PatchesMutation(regional, output, (preserved,)).apply(model)
|
||||
|
||||
patches = model.model_options["transformer_options"]["patches"]
|
||||
assert patches["attn2_patch"] == [regional, preserved]
|
||||
assert patches["attn2_output_patch"] == [output]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("wrapper_type", "key", "wrapper", "message"),
|
||||
[
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Prove Triton remains an optional CUDA acceleration dependency."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.runtime.regional_lora.triton_runtime import (
|
||||
TritonRuntimeResolver,
|
||||
)
|
||||
|
||||
|
||||
def test_node_registration_and_cpu_accumulation_do_not_import_triton() -> None:
|
||||
"""Load public nodes and execute CPU accumulation with Triton blocked."""
|
||||
|
||||
script = """
|
||||
import importlib.abc
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path.cwd().parents[1]))
|
||||
sys.argv = [sys.argv[0], "--cpu"]
|
||||
import comfy.options
|
||||
comfy.options.enable_args_parsing()
|
||||
|
||||
class BlockTriton(importlib.abc.MetaPathFinder):
|
||||
def find_spec(self, fullname, path, target=None):
|
||||
if fullname == "triton" or fullname.startswith("triton."):
|
||||
raise ModuleNotFoundError("blocked optional Triton", name=fullname)
|
||||
return None
|
||||
|
||||
sys.meta_path.insert(0, BlockTriton())
|
||||
import torch
|
||||
from simple_syrup.nodes_v3 import get_nodes
|
||||
from simple_syrup.runtime.regional_lora import ordered_accumulation
|
||||
|
||||
assert get_nodes()
|
||||
base = torch.tensor([1.0, 2.0])
|
||||
result = ordered_accumulation.OrderedTensorAccumulator().accumulate(
|
||||
base,
|
||||
(torch.tensor([3.0, 4.0]), torch.tensor([5.0, 6.0])),
|
||||
)
|
||||
assert result.tolist() == [9.0, 12.0]
|
||||
assert not any(name == "triton" or name.startswith("triton.") for name in sys.modules)
|
||||
"""
|
||||
completed = subprocess.run(
|
||||
[sys.executable, "-c", script],
|
||||
cwd=Path(__file__).resolve().parents[1],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=60,
|
||||
check=False,
|
||||
)
|
||||
|
||||
assert completed.returncode == 0, completed.stderr
|
||||
|
||||
|
||||
def test_resolver_caches_one_missing_result_and_warns_once(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
"""Treat an absent Triton package as one observable optional miss."""
|
||||
|
||||
calls: list[str] = []
|
||||
|
||||
def missing_import(name: str) -> object:
|
||||
calls.append(name)
|
||||
raise ModuleNotFoundError("missing", name="triton")
|
||||
|
||||
resolver = TritonRuntimeResolver(import_module=missing_import)
|
||||
|
||||
with caplog.at_level("WARNING"):
|
||||
assert resolver.resolve("fake.backend") is None
|
||||
assert resolver.resolve("fake.backend") is None
|
||||
|
||||
assert calls == ["fake.backend"]
|
||||
assert [record.message for record in caplog.records] == [
|
||||
"Triton acceleration is unavailable; using the Torch execution path"
|
||||
]
|
||||
|
||||
|
||||
def test_resolver_exposes_broken_backend_import_with_original_cause() -> None:
|
||||
"""Fail visibly when a present backend cannot initialize correctly."""
|
||||
|
||||
failure = RuntimeError("JIT initialization failed")
|
||||
|
||||
def broken_import(_name: str) -> object:
|
||||
raise failure
|
||||
|
||||
resolver = TritonRuntimeResolver(import_module=broken_import)
|
||||
|
||||
with pytest.raises(RuntimeError, match="failed to initialize") as raised:
|
||||
resolver.resolve("fake.backend")
|
||||
|
||||
assert raised.value.__cause__ is failure
|
||||
|
||||
|
||||
def test_resolver_is_thread_safe_and_returns_one_cached_backend() -> None:
|
||||
"""Publish exactly one imported backend across concurrent callers."""
|
||||
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
backend = object()
|
||||
calls: list[str] = []
|
||||
|
||||
def import_backend(name: str) -> object:
|
||||
calls.append(name)
|
||||
return backend
|
||||
|
||||
resolver = TritonRuntimeResolver(import_module=import_backend)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=8) as executor:
|
||||
results = tuple(executor.map(resolver.resolve, ("fake.backend",) * 32))
|
||||
|
||||
assert all(result is backend for result in results)
|
||||
assert calls == ["fake.backend"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"missing_name",
|
||||
["fake.backend", "unrelated_dependency"],
|
||||
)
|
||||
def test_resolver_does_not_hide_non_triton_module_failures(missing_name: str) -> None:
|
||||
"""Reserve optional fallback exclusively for the Triton package family."""
|
||||
|
||||
def missing_import(_name: str) -> object:
|
||||
raise ModuleNotFoundError("missing", name=missing_name)
|
||||
|
||||
resolver = TritonRuntimeResolver(
|
||||
import_module=missing_import,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="failed to initialize"):
|
||||
resolver.resolve("fake.backend")
|
||||
@@ -0,0 +1,94 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Verify typed PPM NegPiP conditioning and call-local option adaptation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import comfy.conds
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.runtime.ppm_negpip_interop import (
|
||||
PpmNegpipInterop,
|
||||
PpmNegpipSemantics,
|
||||
)
|
||||
|
||||
|
||||
def test_anima_adapter_extracts_exact_typed_value_multiplier() -> None:
|
||||
"""Convert PPM's model condition to context-aligned execution state."""
|
||||
|
||||
interop = _anima_interop()
|
||||
context = torch.zeros((1, 3, 4), dtype=torch.float16)
|
||||
source = torch.tensor([[[1], [-1], [1]]], dtype=torch.int32)
|
||||
|
||||
multiplier = interop.extract_value_multiplier(
|
||||
{"c_ppm_negpip_mask": comfy.conds.CONDRegular(source)},
|
||||
context,
|
||||
)
|
||||
|
||||
assert multiplier is not None
|
||||
assert multiplier.dtype is context.dtype
|
||||
assert multiplier.device == context.device
|
||||
assert multiplier.tolist() == [[[1.0], [-1.0], [1.0]]]
|
||||
|
||||
|
||||
def test_anima_adapter_uses_neutral_multiplier_when_condition_is_absent() -> None:
|
||||
"""Represent an all-positive prompt without leaving branch state partial."""
|
||||
|
||||
context = torch.zeros((2, 3, 4))
|
||||
|
||||
multiplier = _anima_interop().extract_value_multiplier({}, context)
|
||||
|
||||
assert multiplier is not None
|
||||
assert torch.equal(multiplier, torch.ones((2, 3, 1)))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"source",
|
||||
[
|
||||
torch.ones((1, 2, 1)),
|
||||
torch.zeros((1, 3, 1)),
|
||||
torch.ones((1, 3, 2)),
|
||||
],
|
||||
)
|
||||
def test_anima_adapter_rejects_misaligned_or_nonbinary_masks(
|
||||
source: torch.Tensor,
|
||||
) -> None:
|
||||
"""Fail closed before malformed PPM state reaches regional packing."""
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_anima_interop().extract_value_multiplier(
|
||||
{"c_ppm_negpip_mask": SimpleNamespace(cond=source)},
|
||||
torch.zeros((1, 3, 4)),
|
||||
)
|
||||
|
||||
|
||||
def test_anima_adapter_publishes_mask_on_an_isolated_option_copy() -> None:
|
||||
"""Keep the ordinary call options unchanged outside original cross-attention."""
|
||||
|
||||
old_mask = torch.ones((1, 2, 1))
|
||||
packed = torch.tensor([[[1.0], [-1.0]], [[-1.0], [1.0]]])
|
||||
source: dict[str, object] = {
|
||||
"ppm_negpip_mask": old_mask,
|
||||
"preserved": object(),
|
||||
}
|
||||
|
||||
prepared = _anima_interop().prepare_anima_transformer_options(source, packed)
|
||||
|
||||
assert prepared is not source
|
||||
assert prepared["preserved"] is source["preserved"]
|
||||
assert prepared["ppm_negpip_mask"] is packed
|
||||
assert source["ppm_negpip_mask"] is old_mask
|
||||
|
||||
|
||||
def _anima_interop() -> PpmNegpipInterop:
|
||||
"""Return one focused admitted Anima semantic adapter."""
|
||||
|
||||
return PpmNegpipInterop(
|
||||
PpmNegpipSemantics.ANIMA_VALUE_MASK,
|
||||
lambda *args, **kwargs: (args, kwargs),
|
||||
)
|
||||
@@ -26,6 +26,7 @@ from simple_syrup.runtime.comfy_conditioning_model_loader import (
|
||||
from simple_syrup.runtime.comfy_conditioning_processing import (
|
||||
ComfyRegionalConditioningProcessor,
|
||||
)
|
||||
from simple_syrup.runtime.ppm_negpip_interop import PpmNegpipInterop
|
||||
from simple_syrup.runtime.regional_lora_conditioning_adapter import (
|
||||
RegionalLoraConditioningAdapter,
|
||||
)
|
||||
@@ -135,6 +136,7 @@ def test_profiled_collaborators_preserve_arguments_results_and_stage_order(
|
||||
noise: torch.Tensor,
|
||||
device: torch.device,
|
||||
context_validator: RegionalContextValidator,
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> ProcessedRegionalAttentionPlan:
|
||||
calls.append(
|
||||
(
|
||||
@@ -145,6 +147,7 @@ def test_profiled_collaborators_preserve_arguments_results_and_stage_order(
|
||||
"noise": noise,
|
||||
"device": device,
|
||||
"context_validator": context_validator,
|
||||
"negpip": negpip,
|
||||
},
|
||||
)
|
||||
)
|
||||
@@ -190,6 +193,7 @@ def test_profiled_collaborators_preserve_arguments_results_and_stage_order(
|
||||
"noise": noise,
|
||||
"device": device,
|
||||
"context_validator": validator,
|
||||
"negpip": None,
|
||||
}
|
||||
assert _stages(caplog) == [
|
||||
"source_model_load",
|
||||
|
||||
@@ -108,6 +108,56 @@ def test_batching_builds_canonical_regions_with_base_fallback() -> None:
|
||||
]
|
||||
|
||||
|
||||
def test_batching_aligns_value_multipliers_with_cfg_regions_and_fallbacks() -> None:
|
||||
"""Keep each scheduled branch's value semantics in identical chunk order."""
|
||||
|
||||
source = _plan()
|
||||
plan = ProcessedRegionalAttentionPlan(
|
||||
ProcessedRegionalAttentionBranch(
|
||||
_context_with_multiplier(source.positive.base_context, 1.0),
|
||||
(
|
||||
_context_with_multiplier(
|
||||
source.positive.regional_contexts[0],
|
||||
-1.0,
|
||||
),
|
||||
_context_with_multiplier(
|
||||
source.positive.regional_contexts[1],
|
||||
1.0,
|
||||
),
|
||||
),
|
||||
),
|
||||
ProcessedRegionalAttentionBranch(
|
||||
_context_with_multiplier(source.negative.base_context, -1.0),
|
||||
(
|
||||
_context_with_multiplier(
|
||||
source.negative.regional_contexts[0],
|
||||
1.0,
|
||||
),
|
||||
),
|
||||
),
|
||||
source.mask_bank,
|
||||
source.lora_plan,
|
||||
)
|
||||
|
||||
aligned = REGIONAL_ATTENTION_BATCHING_SERVICE.align(
|
||||
plan,
|
||||
base_context=_runtime_base(plan, [1, 0], 2),
|
||||
cond_or_uncond=[1, 0],
|
||||
conditioning_uuids=_uuids_for_selectors(plan, [1, 0]),
|
||||
sigma=0.5,
|
||||
latent_batch_size=2,
|
||||
)
|
||||
|
||||
assert aligned.base_value_multiplier is not None
|
||||
assert aligned.base_value_multiplier[:, 0, 0].tolist() == [-1.0, -1.0, 1.0, 1.0]
|
||||
region_zero = aligned.regions[0].entries[0].cross_attention_value_multiplier
|
||||
region_one = aligned.regions[1].entries[0].cross_attention_value_multiplier
|
||||
assert region_zero is not None
|
||||
assert region_one is not None
|
||||
assert region_zero[:, 0, 0].tolist() == [1.0, 1.0, -1.0, -1.0]
|
||||
assert region_one[:, 0, 0].tolist() == [-1.0, -1.0, 1.0, 1.0]
|
||||
|
||||
|
||||
def test_batching_preserves_all_regional_entries_and_per_sample_strengths() -> None:
|
||||
"""Align simultaneous regional entries without collapsing their order."""
|
||||
|
||||
@@ -409,6 +459,34 @@ def _multi_entry_context(
|
||||
)
|
||||
|
||||
|
||||
def _context_with_multiplier(
|
||||
context: ProcessedRegionalAttentionContext,
|
||||
value: float,
|
||||
) -> ProcessedRegionalAttentionContext:
|
||||
"""Copy one context with a uniform sequence-aligned value multiplier."""
|
||||
|
||||
return ProcessedRegionalAttentionContext(
|
||||
context.conditioning_index,
|
||||
context.region_index,
|
||||
tuple(
|
||||
ProcessedRegionalAttentionEntry(
|
||||
entry.entry_index,
|
||||
entry.uuid,
|
||||
entry.schedule,
|
||||
entry.cross_attention,
|
||||
entry.strength,
|
||||
torch.full(
|
||||
(*entry.cross_attention.shape[:2], 1),
|
||||
value,
|
||||
dtype=entry.cross_attention.dtype,
|
||||
device=entry.cross_attention.device,
|
||||
),
|
||||
)
|
||||
for entry in context.entries
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _entry(
|
||||
entry_index: int,
|
||||
tensor: torch.Tensor,
|
||||
|
||||
@@ -30,6 +30,7 @@ from simple_syrup.domain.regional_model_capabilities import (
|
||||
RegionalReferenceLatentPolicy,
|
||||
RegionalSpatialPatchSupport,
|
||||
)
|
||||
from simple_syrup.runtime.ppm_negpip_interop import PpmNegpipSemantics
|
||||
from simple_syrup.runtime.regional_model_patch_interop import (
|
||||
REGIONAL_MODEL_PATCH_INTEROP_VALIDATOR,
|
||||
RegionalPreservedModelModifier,
|
||||
@@ -46,6 +47,11 @@ class _FixtureModel(torch.nn.Module):
|
||||
self.projection = torch.nn.Linear(1, 1)
|
||||
self.latent_format = SimpleNamespace(latent_channels=4)
|
||||
|
||||
def extra_conds(self, **_kwargs: object) -> dict[str, object]:
|
||||
"""Expose the object path patched by Anima NegPiP."""
|
||||
|
||||
return {}
|
||||
|
||||
|
||||
def test_validator_preserves_easycache_and_unrelated_model_state() -> None:
|
||||
"""Accept EasyCache while retaining every collaborator-owned surface."""
|
||||
@@ -186,28 +192,103 @@ def test_validator_rejects_both_core_caches_without_mutating_them() -> None:
|
||||
assert combined.wrappers == before_wrappers
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"family",
|
||||
[RegionalModelFamily.ANIMA, RegionalModelFamily.STANDARD_UNET],
|
||||
)
|
||||
def test_validator_rejects_named_negpip_before_generic_attn2_collision(
|
||||
family: RegionalModelFamily,
|
||||
) -> None:
|
||||
"""Report the installed modifier and regional mask misalignment by name."""
|
||||
def test_validator_admits_exact_standard_unet_negpip_without_mutation() -> None:
|
||||
"""Retain PPM's exact split-K/V callback as typed interop evidence."""
|
||||
|
||||
model = _patcher()
|
||||
callback = _identity_callback(
|
||||
"custom_nodes.ComfyUI-ppm.src.negpip.unet_negpip",
|
||||
"sdxl_attn2_negpip",
|
||||
)
|
||||
model.model_options["ppm_negpip"] = True
|
||||
model.set_model_attn2_patch(callback)
|
||||
|
||||
report = REGIONAL_MODEL_PATCH_INTEROP_VALIDATOR.validate(
|
||||
model,
|
||||
_capabilities(RegionalModelFamily.STANDARD_UNET),
|
||||
)
|
||||
|
||||
assert report.negpip is not None
|
||||
assert report.negpip.semantics is PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE
|
||||
assert report.negpip.attention_patch is callback
|
||||
assert model.model_options["transformer_options"]["patches"]["attn2_patch"] == [
|
||||
callback
|
||||
]
|
||||
|
||||
|
||||
def test_validator_admits_exact_anima_negpip_without_mutation() -> None:
|
||||
"""Retain PPM's complete Anima callback, wrapper, and object-patch family."""
|
||||
|
||||
model = _patcher()
|
||||
callback = _identity_callback(
|
||||
"custom_nodes.ComfyUI-ppm.src.negpip.anima_negpip",
|
||||
"cosmos_attn2_negpip",
|
||||
)
|
||||
wrapper = _identity_callback(
|
||||
"custom_nodes.ComfyUI-ppm.src.negpip.anima_negpip",
|
||||
"cosmos_diffusion_negpip_wrapper",
|
||||
)
|
||||
extra_conds = _identity_callback(
|
||||
"custom_nodes.ComfyUI-ppm.src.negpip.anima_negpip",
|
||||
("anima_extra_conds_negpip_wrapper.<locals>._anima_extra_conds_negpip_wrapper"),
|
||||
)
|
||||
model.model_options["ppm_negpip"] = True
|
||||
model.set_model_attn2_patch(callback)
|
||||
model.add_wrapper_with_key(
|
||||
WrappersMP.DIFFUSION_MODEL,
|
||||
"ppm_negpip_anima",
|
||||
lambda executor, *args, **kwargs: executor(*args, **kwargs),
|
||||
wrapper,
|
||||
)
|
||||
model.set_model_attn2_patch(lambda q, k, v, **kwargs: {"q": q, "k": k, "v": v})
|
||||
model.add_object_patch("extra_conds", extra_conds)
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="NegPiP.*ordinary conditioning batch.*regional branch batch",
|
||||
):
|
||||
report = REGIONAL_MODEL_PATCH_INTEROP_VALIDATOR.validate(
|
||||
model,
|
||||
_capabilities(RegionalModelFamily.ANIMA),
|
||||
)
|
||||
|
||||
assert report.negpip is not None
|
||||
assert report.negpip.semantics is PpmNegpipSemantics.ANIMA_VALUE_MASK
|
||||
assert report.negpip.attention_patch is callback
|
||||
assert model.wrappers[WrappersMP.DIFFUSION_MODEL]["ppm_negpip_anima"] == [wrapper]
|
||||
assert model.object_patches["extra_conds"] is extra_conds
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("family", "configure", "message"),
|
||||
[
|
||||
(
|
||||
RegionalModelFamily.STANDARD_UNET,
|
||||
lambda model: model.model_options.__setitem__("ppm_negpip", True),
|
||||
"requires exactly its PPM split-K/V",
|
||||
),
|
||||
(
|
||||
RegionalModelFamily.STANDARD_UNET,
|
||||
lambda model: model.set_model_attn2_patch(
|
||||
_identity_callback(
|
||||
"custom_nodes.ComfyUI-ppm.src.negpip.unet_negpip",
|
||||
"sdxl_attn2_negpip",
|
||||
)
|
||||
),
|
||||
"incomplete NegPiP patch family",
|
||||
),
|
||||
(
|
||||
RegionalModelFamily.ANIMA,
|
||||
lambda model: model.model_options.__setitem__("ppm_negpip", True),
|
||||
"requires exactly its PPM attention patch",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_validator_rejects_partial_or_foreign_negpip_families(
|
||||
family: RegionalModelFamily,
|
||||
configure: Callable[[ModelPatcher], object],
|
||||
message: str,
|
||||
) -> None:
|
||||
"""Fail closed before partial or identity-foreign NegPiP state is composed."""
|
||||
|
||||
model = _patcher()
|
||||
configure(model)
|
||||
|
||||
with pytest.raises(ValueError, match=message):
|
||||
REGIONAL_MODEL_PATCH_INTEROP_VALIDATOR.validate(
|
||||
model,
|
||||
_capabilities(family),
|
||||
@@ -313,6 +394,19 @@ def _patcher() -> ModelPatcher:
|
||||
)
|
||||
|
||||
|
||||
def _identity_callback(module: str, qualname: str) -> Callable[..., object]:
|
||||
"""Build one executable callback carrying a stable PPM definition identity."""
|
||||
|
||||
def callback(*args: object, **_kwargs: object) -> object:
|
||||
"""Return callback inputs for model-state-only admission tests."""
|
||||
|
||||
return args
|
||||
|
||||
callback.__module__ = module
|
||||
callback.__qualname__ = qualname
|
||||
return callback
|
||||
|
||||
|
||||
def _capabilities(family: RegionalModelFamily) -> RegionalModelCapabilities:
|
||||
"""Build the exact family contract consumed by modifier admission."""
|
||||
|
||||
|
||||
@@ -100,7 +100,7 @@ def test_regional_patch_stack_preserves_state_lineage_and_runtime_nesting() -> N
|
||||
|
||||
assert stack.user_model is source
|
||||
assert stack.attention_model.parent is source
|
||||
assert stack.sampling_model.parent is stack.attention_model
|
||||
assert stack.sampling_model.parent is source
|
||||
assert stack.sampling_model.patches["weight"] == [global_lora_patch]
|
||||
assert (
|
||||
source.get_wrappers("diffusion_model", "simple_syrup.attention_coupling") == []
|
||||
|
||||
@@ -41,6 +41,9 @@ def test_every_matrix_case_requires_exact_modifier_and_terminal_evidence(
|
||||
observed: RegionalPatchInteropHistory
|
||||
if case.expect_success:
|
||||
model_call_count = STEPS
|
||||
diagnostic_record_count = model_call_count * (
|
||||
2 if case.model_family is PatchInteropModelFamily.SDXL else 1
|
||||
)
|
||||
observed = RegionalPatchInteropSuccess(
|
||||
snapshot,
|
||||
{
|
||||
@@ -49,7 +52,7 @@ def test_every_matrix_case_requires_exact_modifier_and_terminal_evidence(
|
||||
"runtime_ms": 10.0,
|
||||
"peak_vram_bytes": 1,
|
||||
},
|
||||
_diagnostics(case, workflow, record_count=model_call_count),
|
||||
_diagnostics(case, workflow, record_count=diagnostic_record_count),
|
||||
ImageReference("result.png", "", "output"),
|
||||
None,
|
||||
)
|
||||
@@ -69,7 +72,7 @@ def test_every_matrix_case_requires_exact_modifier_and_terminal_evidence(
|
||||
assert validated.model_call_count == (STEPS if case.expect_success else 0)
|
||||
|
||||
|
||||
def test_success_requires_one_diagnostic_record_per_actual_model_call() -> None:
|
||||
def test_success_requires_family_specific_diagnostic_records_per_model_call() -> None:
|
||||
"""Reject cache evidence whose diagnostics omit an executed model call."""
|
||||
|
||||
case = next(
|
||||
@@ -89,7 +92,7 @@ def test_success_requires_one_diagnostic_record_per_actual_model_call() -> None:
|
||||
None,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="one diagnostic record per model call"):
|
||||
with pytest.raises(ValueError, match="model-family execution shape"):
|
||||
validate_case(case, workflow, observed)
|
||||
|
||||
|
||||
@@ -216,7 +219,8 @@ def _diagnostics(
|
||||
PatchInteropSpatialMode.CONTEXTUAL: ("tile", "contextual_global"),
|
||||
}[case.spatial_mode]
|
||||
snapshots = [
|
||||
_diagnostic_snapshot(modes[index % len(modes)]) for index in range(record_count)
|
||||
_diagnostic_snapshot(case, modes[index % len(modes)])
|
||||
for index in range(record_count)
|
||||
]
|
||||
return {
|
||||
"run_id": workflow.diagnostics_run_id,
|
||||
@@ -225,8 +229,25 @@ def _diagnostics(
|
||||
}
|
||||
|
||||
|
||||
def _diagnostic_snapshot(mode: str) -> JsonObject:
|
||||
"""Return one exact static-PRIMARY_ADAPTER diagnostic record."""
|
||||
def _diagnostic_snapshot(
|
||||
case: RegionalPatchInteropCase,
|
||||
mode: str,
|
||||
) -> JsonObject:
|
||||
"""Return one exact family-specific regional diagnostic record."""
|
||||
|
||||
if case.model_family is PatchInteropModelFamily.SDXL:
|
||||
return {
|
||||
"strategy": "attention_coupling",
|
||||
"backend": "comfy.ldm.modules.diffusionmodules.openaimodel.UNetModel",
|
||||
"spatial_mode": mode,
|
||||
"region_count": 2,
|
||||
"active_region_indices": [0, 1],
|
||||
"estimated_work": {
|
||||
"cross_attention_branch_multiplier": 3.0,
|
||||
"cross_attention_formula": "base_plus_region_count",
|
||||
"denoiser_call_multiplier": 1.0,
|
||||
},
|
||||
}
|
||||
|
||||
return {
|
||||
"strategy": "attention_coupling",
|
||||
|
||||
@@ -23,13 +23,13 @@ from tools.regional_patch_interop_integration.workflow import (
|
||||
)
|
||||
|
||||
|
||||
def test_matrix_contains_five_acceptances_and_six_exact_rejections() -> None:
|
||||
def test_matrix_contains_seven_acceptances_and_four_exact_rejections() -> None:
|
||||
"""Keep every required modifier, spatial, scheduled, and family case."""
|
||||
|
||||
definitions = cases()
|
||||
|
||||
assert len(definitions) == 11
|
||||
assert sum(case.expect_success for case in definitions) == 5
|
||||
assert sum(case.expect_success for case in definitions) == 7
|
||||
assert {case.modifier for case in definitions} == set(PatchInteropModifier)
|
||||
assert {case.spatial_mode for case in definitions} == set(PatchInteropSpatialMode)
|
||||
assert {case.model_family for case in definitions} == set(PatchInteropModelFamily)
|
||||
@@ -154,11 +154,9 @@ def test_scheduled_cache_graph_authors_the_exact_regional_adapter_interval() ->
|
||||
|
||||
|
||||
def test_sdxl_negpip_graph_uses_public_modifier_snapshot_and_sampler() -> None:
|
||||
"""Submit SDXL NegPiP through the same evidence and rejection boundary."""
|
||||
"""Submit SDXL NegPiP through the same evidence and execution boundary."""
|
||||
|
||||
definition = next(
|
||||
case for case in cases() if case.case_id == "sdxl-negpip-rejected"
|
||||
)
|
||||
definition = next(case for case in cases() if case.case_id == "sdxl-negpip")
|
||||
workflow = RegionalPatchInteropWorkflowBuilder().build(
|
||||
definition,
|
||||
run_id="run",
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
@@ -14,6 +14,10 @@ import torch
|
||||
from simple_syrup.runtime.attention_coupling.unet_attn2_execution import (
|
||||
UnetAttn2Execution,
|
||||
)
|
||||
from simple_syrup.runtime.ppm_negpip_interop import (
|
||||
PpmNegpipInterop,
|
||||
PpmNegpipSemantics,
|
||||
)
|
||||
from simple_syrup.runtime.regional_lora.standard_unet_variant_base_attention import (
|
||||
StandardUnetVariantBaseAttention,
|
||||
)
|
||||
@@ -67,3 +71,32 @@ def test_prepare_rejects_any_preexisting_attn2_callback_surface(key: str) -> Non
|
||||
|
||||
with pytest.raises(ValueError, match="already contains"):
|
||||
StandardUnetVariantBaseAttention(_Resolver()).prepare({"patches": {key: []}})
|
||||
|
||||
|
||||
def test_prepare_places_coupling_before_exact_preserved_negpip_callback() -> None:
|
||||
"""Pack regional alternating tokens before PPM selects K and V views."""
|
||||
|
||||
def negpip(*args: object, **_kwargs: object) -> tuple[object, ...]:
|
||||
"""Represent the identity-validated PPM split callback."""
|
||||
|
||||
return args
|
||||
|
||||
interop = PpmNegpipInterop(
|
||||
PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE,
|
||||
negpip,
|
||||
)
|
||||
source: dict[str, object] = {"patches": {"attn2_patch": [negpip]}}
|
||||
|
||||
prepared = StandardUnetVariantBaseAttention(
|
||||
_Resolver(),
|
||||
negpip=interop,
|
||||
).prepare(source)
|
||||
|
||||
patches = prepared["patches"]
|
||||
assert isinstance(patches, dict)
|
||||
installed = patches["attn2_patch"]
|
||||
assert isinstance(installed, list)
|
||||
assert len(installed) == 2
|
||||
assert installed[1] is negpip
|
||||
assert callable(installed[0])
|
||||
assert source == {"patches": {"attn2_patch": [negpip]}}
|
||||
|
||||
@@ -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(show_downloadable_models=False),
|
||||
)
|
||||
|
||||
assert service.model_choices() == ["bbox/face.pt", "root.pt", "segm/person.pt"]
|
||||
|
||||
@@ -43,7 +56,75 @@ 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(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(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_model_choices_list_installed_models_before_downloadable_entries(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Installed choices precede curated models that still require a download."""
|
||||
|
||||
models_dir = tmp_path / "models"
|
||||
bbox_dir = models_dir / "ultralytics" / "bbox"
|
||||
bbox_dir.mkdir(parents=True)
|
||||
(bbox_dir / "face_yolov8n_v2.pt").write_bytes(b"checkpoint")
|
||||
(bbox_dir / "local-detector.pt").write_bytes(b"checkpoint")
|
||||
folder_paths = _folder_paths(models_dir)
|
||||
service = UltralyticsLoaderService(
|
||||
folder_paths_module=folder_paths,
|
||||
choice_service=_choice_service(show_downloadable_models=True),
|
||||
)
|
||||
|
||||
choices = service.model_choices()
|
||||
|
||||
assert choices[:2] == [
|
||||
"bbox/local-detector.pt",
|
||||
"Bingsu Face YOLOv8n v2 (6.23MB)",
|
||||
]
|
||||
assert choices[2] == "Anzhc Face -seg (6.52MB)"
|
||||
assert "bbox/face_yolov8n_v2.pt" not in choices
|
||||
|
||||
|
||||
def test_hidden_catalog_choices_exclude_installed_curated_model(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Hidden catalog mode excludes installed curated model files."""
|
||||
|
||||
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(show_downloadable_models=False),
|
||||
)
|
||||
|
||||
assert service.model_choices() == [NO_LOCAL_ULTRALYTICS_MODELS]
|
||||
|
||||
@@ -102,6 +183,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(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(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(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(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 +420,57 @@ 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(
|
||||
*,
|
||||
show_downloadable_models: bool,
|
||||
) -> ModelChoiceService:
|
||||
"""Build an Ultralytics choice service with deterministic settings."""
|
||||
|
||||
return ModelChoiceService(_FakeSettingsRepository(show_downloadable_models))
|
||||
|
||||
|
||||
def _folder_paths(models_dir: Path) -> ModuleType:
|
||||
"""Build a minimal fake ComfyUI folder_paths module."""
|
||||
|
||||
|
||||
@@ -29,6 +29,13 @@ from simple_syrup.runtime.attention_coupling.unet_attn2_execution_resolver impor
|
||||
from simple_syrup.runtime.attention_coupling.unet_attn2_patch import (
|
||||
UnetAttn2PatchPair,
|
||||
)
|
||||
from simple_syrup.runtime.ppm_negpip_interop import (
|
||||
PpmNegpipInterop,
|
||||
PpmNegpipSemantics,
|
||||
)
|
||||
from simple_syrup.runtime.regional_lora.standard_unet_variant_base_attention import (
|
||||
StandardUnetVariantBaseAttention,
|
||||
)
|
||||
|
||||
|
||||
class _ZeroAttention(nn.Module):
|
||||
@@ -81,6 +88,32 @@ class _RegionalAttention(nn.Module):
|
||||
return query + context.mean(dim=1, keepdim=True)
|
||||
|
||||
|
||||
class _RecordingKeyValueAttention(nn.Module):
|
||||
"""Record exact post-patch key/value inputs and return packed zeros."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize an empty invocation record."""
|
||||
|
||||
super().__init__()
|
||||
self.calls: list[tuple[torch.Tensor, torch.Tensor]] = []
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
*,
|
||||
context: torch.Tensor | None,
|
||||
value: torch.Tensor | None,
|
||||
transformer_options: dict[str, Any],
|
||||
) -> torch.Tensor:
|
||||
"""Retain post-patch K/V sources without performing attention."""
|
||||
|
||||
del transformer_options
|
||||
if context is None or value is None:
|
||||
raise AssertionError("NegPiP test requires explicit K and V tensors.")
|
||||
self.calls.append((context, value))
|
||||
return torch.zeros_like(query)
|
||||
|
||||
|
||||
class _CountingZeroFeedForward(nn.Module):
|
||||
"""Return zero while counting the retained feed-forward trajectory."""
|
||||
|
||||
@@ -165,6 +198,149 @@ def test_unet_patch_clears_callback_state_after_output_failure() -> None:
|
||||
patches.input_patch(query, context, context, options)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("persistent_base_graph", [False, True])
|
||||
def test_negpip_splits_exact_packed_regional_key_and_value_views(
|
||||
persistent_base_graph: bool,
|
||||
) -> None:
|
||||
"""Select even K and odd V tokens after packing in both UNet base routes."""
|
||||
|
||||
contexts = _negpip_contexts()
|
||||
execution = UnetAttn2Execution(
|
||||
contexts,
|
||||
torch.tensor([[[1.0, 0.0]], [[0.0, 1.0]]]),
|
||||
(1.0, 1.0),
|
||||
1,
|
||||
2,
|
||||
)
|
||||
pair = UnetAttn2PatchPair(StaticUnetAttn2ExecutionResolver(execution))
|
||||
|
||||
def split_negpip(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
_extra_options: dict[str, Any],
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Apply PPM's public alternating-token UNet semantics exactly."""
|
||||
|
||||
return query, key[:, 0::2], value[:, 1::2]
|
||||
|
||||
if persistent_base_graph:
|
||||
interop = PpmNegpipInterop(
|
||||
PpmNegpipSemantics.STANDARD_UNET_SPLIT_KEY_VALUE,
|
||||
split_negpip,
|
||||
)
|
||||
options = StandardUnetVariantBaseAttention(
|
||||
StaticUnetAttn2ExecutionResolver(execution),
|
||||
negpip=interop,
|
||||
).prepare({"patches": {"attn2_patch": [split_negpip]}})
|
||||
else:
|
||||
options = {
|
||||
"patches": {
|
||||
"attn2_patch": [pair.input_patch, split_negpip],
|
||||
"attn2_output_patch": [pair.output_patch],
|
||||
}
|
||||
}
|
||||
recording = _RecordingKeyValueAttention()
|
||||
block = _block(recording)
|
||||
|
||||
block(
|
||||
torch.zeros((1, 2, 1)),
|
||||
context=contexts.base_context,
|
||||
transformer_options=options,
|
||||
)
|
||||
|
||||
assert len(recording.calls) == 1
|
||||
key, value = recording.calls[0]
|
||||
assert key[:, :, 0].tolist() == [[30.0, 40.0], [70.0, 80.0]]
|
||||
assert value[:, :, 0].tolist() == [[31.0, 41.0], [71.0, 81.0]]
|
||||
|
||||
|
||||
def test_persistent_regional_graph_keeps_native_negpip_split_semantics() -> None:
|
||||
"""Split one manually selected regional graph without adding branch packing."""
|
||||
|
||||
recording = _RecordingKeyValueAttention()
|
||||
block = _block(recording)
|
||||
context = torch.tensor([[[30.0], [31.0], [40.0], [41.0]]])
|
||||
|
||||
def split_negpip(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
_extra_options: dict[str, Any],
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Apply PPM's public alternating-token UNet semantics exactly."""
|
||||
|
||||
return query, key[:, 0::2], value[:, 1::2]
|
||||
|
||||
block(
|
||||
torch.zeros((1, 2, 1)),
|
||||
context=context,
|
||||
transformer_options={"patches": {"attn2_patch": [split_negpip]}},
|
||||
)
|
||||
|
||||
key, value = recording.calls[0]
|
||||
assert key[:, :, 0].tolist() == [[30.0, 40.0]]
|
||||
assert value[:, :, 0].tolist() == [[31.0, 41.0]]
|
||||
|
||||
|
||||
def _block(cross_attention: nn.Module) -> BasicTransformerBlock:
|
||||
"""Build one deterministic block around a supplied cross-attention owner."""
|
||||
|
||||
block = BasicTransformerBlock(
|
||||
dim=1,
|
||||
n_heads=1,
|
||||
d_head=1,
|
||||
context_dim=1,
|
||||
checkpoint=False,
|
||||
)
|
||||
block.norm1 = nn.Identity()
|
||||
block.attn1 = _ZeroAttention()
|
||||
block.norm2 = nn.Identity()
|
||||
block.attn2 = cross_attention
|
||||
block.norm3 = nn.Identity()
|
||||
block.ff = _CountingZeroFeedForward()
|
||||
return block
|
||||
|
||||
|
||||
def _negpip_contexts() -> BatchedRegionalAttentionContexts:
|
||||
"""Return alternating K/V token pairs for base and two regions."""
|
||||
|
||||
return BatchedRegionalAttentionContexts(
|
||||
latent_batch_size=1,
|
||||
chunks=(
|
||||
RegionalAttentionChunkBatch(
|
||||
0,
|
||||
RegionalAttentionBranch.POSITIVE,
|
||||
0,
|
||||
1,
|
||||
),
|
||||
),
|
||||
base_context=torch.tensor([[[10.0], [11.0], [20.0], [21.0]]]),
|
||||
regions=(
|
||||
BatchedRegionalAttentionRegion(
|
||||
0,
|
||||
(
|
||||
BatchedRegionalAttentionEntry(
|
||||
0,
|
||||
torch.tensor([[[30.0], [31.0], [40.0], [41.0]]]),
|
||||
(1.0,),
|
||||
),
|
||||
),
|
||||
),
|
||||
BatchedRegionalAttentionRegion(
|
||||
1,
|
||||
(
|
||||
BatchedRegionalAttentionEntry(
|
||||
0,
|
||||
torch.tensor([[[70.0], [71.0], [80.0], [81.0]]]),
|
||||
(1.0,),
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _contexts() -> BatchedRegionalAttentionContexts:
|
||||
"""Return one base and two single-entry regional contexts."""
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -23,6 +23,7 @@ from simple_syrup.runtime.comfy_conditioning_model_loader import (
|
||||
from simple_syrup.runtime.comfy_conditioning_processing import (
|
||||
ComfyRegionalConditioningProcessor,
|
||||
)
|
||||
from simple_syrup.runtime.ppm_negpip_interop import PpmNegpipInterop
|
||||
from simple_syrup.runtime.regional_lora_conditioning_adapter import (
|
||||
RegionalLoraConditioningAdapter,
|
||||
)
|
||||
@@ -101,6 +102,7 @@ class ProfiledComfyRegionalConditioningProcessor(ComfyRegionalConditioningProces
|
||||
noise: torch.Tensor,
|
||||
device: torch.device,
|
||||
context_validator: RegionalContextValidator,
|
||||
negpip: PpmNegpipInterop | None = None,
|
||||
) -> ProcessedRegionalAttentionPlan:
|
||||
"""Delegate conditioning processing with synchronized device timing."""
|
||||
|
||||
@@ -114,6 +116,7 @@ class ProfiledComfyRegionalConditioningProcessor(ComfyRegionalConditioningProces
|
||||
noise=noise,
|
||||
device=device,
|
||||
context_validator=context_validator,
|
||||
negpip=negpip,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -109,11 +109,6 @@ class RegionalPatchInteropCase:
|
||||
def cases() -> tuple[RegionalPatchInteropCase, ...]:
|
||||
"""Return accepted and rejected cases in authoritative evidence order."""
|
||||
|
||||
negpip_error = (
|
||||
"does not support NegPiP",
|
||||
"ordinary conditioning batch",
|
||||
"regional branch batch",
|
||||
)
|
||||
return (
|
||||
_accepted("anima-full-baseline", "Anima full static PRIMARY_ADAPTER baseline"),
|
||||
_accepted(
|
||||
@@ -183,24 +178,22 @@ def cases() -> tuple[RegionalPatchInteropCase, ...]:
|
||||
("easycache", "Contextual spatial views", "view coordinates"),
|
||||
),
|
||||
RegionalPatchInteropCase(
|
||||
"anima-negpip-rejected",
|
||||
"Reject Anima NegPiP before regional branch packing",
|
||||
"anima-negpip",
|
||||
"Anima NegPiP with aligned regional value masks",
|
||||
PatchInteropModelFamily.ANIMA,
|
||||
PatchInteropSpatialMode.FULL,
|
||||
PatchInteropModifier.NEGPIP,
|
||||
PatchInteropOutcome.REJECTED,
|
||||
PatchInteropOutcome.ACCEPTED,
|
||||
False,
|
||||
negpip_error,
|
||||
),
|
||||
RegionalPatchInteropCase(
|
||||
"sdxl-negpip-rejected",
|
||||
"Reject SDXL NegPiP before paired regional attention patches",
|
||||
"sdxl-negpip",
|
||||
"SDXL NegPiP with packed regional split-K/V conditioning",
|
||||
PatchInteropModelFamily.SDXL,
|
||||
PatchInteropSpatialMode.FULL,
|
||||
PatchInteropModifier.NEGPIP,
|
||||
PatchInteropOutcome.REJECTED,
|
||||
PatchInteropOutcome.ACCEPTED,
|
||||
False,
|
||||
negpip_error,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -133,7 +133,7 @@ def _validate_success(
|
||||
workflow: BuiltRegionalPatchInteropWorkflow,
|
||||
observed: RegionalPatchInteropSuccess,
|
||||
) -> ValidatedRegionalPatchInterop:
|
||||
"""Require one exact single-trajectory regional PRIMARY_ADAPTER execution."""
|
||||
"""Require one exact single-trajectory regional execution."""
|
||||
|
||||
metrics = observed.metrics
|
||||
if metrics.get("run_id") != workflow.metrics_run_id:
|
||||
@@ -150,62 +150,26 @@ def _validate_success(
|
||||
raise ValueError("P9.7 diagnostics identity changed.")
|
||||
snapshots = _array(diagnostics.get("snapshots"), "diagnostic snapshots")
|
||||
record_count = _integer(diagnostics.get("record_count"), "record count")
|
||||
if record_count != model_calls or len(snapshots) != record_count:
|
||||
raise ValueError("P9.7 requires one diagnostic record per model call.")
|
||||
expected_records = model_calls * _diagnostic_records_per_model_call(case)
|
||||
if record_count != expected_records or len(snapshots) != record_count:
|
||||
raise ValueError(
|
||||
"P9.7 diagnostic records do not match the model-family execution shape."
|
||||
)
|
||||
if record_count < 1:
|
||||
raise ValueError("P9.7 diagnostics must contain aligned model-call records.")
|
||||
spatial_modes: set[str] = set()
|
||||
adapter_tokens: set[str] = set()
|
||||
for item in snapshots:
|
||||
snapshot = _object(item, "diagnostic snapshot")
|
||||
if (
|
||||
snapshot.get("strategy") != "attention_coupling"
|
||||
or snapshot.get("backend") != "comfy.ldm.anima.model.Anima"
|
||||
):
|
||||
if snapshot.get("strategy") != "attention_coupling" or snapshot.get(
|
||||
"backend"
|
||||
) != _expected_backend(case):
|
||||
raise ValueError("P9.7 diagnostic strategy or backend changed.")
|
||||
spatial_modes.add(_string(snapshot.get("spatial_mode"), "spatial mode"))
|
||||
uses = _array(snapshot.get("adapter_uses"), "adapter uses")
|
||||
if len(uses) != 2:
|
||||
raise ValueError(
|
||||
"P9.7 accepted cases require paired positive/negative adapter uses."
|
||||
)
|
||||
normalized_uses = tuple(_object(use, "adapter use") for use in uses)
|
||||
if {
|
||||
(_string(use.get("branch"), "adapter branch"), use.get("composition_index"))
|
||||
for use in normalized_uses
|
||||
} != {("positive", 0), ("negative", 1)}:
|
||||
raise ValueError(
|
||||
"P9.7 paired regional PRIMARY_ADAPTER branch ownership changed."
|
||||
)
|
||||
for use in normalized_uses:
|
||||
if (
|
||||
use.get("active") is not True
|
||||
or use.get("region_index") != 0
|
||||
or use.get("target_count") != 448
|
||||
):
|
||||
raise ValueError(
|
||||
"P9.7 exact regional PRIMARY_ADAPTER execution changed."
|
||||
)
|
||||
if not math.isclose(
|
||||
_number(use.get("effective_strength"), "effective strength"),
|
||||
0.75,
|
||||
abs_tol=1e-8,
|
||||
):
|
||||
raise ValueError("P9.7 regional PRIMARY_ADAPTER strength changed.")
|
||||
adapter_tokens.add(_string(use.get("adapter_token"), "adapter token"))
|
||||
work = _object(snapshot.get("estimated_work"), "estimated work")
|
||||
if (
|
||||
work.get("active_adapter_uses") != 2
|
||||
or work.get("active_target_count") != 448
|
||||
or work.get("target_use_count") != 896
|
||||
):
|
||||
raise ValueError("P9.7 paired LoRA target-use accounting changed.")
|
||||
if not math.isclose(
|
||||
_number(work.get("denoiser_call_multiplier"), "denoiser multiplier"),
|
||||
1.0,
|
||||
abs_tol=1e-8,
|
||||
):
|
||||
raise ValueError("P9.7 denoiser trajectory multiplier changed.")
|
||||
if case.model_family is PatchInteropModelFamily.ANIMA:
|
||||
_validate_anima_snapshot(snapshot, adapter_tokens)
|
||||
else:
|
||||
_validate_sdxl_snapshot(snapshot)
|
||||
expected_modes = {
|
||||
PatchInteropSpatialMode.FULL: {"full"},
|
||||
PatchInteropSpatialMode.TILED: {"tile"},
|
||||
@@ -213,7 +177,7 @@ def _validate_success(
|
||||
}[case.spatial_mode]
|
||||
if not expected_modes <= spatial_modes:
|
||||
raise ValueError("P9.7 spatial diagnostics are incomplete.")
|
||||
if len(adapter_tokens) != 1:
|
||||
if case.model_family is PatchInteropModelFamily.ANIMA and len(adapter_tokens) != 1:
|
||||
raise ValueError(
|
||||
"P9.7 regional PRIMARY_ADAPTER identity changed during sampling."
|
||||
)
|
||||
@@ -225,6 +189,93 @@ def _validate_success(
|
||||
)
|
||||
|
||||
|
||||
def _validate_anima_snapshot(
|
||||
snapshot: JsonObject,
|
||||
adapter_tokens: set[str],
|
||||
) -> None:
|
||||
"""Require exact Anima regional-LoRA execution evidence."""
|
||||
|
||||
uses = _array(snapshot.get("adapter_uses"), "adapter uses")
|
||||
if len(uses) != 2:
|
||||
raise ValueError(
|
||||
"P9.7 accepted cases require paired positive/negative adapter uses."
|
||||
)
|
||||
normalized_uses = tuple(_object(use, "adapter use") for use in uses)
|
||||
if {
|
||||
(_string(use.get("branch"), "adapter branch"), use.get("composition_index"))
|
||||
for use in normalized_uses
|
||||
} != {("positive", 0), ("negative", 1)}:
|
||||
raise ValueError(
|
||||
"P9.7 paired regional PRIMARY_ADAPTER branch ownership changed."
|
||||
)
|
||||
for use in normalized_uses:
|
||||
if (
|
||||
use.get("active") is not True
|
||||
or use.get("region_index") != 0
|
||||
or use.get("target_count") != 448
|
||||
):
|
||||
raise ValueError("P9.7 exact regional PRIMARY_ADAPTER execution changed.")
|
||||
if not math.isclose(
|
||||
_number(use.get("effective_strength"), "effective strength"),
|
||||
0.75,
|
||||
abs_tol=1e-8,
|
||||
):
|
||||
raise ValueError("P9.7 regional PRIMARY_ADAPTER strength changed.")
|
||||
adapter_tokens.add(_string(use.get("adapter_token"), "adapter token"))
|
||||
work = _object(snapshot.get("estimated_work"), "estimated work")
|
||||
if (
|
||||
work.get("active_adapter_uses") != 2
|
||||
or work.get("active_target_count") != 448
|
||||
or work.get("target_use_count") != 896
|
||||
):
|
||||
raise ValueError("P9.7 paired LoRA target-use accounting changed.")
|
||||
_validate_single_denoiser_trajectory(work)
|
||||
|
||||
|
||||
def _validate_sdxl_snapshot(snapshot: JsonObject) -> None:
|
||||
"""Require exact SDXL regional cross-attention execution evidence."""
|
||||
|
||||
if snapshot.get("region_count") != 2 or snapshot.get("active_region_indices") != [
|
||||
0,
|
||||
1,
|
||||
]:
|
||||
raise ValueError("P9.7 SDXL regional branch execution changed.")
|
||||
work = _object(snapshot.get("estimated_work"), "estimated work")
|
||||
if (
|
||||
work.get("cross_attention_branch_multiplier") != 3.0
|
||||
or work.get("cross_attention_formula") != "base_plus_region_count"
|
||||
):
|
||||
raise ValueError("P9.7 SDXL regional attention accounting changed.")
|
||||
_validate_single_denoiser_trajectory(work)
|
||||
|
||||
|
||||
def _validate_single_denoiser_trajectory(work: JsonObject) -> None:
|
||||
"""Require regional work to retain one denoiser trajectory."""
|
||||
|
||||
if not math.isclose(
|
||||
_number(work.get("denoiser_call_multiplier"), "denoiser multiplier"),
|
||||
1.0,
|
||||
abs_tol=1e-8,
|
||||
):
|
||||
raise ValueError("P9.7 denoiser trajectory multiplier changed.")
|
||||
|
||||
|
||||
def _diagnostic_records_per_model_call(case: RegionalPatchInteropCase) -> int:
|
||||
"""Return the baseline diagnostic cardinality for one model family."""
|
||||
|
||||
if case.model_family is PatchInteropModelFamily.SDXL:
|
||||
return 2
|
||||
return 1
|
||||
|
||||
|
||||
def _expected_backend(case: RegionalPatchInteropCase) -> str:
|
||||
"""Return the exact Comfy denoiser backend identity for one family."""
|
||||
|
||||
if case.model_family is PatchInteropModelFamily.SDXL:
|
||||
return "comfy.ldm.modules.diffusionmodules.openaimodel.UNetModel"
|
||||
return "comfy.ldm.anima.model.Anima"
|
||||
|
||||
|
||||
def _validate_model_call_count(
|
||||
case: RegionalPatchInteropCase,
|
||||
model_calls: int,
|
||||
|
||||
@@ -171,7 +171,7 @@ class RegionalPatchInteropWorkflowBuilder:
|
||||
mask_names: tuple[str, ...],
|
||||
checkpoint_name: str,
|
||||
) -> BuiltRegionalPatchInteropWorkflow:
|
||||
"""Build the focused SDXL NegPiP rejection graph."""
|
||||
"""Build the focused SDXL NegPiP execution graph."""
|
||||
|
||||
graph = AnimaWorkflowGraph()
|
||||
loader = graph.add("CheckpointLoaderSimple", ckpt_name=checkpoint_name)
|
||||
|
||||
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 curated downloadable SAM, GroundingDINO, ViTMatte, WD14 tagger, and Ultralytics models in loader dropdowns.";
|
||||
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 curated downloadable SAM, GroundingDINO, ViTMatte, WD14 tagger, and Ultralytics models in loader dropdowns.";
|
||||
|
||||
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