From ffeb7fe4f3c274810d2b3dbe1fc770f55b0af7f0 Mon Sep 17 00:00:00 2001 From: Reithan Date: Thu, 13 Aug 2026 20:46:47 -0700 Subject: [PATCH] =?UTF-8?q?PR-1:=20cleanup=20=E2=80=94=20remove=20mangled?= =?UTF-8?q?=20guard,=20amputate=20dead=20operation-space=20arms,=20add=20d?= =?UTF-8?q?etection=20tests=20(#34)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## 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. --- NRS/nodes_NRS.py | 103 ++++++------------------------------- tests/test_pred_type.py | 111 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 127 insertions(+), 87 deletions(-) create mode 100644 tests/test_pred_type.py diff --git a/NRS/nodes_NRS.py b/NRS/nodes_NRS.py index 36db2a0..c58bfb3 100644 --- a/NRS/nodes_NRS.py +++ b/NRS/nodes_NRS.py @@ -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) diff --git a/tests/test_pred_type.py b/tests/test_pred_type.py new file mode 100644 index 0000000..531e2ff --- /dev/null +++ b/tests/test_pred_type.py @@ -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