batching
This commit is contained in:
@@ -14,6 +14,7 @@ from antlr4 import InputStream, CommonTokenStream
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
from .Parser.MathExprParser import MathExprParser
|
||||
import re
|
||||
import torch
|
||||
from .Stack import MrmthStack
|
||||
|
||||
class AudioMathNode(io.ComfyNode):
|
||||
@@ -43,16 +44,17 @@ class AudioMathNode(io.ComfyNode):
|
||||
default="error",
|
||||
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
io.Int.Input(id="batching", default=0),
|
||||
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.Audio.Output(),
|
||||
io.Audio.Output(is_output_list=True),
|
||||
MrmthStack.Output(),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack={}):
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",batching=0,stack={}):
|
||||
|
||||
input_stream = InputStream(Expression)
|
||||
lexer = MathExprLexer(input_stream)
|
||||
@@ -84,7 +86,7 @@ class AudioMathNode(io.ComfyNode):
|
||||
return needed1
|
||||
|
||||
@classmethod
|
||||
def execute(cls, V, F, Expression, length_mismatch="tile",stack={}):
|
||||
def execute(cls, V, F, Expression, length_mismatch="tile",batching=0,stack={}):
|
||||
# Identify all present audio inputs and their keys
|
||||
tensor_keys = [k for k, v in V.items() if v is not None and isinstance(v, dict) and "waveform" in v]
|
||||
if not tensor_keys:
|
||||
@@ -152,4 +154,12 @@ class AudioMathNode(io.ComfyNode):
|
||||
visitor = UnifiedMathVisitor(variables, a_w.shape,a_w.device,state_storage=stack)
|
||||
result = visitor.visit(tree)
|
||||
result = as_tensor(result, a_w.shape)
|
||||
return ({"waveform":result,"sample_rate":sample_rate},stack)
|
||||
|
||||
if batching and batching > 0:
|
||||
res = torch.split(result, batching, dim=0)
|
||||
res_list = []
|
||||
for result_chunk in res:
|
||||
res_list.append({"waveform": result_chunk, "sample_rate": sample_rate})
|
||||
return (res_list, stack)
|
||||
else:
|
||||
return ([{"waveform": result, "sample_rate": sample_rate}], stack)
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import torch
|
||||
import torch
|
||||
from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like, get_v_variable, get_f_variable
|
||||
from .Parser.UnifiedMathVisitor import UnifiedMathVisitor
|
||||
from comfy_api.latest import io
|
||||
|
||||
@@ -5,6 +5,7 @@ from antlr4 import InputStream, CommonTokenStream
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
from .Parser.MathExprParser import MathExprParser
|
||||
import re
|
||||
import torch
|
||||
from .Stack import MrmthStack
|
||||
|
||||
class ImageMathNode(io.ComfyNode):
|
||||
@@ -26,7 +27,7 @@ class ImageMathNode(io.ComfyNode):
|
||||
inputs=[
|
||||
io.Autogrow.Input(id="V",template=io.Autogrow.TemplatePrefix(io.Image.Input("values"), prefix="V", min=1, max=50)),
|
||||
io.Autogrow.Input(id="F", template=io.Autogrow.TemplatePrefix(io.Float.Input("float", default=0.0, optional=True, lazy=True, force_input=True), prefix="F", min=1, max=50)),
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on input images"), # Changed ID to Expression to match AudioMathNode pattern, or keep Image? AudioMathNode used "Expression".
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on input images"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
@@ -34,16 +35,17 @@ class ImageMathNode(io.ComfyNode):
|
||||
default="error",
|
||||
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
io.Int.Input(id="batching", default=0),
|
||||
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.Image.Output(),
|
||||
io.Image.Output(is_output_list=True),
|
||||
MrmthStack.Output(),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack={}):
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",batching=0,stack={}):
|
||||
|
||||
input_stream = InputStream(Expression)
|
||||
lexer = MathExprLexer(input_stream)
|
||||
@@ -75,7 +77,7 @@ class ImageMathNode(io.ComfyNode):
|
||||
return needed1
|
||||
|
||||
@classmethod
|
||||
def execute(cls, V, F, Expression, length_mismatch="error",stack={}):
|
||||
def execute(cls, V, F, Expression, length_mismatch="error",batching=0,stack={}):
|
||||
# I and F are Autogrow.Type which is dict[str, Any]
|
||||
|
||||
# Identify all present tensors and their keys
|
||||
@@ -151,4 +153,12 @@ class ImageMathNode(io.ComfyNode):
|
||||
visitor = UnifiedMathVisitor(variables, ae.shape,ae.device,state_storage=stack)
|
||||
result = visitor.visit(tree)
|
||||
result = as_tensor(result, ae.shape)
|
||||
return (result,stack)
|
||||
|
||||
if batching and batching > 0:
|
||||
res = torch.split(result, batching, dim=0)
|
||||
res_list = []
|
||||
for result_chunk in res:
|
||||
res_list.append(result_chunk)
|
||||
return (res_list, stack)
|
||||
else:
|
||||
return ([result], stack)
|
||||
|
||||
@@ -5,6 +5,7 @@ from antlr4 import InputStream, CommonTokenStream
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
from .Parser.MathExprParser import MathExprParser
|
||||
import re
|
||||
import torch
|
||||
from .Stack import MrmthStack
|
||||
|
||||
|
||||
@@ -35,16 +36,17 @@ class MaskMathNode(io.ComfyNode):
|
||||
default="error",
|
||||
tooltip="How to handle mismatched mask batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
io.Int.Input(id="batching", default=0),
|
||||
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.Mask.Output(),
|
||||
io.Mask.Output(is_output_list=True),
|
||||
MrmthStack.Output(),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack={}):
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",batching=0,stack={}):
|
||||
|
||||
input_stream = InputStream(Expression)
|
||||
lexer = MathExprLexer(input_stream)
|
||||
@@ -76,7 +78,7 @@ class MaskMathNode(io.ComfyNode):
|
||||
return needed1
|
||||
|
||||
@classmethod
|
||||
def execute(cls, V, F, Expression, length_mismatch="tile",stack={}):
|
||||
def execute(cls, V, F, Expression, length_mismatch="tile",batching=0,stack={}):
|
||||
# Identify all present tensors and their keys
|
||||
tensor_keys = [k for k, v in V.items() if v is not None]
|
||||
if not tensor_keys:
|
||||
@@ -146,4 +148,12 @@ class MaskMathNode(io.ComfyNode):
|
||||
visitor = UnifiedMathVisitor(variables, ae.shape,ae.device,state_storage=stack)
|
||||
result = visitor.visit(tree)
|
||||
result = as_tensor(result, ae.shape)
|
||||
return (result,stack)
|
||||
|
||||
if batching and batching > 0:
|
||||
res = torch.split(result, batching, dim=0)
|
||||
res_list = []
|
||||
for result_chunk in res:
|
||||
res_list.append(result_chunk)
|
||||
return (res_list, stack)
|
||||
else:
|
||||
return ([result], stack)
|
||||
|
||||
Reference in New Issue
Block a user