massive refactor (AI)

This commit is contained in:
mcDandy
2025-12-25 19:43:09 +01:00
parent 7fad03daf1
commit be29ad7f67
10 changed files with 391 additions and 524 deletions
+43 -55
View File
@@ -1,29 +1,21 @@
import torch
from antlr4 import CommonTokenStream, InputStream
from .Parser.MathExprParser import MathExprParser
from .Parser.MathExprLexer import MathExprLexer
from .Parser.TensorEvalVisitor import TensorEvalVisitor
from .helper_functions import getIndexTensorAlongDim, comonLazy
from .helper_functions import getIndexTensorAlongDim, comonLazy, eval_tensor_expr, make_zero_like
from comfy_api.latest import io
class AudioMathNode(io.ComfyNode):
"""
This node enables the use of math expressions on AUDIO tensors.
inputs:
a, b, c, d:
AUDIO, bound to variables with the same name. Defaults to zero AUDIO if not provided.
w, x, y, z:
Floats, bound to variables of the expression. Defaults to 0.0 if not provided.
Audio expression:
String, describing expression to apply to audio tensors.
outputs:
AUDIO:
Returns an AUDIO object that contains the result of the math expression applied to the input audio tensors.
Enables math expressions on Audio tensors.
Inputs:
a, b, c, d: Audio inputs (b, c, d default to zero if not provided)
w, x, y, z: Float variables for expressions
AudioExpr: Expression to apply on audio tensors
Outputs:
AUDIO: Result of applying expression to input audio
"""
def __init__(self):
pass
@classmethod
def define_schema(cls) -> io.Schema:
@@ -33,56 +25,52 @@ class AudioMathNode(io.ComfyNode):
display_name="Audio math",
inputs=[
io.Audio.Input(id="a", tooltip="Input audio tensor"),
io.Audio.Input(id="b", optional=True,lazy=True, tooltip="Second input audio tensor"),
io.Audio.Input(id="c", optional=True,lazy=True, tooltip="Third input audio tensor"),
io.Audio.Input(id="d", optional=True,lazy=True, tooltip="Fourth input audio tensor"),
io.Float.Input(id="w", default=0.0, optional=True,lazy=True, force_input=True),
io.Float.Input(id="x", default=0.0, optional=True,lazy=True, force_input=True),
io.Float.Input(id="y", default=0.0, optional=True,lazy=True, force_input=True),
io.Float.Input(id="z", default=0.0, optional=True,lazy=True, force_input=True),
io.Audio.Input(id="b", optional=True, lazy=True, tooltip="Second input audio tensor"),
io.Audio.Input(id="c", optional=True, lazy=True, tooltip="Third input audio tensor"),
io.Audio.Input(id="d", optional=True, lazy=True, tooltip="Fourth input audio tensor"),
io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True),
io.String.Input(id="AudioExpr", default="a*(1-w)+b*w", tooltip="Expression to apply on input audio tensors"),
],
outputs=[
io.Audio.Output(),
],
)
@classmethod
def check_lazy_status(cls, AudioExpr, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0):
return comonLazy(AudioExpr, a, b, c, d,w,x,y,z)
def check_lazy_status(cls, AudioExpr, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0):
return comonLazy(AudioExpr, a, b, c, d, w, x, y, z)
@classmethod
def execute(cls, a, AudioExpr, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0):
waveform = a['waveform']
sample_rate = a['sample_rate']
b = b if b else make_zero_like(a)
c = c if c else make_zero_like(a)
d = d if d else make_zero_like(a)
bv = b if b else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']}
cv = c if c else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']}
dv = d if d else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']}
B = getIndexTensorAlongDim(a['waveform'], 0)
C = getIndexTensorAlongDim(a['waveform'], 1)
S = getIndexTensorAlongDim(a['waveform'], 2)
R = torch.full_like(S, a['sample_rate'], dtype=torch.float32)
T = torch.full_like(S, a['waveform'].shape[2], dtype=torch.float32)
bv, cv, dv = b['waveform'], c['waveform'], d['waveform']
variables = {
'a': a['waveform'], 'b': bv['waveform'], 'c': cv['waveform'], 'd': dv['waveform'],
'a': waveform, 'b': bv, 'c': cv, 'd': dv,
'w': w, 'x': x, 'y': y, 'z': z,
'B': B, 'C': C, 'S': S,'R': R, 'T' : T, 'N': a['waveform'].shape[1],
'batch': B, 'channel': C, 'sample': S, 'sample_rate': R, 'sample_count': T,'channel_count': a['waveform'].shape[1]
'B': getIndexTensorAlongDim(waveform, 0),
'C': getIndexTensorAlongDim(waveform, 1),
'S': getIndexTensorAlongDim(waveform, 2),
'R': torch.full_like(waveform, sample_rate, dtype=torch.float32),
'T': torch.full_like(waveform, waveform.shape[2], dtype=torch.float32),
'N': waveform.shape[1],
'batch': getIndexTensorAlongDim(waveform, 0),
'channel': getIndexTensorAlongDim(waveform, 1),
'sample': getIndexTensorAlongDim(waveform, 2),
'sample_rate': torch.full_like(waveform, sample_rate, dtype=torch.float32),
'sample_count': torch.full_like(waveform, waveform.shape[2], dtype=torch.float32),
'channel_count': waveform.shape[1],
}
input_stream = InputStream(AudioExpr)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
tree = parser.expr()
result_tensor = eval_tensor_expr(AudioExpr, variables, waveform.shape)
visitor = TensorEvalVisitor(variables, a['waveform'].shape)
result_tensor = visitor.visit(tree)
# Create output dictionary with the same sample rate
output = {
'waveform': result_tensor,
'sample_rate': a['sample_rate']
}
return (output,)
return ({'waveform': result_tensor, 'sample_rate': sample_rate},)
+47 -92
View File
@@ -1,34 +1,25 @@
from inspect import cleandoc
from antlr4 import CommonTokenStream, InputStream
import torch
from .Parser.MathExprParser import MathExprParser
from .Parser.MathExprLexer import MathExprLexer
from .Parser.TensorEvalVisitor import TensorEvalVisitor
from .helper_functions import ThrowingErrorListener, comonLazy
from .helper_functions import comonLazy, eval_tensor_expr, make_zero_like
from comfy_api.latest import io
class ConditioningMathNode(io.ComfyNode):
"""
This node enables the use of math on conditionings. It is recommended to keep the Tensor expression the same as the pooled_output expression.
inputs:
a, b, c, d:
Conditioning, bound to variables with the same name. Defaults to zero conditioning if not provided.
w, x, y, z:
Floats, bound to variables of the expression. Defaults to 0.0 if not provided.
Tensor expression:
String, describing expression to mix tensor part of conditioning. Tensor probably describes the composition of the image as it has the say what is on the image.
Pooled output expression:
String, describing expression to mix pooled_output part of conditioning. pooled_output is Tensor compressed into 1 token.
outputs:
CONDITIONING:
Returns a CONDITIONING object that contains the result of the math expression applied to the input conditionings.
Enables math operations on conditionings.
Inputs:
a, b, c, d: Conditioning inputs (b, c, d default to zero if not provided)
w, x, y, z: Float variables for expressions
Tensor: Expression for the tensor part (describes image composition)
pooled_output: Expression for the pooled output (condensed representation)
Outputs:
CONDITIONING: Result of applying expressions to input conditionings
"""
def __init__(self):
pass
@classmethod
def define_schema(cls) -> io.Schema:
return io.Schema(
@@ -37,15 +28,15 @@ class ConditioningMathNode(io.ComfyNode):
category="More math",
inputs=[
io.Conditioning.Input(id="a"),
io.Conditioning.Input(id="b", optional=True,lazy=True),
io.Conditioning.Input(id="c", optional=True,lazy=True),
io.Conditioning.Input(id="d", optional=True,lazy=True),
io.Float.Input(id="w", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="x", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="y", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="z", default=0.0,optional=True,lazy=True, force_input=True),
io.String.Input(id="Tensor", default="a*(1-w)+b*w", tooltip="Describes composition of the image."),
io.String.Input(id="pooled_output", default="a*(1-w)+b*w", tooltip="Composition of the image condensed into one vector"),
io.Conditioning.Input(id="b", optional=True, lazy=True),
io.Conditioning.Input(id="c", optional=True, lazy=True),
io.Conditioning.Input(id="d", optional=True, lazy=True),
io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True),
io.String.Input(id="Tensor", default="a*(1-w)+b*w", tooltip="Expression for tensor part (image composition)"),
io.String.Input(id="pooled_output", default="a*(1-w)+b*w", tooltip="Expression for pooled output (condensed representation)"),
],
outputs=[
io.Conditioning.Output(),
@@ -54,71 +45,35 @@ class ConditioningMathNode(io.ComfyNode):
tooltip = cleandoc(__doc__)
#OUTPUT_NODE = False
#OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node
@classmethod
def check_lazy_status(cls, Tensor,pooled_output, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0):
return list(set(comonLazy(Tensor, a, b, c, d)).union(comonLazy(pooled_output, a, b, c, d)))
def check_lazy_status(cls, Tensor, pooled_output, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0):
tensor_needs = set(comonLazy(Tensor, a, b, c, d))
pooled_needs = set(comonLazy(pooled_output, a, b, c, d))
return list(tensor_needs.union(pooled_needs))
@classmethod
def execute(cls, Tensor,pooled_output, a, b=None, c=None, d=None,w=0.0,x=0.0,y=0.0,z=0.0):
if b is None:
b = [[torch.zeros_like(a[0][0]), {"pooled_output": torch.zeros_like(a[0][1]["pooled_output"]) if a[0][1]["pooled_output"] is not None else None}]]
if c is None:
c = [[torch.zeros_like(a[0][0]), {"pooled_output": torch.zeros_like(a[0][1]["pooled_output"]) if a[0][1]["pooled_output"] is not None else None}]]
if d is None:
d = [[torch.zeros_like(a[0][0]), {"pooled_output": torch.zeros_like(a[0][1]["pooled_output"]) if a[0][1]["pooled_output"] is not None else None}]]
ta = a[0][0].clone()
tb = b[0][0].clone()
tc = c[0][0].clone()
td = d[0][0].clone()
pa = None
pb = None
pc = None
pd = None
if a[0][1]["pooled_output"] is not None:
pa = a[0][1]["pooled_output"].clone()
pb = b[0][1]["pooled_output"].clone()
pc = c[0][1]["pooled_output"].clone()
pd = d[0][1]["pooled_output"].clone()
def execute(cls, Tensor, pooled_output, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0):
# Default missing conditionings to zero
b = b if b is not None else make_zero_like(a)
c = c if c is not None else make_zero_like(a)
d = d if d is not None else make_zero_like(a)
# Extract tensors
ta, tb, tc, td = a[0][0], b[0][0], c[0][0], d[0][0]
# Evaluate tensor expression
variables = {'a': ta, 'b': tb, 'c': tc, 'd': td, 'w': w, 'x': x, 'y': y, 'z': z}
result_tensor = eval_tensor_expr(Tensor, variables, ta.shape)
input_stream = InputStream(Tensor)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
parser.addErrorListener(ThrowingErrorListener())
tree = parser.expr()
visitor = TensorEvalVisitor(variables,ta.shape)
result1 = visitor.visit(tree)
if a[0][1]["pooled_output"] is not None:
# Evaluate pooled_output expression if available
pa = a[0][1].get("pooled_output")
if pa is not None:
pb = b[0][1].get("pooled_output")
pc = c[0][1].get("pooled_output")
pd = d[0][1].get("pooled_output")
variables = {'a': pa, 'b': pb, 'c': pc, 'd': pd, 'w': w, 'x': x, 'y': y, 'z': z}
input_stream = InputStream(Tensor)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
tree = parser.expr()
visitor = TensorEvalVisitor(variables,pa.shape)
result2 = visitor.visit(tree)
result_pooled = eval_tensor_expr(pooled_output, variables, pa.shape)
else:
result2 = None
print("No pooled_output found in input conditioning, skipping pooled_output calculation.")
result_pooled = None
result = [[result1, {"pooled_output": result2}]]
return (result,)
"""
The node will always be re executed if any of the inputs change but
this method can be used to force the node to execute again even when the inputs don't change.
You can make this node return a number or a string. This value will be compared to the one returned the last time the node was
executed, if it is different the node will be executed again.
This method is used in the core repo for the LoadImage node where they return the image hash as a string, if the image hash
changes between executions the LoadImage node is executed again.
"""
#@classmethod
#def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen):
# return ""
return ([[result_tensor, {"pooled_output": result_pooled}]],)
+23 -58
View File
@@ -1,51 +1,38 @@
from inspect import cleandoc
from antlr4 import CommonTokenStream, InputStream
from .helper_functions import ThrowingErrorListener, comonLazy
from .Parser.MathExprParser import MathExprParser
from .Parser.MathExprLexer import MathExprLexer
from .Parser.FloatEvalVisitor import FloatEvalVisitor
from .helper_functions import comonLazy, eval_float_expr
from comfy_api.latest import io
class FloatMathNode(io.ComfyNode):
"""
This node enables the use of math expressions on Latents.
inputs:
a, b, c, d:
Floats, bound to variables with the same name. Defaults to 0.0 if not provided.
w, x, y, z:
Floats, bound to variables of the expression. Defaults to 0.0 if not provided.
Latent expression:
String, describing expression to mix latents. Valid functions are sin, cos, tan, abs, sqrt, min, max, norm. Valid operators are +, -, *, /, ^, %. Usable constants are e and pi.
outputs:
LATENT:
Returns a LATENT object that contains the result of the math expression applied to the input conditionings.
This node enables the use of math expressions on Floats.
Inputs:
a, b, c, d: Floats, bound to variables with the same name.
w, x, y, z: Floats, bound to variables of the expression.
FloatFunc: String, describing math expression.
Valid functions: sin, cos, tan, abs, sqrt, min, max, norm, etc.
Operators: +, -, *, /, ^, %.
Constants: e, pi.
Outputs:
FLOAT: The result of evaluating the math expression.
"""
def __init__(self):
pass
@classmethod
def check_lazy_status(cls, Model, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0):
return comonLazy(Model, a, b, c, d,w,x,y,z)
@classmethod
def define_schema(cls) -> io.Schema:
"""
"""
return io.Schema(
node_id="mrmth_FloatMathNode",
category="More math",
display_name="Float math",
inputs=[
io.Float.Input(id="a", force_input=True),
io.Float.Input(id="b", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="c", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="d", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="w", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="x", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="y", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="z", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="b", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="c", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="d", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True),
io.String.Input(id="FloatFunc", default="a*(1-w)+b*w", tooltip="Expression to use on inputs"),
],
outputs=[
@@ -53,36 +40,14 @@ class FloatMathNode(io.ComfyNode):
],
)
#RETURN_NAMES = ("image_output_name",)
tooltip = cleandoc(__doc__)
#OUTPUT_NODE = False
#OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node
@classmethod
def check_lazy_status(cls, FloatFunc, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0):
def check_lazy_status(cls, FloatFunc, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0):
return comonLazy(FloatFunc, a, b, c, d)
@classmethod
def execute(cls, FloatFunc, a, b=0.0, c=0.0, d=0.0, w=0.0, x=0.0, y=0.0, z=0.0):
variables = {'a': a, 'b': b, 'c': c, 'd': d, 'w': w, 'x': x, 'y': y, 'z': z}
input_stream = InputStream(FloatFunc)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
parser.addErrorListener(ThrowingErrorListener())
tree = parser.expr()
visitor = FloatEvalVisitor(variables)
result = visitor.visit(tree)
result = eval_float_expr(FloatFunc, variables)
return (result,)
"""
The node will always be re executed if any of the inputs change but
this method can be used to force the node to execute again even when the inputs don't change.
You can make this node return a number or a string. This value will be compared to the one returned the last time the node was
executed, if it is different the node will be executed again.
This method is used in the core repo for the LoadImage node where they return the image hash as a string, if the image hash
changes between executions the LoadImage node is executed again.
"""
#@classmethod
#def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen):
# return ""
+36 -70
View File
@@ -1,109 +1,75 @@
from antlr4 import CommonTokenStream
from antlr4.atn.LexerActionExecutor import InputStream
import torch
from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim, comonLazy
from .Parser.MathExprParser import MathExprParser
from .Parser.MathExprLexer import MathExprLexer
from .Parser.TensorEvalVisitor import TensorEvalVisitor
from .helper_functions import getIndexTensorAlongDim, comonLazy, eval_tensor_expr, make_zero_like
from comfy_api.latest import io
class ImageMathNode(io.ComfyNode):
"""
This node enables the use of math expressions on Latents.
inputs:
a, b, c, d:
Latent, bound to variables with the same name. Defaults to zero latent if not provided.
w, x, y, z:
Floats, bound to variables of the expression. Defaults to 0.0 if not provided.
Image expression:
String, describing expression to mix images.
outputs:
LATENT:
Returns a LATENT object that contains the result of the math expression applied to the input conditionings.
Enables math expressions on Images.
Inputs:
a, b, c, d: Image inputs (b, c, d default to zero if not provided)
w, x, y, z: Float variables for expressions
Image: Expression to apply on input images
Outputs:
IMAGE: Result of applying expression to input images
"""
def __init__(self):
pass
@classmethod
def define_schema(cls) -> io.Schema:
"""
"""
return io.Schema(
node_id="mrmth_ImageMathNode",
category="More math",
display_name="Image math",
inputs=[
io.Image.Input(id="a"),
io.Image.Input(id="b", optional=True,lazy=True),
io.Image.Input(id="c", optional=True,lazy=True),
io.Image.Input(id="d", optional=True,lazy=True),
io.Float.Input(id="w", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="x", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="y", default=0.0,optional=True,lazy=True, force_input=True),
io.Float.Input(id="z", default=0.0,optional=True,lazy=True, force_input=True),
io.Image.Input(id="b", optional=True, lazy=True),
io.Image.Input(id="c", optional=True, lazy=True),
io.Image.Input(id="d", optional=True, lazy=True),
io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True),
io.String.Input(id="Image", default="a*(1-w)+b*w", tooltip="Expression to apply on input images"),
],
outputs=[
io.Image.Output(),
],
)
@classmethod
def check_lazy_status(cls, Image, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0):
return comonLazy(Image, a, b, c, d,w,x,y,z)
@classmethod
def execute(scls, Image, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0):
b = torch.zeros_like(a) if b is None else b
c = torch.zeros_like(a) if c is None else c
d = torch.zeros_like(a) if d is None else d
# permute to B, C, H, W
@classmethod
def check_lazy_status(cls, Image, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0):
return comonLazy(Image, a, b, c, d, w, x, y, z)
@classmethod
def execute(cls, Image, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0):
b = make_zero_like(a) if b is None else b
c = make_zero_like(a) if c is None else c
d = make_zero_like(a) if d is None else d
# Permute to B, C, H, W for processing
a = a.permute(0, 3, 1, 2)
b = b.permute(0, 3, 1, 2)
c = c.permute(0, 3, 1, 2)
d = d.permute(0, 3, 1, 2)
B = getIndexTensorAlongDim(a, 0)
C = getIndexTensorAlongDim(a, 1)
H = getIndexTensorAlongDim(a, 2)
W = getIndexTensorAlongDim(a, 3)
variables = {
'a': a, 'b': b, 'c': c, 'd': d,
'w': w, 'x': x, 'y': y, 'z': z,
'X': W, 'Y': H,
'B': B,'batch': B,
'C': C,'channel': C,
'X': getIndexTensorAlongDim(a, 3),
'Y': getIndexTensorAlongDim(a, 2),
'B': getIndexTensorAlongDim(a, 0), 'batch': getIndexTensorAlongDim(a, 0),
'C': getIndexTensorAlongDim(a, 1), 'channel': getIndexTensorAlongDim(a, 1),
'W': a.shape[3], 'width': a.shape[3],
'H': a.shape[2], 'height': a.shape[2],
'T': a.shape[0], 'batch_count': a.shape[0],
'N': a.shape[1], 'channel_count': a.shape[1],
}
input_stream = InputStream(Image)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
parser.addErrorListener(ThrowingErrorListener())
tree = parser.expr()
visitor = TensorEvalVisitor(variables, a.shape)
result = visitor.visit(tree)
# permute back to B, H, W, C
result = eval_tensor_expr(Image, variables, a.shape)
# Permute back to B, H, W, C
result = result.permute(0, 2, 3, 1)
return (result,)
"""
The node will always be re executed if any of the inputs change but
this method can be used to force the node to execute again even when the inputs don't change.
You can make this node return a number or a string. This value will be compared to the one returned the last time the node was
executed, if it is different the node will be executed again.
This method is used in the core repo for the LoadImage node where they return the image hash as a string, if the image hash
changes between executions the LoadImage node is executed again.
"""
#@classmethod
#def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen):
# return ""
+9 -18
View File
@@ -2,14 +2,9 @@ from inspect import cleandoc
from comfy_api.latest import io
from antlr4 import CommonTokenStream, InputStream
import torch
from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim,comonLazy
from .Parser.MathExprParser import MathExprParser
from .Parser.MathExprLexer import MathExprLexer
from .Parser.TensorEvalVisitor import TensorEvalVisitor
from .helper_functions import getIndexTensorAlongDim, comonLazy, parse_expr, eval_tensor_expr_with_tree, make_zero_like
# try to import NestedTensor type if available
try:
@@ -19,6 +14,8 @@ except Exception:
_nested_tensor_module = None
_NESTED_TENSOR_AVAILABLE = False
class LatentMathNode(io.ComfyNode):
"""
This node enables the use of math expressions on Latents.
@@ -78,12 +75,7 @@ class LatentMathNode(io.ComfyNode):
d_in = None if d is None else d["samples"]
# parse expression once
input_stream = InputStream(Latent)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
parser.addErrorListener(ThrowingErrorListener())
tree = parser.expr()
tree = parse_expr(Latent)
# Helper to evaluate for a single tensor
def eval_single_tensor(a_t, b_t, c_t, d_t):
@@ -131,8 +123,7 @@ class LatentMathNode(io.ComfyNode):
F = getIndexTensorAlongDim(a_t, time_dim)
variables.update({'frame_idx': F, 'frame': F, 'frame_count': frame_count})
visitor = TensorEvalVisitor(variables, a_t.shape)
return visitor.visit(tree)
return eval_tensor_expr_with_tree(tree, variables, a_t.shape)
# If input is a NestedTensor (from comfy), evaluate per-subtensor and return NestedTensor result
if hasattr(a_in, 'is_nested') and getattr(a_in, 'is_nested'):
@@ -145,7 +136,7 @@ class LatentMathNode(io.ComfyNode):
def merge_to_tensor(val, ref):
# ref is merged_a
if val is None:
return torch.zeros_like(ref)
return make_zero_like(ref)
if hasattr(val, 'is_nested') and getattr(val, 'is_nested'):
lst = val.unbind()
return torch.cat(lst, dim=0)
@@ -166,7 +157,7 @@ class LatentMathNode(io.ComfyNode):
if val.shape[0] == sum(sizes):
return val
# fallback
return torch.zeros_like(ref)
return make_zero_like(ref)
merged_b = merge_to_tensor(b_in, merged_a)
merged_c = merge_to_tensor(c_in, merged_a)
@@ -187,7 +178,7 @@ class LatentMathNode(io.ComfyNode):
# ensure b/c/d are set appropriately (zeros_like if None)
def to_tensor(val, ref):
if val is None:
return torch.zeros_like(ref)
return make_zero_like(ref)
if hasattr(val, 'is_nested') and getattr(val, 'is_nested'):
lst = val.unbind()
elif isinstance(val, (list, tuple)):
@@ -196,7 +187,7 @@ class LatentMathNode(io.ComfyNode):
return val
if len(lst) == 0:
return torch.zeros_like(ref)
return make_zero_like(ref)
if len(lst) == 1:
return lst[0]
try:
+10 -24
View File
@@ -1,17 +1,11 @@
from inspect import cleandoc
from antlr4 import CommonTokenStream, InputStream
import torch
from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim,comonLazy
from .Parser.MathExprParser import MathExprParser
from .Parser.MathExprLexer import MathExprLexer
from .Parser.TensorEvalVisitor import TensorEvalVisitor
from .helper_functions import getIndexTensorAlongDim, comonLazy, parse_expr, eval_tensor_expr_with_tree, make_zero_like
from comfy_api.latest import io
# try to import NestedTensor type if available
import comfy.nested_tensor as _nested_tensor_module
@@ -95,27 +89,20 @@ class NoiseExecutor():
self.y = y
self.z = z
self.expr = Noise
# parse expression once
input_stream = InputStream(Noise)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
parser.addErrorListener(ThrowingErrorListener())
self.tree = parser.expr()
self.tree = parse_expr(Noise)
seed = -1;
def generate_noise(self, input_latent:torch.Tensor) -> torch.Tensor:
samples = input_latent["samples"]
# evaluate generators / default zeros
a_val = self.a.generate_noise(input_latent) if self.a is not None else torch.zeros_like(samples)
b_val = self.b.generate_noise(input_latent) if self.b is not None else torch.zeros_like(samples)
c_val = self.c.generate_noise(input_latent) if self.c is not None else torch.zeros_like(samples)
d_val = self.d.generate_noise(input_latent) if self.d is not None else torch.zeros_like(samples)
a_val = self.a.generate_noise(input_latent) if self.a is not None else make_zero_like(samples)
b_val = self.b.generate_noise(input_latent) if self.b is not None else make_zero_like(samples)
c_val = self.c.generate_noise(input_latent) if self.c is not None else make_zero_like(samples)
d_val = self.d.generate_noise(input_latent) if self.d is not None else make_zero_like(samples)
# helper to convert a returned value into a list matching ref_list
def to_list(val, ref_list):
if val is None:
return [torch.zeros_like(r) for r in ref_list]
return [make_zero_like(r) for r in ref_list]
# If val is a NestedTensor-like, return underlying list
if hasattr(val, 'is_nested') and getattr(val, 'is_nested'):
return val.unbind()
@@ -141,7 +128,7 @@ class NoiseExecutor():
def merge_to_tensor(val, ref):
if val is None:
return torch.zeros_like(ref)
return make_zero_like(ref)
if hasattr(val, 'is_nested') and getattr(val, 'is_nested'):
lst = val.unbind()
return torch.cat(lst, dim=0)
@@ -158,7 +145,7 @@ class NoiseExecutor():
return torch.cat([val[i].unsqueeze(0).expand(sample_list[i].shape[0], *val.shape[1:]) for i in range(len(sample_list))], dim=0)
except Exception:
pass
return torch.zeros_like(ref)
return make_zero_like(ref)
merged_a = merge_to_tensor(a_val, merged_samples)
merged_b = merge_to_tensor(b_val, merged_samples)
@@ -201,8 +188,7 @@ class NoiseExecutor():
F = getIndexTensorAlongDim(merged_samples, time_dim)
variables.update({'frame': F, 'frame_count': frame_count})
visitor = TensorEvalVisitor(variables, variables['a'].shape)
merged_result = visitor.visit(self.tree)
merged_result = eval_tensor_expr_with_tree(self.tree, variables, variables['a'].shape)
if hasattr(samples, 'is_nested') and getattr(samples, 'is_nested'):
split_results = list(merged_result.split(sizes, dim=0))
+60 -118
View File
@@ -4,38 +4,27 @@ from comfy_api.latest import io
from comfy_api.input_impl import VideoFromComponents
from comfy_api.util import VideoComponents
from antlr4 import CommonTokenStream, InputStream
import torch
from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim
from .helper_functions import getIndexTensorAlongDim, eval_tensor_expr, make_zero_like
from .Parser.MathExprParser import MathExprParser
from .Parser.MathExprLexer import MathExprLexer
from .Parser.TensorEvalVisitor import TensorEvalVisitor
class VideoMathNode(io.ComfyNode):
"""
This node enables the use of math expressions on Latents.
inputs:
a, b, c, d:
Latent, bound to variables with the same name. Defaults to zero latent if not provided.
w, x, y, z:
Floats, bound to variables of the expression. Defaults to 0.0 if not provided.
Latent expression:
String, describing expression to aply to latents.
outputs:
LATENT:
Returns a LATENT object that contains the result of the math expression applied to the input conditionings.
Enables math expressions on Video (images + audio).
Inputs:
a, b, c, d: Video inputs (b, c, d default to zero if not provided)
w, x, y, z: Float variables for expressions
Audio: Expression for audio component
Images: Expression for image component
Outputs:
VIDEO: Result of applying expressions to input videos
"""
def __init__(self):
pass
@classmethod
def define_schema(cls) -> io.Schema:
"""
"""
return io.Schema(
node_id="mrmth_VideoMathNode",
display_name="Video math",
@@ -45,10 +34,10 @@ class VideoMathNode(io.ComfyNode):
io.Video.Input(id="b", optional=True),
io.Video.Input(id="c", optional=True),
io.Video.Input(id="d", optional=True),
io.Float.Input(id="w", default=0.0,optional=True, force_input=True),
io.Float.Input(id="x", default=0.0,optional=True, force_input=True),
io.Float.Input(id="y", default=0.0,optional=True, force_input=True),
io.Float.Input(id="z", default=0.0,optional=True, force_input=True),
io.Float.Input(id="w", default=0.0, optional=True, force_input=True),
io.Float.Input(id="x", default=0.0, optional=True, force_input=True),
io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True),
io.Float.Input(id="z", default=0.0, optional=True, force_input=True),
io.String.Input(id="Audio", default="a*(1-w)+b*w", tooltip="Expression to apply on audio part of video"),
io.String.Input(id="Images", default="a*(1-w)+b*w", tooltip="Expression to apply on image part of video"),
],
@@ -57,109 +46,62 @@ class VideoMathNode(io.ComfyNode):
],
)
#RETURN_NAMES = ("image_output_name",)
tooltip = cleandoc(__doc__)
#OUTPUT_NODE = False
#OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node
@classmethod
def execute(cls, Audio,Images, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0) -> io.NodeOutput:
def execute(cls, Audio, Images, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0) -> io.NodeOutput:
ac = a.get_components()
bc = b.get_components() if b is not None else VideoComponents(images=torch.zeros_like(ac.images), audio={'waveform':torch.zeros_like(ac.audio['waveform']),'sample_rate':ac.audio['sample_rate']}, frame_rate=ac.frame_rate,metadata=None)
cc = c.get_components() if c is not None else VideoComponents(images=torch.zeros_like(ac.images), audio={'waveform':torch.zeros_like(ac.audio['waveform']),'sample_rate':ac.audio['sample_rate']}, frame_rate=ac.frame_rate,metadata=None)
dc = d.get_components() if d is not None else VideoComponents(images=torch.zeros_like(ac.images), audio={'waveform':torch.zeros_like(ac.audio['waveform']),'sample_rate':ac.audio['sample_rate']}, frame_rate=ac.frame_rate,metadata=None)
bc = b.get_components() if b is not None else make_zero_like(ac)
cc = c.get_components() if c is not None else make_zero_like(ac)
dc = d.get_components() if d is not None else make_zero_like(ac)
# Process images (permute to B, C, H, W)
imgs_a = ac.images.permute(0, 3, 1, 2)
imgs_b = bc.images.permute(0, 3, 1, 2)
imgs_c = cc.images.permute(0, 3, 1, 2)
imgs_d = dc.images.permute(0, 3, 1, 2)
# permute images to B, C, H, W
ac.images = ac.images.permute(0, 3, 1, 2)
bc.images = bc.images.permute(0, 3, 1, 2)
cc.images = cc.images.permute(0, 3, 1, 2)
dc.images = dc.images.permute(0, 3, 1, 2)
B = getIndexTensorAlongDim(ac.images, 0)
C = getIndexTensorAlongDim(ac.images, 1)
X = getIndexTensorAlongDim(ac.images, 3) # W
Y = getIndexTensorAlongDim(ac.images, 2) # H
W = torch.full_like(Y, ac.images.shape[3], dtype=torch.float32)
H = torch.full_like(Y, ac.images.shape[2], dtype=torch.float32)
R = torch.full_like(Y, float(ac.frame_rate), dtype=torch.float32)
T = torch.full_like(Y, ac.images.shape[0], dtype=torch.float32)
variables = {'a': ac.images, 'b': bc.images, 'c': cc.images, 'd': dc.images, 'w': w, 'x': x, 'y': y, 'z': z,
'X':X,'Y':Y,
'B':B,'frame':B,
'W':W,'width':W,
'H':H,'height':H,
'C':C,'channel':C,
'R':R,'frame_rate':R,
'frame_count':ac.images.shape[0],
'N':ac.images.shape[1],'channel_count':ac.images.shape[1]}
input_stream = InputStream(Images)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
parser.addErrorListener(ThrowingErrorListener())
tree = parser.expr()
visitor = TensorEvalVisitor(variables, ac.images.shape)
imgs = visitor.visit(tree)
# permute back to B, H, W, C
imgs = imgs.permute(0, 2, 3, 1)
B = getIndexTensorAlongDim(ac.audio['waveform'], 0)
C = getIndexTensorAlongDim(ac.audio['waveform'], 1)
S = getIndexTensorAlongDim(ac.audio['waveform'], 2)
R = torch.full_like(S, ac.audio['sample_rate'], dtype=torch.float32)
T = torch.full_like(S, ac.audio['waveform'].shape[2], dtype=torch.float32)
N= ac.audio['waveform'].shape[1]
variables = {
'a': ac.audio['waveform'], 'b': bc.audio['waveform'], 'c': cc.audio['waveform'], 'd': dc.audio['waveform'],
img_vars = {
'a': imgs_a, 'b': imgs_b, 'c': imgs_c, 'd': imgs_d,
'w': w, 'x': x, 'y': y, 'z': z,
'B': B, 'batch': B,
'C': C, 'channel': C,
'S': S, 'sample': S,
'R': R, 'sample_rate': R,
'T': T, 'sample_count': T,
'N': N, 'channel_count': N
'X': getIndexTensorAlongDim(imgs_a, 3),
'Y': getIndexTensorAlongDim(imgs_a, 2),
'B': getIndexTensorAlongDim(imgs_a, 0), 'frame': getIndexTensorAlongDim(imgs_a, 0),
'C': getIndexTensorAlongDim(imgs_a, 1), 'channel': getIndexTensorAlongDim(imgs_a, 1),
'W': imgs_a.shape[3], 'width': imgs_a.shape[3],
'H': imgs_a.shape[2], 'height': imgs_a.shape[2],
'R': float(ac.frame_rate), 'frame_rate': float(ac.frame_rate),
'T': imgs_a.shape[0], 'frame_count': imgs_a.shape[0],
'N': imgs_a.shape[1], 'channel_count': imgs_a.shape[1],
}
input_stream = InputStream(Audio)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
tree = parser.expr()
result_imgs = eval_tensor_expr(Images, img_vars, imgs_a.shape)
result_imgs = result_imgs.permute(0, 2, 3, 1) # Back to B, H, W, C
visitor = TensorEvalVisitor(variables, ac.audio['waveform'].shape)
result_tensor = visitor.visit(tree)
# Process audio
audio_a = ac.audio['waveform']
audio_b = bc.audio['waveform']
audio_c = cc.audio['waveform']
audio_d = dc.audio['waveform']
# Create output dictionary with the same sample rate
audioo = {
'waveform': result_tensor,
'sample_rate': ac.audio['sample_rate']
audio_vars = {
'a': audio_a, 'b': audio_b, 'c': audio_c, 'd': audio_d,
'w': w, 'x': x, 'y': y, 'z': z,
'B': getIndexTensorAlongDim(audio_a, 0), 'batch': getIndexTensorAlongDim(audio_a, 0),
'C': getIndexTensorAlongDim(audio_a, 1), 'channel': getIndexTensorAlongDim(audio_a, 1),
'S': getIndexTensorAlongDim(audio_a, 2), 'sample': getIndexTensorAlongDim(audio_a, 2),
'R': ac.audio['sample_rate'], 'sample_rate': ac.audio['sample_rate'],
'T': audio_a.shape[2], 'sample_count': audio_a.shape[2],
'N': audio_a.shape[1], 'channel_count': audio_a.shape[1],
}
result_audio = eval_tensor_expr(Audio, audio_vars, audio_a.shape)
out = VideoFromComponents(VideoComponents(images=imgs, audio=audioo, frame_rate=ac.frame_rate,metadata=ac.metadata))
return (out,)
"""
The node will always be re executed if any of the inputs change but
this method can be used to force the node to execute again even when the inputs don't change.
You can make this node return a number or a string. This value will be compared to the one returned the last time the node was
executed, if it is different the node will be executed again.
This method is used in the core repo for the LoadImage node where they return the image hash as a string, if the image hash
changes between executions the LoadImage node is executed again.
"""
#@classmethod
#def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen):
# return ""
output = VideoFromComponents(VideoComponents(
images=result_imgs,
audio={'waveform': result_audio, 'sample_rate': ac.audio['sample_rate']},
frame_rate=ac.frame_rate,
metadata=ac.metadata
))
return (output,)
+130 -30
View File
@@ -1,51 +1,151 @@
from antlr4.error.ErrorListener import ErrorListener
from antlr4 import InputStream
from .Parser.MathExprLexer import MathExprLexer
from .Parser.MathExprParser import MathExprParser
from antlr4 import CommonTokenStream
from antlr4 import InputStream, CommonTokenStream
import torch
def getIndexTensorAlongDim(tensor, dim):
shape = tensor.shape
from .Parser.MathExprLexer import MathExprLexer
from .Parser.MathExprParser import MathExprParser
# Create values: shape (size of dim)
values = torch.arange(shape[dim], dtype=torch.float32)
# Reshape values to align with the target dimension
view_shape = [1] * len(shape)
view_shape[dim] = shape[dim]
values = values.view(*view_shape)
# Broadcast to full shape
return values.expand(*shape)
def time_to_freq(element: torch.Tensor) -> torch.Tensor:
if element.ndim < 2:
raise ValueError("FFT requires at least 2 dimensions (Batch, Channel)")
dims = tuple(range(2, element.ndim))
return torch.fft.fftn(element, dim=dims)
def freq_to_time(element: torch.Tensor) -> torch.Tensor:
if element.ndim < 2:
raise ValueError("IFFT requires at least 2 dimensions (Batch, Channel)")
dims = tuple(range(2, element.ndim))
return torch.fft.ifftn(element, dim=dims).real
class ThrowingErrorListener(ErrorListener):
"""Error listener that raises ValueError on syntax errors."""
def syntaxError(self, recognizer, offendingSymbol, line, column, msg, e):
raise ValueError(f"Syntax error in expression at line {line}, col {column}: {msg}")
def parse_expr(expr: str):
"""Parse a math expression and return the parse tree."""
input_stream = InputStream(expr)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
parser.addErrorListener(ThrowingErrorListener())
return parser.expr()
def eval_tensor_expr(expr: str, variables: dict, shape: tuple, device=None):
"""Parse and evaluate a tensor math expression.
Args:
expr: Math expression string
variables: Dict of variable names to tensor/scalar values
shape: Shape tuple for the TensorEvalVisitor
device: Optional device override
Returns:
Result tensor from evaluating the expression
"""
from .Parser.TensorEvalVisitor import TensorEvalVisitor
tree = parse_expr(expr)
visitor = TensorEvalVisitor(variables, shape, device=device)
return visitor.visit(tree)
def eval_tensor_expr_with_tree(tree, variables: dict, shape: tuple, device=None):
"""Evaluate a pre-parsed expression tree with TensorEvalVisitor."""
from .Parser.TensorEvalVisitor import TensorEvalVisitor
visitor = TensorEvalVisitor(variables, shape, device=device)
return visitor.visit(tree)
def eval_float_expr(expr: str, variables: dict):
"""Parse and evaluate a float math expression."""
from .Parser.FloatEvalVisitor import FloatEvalVisitor
tree = parse_expr(expr)
visitor = FloatEvalVisitor(variables)
return visitor.visit(tree)
def eval_float_expr_with_tree(tree, variables: dict):
"""Evaluate a pre-parsed expression tree with FloatEvalVisitor."""
from .Parser.FloatEvalVisitor import FloatEvalVisitor
visitor = FloatEvalVisitor(variables)
return visitor.visit(tree)
def getIndexTensorAlongDim(tensor, dim):
"""Create a tensor of indices along a dimension, broadcasted to full shape."""
shape = tensor.shape
values = torch.arange(shape[dim], dtype=torch.float32, device=tensor.device)
view_shape = [1] * len(shape)
view_shape[dim] = shape[dim]
values = values.view(*view_shape)
return values.expand(*shape)
def comonLazy(expr, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0):
variables = {'a':a,'b':b,'c':c,'d':d,'w':w,'x':x,'y':y,'z':z}
"""Determine which lazy inputs are needed based on expression variables."""
variables = {'a': a, 'b': b, 'c': c, 'd': d, 'w': w, 'x': x, 'y': y, 'z': z}
need_eval = []
input_stream = InputStream(expr)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
stream.fill()
for token in filter(lambda t: t.type == MathExprParser.VARIABLE, stream.tokens):
if token.text in variables and variables[token.text] is None:
need_eval.append(token.text)
return need_eval
return need_eval
def make_zero_like(ref):
"""
Create a zero-initialized version of the reference object, maintaining its structure.
Handles Conditioning, Audio, Latent, VideoComponents, and Tensors.
"""
if ref is None:
return None
# raw torch tensor
if torch.is_tensor(ref):
return torch.zeros_like(ref)
# Conditioning: list of lists [[tensor, dict]]
if isinstance(ref, list) and len(ref) > 0 and isinstance(ref[0], list) and len(ref[0]) >= 2:
ref_tensor = ref[0][0]
# Ensure it's a tensor-like structure
if torch.is_tensor(ref_tensor):
ref_pooled = ref[0][1].get("pooled_output")
return [[
torch.zeros_like(ref_tensor),
{"pooled_output": torch.zeros_like(ref_pooled) if ref_pooled is not None else None}
]]
# Audio or Latent: dict
if isinstance(ref, dict):
if 'waveform' in ref: # Audio
return {
'waveform': torch.zeros_like(ref['waveform']),
'sample_rate': ref['sample_rate']
}
if 'samples' in ref: # Latent
return {
'samples': torch.zeros_like(ref['samples'])
}
# VideoComponents or other objects with images/audio attributes
if hasattr(ref, 'images') and hasattr(ref, 'audio'):
# Dynamically create same type (e.g. VideoComponents)
return type(ref)(
images=torch.zeros_like(ref.images),
audio={'waveform': torch.zeros_like(ref.audio['waveform']), 'sample_rate': ref.audio['sample_rate']},
frame_rate=ref.frame_rate,
metadata=None
)
return None
# Legacy FFT functions (kept for backward compatibility, but now unused)
def time_to_freq(element: torch.Tensor) -> torch.Tensor:
if element.ndim < 2:
raise ValueError("FFT requires at least 2 dimensions (Batch, Channel)")
dims = tuple(range(2, element.ndim))
return torch.fft.fftn(element, dim=dims)
def freq_to_time(element: torch.Tensor) -> torch.Tensor:
if element.ndim < 2:
raise ValueError("IFFT requires at least 2 dimensions (Batch, Channel)")
dims = tuple(range(2, element.ndim))
return torch.fft.ifftn(element, dim=dims).real
+25 -47
View File
@@ -1,21 +1,13 @@
from antlr4.error.ErrorListener import ErrorListener
from antlr4 import InputStream
from .Parser.MathExprLexer import MathExprLexer
from .Parser.MathExprParser import MathExprParser
from .Parser.TensorEvalVisitor import TensorEvalVisitor
from antlr4 import CommonTokenStream
import torch
from .helper_functions import getIndexTensorAlongDim, ThrowingErrorListener
import comfy.utils
from .helper_functions import getIndexTensorAlongDim, parse_expr, eval_tensor_expr_with_tree
def calculate_patches(Model, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0):
# Parse expression
input_stream = InputStream(Model)
lexer = MathExprLexer(input_stream)
stream = CommonTokenStream(lexer)
parser = MathExprParser(stream)
parser.addErrorListener(ThrowingErrorListener())
tree = parser.expr()
"""Calculate model weight patches by applying math expression to state dicts."""
# Parse expression once
tree = parse_expr(Model)
sd_a = a.model.state_dict()
sd_b = b.model.state_dict() if b is not None else {}
@@ -24,55 +16,41 @@ def calculate_patches(Model, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0
patches = {}
layer_count = len(sd_a)
pbar = comfy.utils.ProgressBar(layer_count)
# Iterate over all keys in the main model 'a'
for i, (key, tens_a) in enumerate(sd_a.items()):
# Get corresponding tensors from other models, defaulting to zeros if missing or models not provided
tens_b = sd_b.get(key, None)
if tens_b is None: tens_b = torch.zeros_like(tens_a,device=tens_a.device)
else: tens_b = tens_b.to(tens_a.device)
# Get tensors from other models, default to zeros
tens_b = sd_b.get(key)
tens_b = tens_b.to(tens_a.device) if tens_b is not None else torch.zeros_like(tens_a)
tens_c = sd_c.get(key, None)
if tens_c is None: tens_c = torch.zeros_like(tens_a,device=tens_a.device)
else: tens_c = tens_c.to(tens_a.device)
tens_c = sd_c.get(key)
tens_c = tens_c.to(tens_a.device) if tens_c is not None else torch.zeros_like(tens_a)
tens_d = sd_d.get(key, None)
if tens_d is None: tens_d = torch.zeros_like(tens_a,device=tens_a.device)
else: tens_d = tens_d.to(tens_a.device)
tens_d = sd_d.get(key)
tens_d = tens_d.to(tens_a.device) if tens_d is not None else torch.zeros_like(tens_a)
# Variables for the visitor
# Build variables
variables = {
'a': tens_a,
'b': tens_b,
'c': tens_c,
'd': tens_d,
'a': tens_a, 'b': tens_b, 'c': tens_c, 'd': tens_d,
'w': w, 'x': x, 'y': y, 'z': z,
'L': i, 'layer': i,
'LC': layer_count, 'layer_count': layer_count
'LC': layer_count, 'layer_count': layer_count,
}
# Add dimension index tensors
for dim_idx in range(tens_a.ndim):
idx_tensor = getIndexTensorAlongDim(tens_a, dim_idx)
idx_tensor = idx_tensor.to(tens_a.device)
variables[f'D{dim_idx}'] = idx_tensor
variables[f'dim_{dim_idx}'] = idx_tensor
idx_tensor = getIndexTensorAlongDim(tens_a, dim_idx)
variables[f'D{dim_idx}'] = idx_tensor
variables[f'dim_{dim_idx}'] = idx_tensor
visitor = TensorEvalVisitor(variables, tens_a.shape)
result_tensor = visitor.visit(tree)
result_tensor = eval_tensor_expr_with_tree(tree, variables, tens_a.shape)
# Calculate difference for patching
# The patch should be: result - original
# Because ComfyUI applies: original + patch
# Calculate patch (diff from original)
diff = result_tensor - tens_a
# Allow skipping zero patches to save memory
if torch.all(diff == 0):
continue
# Store patch. ComfyUI expects { key: (tensor,) } usually
patches[key] = (diff,)
# Skip zero patches to save memory
if not torch.all(diff == 0):
patches[key] = (diff,)
pbar.update(1)
+8 -12
View File
@@ -1,21 +1,17 @@
import os
import sys
# Ensure test discovery (Visual Studio / pytest) can import the package.
# conftest.py is imported during collection, so top-level path changes affect discovery.
_here = os.path.abspath(os.path.dirname(__file__))
_project_root = os.path.abspath(os.path.join(_here, os.pardir))
_src_path = os.path.join(_project_root, "src")
_comfy_root = os.path.abspath(os.path.join(_project_root, os.pardir, os.pardir))
if os.path.isdir(_src_path):
_path_to_add = _src_path
else:
_path_to_add = _project_root
if _path_to_add not in sys.path:
sys.path.insert(0, _path_to_add)
# also make it visible to subprocesses that inspect PYTHONPATH
os.environ["PYTHONPATH"] = _path_to_add + os.pathsep + os.environ.get("PYTHONPATH", "")
# Add project root and ComfyUI root to sys.path
# This ensures import comfy_api and import comfy work.
for p in [_project_root, _comfy_root]:
if p not in sys.path:
sys.path.insert(0, p)
# also make it visible to subprocesses that inspect PYTHONPATH
os.environ["PYTHONPATH"] = p + os.pathsep + os.environ.get("PYTHONPATH", "")
import pytest