fix(runtime): centralize Comfy patcher lifecycle
This commit is contained in:
@@ -9,8 +9,13 @@ from __future__ import annotations
|
||||
import importlib
|
||||
from collections.abc import Iterable
|
||||
from types import ModuleType
|
||||
from typing import Any, Protocol, cast, runtime_checkable
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
from .patcher_lifecycle import (
|
||||
PATCHER_LIFECYCLE,
|
||||
ClipLayerMutation,
|
||||
ComfyPatcherLifecycle,
|
||||
)
|
||||
from .vae_loader import VaeLoaderService
|
||||
|
||||
USE_CHECKPOINT_VAE_CHOICE = "Use Checkpoint VAE"
|
||||
@@ -25,11 +30,13 @@ class CheckpointLoaderService:
|
||||
self,
|
||||
folder_paths_module: ModuleType | None = None,
|
||||
vae_loader: VaeLoaderBoundary | None = None,
|
||||
patcher_lifecycle: ComfyPatcherLifecycle | None = None,
|
||||
) -> None:
|
||||
"""Create a checkpoint loader with injectable runtime boundaries."""
|
||||
|
||||
self._folder_paths_module = folder_paths_module
|
||||
self._vae_loader = vae_loader or VaeLoaderService(folder_paths_module)
|
||||
self._patcher_lifecycle = patcher_lifecycle or PATCHER_LIFECYCLE
|
||||
|
||||
def load_checkpoint(
|
||||
self,
|
||||
@@ -57,12 +64,32 @@ class CheckpointLoaderService:
|
||||
model = loaded[0]
|
||||
clip = loaded[1]
|
||||
checkpoint_vae = loaded[2]
|
||||
selected_clip = _selected_clip(clip, clip_skip)
|
||||
selected_clip = self._selected_clip(clip, clip_skip)
|
||||
|
||||
if vae_name == USE_CHECKPOINT_VAE_CHOICE:
|
||||
return model, selected_clip, checkpoint_vae
|
||||
selected_vae = checkpoint_vae
|
||||
else:
|
||||
selected_vae = self._vae_loader.load_vae(vae_name)
|
||||
|
||||
return model, selected_clip, self._vae_loader.load_vae(vae_name)
|
||||
return (
|
||||
model,
|
||||
selected_clip,
|
||||
self._patcher_lifecycle.preserve_vae(
|
||||
selected_vae,
|
||||
operation="SimpleSyrup checkpoint loading",
|
||||
),
|
||||
)
|
||||
|
||||
def _selected_clip(self, clip: object, clip_skip: bool) -> object:
|
||||
"""Return the loaded CLIP or a lifecycle-owned clip-skip derivation."""
|
||||
|
||||
if not clip_skip:
|
||||
return clip
|
||||
return self._patcher_lifecycle.derive_clip(
|
||||
clip,
|
||||
(ClipLayerMutation(CLIP_SKIP_LAYER),),
|
||||
operation="SimpleSyrup checkpoint clip skip",
|
||||
)
|
||||
|
||||
def _folder_paths(self) -> ModuleType:
|
||||
"""Return the ComfyUI folder_paths module."""
|
||||
@@ -83,17 +110,6 @@ class VaeLoaderBoundary(Protocol):
|
||||
"""Load the named external VAE."""
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ClipLayerBoundary(Protocol):
|
||||
"""CLIP interface required to apply the ComfyUI clip-skip layer."""
|
||||
|
||||
def clone(self) -> ClipLayerBoundary:
|
||||
"""Return an independent CLIP object."""
|
||||
|
||||
def clip_layer(self, layer_idx: int) -> None:
|
||||
"""Set the CLIP layer index used during prompt encoding."""
|
||||
|
||||
|
||||
def _validate_clip_skip(clip_skip: object) -> None:
|
||||
"""Reject non-boolean clip-skip selections from runtime callers."""
|
||||
|
||||
@@ -101,22 +117,6 @@ def _validate_clip_skip(clip_skip: object) -> None:
|
||||
raise TypeError("clip_skip must be a boolean.")
|
||||
|
||||
|
||||
def _selected_clip(clip: object, clip_skip: bool) -> object:
|
||||
"""Return the loaded CLIP or a cloned CLIP with clip skip applied."""
|
||||
|
||||
if not clip_skip:
|
||||
return clip
|
||||
|
||||
if not isinstance(clip, ClipLayerBoundary):
|
||||
raise TypeError(
|
||||
"clip_skip requires a CLIP object with clone() and clip_layer()."
|
||||
)
|
||||
|
||||
selected_clip = clip.clone()
|
||||
selected_clip.clip_layer(CLIP_SKIP_LAYER)
|
||||
return selected_clip
|
||||
|
||||
|
||||
def _comfy_sd() -> Any:
|
||||
"""Import ComfyUI's stable diffusion loading module lazily."""
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ from ..domain.contextual_diffusion import (
|
||||
)
|
||||
from ..shared.logging import get_logger
|
||||
from . import sampling_samplers, sampling_schedulers
|
||||
from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelUnetWrapperMutation
|
||||
from .tiled_sampling import (
|
||||
ApplyModel,
|
||||
Latent,
|
||||
@@ -156,10 +157,9 @@ def clone_model_with_contextual_diffusion(
|
||||
sigmas: torch.Tensor,
|
||||
diffusion_mode: str,
|
||||
) -> Any:
|
||||
"""Clone a model and install one pre-CFG contextual prediction wrapper."""
|
||||
"""Derive a model with one pre-CFG contextual prediction wrapper."""
|
||||
|
||||
cloned_model = model.clone()
|
||||
old_wrapper = cloned_model.model_options.get("model_function_wrapper")
|
||||
old_wrapper = model.model_options.get("model_function_wrapper")
|
||||
if old_wrapper is not None and not callable(old_wrapper):
|
||||
raise ValueError("Existing model_function_wrapper is not callable.")
|
||||
wrapper = ContextualDiffusionModelWrapper(
|
||||
@@ -169,8 +169,11 @@ def clone_model_with_contextual_diffusion(
|
||||
diffusion_mode=diffusion_mode,
|
||||
existing_wrapper=cast(ModelFunctionWrapper | None, old_wrapper),
|
||||
)
|
||||
cloned_model.set_model_unet_function_wrapper(wrapper)
|
||||
return cloned_model
|
||||
return PATCHER_LIFECYCLE.derive_model(
|
||||
model,
|
||||
(ModelUnetWrapperMutation(wrapper),),
|
||||
operation="SimpleSyrup contextual diffusion",
|
||||
)
|
||||
|
||||
|
||||
class ContextualDiffusionModelWrapper:
|
||||
|
||||
@@ -9,6 +9,8 @@ from __future__ import annotations
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelDenoiseMaskMutation
|
||||
|
||||
|
||||
def has_denoise_mask_function(model: Any) -> bool:
|
||||
"""Return whether a model patcher already has denoise-mask behavior."""
|
||||
@@ -20,31 +22,29 @@ def has_denoise_mask_function(model: Any) -> bool:
|
||||
|
||||
|
||||
def clone_with_differential_diffusion(model: Any, strength: float = 1.0) -> Any:
|
||||
"""Return a clone patched with ComfyUI differential denoise masks."""
|
||||
"""Return a lifecycle-owned MODEL with differential denoise masks."""
|
||||
|
||||
if has_denoise_mask_function(model):
|
||||
return model
|
||||
cloned_model = model.clone()
|
||||
install_differential_diffusion(cloned_model, strength=strength)
|
||||
return cloned_model
|
||||
return PATCHER_LIFECYCLE.derive_model(
|
||||
model,
|
||||
(differential_diffusion_mutation(strength=strength),),
|
||||
operation="SimpleSyrup differential diffusion",
|
||||
)
|
||||
|
||||
|
||||
def install_differential_diffusion(model: Any, strength: float = 1.0) -> Any:
|
||||
"""Install ComfyUI differential denoise-mask behavior on a model patcher."""
|
||||
def differential_diffusion_mutation(
|
||||
strength: float = 1.0,
|
||||
) -> ModelDenoiseMaskMutation:
|
||||
"""Build the ComfyUI differential denoise-mask mutation."""
|
||||
|
||||
if has_denoise_mask_function(model):
|
||||
return model
|
||||
set_mask_function = getattr(model, "set_model_denoise_mask_function", None)
|
||||
if not callable(set_mask_function):
|
||||
raise ValueError("Model does not support differential diffusion denoise masks.")
|
||||
differential_diffusion = import_module(
|
||||
"comfy_extras.nodes_differential_diffusion"
|
||||
).DifferentialDiffusion
|
||||
set_mask_function(
|
||||
return ModelDenoiseMaskMutation(
|
||||
lambda *args, **kwargs: differential_diffusion.forward(
|
||||
*args,
|
||||
**kwargs,
|
||||
strength=strength,
|
||||
)
|
||||
)
|
||||
return model
|
||||
|
||||
@@ -23,7 +23,15 @@ from ..domain.tiled_diffusion import (
|
||||
from ..shared.logging import get_logger
|
||||
from . import sampling_samplers, sampling_schedulers
|
||||
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
|
||||
from .differential_diffusion import install_differential_diffusion
|
||||
from .differential_diffusion import (
|
||||
differential_diffusion_mutation,
|
||||
has_denoise_mask_function,
|
||||
)
|
||||
from .patcher_lifecycle import (
|
||||
PATCHER_LIFECYCLE,
|
||||
ModelMutation,
|
||||
ModelUnetWrapperMutation,
|
||||
)
|
||||
from .tiled_sampling import (
|
||||
ApplyModel,
|
||||
Latent,
|
||||
@@ -176,7 +184,7 @@ def clone_model_with_mixture_of_diffusers(
|
||||
differential_diffusion: bool = False,
|
||||
tiled_plan: TiledDiffusionPlan | None = None,
|
||||
) -> tuple[Any, TiledDiffusionPlan]:
|
||||
"""Return a model clone patched with a pre-CFG Mixture wrapper."""
|
||||
"""Return a derived model patched with a pre-CFG Mixture wrapper."""
|
||||
|
||||
plan = tiled_plan or build_tiled_diffusion_plan(
|
||||
latent_width=latent_width,
|
||||
@@ -187,10 +195,7 @@ def clone_model_with_mixture_of_diffusers(
|
||||
tile_batch_size=tile_batch_size,
|
||||
)
|
||||
_validate_supplied_plan(plan, latent_width, latent_height)
|
||||
cloned_model = model.clone()
|
||||
if differential_diffusion:
|
||||
install_differential_diffusion(cloned_model)
|
||||
old_wrapper = cloned_model.model_options.get("model_function_wrapper")
|
||||
old_wrapper = model.model_options.get("model_function_wrapper")
|
||||
if old_wrapper is not None and not callable(old_wrapper):
|
||||
raise ValueError("Existing model_function_wrapper is not callable.")
|
||||
|
||||
@@ -198,8 +203,16 @@ def clone_model_with_mixture_of_diffusers(
|
||||
plan=plan,
|
||||
existing_wrapper=cast(ModelFunctionWrapper | None, old_wrapper),
|
||||
)
|
||||
cloned_model.set_model_unet_function_wrapper(wrapper)
|
||||
return cloned_model, plan
|
||||
mutations: list[ModelMutation] = []
|
||||
if differential_diffusion and not has_denoise_mask_function(model):
|
||||
mutations.append(differential_diffusion_mutation())
|
||||
mutations.append(ModelUnetWrapperMutation(wrapper))
|
||||
derived_model = PATCHER_LIFECYCLE.derive_model(
|
||||
model,
|
||||
mutations,
|
||||
operation="SimpleSyrup Mixture of Diffusers",
|
||||
)
|
||||
return derived_model, plan
|
||||
|
||||
|
||||
class MixtureOfDiffusersModelWrapper:
|
||||
|
||||
@@ -23,7 +23,15 @@ from ..domain.tiled_diffusion import (
|
||||
from ..shared.logging import get_logger
|
||||
from . import sampling_samplers, sampling_schedulers
|
||||
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
|
||||
from .differential_diffusion import install_differential_diffusion
|
||||
from .differential_diffusion import (
|
||||
differential_diffusion_mutation,
|
||||
has_denoise_mask_function,
|
||||
)
|
||||
from .patcher_lifecycle import (
|
||||
PATCHER_LIFECYCLE,
|
||||
ModelMutation,
|
||||
ModelUnetWrapperMutation,
|
||||
)
|
||||
from .tiled_sampling import (
|
||||
ApplyModel,
|
||||
Latent,
|
||||
@@ -179,7 +187,7 @@ def clone_model_with_multidiffusion(
|
||||
differential_diffusion: bool = False,
|
||||
tiled_plan: TiledDiffusionPlan | None = None,
|
||||
) -> tuple[Any, TiledDiffusionPlan]:
|
||||
"""Return a model clone patched with a pre-CFG MultiDiffusion wrapper."""
|
||||
"""Return a derived model patched with a pre-CFG MultiDiffusion wrapper."""
|
||||
|
||||
plan = tiled_plan or build_tiled_diffusion_plan(
|
||||
latent_width=latent_width,
|
||||
@@ -190,10 +198,7 @@ def clone_model_with_multidiffusion(
|
||||
tile_batch_size=tile_batch_size,
|
||||
)
|
||||
_validate_supplied_plan(plan, latent_width, latent_height)
|
||||
cloned_model = model.clone()
|
||||
if differential_diffusion:
|
||||
install_differential_diffusion(cloned_model)
|
||||
old_wrapper = cloned_model.model_options.get("model_function_wrapper")
|
||||
old_wrapper = model.model_options.get("model_function_wrapper")
|
||||
if old_wrapper is not None and not callable(old_wrapper):
|
||||
raise ValueError("Existing model_function_wrapper is not callable.")
|
||||
|
||||
@@ -201,8 +206,16 @@ def clone_model_with_multidiffusion(
|
||||
plan=plan,
|
||||
existing_wrapper=cast(ModelFunctionWrapper | None, old_wrapper),
|
||||
)
|
||||
cloned_model.set_model_unet_function_wrapper(wrapper)
|
||||
return cloned_model, plan
|
||||
mutations: list[ModelMutation] = []
|
||||
if differential_diffusion and not has_denoise_mask_function(model):
|
||||
mutations.append(differential_diffusion_mutation())
|
||||
mutations.append(ModelUnetWrapperMutation(wrapper))
|
||||
derived_model = PATCHER_LIFECYCLE.derive_model(
|
||||
model,
|
||||
mutations,
|
||||
operation="SimpleSyrup MultiDiffusion",
|
||||
)
|
||||
return derived_model, plan
|
||||
|
||||
|
||||
class MultiDiffusionModelWrapper:
|
||||
|
||||
@@ -0,0 +1,216 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Own lifecycle-safe derivation of ComfyUI MODEL and CLIP values."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Iterable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol, TypeVar, cast
|
||||
|
||||
|
||||
class ModelMutation(Protocol):
|
||||
"""Apply one supported mutation to an already-derived MODEL patcher."""
|
||||
|
||||
def apply(self, model: object) -> None:
|
||||
"""Apply the mutation through ComfyUI's public patcher API."""
|
||||
|
||||
|
||||
class ClipMutation(Protocol):
|
||||
"""Apply one supported mutation to an already-derived CLIP value."""
|
||||
|
||||
def apply(self, clip: object) -> None:
|
||||
"""Apply the mutation through ComfyUI's public CLIP API."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelDenoiseMaskMutation:
|
||||
"""Install a denoise-mask function through the MODEL patcher API."""
|
||||
|
||||
function: Callable[..., object]
|
||||
|
||||
def apply(self, model: object) -> None:
|
||||
"""Install the configured denoise-mask function."""
|
||||
|
||||
setter = getattr(model, "set_model_denoise_mask_function", None)
|
||||
if not callable(setter):
|
||||
raise TypeError("MODEL does not support denoise-mask functions.")
|
||||
setter(self.function)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelUnetWrapperMutation:
|
||||
"""Install a model-function wrapper through the MODEL patcher API."""
|
||||
|
||||
wrapper: Callable[..., object]
|
||||
|
||||
def apply(self, model: object) -> None:
|
||||
"""Install the configured model-function wrapper."""
|
||||
|
||||
setter = getattr(model, "set_model_unet_function_wrapper", None)
|
||||
if not callable(setter):
|
||||
raise TypeError("MODEL does not support model-function wrappers.")
|
||||
setter(self.wrapper)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelCalcCondBatchMutation:
|
||||
"""Install a calc-cond-batch function through the MODEL patcher API."""
|
||||
|
||||
function: Callable[..., object]
|
||||
|
||||
def apply(self, model: object) -> None:
|
||||
"""Install the configured calc-cond-batch function."""
|
||||
|
||||
setter = getattr(model, "set_model_sampler_calc_cond_batch_function", None)
|
||||
if not callable(setter):
|
||||
raise TypeError("MODEL does not support calc-cond-batch functions.")
|
||||
setter(self.function)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ClipLayerMutation:
|
||||
"""Select the text-encoder layer used by a derived CLIP value."""
|
||||
|
||||
layer_index: int
|
||||
|
||||
def apply(self, clip: object) -> None:
|
||||
"""Apply the configured CLIP layer index."""
|
||||
|
||||
select_layer = getattr(clip, "clip_layer", None)
|
||||
if not callable(select_layer):
|
||||
raise TypeError("CLIP does not support layer selection.")
|
||||
select_layer(self.layer_index)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ClipHookScheduleMutation:
|
||||
"""Install a native ComfyUI hook schedule on a derived CLIP value."""
|
||||
|
||||
hooks: object
|
||||
target: object
|
||||
|
||||
def apply(self, clip: object) -> None:
|
||||
"""Clone and register hooks through the derived CLIP patcher."""
|
||||
|
||||
patcher = _required_attribute(clip, "patcher", value_name="CLIP")
|
||||
clone_hooks = getattr(self.hooks, "clone", None)
|
||||
if not callable(clone_hooks):
|
||||
raise TypeError("HOOKS does not support lifecycle-safe cloning.")
|
||||
register_hooks = getattr(patcher, "register_all_hook_patches", None)
|
||||
if not callable(register_hooks):
|
||||
raise TypeError("CLIP patcher does not support hook registration.")
|
||||
|
||||
patcher_boundary = cast(Any, patcher)
|
||||
clip_boundary = cast(Any, clip)
|
||||
patcher_boundary.forced_hooks = clone_hooks()
|
||||
clip_boundary.use_clip_schedule = True
|
||||
register_hooks(self.hooks, self.target)
|
||||
|
||||
|
||||
PatcherValue = TypeVar("PatcherValue")
|
||||
|
||||
|
||||
class ComfyPatcherLifecycle:
|
||||
"""Derive Comfy patchers while preserving their source lineage."""
|
||||
|
||||
def derive_model(
|
||||
self,
|
||||
source: PatcherValue,
|
||||
mutations: Iterable[ModelMutation],
|
||||
*,
|
||||
operation: str,
|
||||
) -> PatcherValue:
|
||||
"""Clone one MODEL, verify its lineage, and apply all mutations."""
|
||||
|
||||
derived = self._clone(source, operation=operation)
|
||||
self._require_direct_parent(source, derived, operation=operation)
|
||||
for mutation in mutations:
|
||||
mutation.apply(derived)
|
||||
return derived
|
||||
|
||||
def derive_clip(
|
||||
self,
|
||||
source: PatcherValue,
|
||||
mutations: Iterable[ClipMutation],
|
||||
*,
|
||||
operation: str,
|
||||
disable_dynamic: bool = False,
|
||||
) -> PatcherValue:
|
||||
"""Clone one CLIP, verify patcher lineage, and apply all mutations."""
|
||||
|
||||
clone = getattr(source, "clone", None)
|
||||
if not callable(clone):
|
||||
raise TypeError(f"{operation} requires a cloneable CLIP value.")
|
||||
derived = cast(
|
||||
PatcherValue,
|
||||
clone(disable_dynamic=disable_dynamic) if disable_dynamic else clone(),
|
||||
)
|
||||
if derived is source:
|
||||
raise RuntimeError(f"{operation} returned the source CLIP from clone().")
|
||||
|
||||
source_patcher = _required_attribute(
|
||||
source,
|
||||
"patcher",
|
||||
value_name="source CLIP",
|
||||
)
|
||||
derived_patcher = _required_attribute(
|
||||
derived,
|
||||
"patcher",
|
||||
value_name="derived CLIP",
|
||||
)
|
||||
self._require_direct_parent(
|
||||
source_patcher,
|
||||
derived_patcher,
|
||||
operation=operation,
|
||||
)
|
||||
for mutation in mutations:
|
||||
mutation.apply(derived)
|
||||
return derived
|
||||
|
||||
def preserve_vae(self, vae: PatcherValue, *, operation: str) -> PatcherValue:
|
||||
"""Return an unmodified VAE and make the no-derivation contract explicit."""
|
||||
|
||||
if not operation.strip():
|
||||
raise ValueError("VAE lifecycle operations require a descriptive name.")
|
||||
return vae
|
||||
|
||||
@staticmethod
|
||||
def _clone(source: PatcherValue, *, operation: str) -> PatcherValue:
|
||||
"""Clone one MODEL through its ComfyUI boundary."""
|
||||
|
||||
clone = getattr(source, "clone", None)
|
||||
if not callable(clone):
|
||||
raise TypeError(f"{operation} requires a cloneable MODEL value.")
|
||||
derived = cast(PatcherValue, clone())
|
||||
if derived is source:
|
||||
raise RuntimeError(f"{operation} returned the source MODEL from clone().")
|
||||
return derived
|
||||
|
||||
@staticmethod
|
||||
def _require_direct_parent(
|
||||
source: object,
|
||||
derived: object,
|
||||
*,
|
||||
operation: str,
|
||||
) -> None:
|
||||
"""Require Comfy's parent link used by loaded-model cleanup."""
|
||||
|
||||
if getattr(derived, "parent", None) is not source:
|
||||
raise RuntimeError(
|
||||
f"{operation} produced a derived patcher without its source as parent."
|
||||
)
|
||||
|
||||
|
||||
PATCHER_LIFECYCLE = ComfyPatcherLifecycle()
|
||||
|
||||
|
||||
def _required_attribute(value: object, name: str, *, value_name: str) -> object:
|
||||
"""Return a required dynamic ComfyUI boundary attribute."""
|
||||
|
||||
attribute = getattr(value, name, None)
|
||||
if attribute is None:
|
||||
raise TypeError(f"{value_name} does not expose {name}.")
|
||||
return attribute
|
||||
@@ -10,6 +10,8 @@ import logging
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
from .patcher_lifecycle import PATCHER_LIFECYCLE, ClipHookScheduleMutation
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -32,12 +34,18 @@ def prepare_regional_lora_clip(clip: Any, hooks: object) -> tuple[Any, object]:
|
||||
"text-encoder hook group entries.",
|
||||
matching_hook_count,
|
||||
)
|
||||
prepared_clip = clip.clone(disable_dynamic=True)
|
||||
prepared_clip.patcher.forced_hooks = hooks.clone()
|
||||
prepared_clip.use_clip_schedule = True
|
||||
prepared_clip.patcher.register_all_hook_patches(
|
||||
hooks,
|
||||
comfy_hooks.create_target_dict(comfy_hooks.EnumWeightTarget.Clip),
|
||||
prepared_clip = PATCHER_LIFECYCLE.derive_clip(
|
||||
clip,
|
||||
(
|
||||
ClipHookScheduleMutation(
|
||||
hooks=hooks,
|
||||
target=comfy_hooks.create_target_dict(
|
||||
comfy_hooks.EnumWeightTarget.Clip
|
||||
),
|
||||
),
|
||||
),
|
||||
operation="SimpleSyrup regional LoRA CLIP preparation",
|
||||
disable_dynamic=True,
|
||||
)
|
||||
return prepared_clip, hooks
|
||||
|
||||
|
||||
@@ -22,7 +22,15 @@ from ..domain.regional_detailing import LatentRegion
|
||||
from ..shared.logging import get_logger
|
||||
from . import sampling_samplers, sampling_schedulers
|
||||
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
|
||||
from .differential_diffusion import install_differential_diffusion
|
||||
from .differential_diffusion import (
|
||||
differential_diffusion_mutation,
|
||||
has_denoise_mask_function,
|
||||
)
|
||||
from .patcher_lifecycle import (
|
||||
PATCHER_LIFECYCLE,
|
||||
ModelCalcCondBatchMutation,
|
||||
ModelMutation,
|
||||
)
|
||||
from .tiled_sampling import (
|
||||
Latent,
|
||||
reject_unsupported_conditioning,
|
||||
@@ -168,7 +176,7 @@ def clone_model_with_regional_multidiffusion(
|
||||
global_prompt_weight: float = 0.0,
|
||||
differential_diffusion: bool = False,
|
||||
) -> tuple[Any, RegionalMultiDiffusionSummary]:
|
||||
"""Return a model clone patched with regional calc-cond-batch blending."""
|
||||
"""Return a derived model patched with regional calc-cond-batch blending."""
|
||||
|
||||
_validate_sampling_controls(
|
||||
steps=1,
|
||||
@@ -180,10 +188,7 @@ def clone_model_with_regional_multidiffusion(
|
||||
latent_height=latent_height,
|
||||
regions=regions,
|
||||
)
|
||||
cloned_model = model.clone()
|
||||
if differential_diffusion:
|
||||
install_differential_diffusion(cloned_model)
|
||||
old_wrapper = cloned_model.model_options.get("sampler_calc_cond_batch_function")
|
||||
old_wrapper = model.model_options.get("sampler_calc_cond_batch_function")
|
||||
if old_wrapper is not None and not callable(old_wrapper):
|
||||
raise ValueError("Existing sampler_calc_cond_batch_function is not callable.")
|
||||
|
||||
@@ -194,7 +199,15 @@ def clone_model_with_regional_multidiffusion(
|
||||
existing_calc_cond_batch=cast(CalcCondBatchFunction | None, old_wrapper),
|
||||
global_prompt_weight=global_prompt_weight,
|
||||
)
|
||||
cloned_model.set_model_sampler_calc_cond_batch_function(wrapper)
|
||||
mutations: list[ModelMutation] = []
|
||||
if differential_diffusion and not has_denoise_mask_function(model):
|
||||
mutations.append(differential_diffusion_mutation())
|
||||
mutations.append(ModelCalcCondBatchMutation(wrapper))
|
||||
derived_model = PATCHER_LIFECYCLE.derive_model(
|
||||
model,
|
||||
mutations,
|
||||
operation="SimpleSyrup regional MultiDiffusion",
|
||||
)
|
||||
summary = RegionalMultiDiffusionSummary(
|
||||
latent_width=latent_width,
|
||||
latent_height=latent_height,
|
||||
@@ -208,7 +221,7 @@ def clone_model_with_regional_multidiffusion(
|
||||
default=0,
|
||||
),
|
||||
)
|
||||
return cloned_model, summary
|
||||
return derived_model, summary
|
||||
|
||||
|
||||
class RegionalMultiDiffusionCalcCondBatch:
|
||||
|
||||
@@ -14,6 +14,8 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from .patcher_lifecycle import PATCHER_LIFECYCLE
|
||||
|
||||
VIDEO_TAES = ("taehv", "lighttaew2_2", "lighttaew2_1", "lighttaehy1_5", "taeltx_2")
|
||||
IMAGE_TAES = ("taesd", "taesdxl", "taesd3", "taef1")
|
||||
|
||||
@@ -156,7 +158,10 @@ def _build_vae(sd: dict[str, object], metadata: object | None) -> object:
|
||||
comfy_sd = _comfy_sd()
|
||||
vae = comfy_sd.VAE(sd=sd, metadata=metadata)
|
||||
vae.throw_exception_if_invalid()
|
||||
return vae
|
||||
return PATCHER_LIFECYCLE.preserve_vae(
|
||||
vae,
|
||||
operation="SimpleSyrup VAE loading",
|
||||
)
|
||||
|
||||
|
||||
def _folder_paths() -> ModuleType:
|
||||
|
||||
@@ -9,7 +9,7 @@ from __future__ import annotations
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
from types import ModuleType, SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -24,18 +24,23 @@ from simple_syrup.runtime.checkpoint_loader import (
|
||||
class FakeClip:
|
||||
"""CLIP double that records clone and layer selection behavior."""
|
||||
|
||||
def __init__(self, name: str = "checkpoint_clip") -> None:
|
||||
def __init__(
|
||||
self,
|
||||
name: str = "checkpoint_clip",
|
||||
parent_patcher: object | None = None,
|
||||
) -> None:
|
||||
"""Create a CLIP double with no selected layer."""
|
||||
|
||||
self.name = name
|
||||
self.layer: int | None = None
|
||||
self.clone_count = 0
|
||||
self.patcher = SimpleNamespace(parent=parent_patcher)
|
||||
|
||||
def clone(self) -> FakeClip:
|
||||
"""Return an independent CLIP double and record the clone call."""
|
||||
|
||||
self.clone_count += 1
|
||||
return FakeClip(f"{self.name}_clone")
|
||||
return FakeClip(f"{self.name}_clone", parent_patcher=self.patcher)
|
||||
|
||||
def clip_layer(self, layer: int) -> None:
|
||||
"""Record the selected CLIP layer."""
|
||||
|
||||
@@ -0,0 +1,239 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Characterize ComfyUI's patcher lifecycle leak detector."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import logging
|
||||
import weakref
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.runtime.patcher_lifecycle import (
|
||||
ClipLayerMutation,
|
||||
ComfyPatcherLifecycle,
|
||||
ModelCalcCondBatchMutation,
|
||||
ModelDenoiseMaskMutation,
|
||||
ModelUnetWrapperMutation,
|
||||
)
|
||||
|
||||
|
||||
class AnimaTEModel_(torch.nn.Module):
|
||||
"""Represent a tiny Anima text encoder without production weights."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create one parameter so ComfyUI sees a normal torch module."""
|
||||
|
||||
super().__init__()
|
||||
self.projection = torch.nn.Linear(1, 1)
|
||||
|
||||
|
||||
class _AnimaClip:
|
||||
"""Expose an Anima-style encoder through the native CLIP clone shape."""
|
||||
|
||||
def __init__(self, patcher: Any, encoder: AnimaTEModel_) -> None:
|
||||
"""Create a CLIP value around one real Comfy ModelPatcher."""
|
||||
|
||||
self.patcher = patcher
|
||||
self.cond_stage_model = encoder
|
||||
self.layer_index: int | None = None
|
||||
|
||||
def clone(self, disable_dynamic: bool = False) -> _AnimaClip:
|
||||
"""Clone the patcher with the same behavior as Comfy's CLIP wrapper."""
|
||||
|
||||
return _AnimaClip(
|
||||
self.patcher.clone(disable_dynamic=disable_dynamic),
|
||||
self.cond_stage_model,
|
||||
)
|
||||
|
||||
def clip_layer(self, layer_index: int) -> None:
|
||||
"""Record the selected text-encoder layer."""
|
||||
|
||||
self.layer_index = layer_index
|
||||
|
||||
|
||||
def test_real_comfy_anima_lifecycle_regression(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Exercise the unsafe predicate and every safe owner path in one worker."""
|
||||
|
||||
_assert_comfy_marks_rootless_anima_patcher_dead()
|
||||
_assert_comfy_returns_clone_to_live_source()
|
||||
_assert_lifecycle_owned_anima_clip_is_safe(caplog, monkeypatch)
|
||||
_assert_supported_model_mutations_share_one_clone()
|
||||
|
||||
|
||||
def _assert_comfy_marks_rootless_anima_patcher_dead() -> None:
|
||||
"""Capture the exact condition behind ComfyUI's memory-leak warning."""
|
||||
|
||||
from comfy.model_management import LoadedModel
|
||||
|
||||
encoder = AnimaTEModel_()
|
||||
patcher = _patcher(encoder)
|
||||
loaded = LoadedModel(patcher)
|
||||
loaded.real_model = weakref.ref(encoder)
|
||||
|
||||
del patcher
|
||||
gc.collect()
|
||||
|
||||
assert loaded.model is None
|
||||
assert loaded.real_model() is encoder
|
||||
assert loaded.is_dead() is True
|
||||
|
||||
|
||||
def _assert_comfy_returns_clone_to_live_source() -> None:
|
||||
"""Prove valid clone lineage cannot satisfy the stale-patcher condition."""
|
||||
|
||||
from comfy.model_management import LoadedModel
|
||||
|
||||
encoder = AnimaTEModel_()
|
||||
source = _patcher(encoder)
|
||||
derived = source.clone()
|
||||
loaded = LoadedModel(derived)
|
||||
loaded.real_model = weakref.ref(encoder)
|
||||
|
||||
del derived
|
||||
gc.collect()
|
||||
|
||||
assert loaded.model is source
|
||||
assert loaded.is_dead() is False
|
||||
|
||||
|
||||
def _assert_lifecycle_owned_anima_clip_is_safe(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""A released first-party Anima CLIP derivation resolves to its live source."""
|
||||
|
||||
import comfy.model_management
|
||||
from comfy.model_management import LoadedModel
|
||||
|
||||
encoder = AnimaTEModel_()
|
||||
source_clip = _AnimaClip(_patcher(encoder), encoder)
|
||||
derived_clip = ComfyPatcherLifecycle().derive_clip(
|
||||
source_clip,
|
||||
(ClipLayerMutation(-2),),
|
||||
operation="Anima CLIP regression",
|
||||
)
|
||||
assert isinstance(derived_clip, _AnimaClip)
|
||||
assert derived_clip.layer_index == -2
|
||||
loaded = LoadedModel(derived_clip.patcher)
|
||||
loaded.real_model = weakref.ref(encoder)
|
||||
monkeypatch.setattr(comfy.model_management, "current_loaded_models", [loaded])
|
||||
|
||||
del derived_clip
|
||||
gc.collect()
|
||||
with caplog.at_level(logging.INFO):
|
||||
comfy.model_management.cleanup_models_gc()
|
||||
|
||||
assert loaded.model is source_clip.patcher
|
||||
assert loaded.is_dead() is False
|
||||
assert "Potential memory leak detected" not in caplog.text
|
||||
assert "WARNING, memory leak" not in caplog.text
|
||||
|
||||
caplog.clear()
|
||||
released_encoder = AnimaTEModel_()
|
||||
released_source = _AnimaClip(_patcher(released_encoder), released_encoder)
|
||||
released_clip = ComfyPatcherLifecycle().derive_clip(
|
||||
released_source,
|
||||
(ClipLayerMutation(-2),),
|
||||
operation="released Anima CLIP regression",
|
||||
)
|
||||
released = LoadedModel(released_clip.patcher)
|
||||
released.real_model = weakref.ref(released_encoder)
|
||||
monkeypatch.setattr(comfy.model_management, "current_loaded_models", [released])
|
||||
|
||||
del released_clip, released_source, released_encoder
|
||||
gc.collect()
|
||||
with caplog.at_level(logging.INFO):
|
||||
comfy.model_management.cleanup_models_gc()
|
||||
|
||||
assert released.model is None
|
||||
assert released.real_model() is None
|
||||
assert released.is_dead() is False
|
||||
assert "Potential memory leak detected" not in caplog.text
|
||||
assert "WARNING, memory leak" not in caplog.text
|
||||
|
||||
|
||||
def _assert_supported_model_mutations_share_one_clone() -> None:
|
||||
"""All supported first-party MODEL changes share one verified derivation."""
|
||||
|
||||
source = _patcher(torch.nn.Linear(1, 1))
|
||||
|
||||
def denoise_mask(*args: object, **kwargs: object) -> object:
|
||||
"""Return a stable test sentinel."""
|
||||
|
||||
del args, kwargs
|
||||
return object()
|
||||
|
||||
def model_wrapper(args: object) -> object:
|
||||
"""Return the supplied model-wrapper arguments."""
|
||||
|
||||
return args
|
||||
|
||||
def calc_cond_batch(args: object) -> object:
|
||||
"""Return the supplied calc-cond-batch arguments."""
|
||||
|
||||
return args
|
||||
|
||||
derived = ComfyPatcherLifecycle().derive_model(
|
||||
source,
|
||||
(
|
||||
ModelDenoiseMaskMutation(denoise_mask),
|
||||
ModelUnetWrapperMutation(model_wrapper),
|
||||
ModelCalcCondBatchMutation(calc_cond_batch),
|
||||
),
|
||||
operation="MODEL mutation regression",
|
||||
)
|
||||
|
||||
assert derived.parent is source
|
||||
assert derived.model_options["denoise_mask_function"] is denoise_mask
|
||||
assert derived.model_options["model_function_wrapper"] is model_wrapper
|
||||
assert derived.model_options["sampler_calc_cond_batch_function"] is calc_cond_batch
|
||||
|
||||
|
||||
def test_lifecycle_rejects_a_clone_without_comfy_parent_lineage() -> None:
|
||||
"""A non-Comfy clone cannot silently enter a first-party lifecycle path."""
|
||||
|
||||
class BrokenModel:
|
||||
"""Return an unrelated object from clone()."""
|
||||
|
||||
def clone(self) -> BrokenModel:
|
||||
"""Return a clone without a parent link."""
|
||||
|
||||
return BrokenModel()
|
||||
|
||||
with pytest.raises(RuntimeError, match="without its source as parent"):
|
||||
ComfyPatcherLifecycle().derive_model(
|
||||
BrokenModel(),
|
||||
(),
|
||||
operation="broken regression",
|
||||
)
|
||||
|
||||
|
||||
def test_lifecycle_preserves_vae_identity() -> None:
|
||||
"""The first-party VAE contract never clones or mutates the supplied value."""
|
||||
|
||||
vae = object()
|
||||
|
||||
result = ComfyPatcherLifecycle().preserve_vae(
|
||||
vae,
|
||||
operation="VAE regression",
|
||||
)
|
||||
|
||||
assert result is vae
|
||||
|
||||
|
||||
def _patcher(model: torch.nn.Module) -> Any:
|
||||
"""Create a CPU patcher without loading model weights."""
|
||||
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
device = torch.device("cpu")
|
||||
return ModelPatcher(model, load_device=device, offload_device=device)
|
||||
@@ -323,17 +323,22 @@ def _controls(
|
||||
class _FakeModel:
|
||||
"""Provide the ModelPatcher surface used by the semantic runtime."""
|
||||
|
||||
def __init__(self, model_options: dict[str, Any] | None = None) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
model_options: dict[str, Any] | None = None,
|
||||
parent: _FakeModel | None = None,
|
||||
) -> None:
|
||||
"""Create a CPU-backed fake model patcher."""
|
||||
|
||||
self.load_device = torch.device("cpu")
|
||||
self.model_options = {} if model_options is None else model_options
|
||||
self.wrapper: object | None = None
|
||||
self.parent = parent
|
||||
|
||||
def clone(self) -> _FakeModel:
|
||||
"""Return a clone with copied model options."""
|
||||
|
||||
return _FakeModel(self.model_options.copy())
|
||||
return _FakeModel(self.model_options.copy(), parent=self)
|
||||
|
||||
def set_model_unet_function_wrapper(self, wrapper: object) -> None:
|
||||
"""Capture the installed wrapper."""
|
||||
|
||||
@@ -0,0 +1,200 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Enforce the single first-party Comfy patcher lifecycle boundary."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
SOURCE_ROOT = PROJECT_ROOT / "simple_syrup"
|
||||
LIFECYCLE_MODULE = "simple_syrup/runtime/patcher_lifecycle.py"
|
||||
FORBIDDEN_PATCHER_CALLS = frozenset(
|
||||
{
|
||||
"add_object_patch",
|
||||
"add_patches",
|
||||
"clip_layer",
|
||||
"register_all_hook_patches",
|
||||
"set_tokenizer_option",
|
||||
}
|
||||
)
|
||||
FORBIDDEN_PATCHER_WRITES = frozenset({"forced_hooks", "use_clip_schedule"})
|
||||
APPROVED_VALUE_CLONES = Counter(
|
||||
{
|
||||
("simple_syrup/domain/segs_tiled_diffusion.py", "mask"): 2,
|
||||
("simple_syrup/image/crop_composite.py", "image"): 1,
|
||||
(
|
||||
"simple_syrup/image/resize_service.py",
|
||||
"values.expand(cropped.shape[0], cropped.shape[1], "
|
||||
"plan.output_height, plan.output_width)",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/masking/prompt_segs_with_sam_service.py",
|
||||
"crop_image(image_tensor, crop_region).detach()",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/masking/prompt_segs_with_sam_service.py",
|
||||
"crop_mask(final_mask, crop_region).detach()",
|
||||
): 1,
|
||||
("simple_syrup/runtime/sam_region_overlay_renderer.py", "image.detach()"): 2,
|
||||
(
|
||||
"simple_syrup/services/detail_segs_as_regions_service.py",
|
||||
"image_tensor",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/services/detail_segs_by_scale_factor_service.py",
|
||||
"image_tensor",
|
||||
): 2,
|
||||
(
|
||||
"simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py",
|
||||
"image_tensor",
|
||||
): 2,
|
||||
(
|
||||
"simple_syrup/services/mask_to_segs_service.py",
|
||||
"crop_image(image_tensor, crop_region).detach()",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/services/mask_to_segs_service.py",
|
||||
"cropped_segment_mask",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/services/segs_detection_service.py",
|
||||
"crop_mask(mask, crop_region).detach()",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/services/segs_detection_service.py",
|
||||
"crop_image(image_tensor, crop_region).detach()",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/services/segs_from_sam_output_service.py",
|
||||
"image[:, crop_region.top:crop_region.bottom, "
|
||||
"crop_region.left:crop_region.right, :].detach()",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/services/segs_from_sam_output_service.py",
|
||||
"local_mask.unsqueeze(0).detach()",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/services/segs_output_service.py",
|
||||
"crop_mask(combined_mask, crop_region).detach()",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/services/segs_output_service.py",
|
||||
"crop_image(image_tensor, crop_region).detach()",
|
||||
): 1,
|
||||
(
|
||||
"simple_syrup/services/simple_preview_segs_service.py",
|
||||
"image.detach().cpu()",
|
||||
): 1,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_first_party_patcher_lifecycle_has_one_authoritative_owner() -> None:
|
||||
"""Reject every first-party clone or patcher mutation outside the owner."""
|
||||
|
||||
observed_clones: Counter[tuple[str, str]] = Counter()
|
||||
violations: list[str] = []
|
||||
for path in sorted(SOURCE_ROOT.rglob("*.py")):
|
||||
relative_path = path.relative_to(PROJECT_ROOT).as_posix()
|
||||
if "third_party" in path.parts or relative_path == LIFECYCLE_MODULE:
|
||||
continue
|
||||
source = path.read_text(encoding="utf-8")
|
||||
tree = ast.parse(source, filename=str(path))
|
||||
file_violations, clones = _lifecycle_bypasses(tree, relative_path)
|
||||
violations.extend(file_violations)
|
||||
observed_clones.update(clones)
|
||||
|
||||
unexpected_clones = observed_clones - APPROVED_VALUE_CLONES
|
||||
missing_clones = APPROVED_VALUE_CLONES - observed_clones
|
||||
assert violations == []
|
||||
assert unexpected_clones == Counter()
|
||||
assert missing_clones == Counter()
|
||||
|
||||
|
||||
def test_policy_detects_a_new_direct_patcher_code_path() -> None:
|
||||
"""Prove the policy rejects clone, mutation, and object-patch bypasses."""
|
||||
|
||||
source = """
|
||||
def unsafe(model):
|
||||
cloned = model.clone()
|
||||
cloned.set_model_unet_function_wrapper(lambda args: args)
|
||||
cloned.add_object_patch("encode", model.encode)
|
||||
cloned.forced_hooks = object()
|
||||
return cloned
|
||||
"""
|
||||
violations, clones = _lifecycle_bypasses(
|
||||
ast.parse(source),
|
||||
"simple_syrup/runtime/new_feature.py",
|
||||
)
|
||||
|
||||
assert len(violations) == 3
|
||||
assert clones == Counter({("simple_syrup/runtime/new_feature.py", "model"): 1})
|
||||
|
||||
|
||||
def _lifecycle_bypasses(
|
||||
tree: ast.AST,
|
||||
relative_path: str,
|
||||
) -> tuple[list[str], Counter[tuple[str, str]]]:
|
||||
"""Return direct patcher mutations and all clone callsites in one tree."""
|
||||
|
||||
violations: list[str] = []
|
||||
clones: Counter[tuple[str, str]] = Counter()
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Attribute):
|
||||
if node.attr == "clone":
|
||||
clones[(relative_path, ast.unparse(node.value))] += 1
|
||||
elif _is_forbidden_patcher_call(node.attr):
|
||||
violations.append(f"{relative_path}:{node.lineno}: {node.attr}")
|
||||
elif isinstance(node, ast.Call) and _uses_forbidden_dynamic_attribute(node):
|
||||
violations.append(f"{relative_path}:{node.lineno}: dynamic patcher access")
|
||||
elif isinstance(node, (ast.Assign, ast.AnnAssign, ast.AugAssign)):
|
||||
for target in _assignment_targets(node):
|
||||
if (
|
||||
isinstance(target, ast.Attribute)
|
||||
and target.attr in FORBIDDEN_PATCHER_WRITES
|
||||
):
|
||||
violations.append(
|
||||
f"{relative_path}:{node.lineno}: write {target.attr}"
|
||||
)
|
||||
return violations, clones
|
||||
|
||||
|
||||
def _uses_forbidden_dynamic_attribute(node: ast.Call) -> bool:
|
||||
"""Return whether getattr or setattr hides a protected patcher attribute."""
|
||||
|
||||
if not isinstance(node.func, ast.Name) or node.func.id not in {
|
||||
"getattr",
|
||||
"setattr",
|
||||
}:
|
||||
return False
|
||||
if len(node.args) < 2 or not isinstance(node.args[1], ast.Constant):
|
||||
return False
|
||||
attribute_name = node.args[1].value
|
||||
return isinstance(attribute_name, str) and (
|
||||
_is_forbidden_patcher_call(attribute_name)
|
||||
or attribute_name in FORBIDDEN_PATCHER_WRITES
|
||||
)
|
||||
|
||||
|
||||
def _is_forbidden_patcher_call(attribute_name: str) -> bool:
|
||||
"""Return whether an attribute mutates a managed MODEL or CLIP value."""
|
||||
|
||||
return (
|
||||
attribute_name.startswith("set_model_")
|
||||
or attribute_name in FORBIDDEN_PATCHER_CALLS
|
||||
)
|
||||
|
||||
|
||||
def _assignment_targets(
|
||||
node: ast.Assign | ast.AnnAssign | ast.AugAssign,
|
||||
) -> tuple[ast.expr, ...]:
|
||||
"""Normalize assignment node targets for lifecycle-policy inspection."""
|
||||
|
||||
if isinstance(node, ast.Assign):
|
||||
return tuple(node.targets)
|
||||
return (node.target,)
|
||||
@@ -19,13 +19,14 @@ from simple_syrup.runtime.regional_lora_hooks import prepare_regional_lora_clip
|
||||
class _FakeClip:
|
||||
"""Record whether regional preparation clones the text encoder."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, parent_patcher: object | None = None) -> None:
|
||||
"""Create a clip and its minimal patcher collaboration."""
|
||||
|
||||
self.cond_stage_model = object()
|
||||
self.registrations: list[tuple[Any, Any]] = []
|
||||
self.patcher = SimpleNamespace(
|
||||
forced_hooks=None,
|
||||
parent=parent_patcher,
|
||||
register_all_hook_patches=self._register_hooks,
|
||||
)
|
||||
self.use_clip_schedule = False
|
||||
@@ -35,7 +36,7 @@ class _FakeClip:
|
||||
"""Return a distinct clip while recording clone policy."""
|
||||
|
||||
self.clone_calls.append(disable_dynamic)
|
||||
clone = _FakeClip()
|
||||
clone = _FakeClip(parent_patcher=self.patcher)
|
||||
clone.clone_calls = self.clone_calls
|
||||
clone.registrations = self.registrations
|
||||
clone.patcher.register_all_hook_patches = clone._register_hooks
|
||||
|
||||
Reference in New Issue
Block a user