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:
Reithan
2026-08-13 20:46:47 -07:00
committed by GitHub
parent 07f5ea0dbe
commit ffeb7fe4f3
2 changed files with 127 additions and 87 deletions
+16 -87
View File
@@ -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)
+111
View File
@@ -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