diff --git a/more_math/ClipMathNode.py b/more_math/ClipMathNode.py index 234d1b6..016170f 100644 --- a/more_math/ClipMathNode.py +++ b/more_math/ClipMathNode.py @@ -1,5 +1,4 @@ from comfy_api.latest import io -from .modelLikeCommon import calculate_patches from inspect import cleandoc from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer @@ -81,21 +80,21 @@ class CLIPMathNode(io.ComfyNode): raise ValueError("At least one input CLIP is required.") patcher_a = a.patcher - + # Prepare CLIP patchers patchers_V = {} for k, v in V.items(): if v is not None: patchers_V[k] = v.patcher - + # Call autogrow patch calculation from .modelLikeCommon import calculate_patches_autogrow - + # Populate aliases map for backward compatibility in expressions (users might still use a/b/c/d/w/x/y/z) # a=V0, b=V1, etc created by us or expected by user? # The prompt says aliases are supported in check_lazy_status. Variables map in helper handles logic. aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"} - + patches = calculate_patches_autogrow(Expression, V=patchers_V, F=F, mapping=aliases) out_clip = a.clone() diff --git a/more_math/FloatMathNode.py b/more_math/FloatMathNode.py index 04842fc..ce259c4 100644 --- a/more_math/FloatMathNode.py +++ b/more_math/FloatMathNode.py @@ -1,7 +1,7 @@ from inspect import cleandoc import torch -from .helper_functions import parse_expr, as_tensor +from .helper_functions import parse_expr from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io @@ -81,7 +81,7 @@ class FloatMathNode(io.ComfyNode): variables["x"] = V.get("V5", 0.0) variables["y"] = V.get("V6", 0.0) variables["z"] = V.get("V7", 0.0) - + # Populate all V inputs for k, val in V.items(): variables[k] = val if val is not None else 0.0 diff --git a/more_math/LatentMathNode.py b/more_math/LatentMathNode.py index 678a314..518c621 100644 --- a/more_math/LatentMathNode.py +++ b/more_math/LatentMathNode.py @@ -177,7 +177,7 @@ class LatentMathNode(io.ComfyNode): if time_dim is not None: F_idx = getIndexTensorAlongDim(ae, time_dim) variables.update({"frame_idx": F_idx, "frame": F_idx, "frame_count": frame_count}) - + # Add all dynamic inputs for k, v in V.items(): if v is not None: diff --git a/more_math/MaskMathNode.py b/more_math/MaskMathNode.py index c16f91a..acf8181 100644 --- a/more_math/MaskMathNode.py +++ b/more_math/MaskMathNode.py @@ -1,4 +1,4 @@ -from .helper_functions import generate_dim_variables,parse_expr, getIndexTensorAlongDim, as_tensor, commonLazy, normalize_to_common_shape,prepare_inputs, make_zero_like +from .helper_functions import generate_dim_variables,parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape,prepare_inputs, make_zero_like from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io from antlr4 import InputStream, CommonTokenStream @@ -98,7 +98,7 @@ class MaskMathNode(io.ComfyNode): for name, tensor in V.items(): if tensor is not None and tensor.shape[0] != max_length: raise ValueError(f"Input '{name}' has shape {tensor.shape[0]}, expected {max_length} to match largest input.") - + ae, be, ce, de = normalize_to_common_shape(ae, be, ce, de, mode=length_mismatch) variables = { diff --git a/more_math/ModelMathNode.py b/more_math/ModelMathNode.py index 39b1267..537342e 100644 --- a/more_math/ModelMathNode.py +++ b/more_math/ModelMathNode.py @@ -1,7 +1,5 @@ from inspect import cleandoc from comfy_api.latest import io -from .helper_functions import commonLazy -from .modelLikeCommon import calculate_patches from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser @@ -97,26 +95,26 @@ class ModelMathNode(io.ComfyNode): # Looking at Step 34, calculate_patches signature: (Model, a, b, c, d, w, x, y, z) # We need to verify if calculate_patches handles V/F. It probably doesn't. # We should check 'modelLikeCommon.py' to see if update is needed. - + # Assume for now we pass a,b,c,d,w,x,y,z as standard. # But for full autogrow support (more than 4 inputs), calculate_patches needs update. # The prompt didn't explicitly ask to update modelLikeCommon, but "switch to Autogrow" implies full functionality. # I'll check modelLikeCommon.py after this block. # For now, I will pass V and F to calculate_patches if I modify it, or I will stick to legacy args if I don't modify it. # However, to support V4+, I MUST modify calculate_patches. - + # Let's pass the V and F dicts to a modified calculate_patches, or overload it. # I will update modelLikeCommon.py as part of this task. - + from .modelLikeCommon import calculate_patches_autogrow - + # Map inputs to patchers if needed (Model.Input gives Model wrapper, need state_dict source?) # ModelMathNode inputs are Model wrappers (comfy.model_patcher.ModelPatcher). # So V items are ready to be used. - + aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"} patches = calculate_patches_autogrow(Expression, V=V, F=F, mapping=aliases) - + out_model = a.clone() if patches: out_model.add_patches(patches, 1.0, 1.0) diff --git a/more_math/VaeMathNode.py b/more_math/VaeMathNode.py index df7fdf5..166c307 100644 --- a/more_math/VaeMathNode.py +++ b/more_math/VaeMathNode.py @@ -1,7 +1,6 @@ from inspect import cleandoc from comfy_api.latest import io import copy -from .modelLikeCommon import calculate_patches from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser @@ -82,15 +81,15 @@ class VAEMathNode(io.ComfyNode): raise ValueError("At least one input VAE is required.") patcher_a = a.patcher - + # Prepare VAE patchers for calculation # We need to map VAE wrappers to their patchers for `calculate_patches` - + patchers_V = {} for k, v in V.items(): if v is not None: patchers_V[k] = v.patcher - + # Calculate patches using the patchers (weights are in patcher.model.state_dict) from .modelLikeCommon import calculate_patches_autogrow aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"} diff --git a/more_math/modelLikeCommon.py b/more_math/modelLikeCommon.py index 5fdffed..c1d9574 100644 --- a/more_math/modelLikeCommon.py +++ b/more_math/modelLikeCommon.py @@ -1,10 +1,6 @@ from .helper_functions import parse_expr, getIndexTensorAlongDim, as_tensor from .Parser.UnifiedMathVisitor import UnifiedMathVisitor -from .Parser.MathExprParser import MathExprParser -from antlr4 import InputStream, CommonTokenStream -from .Parser.MathExprLexer import MathExprLexer import torch -import comfy.utils def calculate_patches(Model, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0): """Legacy calculate_patches for backward compatibility.""" @@ -28,7 +24,7 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): # Collect all unique keys from all models all_keys = set() models = [v for v in V.values() if v is not None] - + if not models: return {} @@ -54,17 +50,17 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): tree = parse_expr(Expr) patches = {} - + # Progress bar if possible (comfy.utils.ProgressBar might assume unthreaded?) # Just skip for utility or use if substantial. - + for key in all_keys: variables = {} - + # Populate F variables (constants for all keys) for k, val in F.items(): variables[k] = val if val is not None else 0.0 - + # Also populate mapped aliases for F (w, x, y, z) for alias, target in mapping.items(): if target in F: @@ -73,7 +69,7 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): # Inject weights for this key from V models valid_key = False ref_tensor = None - + for v_name, v_val in V.items(): if v_val is not None: w_tensor = get_weight(v_val, key) @@ -84,19 +80,19 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): else: # Missing key in this model will be handled later (zero init) pass - + if not valid_key: continue - + # Find reference shape if ref_tensor is None: continue # Should not happen if valid_key is true - + # Fill missing models with zeros for v_name in V.keys(): if v_name not in variables: variables[v_name] = torch.zeros_like(ref_tensor) - + # Populate aliases for V (a, b, c, d) for alias, target in mapping.items(): if target in variables: @@ -109,31 +105,31 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): idx_tensor = getIndexTensorAlongDim(ref_tensor, dim_idx) variables[f"D{dim_idx}"] = idx_tensor variables[f"dim_{dim_idx}"] = idx_tensor - + # Execute math try: visitor = UnifiedMathVisitor(variables, ref_tensor.shape) res = visitor.visit(tree) res = as_tensor(res, ref_tensor.shape) - + # Calculate patch: Result - Original(V0) # Assumption: We are patching V0. # If V0 doesn't have the key, we assume V0 was zero? # ComfyUI patching mechanism adds patch to original weights. # If we output 'res', we need to return (res - original). - + # Get original weight for V0 (alias 'a' usually) original = variables.get("V0") # Or strictly V.get("V0")'s weight if original is None: original = torch.zeros_like(res) - + diff = res - original - + # Clean up: don't store zero patches if not torch.all(diff == 0): patches[key] = (diff,) - + except Exception: pass - + return patches diff --git a/tests/test_deprecated_nodes.py b/tests/test_deprecated_nodes.py index 183100a..7306531 100644 --- a/tests/test_deprecated_nodes.py +++ b/tests/test_deprecated_nodes.py @@ -44,7 +44,7 @@ class MockPatcherContainer: 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) @@ -84,7 +84,7 @@ def test_deprecated_model_math(): 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()}") @@ -97,7 +97,7 @@ def test_deprecated_model_math(): 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. @@ -107,7 +107,7 @@ def test_deprecated_vae_math(): 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. diff --git a/tests/test_model_math.py b/tests/test_model_math.py index 86691df..61b4c3f 100644 --- a/tests/test_model_math.py +++ b/tests/test_model_math.py @@ -266,7 +266,7 @@ def test_model_math_3_way_merge(): # Expression: a + b + c # Expected: 1 + 2 + 3 = 6. # Diff vs a (1.0) = 5.0. - + result_tuple = ModelMathNode.execute(Expression="a + b + c", V={"V0": a.patcher, "V1": b.patcher, "V2": c.patcher}, F={}) patches = result_tuple[0].patches