PR-2: pack-aware per-stream NRS routing + degeneracy tripwire (#36)
## What Fixes NRS on MiniMax H3 (and any sampler that packs multiple streams). ComfyUI hands the cfg hook a **flat packed latent** `[B,1,N]` whose `dim=1` is a singleton, so every NRS reduction runs over a size-1 axis and the projection geometry collapses — Skew becomes a silent no-op, Stretch degenerates to elementwise CFG, Squash to a Stretch-reverter (proven to float precision in Phase 0). This PR unpacks the flat pack into its real channels-first per-stream tensors, runs the **unchanged** NRS geometry per stream, then repacks. ## Changes (`NRS/nodes_NRS.py`) - Guarded `import comfy.utils as _comfy_utils` (module still imports with no ComfyUI present) + `import math`. - Module-level pure-Python `_unpack_latents`/`_pack_latents` fallbacks matching ComfyUI's pack contract (each stream reshaped to `(B,1,-1)` and concatenated on the last dim; unpack slices `prod(shape[1:])` per stream). - Extracted `_apply_guidance(...)` — the per-stream geometry pipeline (sigma reshape → `_convert_to_v_space` → dot/proj/stretch/skew/squash → `_finalize_from_v_space`), **byte-for-byte identical** to the prior inline math. - Rewired `nrs()`: read `args["model"].latent_shapes`; when present with `len>1`, unpack cond/uncond/input per stream (comfy-preferred, pure-Python fallback), map `_apply_guidance`, repack; otherwise the single-stream path is the **exact prior behavior** (regression no-op for SDXL etc.). - Permanent **degeneracy tripwire**: warns once per `patch()` if any routed stream still has a singleton reduction axis — the silent-failure class that cost the original Skew investigation days. ## Why this is safe - Single-stream models are a verified byte-for-byte no-op (two regression tests compare against a manually-computed `_apply_guidance`). - Not H3 fingerprinting: any multi-stream packed sampler (LTXV AV variants) benefits; older ComfyUI without the API falls through cleanly. - Verified against the real 18MB H3 capture (not committed): unpack→repack is bit-exact for cond/uncond/x_orig; streams recover as `[1,24,72,38,22]` (video) / `[1,32,2,405]` (audio); per-stream guidance produces large live deviations where the flat pack was near-degenerate — Skew is alive again. ## Tests New `tests/test_pack_split.py` (9 tests, real-torch via a module-scoped isolation harness that restores the mock afterward): round-trip pack/unpack, single- & multi-stream fallback regression, tripwire fires/doesn't-fire, per-stream reduced shapes, and non-degenerate-vs-degenerate rejection. - `uvx ruff check .` → All checks passed! - `uvx --with torch pytest -q` → **33 passed** (24 existing + 9 new) Breaking for H3 outputs only. Baseline preserved at tag `pre-flow`.
This commit is contained in:
+102
-31
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user