massive refactor (AI)
This commit is contained in:
+43
-55
@@ -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},)
|
||||
|
||||
@@ -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
@@ -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
@@ -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 ""
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user