diff --git a/src/effective_sampling.py b/src/effective_sampling.py new file mode 100644 index 0000000..b9d98c2 --- /dev/null +++ b/src/effective_sampling.py @@ -0,0 +1,73 @@ +"""Effective model_sampling resolution through a ModelPatcher (v2.16.0). + +The node layer must never read ``model.model.model_sampling`` (the LIVE +attribute on the shared BaseModel) to derive sigma schedules, timestep +conversions or prediction types. ComfyUI's object-patch lifecycle makes that +attribute history-dependent: + +- object patches are applied to the shared BaseModel at load and stay applied + between runs (model_patcher.py patch_model / partially_load); +- ``load_models_gpu`` detaches stale same-base patchers with + ``detach(unpatch_all=False)`` (model_management.py:962) — their patches LEAK; +- restore-to-original only happens through ``object_patches_backup``, which + every ``unpatch_model`` CLEARS (model_patcher.py:1165-1169); +- ``partially_unload`` moves weights only — patches and backup untouched. + +Whichever patcher (a DyPE/SEGA clone, a stock ModelSamplingFlux node clone, +...) was loaded last therefore decides the schedule a live-attr reader sees. +ComfyUI's own sampler resolves through the patcher instead — +``comfy/samplers.py:1425``: ``calculate_sigmas(self.model.get_model_object( +"model_sampling"), ...)``. :func:`effective_model_sampling` mirrors exactly +that semantics (object_patches -> object_patches_backup -> live attr), making +the schedule a deterministic function of the graph's own patch chain. + +:func:`is_stale_dype_leak` detects the residual case this cannot repair: the +patcher carries no ``model_sampling`` patch/backup entry while the live attr is +one of this pack's function-local leak classes — i.e. a stale patch inherited +from a PREVIOUS run's patch node that is no longer in the graph. +""" + +from __future__ import annotations + +# Function-local classes installed by apply_dype_to_model / apply_sega_to_model +# (src/patch_utils.py). Matched by __name__: the classes are defined inside the +# installer functions, so identity comparison across imports is impossible. +STALE_LEAK_CLASS_NAMES = ( + "DypeModelSamplingFlux", + "SegaModelSamplingFlux", + "DefaultModelSamplingFlux", +) + + +def effective_model_sampling(model): + """Resolve the model_sampling object the way ComfyUI's sampler does. + + ``model`` is a ModelPatcher (or a test mock). Resolution order mirrors + ``ModelPatcher.get_model_object`` (model_patcher.py:758-768): the patcher's + own object patch, then its object_patches_backup (the original captured at + patch time), then the live BaseModel attribute. Plain objects without the + method fall back to the live attribute (mock/test safety). + """ + get_model_object = getattr(model, "get_model_object", None) + if callable(get_model_object): + return get_model_object("model_sampling") + return getattr(getattr(model, "model", None), "model_sampling", None) + + +def is_stale_dype_leak(model) -> bool: + """True when the live sampling is a stale patch from a previous run. + + Stale means: the patcher carries NO ``model_sampling`` object patch and NO + backup entry (so the resolution falls through to the live attribute), and + the live attribute's class is one of this pack's installer-local sampling + patches. Only then is the schedule inherited from a run that is no longer + in the graph — the user-fixable-by-cache-clear drift this pack warns about. + """ + if not callable(getattr(model, "get_model_object", None)): + return False + if "model_sampling" in (getattr(model, "object_patches", None) or {}): + return False + if "model_sampling" in (getattr(model, "object_patches_backup", None) or {}): + return False + live = getattr(getattr(model, "model", None), "model_sampling", None) + return type(live).__name__ in STALE_LEAK_CLASS_NAMES diff --git a/tests/test_effective_sampling.py b/tests/test_effective_sampling.py new file mode 100644 index 0000000..4aed35e --- /dev/null +++ b/tests/test_effective_sampling.py @@ -0,0 +1,126 @@ +"""Tests for src/effective_sampling.py (v2.16.0, plan S1). + +Covers the resolution order (patch -> backup -> live), the mock fallback and +the stale-leak detector, using fake patchers that replicate the ModelPatcher +get_model_object semantics (object_patches -> object_patches_backup -> live +attr; comfy/model_patcher.py:758-768). +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import pytest + +sys.path.insert(0, str(Path(__file__).parent.parent)) + +from src.effective_sampling import ( # noqa: E402 + STALE_LEAK_CLASS_NAMES, + effective_model_sampling, + is_stale_dype_leak, +) + + +class FakeBaseModel: + def __init__(self, model_sampling): + self.model_sampling = model_sampling + + +class FakePatcher: + """Minimal ModelPatcher stand-in with the real get_model_object order.""" + + def __init__(self, live, patches=None, backup=None): + self.model = FakeBaseModel(live) + self.object_patches = dict(patches or {}) + self.object_patches_backup = dict(backup or {}) + + def get_model_object(self, name): + if name in self.object_patches: + return self.object_patches[name] + if name in self.object_patches_backup: + return self.object_patches_backup[name] + return getattr(self.model, name) + + +class PlainNode: + """Object WITHOUT get_model_object (the mock-fallback path).""" + + def __init__(self, live): + self.model = FakeBaseModel(live) + + +def _make_leak_class(name: str) -> type: + return type(name, (), {}) + + +ORIG = _make_leak_class("ModelSamplingContinuousFlow")() +LEAK = _make_leak_class("DypeModelSamplingFlux")() +PATCH = _make_leak_class("DypeModelSamplingFlux")() + + +class TestEffectiveModelSampling: + def test_own_patch_wins_over_leaked_live(self): + patcher = FakePatcher(LEAK, patches={"model_sampling": PATCH}) + assert effective_model_sampling(patcher) is PATCH + + def test_backup_used_when_no_patch(self): + patcher = FakePatcher(LEAK, backup={"model_sampling": ORIG}) + assert effective_model_sampling(patcher) is ORIG + + def test_live_fallback_when_patch_and_backup_empty(self): + patcher = FakePatcher(ORIG) + assert effective_model_sampling(patcher) is ORIG + + def test_live_leak_returned_when_nothing_else_available(self): + # Documented residual: patcher without patch/backup sees the leak — + # the same object a stock KSampler would resolve (KSampler semantics). + patcher = FakePatcher(LEAK) + assert effective_model_sampling(patcher) is LEAK + + def test_plain_object_fallback(self): + node = PlainNode(ORIG) + assert effective_model_sampling(node) is ORIG + + def test_plain_object_without_model_attr_returns_none(self): + class Empty: + pass + + assert effective_model_sampling(Empty()) is None + + +class TestIsStaleDypeLeak: + def test_true_for_leak_class_with_empty_patcher(self): + assert is_stale_dype_leak(FakePatcher(LEAK)) is True + + @pytest.mark.parametrize("name", STALE_LEAK_CLASS_NAMES) + def test_true_for_every_known_leak_class(self, name): + leak = _make_leak_class(name)() + assert is_stale_dype_leak(FakePatcher(leak)) is True + + def test_false_when_patcher_carries_own_patch(self): + # A DyPE clone in THIS graph legitimately patches model_sampling — + # the same class name live is not stale there. + patcher = FakePatcher(PATCH, patches={"model_sampling": PATCH}) + assert is_stale_dype_leak(patcher) is False + + def test_false_when_backup_holds_the_key(self): + patcher = FakePatcher(LEAK, backup={"model_sampling": ORIG}) + assert is_stale_dype_leak(patcher) is False + + def test_false_for_foreign_live_class(self): + assert is_stale_dype_leak(FakePatcher(ORIG)) is False + + def test_false_without_get_model_object(self): + # Mock safety: no resolution contract -> never claim a leak. + assert is_stale_dype_leak(PlainNode(LEAK)) is False + + def test_false_when_live_attr_missing(self): + class BarePatcher: + object_patches = {} + object_patches_backup = {} + + def get_model_object(self, name): + return None + + assert is_stale_dype_leak(BarePatcher()) is False