Files
mcDandy-more_math/tests/test_deprecated_nodes.py
T
2026-01-24 23:28:40 +01:00

116 lines
4.3 KiB
Python

import torch
import sys
import os
# Ensure test runner usage (Visual Studio) can import the package regardless of working dir.
_here = os.path.abspath(os.path.dirname(__file__))
_project_root = os.path.abspath(os.path.join(_here, os.pardir))
if _project_root not in sys.path:
sys.path.insert(0, _project_root)
# Add ComfyUI root to path to find 'comfy' and 'comfy_api' packages
_comfy_root = os.path.abspath(os.path.join(_here, "../../.."))
if _comfy_root not in sys.path:
sys.path.insert(0, _comfy_root)
from more_math.deprecated.ImageMathNode import ImageMathNodeOLD
from more_math.deprecated.LatentMathNode import LatentMathNodeOLD
from more_math.deprecated.MaskMathNode import MaskMathNodeOLD
from more_math.deprecated.AudioMathNode import AudioMathNodeOLD
from more_math.deprecated.ModelMathNode import ModelMathNodeOLD
from more_math.deprecated.VaeMathNode import VAEMathNodeOLD
from more_math.deprecated.ClipMathNode import CLIPMathNodeOLD
import comfy_api
# Mock objects for Model/VAE/CLIP tests
class MockModelPatcher:
def __init__(self, state_dict):
self.model = comfy_api.Model()
self.model.state_dict = lambda: state_dict
self.patches = {}
def clone(self):
new_patcher = MockModelPatcher(self.model.state_dict())
return new_patcher
def add_patches(self, patches, strength1, strength2):
self.patches.update(patches)
class MockPatcherContainer:
def __init__(self, state_dict):
self.patcher = MockModelPatcher(state_dict)
def clone(self):
new_obj = MockPatcherContainer(self.patcher.model.state_dict())
new_obj.patcher = self.patcher.clone()
return new_obj
def add_patches(self, patches, s1, s2):
self.patcher.add_patches(patches, s1, s2)
class MockVAE(MockPatcherContainer):
pass # VAE structure is similar enough for simple patching test
def test_deprecated_image_math():
t1 = torch.zeros((1, 32, 32, 3))
t2 = torch.full((1, 32, 32, 3), 1.0)
# OLD signature: Image, a, b=...
res = ImageMathNodeOLD.execute(Image="a+b", a=t1, b=t2)[0]
assert torch.allclose(res, torch.ones_like(res))
def test_deprecated_latent_math():
l1 = {"samples": torch.zeros((1, 4, 32, 32))}
l2 = {"samples": torch.full((1, 4, 32, 32), 1.0)}
# OLD signature: Latent, a, b=...
res = LatentMathNodeOLD.execute(Latent="a+b", a=l1, b=l2)[0]["samples"]
assert torch.allclose(res, torch.ones_like(res))
def test_deprecated_mask_math():
m1 = torch.zeros((1, 32, 32))
m2 = torch.full((1, 32, 32), 1.0)
# OLD signature: Mask, a, b=...
res = MaskMathNodeOLD.execute(Mask="a+b", a=m1, b=m2)[0]
assert torch.allclose(res, torch.ones_like(res))
def test_deprecated_audio_math():
a1 = {"waveform": torch.zeros((1, 1, 100)), "sample_rate": 44100}
a2 = {"waveform": torch.full((1, 1, 100), 1.0), "sample_rate": 44100}
# OLD signature: Audio, a, b=...
res = AudioMathNodeOLD.execute(Audio="a+b", a=a1, b=a2)[0]["waveform"]
assert torch.allclose(res, torch.ones_like(res))
def test_deprecated_model_math():
sd_a = {"w": torch.tensor([1.0])}
sd_b = {"w": torch.tensor([2.0])}
patcher_a = MockModelPatcher(sd_a)
patcher_b = MockModelPatcher(sd_b)
# OLD signature: Model, a, b=...
res = ModelMathNodeOLD.execute(Model="a+b", a=patcher_a, b=patcher_b)[0]
print(f"DEBUG: res.patches keys: {res.patches.keys()}")
if "w" in res.patches:
print(f"DEBUG: res.patches['w']: {res.patches['w']}")
# a+b = 1+2 = 3. diff = 3-1 = 2.
assert "w" in res.patches
assert torch.allclose(res.patches["w"][0], torch.tensor([2.0]))
def test_deprecated_vae_math():
sd_a = {"w": torch.tensor([1.0])}
vae_a = MockVAE(sd_a) # VAE wrapper
# OLD signature: Model, a, b=... (Note: VAE node param name was Model in old schema too)
res = VAEMathNodeOLD.execute(Model="a+1", a=vae_a)[0]
# 1+1=2. diff=1.
assert "w" in res.patcher.patches
assert torch.allclose(res.patcher.patches["w"][0], torch.tensor([1.0]))
def test_deprecated_clip_math():
sd_a = {"w": torch.tensor([1.0])}
clip_a = MockPatcherContainer(sd_a)
# OLD signature: Model, a, b=...
res = CLIPMathNodeOLD.execute(Model="a*2", a=clip_a)[0]
# 1*2=2. diff=1.
assert "w" in res.patcher.patches
assert torch.allclose(res.patcher.patches["w"][0], torch.tensor([1.0]))