feat: effective_model_sampling helper — patcher-resolved sampling (KSampler semantics)
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user