PR-1: cleanup — remove mangled guard, amputate dead operation-space arms, add detection tests (#34)
## Summary Behavior-preserving cleanup ahead of the H3 flow-fix work (see `.claude/plans/nrs_h3_fix_plan.2.md`, fact 8 / Phase 3 / Sequencing PR-1). No output change — pure refactor plus a regression net. ### Changes to `NRS/nodes_NRS.py` - **Removed the broken name-mangling guard.** `hasattr(self, "__pred_type")` checked the literal name while assignment created `_NRS__pred_type`, so the guard never fired. `patch()` now computes `pred_type = self._get_pred_type(model)` unconditionally and passes it as a closure local / explicit parameter into `nrs()`, `_convert_to_v_space`, and `_finalize_from_v_space`. This preserves the (correct) always-redetect behavior and removes a latent cross-model aliasing bug. The `hasattr` string was **not** "repaired" — doing so would introduce a stale-cache bug since ComfyUI reuses node instances across queue runs. - **Amputated dead operation-space arms.** `self.__OPERATION_SPACE` was hardcoded to `PredictionType.V`, making both `match` blocks and the `_convert_to_eps_space`/`_finalize_from_eps_space` helpers unreachable. Removed them; the `nrs()` hook now calls the V-space conversion/finalization directly. The inner V-vs-EPS conversion math inside the V-space helpers is untouched — that is the real per-model algebra, not dead code. FLOW is added in a later PR. ### Tests - New `tests/test_pred_type.py`: parametrized coverage of all 11 `_RAW_TO_ENUM` entries plus `_get_pred_type` walks over stub models (direct-attribute path and enhanced-detection fallback). These pin **current** behavior (flow-family → EPS) as a regression net for the FLOW reclassification PR. ## Verification - `24/24` tests pass (21 new + 3 pre-existing smoke) from a clean checkout. - Greps for `self.__pred_type`, `_NRS__pred_type`, `self.__OPERATION_SPACE` all empty. - Diff: `NRS/nodes_NRS.py` +16/-87, `tests/test_pred_type.py` +111 — behavior-preserving, well under 500 LoC.
This commit is contained in:
+16
-87
@@ -159,57 +159,14 @@ class NRS:
|
||||
)
|
||||
return PredictionType.EPS
|
||||
|
||||
def _convert_to_eps_space(self, x_orig, sig_root, sigma, cond, uncond):
|
||||
x_div = None
|
||||
eps_cond = cond
|
||||
eps_uncond = uncond
|
||||
if self.__pred_type == PredictionType.V:
|
||||
# v → ε conversion
|
||||
logging.debug("NRS._convert_to_eps_space: generating x_div, eps_cond, and eps_uncond for v-pred")
|
||||
x_div = x_orig / (sigma**2 + 1)
|
||||
|
||||
eps_cond = ((x_div - (x_orig - cond)) * sig_root) / (sigma)
|
||||
eps_uncond = ((x_div - (x_orig - uncond)) * sig_root) / (sigma)
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
logging.debug("NRS._convert_to_eps_space: already in eps, no pre-scale needed")
|
||||
pass # already in ε space
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._convert_to_eps_space: x0-prediction not supported yet.")
|
||||
else:
|
||||
# Fallback: treat UNKNOWN as EPS (should not happen with enhanced detection)
|
||||
logging.warning(f"NRS._convert_to_eps_space: Unknown prediction type {self.__pred_type}, treating as EPS")
|
||||
pass
|
||||
|
||||
return x_div, eps_cond, eps_uncond
|
||||
|
||||
def _finalize_from_eps_space(self, x_orig, x_div, x_final, sig_root, sigma):
|
||||
nrs_result = x_final
|
||||
if self.__pred_type == PredictionType.V:
|
||||
# ε → v conversion
|
||||
logging.debug("NRS._finalize_from_eps_space: generating cfg_result for v-pred")
|
||||
nrs_result = x_orig - (x_div - x_final * sigma / sig_root)
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
# already in ε space
|
||||
logging.debug("NRS._finalize_from_eps_space: already in eps, no post-scale needed")
|
||||
pass
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._finalize_from_eps_space: x0-prediction not supported yet.")
|
||||
else:
|
||||
# Fallback: treat UNKNOWN as EPS (should not happen with enhanced detection)
|
||||
logging.warning(
|
||||
f"NRS._finalize_from_eps_space: Unknown prediction type {self.__pred_type}, treating as EPS"
|
||||
)
|
||||
pass
|
||||
return nrs_result
|
||||
|
||||
def _convert_to_v_space(self, x_orig, sig_root, sigma, cond, uncond):
|
||||
def _convert_to_v_space(self, x_orig, sig_root, sigma, cond, uncond, pred_type):
|
||||
x_div = None
|
||||
v_cond = cond
|
||||
v_uncond = uncond
|
||||
if self.__pred_type == PredictionType.V:
|
||||
if pred_type == PredictionType.V:
|
||||
logging.debug("NRS._convert_to_v_space: already in v, no pre-scale needed")
|
||||
pass # already in v space
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
elif pred_type == PredictionType.EPS:
|
||||
# ε → v conversion
|
||||
logging.debug("NRS._convert_to_v_space: generating x_div, v_cond, and v_uncond for eps")
|
||||
x_div = x_orig / (sigma**2 + 1)
|
||||
@@ -217,11 +174,11 @@ class NRS:
|
||||
|
||||
v_cond = x_orig - (x_div - cond * factor)
|
||||
v_uncond = x_orig - (x_div - uncond * factor)
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
elif pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._convert_to_v_space: x0-prediction not supported yet.")
|
||||
else:
|
||||
# Fallback: treat UNKNOWN as EPS and convert to V-space
|
||||
logging.warning(f"NRS._convert_to_v_space: Unknown prediction type {self.__pred_type}, treating as EPS")
|
||||
logging.warning(f"NRS._convert_to_v_space: Unknown prediction type {pred_type}, treating as EPS")
|
||||
logging.debug("NRS._convert_to_v_space: generating x_div, v_cond, and v_uncond for eps (fallback)")
|
||||
x_div = x_orig / (sigma**2 + 1)
|
||||
factor = sigma / sig_root
|
||||
@@ -230,32 +187,30 @@ class NRS:
|
||||
|
||||
return x_div, v_cond, v_uncond
|
||||
|
||||
def _finalize_from_v_space(self, x_orig, x_div, x_final, sig_root, sigma):
|
||||
def _finalize_from_v_space(self, x_orig, x_div, x_final, sig_root, sigma, pred_type):
|
||||
nrs_result = x_final
|
||||
if self.__pred_type == PredictionType.V:
|
||||
if pred_type == PredictionType.V:
|
||||
# already in v space
|
||||
logging.debug("NRS._finalize_from_v_space: already in v, no post-scale needed")
|
||||
pass
|
||||
elif self.__pred_type == PredictionType.EPS:
|
||||
elif pred_type == PredictionType.EPS:
|
||||
# v → ε conversion
|
||||
logging.debug("NRS._finalize_from_v_space: generating cfg_result for eps")
|
||||
nrs_result = (x_div - (x_orig - x_final)) * (sig_root / sigma)
|
||||
elif self.__pred_type == PredictionType.X0:
|
||||
elif pred_type == PredictionType.X0:
|
||||
raise NotImplementedError("NRS._finalize_from_v_space: x0-prediction not supported yet.")
|
||||
else:
|
||||
# Fallback: treat UNKNOWN as EPS and convert from V-space
|
||||
logging.warning(f"NRS._finalize_from_v_space: Unknown prediction type {self.__pred_type}, treating as EPS")
|
||||
logging.warning(f"NRS._finalize_from_v_space: Unknown prediction type {pred_type}, treating as EPS")
|
||||
logging.debug("NRS._finalize_from_v_space: generating cfg_result for eps (fallback)")
|
||||
nrs_result = (x_div - (x_orig - x_final)) * (sig_root / sigma)
|
||||
return nrs_result
|
||||
|
||||
def patch(self, model, skew, stretch, squash):
|
||||
self.__pred_type = self._get_pred_type(model) if not hasattr(self, "__pred_type") else self.__pred_type
|
||||
self.__OPERATION_SPACE = PredictionType.V
|
||||
pred_type = self._get_pred_type(model)
|
||||
|
||||
def nrs(args):
|
||||
logging.debug(f"NRS.nrs: Skew: {skew}, Stretch: {stretch}, Squash: {squash}")
|
||||
# self.__pred_type = self.__pred_type if self.__pred_type is not None else self._get_pred_type(model)
|
||||
cond = args["cond"]
|
||||
uncond = args["uncond"]
|
||||
x_orig = args["input"]
|
||||
@@ -264,22 +219,10 @@ class NRS:
|
||||
sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1))
|
||||
sig_root = (sigma**2 + 1).sqrt()
|
||||
|
||||
x_div, nrs_cond, nrs_uncond = None, None, None
|
||||
match self.__OPERATION_SPACE:
|
||||
case PredictionType.V:
|
||||
x_div, nrs_cond, nrs_uncond = self._convert_to_v_space(x_orig, sig_root, sigma, cond, uncond)
|
||||
case PredictionType.EPS:
|
||||
x_div, nrs_cond, nrs_uncond = self._convert_to_eps_space(x_orig, sig_root, sigma, cond, uncond)
|
||||
case PredictionType.X0:
|
||||
raise RuntimeError("NRS.nrs: x0-prediction not supported yet.")
|
||||
case PredictionType.UNKNOWN:
|
||||
# Fallback: treat as EPS (should not happen with enhanced detection)
|
||||
logging.warning("NRS.nrs: Unknown operation space, treating as EPS for conversion")
|
||||
x_div, nrs_cond, nrs_uncond = self._convert_to_eps_space(x_orig, sig_root, sigma, cond, uncond)
|
||||
case _:
|
||||
raise RuntimeError(
|
||||
f"NRS.nrs: Invalid PredictionType used for operation space: {self.__OPERATION_SPACE}"
|
||||
)
|
||||
# Operation space is hardcoded to V for now; FLOW is added in a later PR.
|
||||
x_div, nrs_cond, nrs_uncond = self._convert_to_v_space(
|
||||
x_orig, sig_root, sigma, cond, uncond, pred_type
|
||||
)
|
||||
|
||||
def _dot(a, b):
|
||||
return (a * b).sum(dim=1, keepdim=True) # [B,C,W,H] => [B,1,W,H]
|
||||
@@ -307,21 +250,7 @@ class NRS:
|
||||
squash_scale = (1 - squash) + (squash * (cond_len / nrs_len))
|
||||
x_final = skewed * squash_scale
|
||||
|
||||
match self.__OPERATION_SPACE:
|
||||
case PredictionType.V:
|
||||
return self._finalize_from_v_space(x_orig, x_div, x_final, sig_root, sigma)
|
||||
case PredictionType.EPS:
|
||||
return self._finalize_from_eps_space(x_orig, x_div, x_final, sig_root, sigma)
|
||||
case PredictionType.X0:
|
||||
raise RuntimeError("NRS.nrs: x0-prediction not supported yet.")
|
||||
case PredictionType.UNKNOWN:
|
||||
# Fallback: treat as EPS (should not happen with enhanced detection)
|
||||
logging.warning("NRS.nrs: Unknown operation space, treating as EPS for finalization")
|
||||
return self._finalize_from_eps_space(x_orig, x_div, x_final, sig_root, sigma)
|
||||
case _:
|
||||
raise RuntimeError(
|
||||
f"NRS.nrs: Invalid PredictionType used for operation space: {self.__OPERATION_SPACE}"
|
||||
)
|
||||
return self._finalize_from_v_space(x_orig, x_div, x_final, sig_root, sigma, pred_type)
|
||||
|
||||
m = model.clone()
|
||||
m.set_model_sampler_cfg_function(nrs, True)
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
"""Regression tests for NRS._get_pred_type and _RAW_TO_ENUM mappings.
|
||||
|
||||
These tests pin CURRENT behavior (flow-family names resolve to EPS) as a
|
||||
safety net ahead of the FLOW reclassification planned for a later PR. If
|
||||
this file needs updating because flow-family names now map to
|
||||
PredictionType.FLOW, that is expected -- it means the reclassification
|
||||
landed and this net did its job.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
# Add project root to path for imports
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
from NRS.nodes_NRS import PredictionType, _RAW_TO_ENUM, NRS
|
||||
|
||||
|
||||
def _make_model_sampling(class_name):
|
||||
"""Build an instance whose type name matches class_name, for class-name fingerprinting."""
|
||||
return type(class_name, (object,), {})()
|
||||
|
||||
|
||||
class _StubModel:
|
||||
"""Minimal stand-in for a model object walked by _get_pred_type."""
|
||||
|
||||
def __init__(self, model_type=None, model_sampling=None, inner_model_type=None):
|
||||
if model_type is not None:
|
||||
self.model_type = model_type
|
||||
if model_sampling is not None:
|
||||
self.model_sampling = model_sampling
|
||||
if inner_model_type is not None:
|
||||
# Emulates model.model.model_type used by the enhanced-detection fallback.
|
||||
self.model = _StubModel(model_type=inner_model_type)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"raw,expected",
|
||||
[
|
||||
("eps", PredictionType.EPS),
|
||||
("epsilon", PredictionType.EPS),
|
||||
("flux", PredictionType.EPS),
|
||||
("chroma", PredictionType.EPS),
|
||||
("flow", PredictionType.EPS),
|
||||
("wan", PredictionType.EPS),
|
||||
("const", PredictionType.EPS),
|
||||
("v", PredictionType.V),
|
||||
("v_prediction", PredictionType.V),
|
||||
("x0", PredictionType.X0),
|
||||
("sample", PredictionType.X0),
|
||||
],
|
||||
)
|
||||
def test_raw_to_enum_mapping(raw, expected):
|
||||
"""Pin the current _RAW_TO_ENUM dict mappings."""
|
||||
assert _RAW_TO_ENUM[raw] == expected
|
||||
|
||||
|
||||
def test_raw_to_enum_unknown_raw_not_present():
|
||||
"""Unrecognized raw strings are not in the dict; callers fall back to UNKNOWN."""
|
||||
assert "totally-unrecognized" not in _RAW_TO_ENUM
|
||||
|
||||
|
||||
class TestGetPredTypeDirectAttribute:
|
||||
"""_get_pred_type's direct-hit path via model_type/prediction_type/parameterization."""
|
||||
|
||||
def test_model_type_v_prediction(self):
|
||||
node = NRS()
|
||||
model = _StubModel(model_type="v_prediction")
|
||||
assert node._get_pred_type(model) == PredictionType.V
|
||||
|
||||
def test_model_type_eps(self):
|
||||
node = NRS()
|
||||
model = _StubModel(model_type="eps")
|
||||
assert node._get_pred_type(model) == PredictionType.EPS
|
||||
|
||||
def test_model_type_x0(self):
|
||||
node = NRS()
|
||||
model = _StubModel(model_type="x0")
|
||||
assert node._get_pred_type(model) == PredictionType.X0
|
||||
|
||||
def test_model_type_flow_family_is_currently_eps(self):
|
||||
"""Flow-family models currently resolve to EPS (pre-reclassification)."""
|
||||
model = _StubModel(model_type="flow")
|
||||
assert NRS()._get_pred_type(model) == PredictionType.EPS
|
||||
|
||||
def test_model_type_wan_is_currently_eps(self):
|
||||
model = _StubModel(model_type="wan")
|
||||
assert NRS()._get_pred_type(model) == PredictionType.EPS
|
||||
|
||||
|
||||
class TestGetPredTypeEnhancedDetectionFallback:
|
||||
"""The model_sampling class-name and model.model.model_type fallback paths."""
|
||||
|
||||
def test_model_sampling_const_class_is_eps(self):
|
||||
model = _StubModel(model_sampling=_make_model_sampling("ModelSamplingContinuousEDMConst"))
|
||||
assert NRS()._get_pred_type(model) == PredictionType.EPS
|
||||
|
||||
def test_model_sampling_v_prediction_class_is_v(self):
|
||||
model = _StubModel(model_sampling=_make_model_sampling("ModelSamplingV_Prediction"))
|
||||
assert NRS()._get_pred_type(model) == PredictionType.V
|
||||
|
||||
def test_model_sampling_eps_class_is_eps(self):
|
||||
model = _StubModel(model_sampling=_make_model_sampling("ModelSamplingEps"))
|
||||
assert NRS()._get_pred_type(model) == PredictionType.EPS
|
||||
|
||||
def test_unrecognized_model_defaults_to_eps(self):
|
||||
"""Fully-unrecognized models fall back to EPS (documented default)."""
|
||||
model = _StubModel()
|
||||
assert NRS()._get_pred_type(model) == PredictionType.EPS
|
||||
Reference in New Issue
Block a user