feat: effective_model_sampling helper — patcher-resolved sampling (KSampler semantics)

This commit is contained in:
WildAi
2026-09-17 22:01:49 +03:00
parent 5cc26ee9b5
commit 3bf3045d99
2 changed files with 199 additions and 0 deletions
+73
View File
@@ -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
+126
View File
@@ -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