diff --git a/NRS/nodes_NRS.py b/NRS/nodes_NRS.py index c58bfb3..b472abd 100644 --- a/NRS/nodes_NRS.py +++ b/NRS/nodes_NRS.py @@ -1,8 +1,41 @@ import logging +import math from enum import Enum, auto import torch +try: + import comfy.utils as _comfy_utils +except Exception: + _comfy_utils = None + + +def _unpack_latents(combined, latent_shapes): + """Split a flat packed latent [B, 1, N] back into its per-stream tensors. + + Mirrors comfy.utils.unpack_latents: for each shape in latent_shapes, take + math.prod(shape[1:]) elements off the last dim and reshape that [B, 1, n] + slice back to `shape`. + """ + streams = [] + offset = 0 + for shape in latent_shapes: + n = math.prod(shape[1:]) + chunk = combined[:, :, offset : offset + n] + streams.append(chunk.reshape(shape)) + offset += n + return streams + + +def _pack_latents(streams): + """Pack a list of per-stream tensors [B, C, ...] into a flat [B, 1, N] tensor. + + Mirrors comfy.utils.pack_latents: each stream is reshaped to (B, 1, -1) + and concatenated on the last dim. + """ + flat = [s.reshape(s.shape[0], 1, -1) for s in streams] + return torch.cat(flat, dim=-1) + # fmt: off class PredictionType(Enum): @@ -206,51 +239,89 @@ class NRS: nrs_result = (x_div - (x_orig - x_final)) * (sig_root / sigma) return nrs_result + def _apply_guidance(self, x_orig, cond, uncond, sigma, skew, stretch, squash, pred_type): + """Run the NRS geometry pipeline on a single (already-unpacked, channels-first) stream.""" + sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1)) + sig_root = (sigma**2 + 1).sqrt() + + # 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] + + def _nrm2(v): + return _dot(v, v) + + eps = torch.finfo(nrs_cond.dtype).eps + c_dot_c = _nrm2(nrs_cond) + eps # [B,1,W,H] + u_dot_c = _dot(nrs_uncond, nrs_cond) # [B,1,W,H] + u_on_c = (u_dot_c / c_dot_c) * nrs_cond # [B,1,W,H] * [B,C,H,W] + + # Amplify Cond based on length compared to projection of uncond + proj_diff = nrs_cond - u_on_c + stretched = nrs_cond + (stretch * proj_diff) + + # Skew/Steer Conf based on rejection of uncond on cond + u_rej_c = nrs_uncond - u_on_c + skewed = stretched - (skew * u_rej_c) + + # Squash final length back down to original length of cond + cond_len = nrs_cond.norm(dim=1, keepdim=True) + nrs_len = skewed.norm(dim=1, keepdim=True) + eps + + squash_scale = (1 - squash) + (squash * (cond_len / nrs_len)) + x_final = skewed * squash_scale + + return self._finalize_from_v_space(x_orig, x_div, x_final, sig_root, sigma, pred_type) + def patch(self, model, skew, stretch, squash): pred_type = self._get_pred_type(model) + warned = {"done": False} def nrs(args): logging.debug(f"NRS.nrs: Skew: {skew}, Stretch: {stretch}, Squash: {squash}") cond = args["cond"] uncond = args["uncond"] x_orig = args["input"] - sigma = args["sigma"] - sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1)) - sig_root = (sigma**2 + 1).sqrt() - # 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 - ) + shapes = getattr(args["model"], "latent_shapes", None) + if shapes and len(shapes) > 1: + if _comfy_utils is not None and hasattr(_comfy_utils, "unpack_latents"): + cond_streams = _comfy_utils.unpack_latents(cond, shapes) + uncond_streams = _comfy_utils.unpack_latents(uncond, shapes) + x_streams = _comfy_utils.unpack_latents(x_orig, shapes) + else: + cond_streams = _unpack_latents(cond, shapes) + uncond_streams = _unpack_latents(uncond, shapes) + x_streams = _unpack_latents(x_orig, shapes) + else: + cond_streams, uncond_streams, x_streams = [cond], [uncond], [x_orig] - def _dot(a, b): - return (a * b).sum(dim=1, keepdim=True) # [B,C,W,H] => [B,1,W,H] + if not warned["done"]: + for stream in x_streams: + if stream.shape[1] == 1: + logging.warning( + f"NRS.nrs: routed stream has a singleton reduction axis {tuple(stream.shape)}; " + "NRS geometry (dot/proj/skew) will degenerate to a no-op on this stream." + ) + warned["done"] = True + break - def _nrm2(v): - return _dot(v, v) + results = [ + self._apply_guidance( + x_streams[i], cond_streams[i], uncond_streams[i], sigma, skew, stretch, squash, pred_type + ) + for i in range(len(cond_streams)) + ] - eps = torch.finfo(nrs_cond.dtype).eps - c_dot_c = _nrm2(nrs_cond) + eps # [B,1,W,H] - u_dot_c = _dot(nrs_uncond, nrs_cond) # [B,1,W,H] - u_on_c = (u_dot_c / c_dot_c) * nrs_cond # [B,1,W,H] * [B,C,H,W] + if len(results) == 1: + return results[0] - # Amplify Cond based on length compared to projection of uncond - proj_diff = nrs_cond - u_on_c - stretched = nrs_cond + (stretch * proj_diff) - - # Skew/Steer Conf based on rejection of uncond on cond - u_rej_c = nrs_uncond - u_on_c - skewed = stretched - (skew * u_rej_c) - - # Squash final length back down to original length of cond - cond_len = nrs_cond.norm(dim=1, keepdim=True) - nrs_len = skewed.norm(dim=1, keepdim=True) + eps - - squash_scale = (1 - squash) + (squash * (cond_len / nrs_len)) - x_final = skewed * squash_scale - - return self._finalize_from_v_space(x_orig, x_div, x_final, sig_root, sigma, pred_type) + if _comfy_utils is not None and hasattr(_comfy_utils, "pack_latents"): + return _comfy_utils.pack_latents(results)[0] + return _pack_latents(results) m = model.clone() m.set_model_sampler_cfg_function(nrs, True) diff --git a/tests/test_pack_split.py b/tests/test_pack_split.py new file mode 100644 index 0000000..8ac85e8 --- /dev/null +++ b/tests/test_pack_split.py @@ -0,0 +1,282 @@ +"""Tests for pack-aware per-stream routing in NRS.nodes_NRS. + +These tests need real torch (tensor math), but tests/conftest.py installs a +MagicMock in sys.modules["torch"] for the whole session so other test modules +can import without the heavy dependency. We swap the real torch module in for +the duration of this module only, then restore the mock so the rest of the +suite is unaffected. +""" + +import sys + +import pytest + +_saved_torch = None +_saved_nodes_nrs = None +torch = None +nrs_module = None + + +def setup_module(module): + # NOTE: we deliberately avoid importlib.reload() here. reload() mutates + # the *existing* NRS.nodes_NRS module dict in place, and other test + # modules (e.g. test_pred_type.py) import PredictionType/NRS at + # collection time and keep those references for the whole session. Their + # methods' __globals__ point at that same dict, so an in-place reload + # would silently swap PredictionType out from under them (new class + # object, same name -> broken identity-based Enum equality). Instead we + # unregister the module from sys.modules and import it fresh: this + # creates an independent module object, leaving the original (still + # cached in other modules' namespaces) untouched. We restore the exact + # original module object on teardown. + global _saved_torch, _saved_nodes_nrs, torch, nrs_module + _saved_torch = sys.modules.get("torch") + sys.modules.pop("torch", None) + try: + import torch as real_torch + except ImportError: + pytest.skip("real torch unavailable", allow_module_level=True) + torch = real_torch + + _saved_nodes_nrs = sys.modules.get("NRS.nodes_NRS") + sys.modules.pop("NRS.nodes_NRS", None) + + import NRS.nodes_NRS as m + + nrs_module = m + + +def teardown_module(module): + if _saved_torch is not None: + sys.modules["torch"] = _saved_torch + else: + sys.modules.pop("torch", None) + + if _saved_nodes_nrs is not None: + sys.modules["NRS.nodes_NRS"] = _saved_nodes_nrs + else: + sys.modules.pop("NRS.nodes_NRS", None) + + +class _StubModelSampling: + """Minimal stand-in that makes _get_pred_type fall back to EPS quickly.""" + + +class _StubInnerModel: + def __init__(self, latent_shapes=None): + self.model_sampling = _StubModelSampling() + if latent_shapes is not None: + self.latent_shapes = latent_shapes + + +class _StubModel: + """Stub for the outer ComfyUI ModelPatcher passed to NRS.patch().""" + + def __init__(self, latent_shapes=None): + self.model = _StubInnerModel(latent_shapes) + self._captured_fn = None + + def clone(self): + return self + + def set_model_sampler_cfg_function(self, fn, flag): + self._captured_fn = fn + + +def _make_args(model, cond, uncond, x_orig, sigma): + return { + "model": model.model, # args["model"] is the inner model carrying latent_shapes + "cond": cond, + "uncond": uncond, + "input": x_orig, + "sigma": sigma, + } + + +# --------------------------------------------------------------------------- +# Phase 1: round-trip pack/unpack correctness +# --------------------------------------------------------------------------- + + +def test_roundtrip_unpack_repack_two_streams(): + video = torch.randn(1, 4, 3, 2) + audio = torch.randn(1, 6, 5) + shapes = [video.shape, audio.shape] + + packed = nrs_module._pack_latents([video, audio]) + assert packed.shape == (1, 1, video.numel() + audio.numel()) + + unpacked = nrs_module._unpack_latents(packed, shapes) + assert len(unpacked) == 2 + assert torch.allclose(unpacked[0], video) + assert torch.allclose(unpacked[1], audio) + + repacked = nrs_module._pack_latents(unpacked) + assert torch.allclose(repacked, packed) + + +def test_roundtrip_single_stream(): + x = torch.randn(1, 4, 8, 8) + shapes = [x.shape] + + packed = nrs_module._pack_latents([x]) + unpacked = nrs_module._unpack_latents(packed, shapes) + assert len(unpacked) == 1 + assert torch.allclose(unpacked[0], x) + + +# --------------------------------------------------------------------------- +# Phase 2: fallback path (no latent_shapes) is a byte-for-byte regression no-op +# --------------------------------------------------------------------------- + + +def _run_nrs(model, cond, uncond, x_orig, sigma, skew=2.0, stretch=5.0, squash=0.75): + node = nrs_module.NRS() + (patched_model,) = node.patch(model, skew, stretch, squash) + fn = patched_model._captured_fn + args = _make_args(model, cond, uncond, x_orig, sigma) + return fn(args) + + +def test_fallback_no_latent_shapes_matches_single_stream_shape(): + model = _StubModel(latent_shapes=None) + cond = torch.randn(2, 4, 8, 8) + uncond = torch.randn(2, 4, 8, 8) + x_orig = torch.randn(2, 4, 8, 8) + sigma = torch.rand(2) + 0.1 + + result = _run_nrs(model, cond, uncond, x_orig, sigma) + assert result.shape == x_orig.shape + + # Regression check: manually compute the single-stream result the same + # way the pre-split code path did, and confirm equality. + node = nrs_module.NRS() + expected = node._apply_guidance( + x_orig, cond, uncond, sigma, 2.0, 5.0, 0.75, nrs_module.PredictionType.EPS + ) + assert torch.allclose(result, expected) + + +def test_single_stream_latent_shapes_also_matches(): + """A model.latent_shapes list of length 1 must take the same code path.""" + cond = torch.randn(1, 4, 5, 5) + uncond = torch.randn(1, 4, 5, 5) + x_orig = torch.randn(1, 4, 5, 5) + sigma = torch.rand(1) + 0.1 + + model = _StubModel(latent_shapes=[cond.shape]) + result = _run_nrs(model, cond, uncond, x_orig, sigma) + + node = nrs_module.NRS() + expected = node._apply_guidance( + x_orig, cond, uncond, sigma, 2.0, 5.0, 0.75, nrs_module.PredictionType.EPS + ) + assert torch.allclose(result, expected) + + +# --------------------------------------------------------------------------- +# Phase 3: degeneracy tripwire +# --------------------------------------------------------------------------- + + +def test_tripwire_fires_on_flat_pack_without_latent_shapes(caplog): + model = _StubModel(latent_shapes=None) + cond = torch.randn(1, 1, 100) + uncond = torch.randn(1, 1, 100) + x_orig = torch.randn(1, 1, 100) + sigma = torch.rand(1) + 0.1 + + with caplog.at_level("WARNING"): + _run_nrs(model, cond, uncond, x_orig, sigma) + + assert any("singleton reduction axis" in rec.message for rec in caplog.records) + + +def test_tripwire_does_not_fire_for_normal_single_stream(caplog): + model = _StubModel(latent_shapes=None) + cond = torch.randn(1, 4, 8, 8) + uncond = torch.randn(1, 4, 8, 8) + x_orig = torch.randn(1, 4, 8, 8) + sigma = torch.rand(1) + 0.1 + + with caplog.at_level("WARNING"): + _run_nrs(model, cond, uncond, x_orig, sigma) + + assert not any("singleton reduction axis" in rec.message for rec in caplog.records) + + +# --------------------------------------------------------------------------- +# Phase 4: per-stream reduced shapes after unpack (H3-like video + audio) +# --------------------------------------------------------------------------- + + +def test_per_stream_reduced_shapes_after_unpack(): + video = torch.randn(1, 24, 4, 3, 2) + audio = torch.randn(1, 32, 2, 5) + shapes = [video.shape, audio.shape] + + packed = nrs_module._pack_latents([video, audio]) + unpacked = nrs_module._unpack_latents(packed, shapes) + + video_u, audio_u = unpacked + assert video_u.shape == video.shape + assert audio_u.shape == audio.shape + + # Channels sit at dim 1 for both streams. + assert video_u.shape[1] == 24 + assert audio_u.shape[1] == 32 + + video_reduced = video_u.sum(dim=1, keepdim=True) + audio_reduced = audio_u.sum(dim=1, keepdim=True) + assert video_reduced.shape == (1, 1, 4, 3, 2) + assert audio_reduced.shape == (1, 1, 2, 5) + + +# --------------------------------------------------------------------------- +# Phase 5: split restores non-degenerate rejection (proves Skew is alive) +# --------------------------------------------------------------------------- + + +def test_split_restores_nondegenerate_rejection(): + """On a real multi-channel stream, uncond's rejection on cond must not + collapse to ~0 -- this is the geometry that was silently dead on the flat + [B,1,N] pack before the unpack/repack fix. + """ + torch.manual_seed(0) + cond = torch.randn(1, 8, 4, 4) + # Make uncond non-parallel to cond so the rejection component is nonzero. + uncond = torch.randn(1, 8, 4, 4) + + def _dot(a, b): + return (a * b).sum(dim=1, keepdim=True) + + eps = torch.finfo(cond.dtype).eps + c_dot_c = _dot(cond, cond) + eps + u_dot_c = _dot(uncond, cond) + u_on_c = (u_dot_c / c_dot_c) * cond + u_rej_c = uncond - u_on_c + + assert u_rej_c.abs().max().item() > 1e-4 + + +def test_flat_pack_rejection_is_degenerate_without_split(): + """Sanity check for the bug this PR fixes: reducing over the flat pack's + singleton dim=1 axis collapses the rejection to exactly zero (up to + floating point noise from the eps regularization term). + """ + # float64 keeps the residual from the eps regularizer near the true + # machine epsilon instead of float32 accumulation noise, so the + # collapse-to-zero identity is exact enough to assert tightly. + packed_cond = torch.randn(1, 1, 100, dtype=torch.float64) + packed_uncond = torch.randn(1, 1, 100, dtype=torch.float64) + + def _dot(a, b): + return (a * b).sum(dim=1, keepdim=True) + + eps = torch.finfo(packed_cond.dtype).eps + c_dot_c = _dot(packed_cond, packed_cond) + eps + u_dot_c = _dot(packed_uncond, packed_cond) + u_on_c = (u_dot_c / c_dot_c) * packed_cond + u_rej_c = packed_uncond - u_on_c + + assert u_rej_c.abs().max().item() < 1e-8