Files
mcDandy-more_math/more_math/NoiseMathNode.py
T
2026-03-26 14:23:33 +01:00

139 lines
5.6 KiB
Python

from .helper_functions import generate_dim_variables, as_tensor, parse_expr, getIndexTensorAlongDim, make_zero_like, get_v_variable, get_f_variable, checkLazyNew
from comfy_api.latest import io
import torch
from .Parser.UnifiedMathVisitor import UnifiedMathVisitor
from .Stack import MrmthStack
from .ParseTree import MrmthParseTree
import copy
class NoiseMathNode(io.ComfyNode):
"""
This node enables the use of math expressions on noise generators.
inputs:
a, b, c, d:
Noise generators.
w, x, y, z:
Floats.
Noise expression:
The expression to apply on those noise generators.
Note that variables X, Y, W, H, C, batch, batch_count, input_latent refer to input_latent.
outputs:
NOISE:
The resulting noise generator.
"""
@classmethod
def define_schema(cls) -> io.Schema:
return io.Schema(
node_id="mrmth_ag_NoiseMathNode",
display_name="Noise math",
category="More math",
inputs=[
io.Autogrow.Input(id="V",template=io.Autogrow.TemplatePrefix(io.Noise.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.MultiType.Input(
io.String.Input("Noise", default="a*(1-w)+b*w", multiline=False),
types=[io.String,MrmthParseTree],
tooltip="Expression for noise",
),
io.Boolean.Input(
id="remember_stack",
default=False,
display_name="Remember stack across batch",
tooltip=(
"If enabled, stack is copied at output leading to changes being remembered during batch operations (node runs multiple times in sucession). If disabled each batch gets it's own copy of the stack."
),
),
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
],
outputs=[
io.Noise.Output(),
MrmthStack.Output(),
],
)
@classmethod
def check_lazy_status(cls, Noise, V, F,remember_stack=False,stack={}):
return checkLazyNew(Noise,V,F)
@classmethod
def execute(cls, Noise, V,F,remember_stack=False,stack={}):
stack = stack if remember_stack else (copy.deepcopy(stack) if stack is not None else {})
executer = NoiseExecutor(V,F, Noise,stack)
stack = stack if remember_stack else copy.deepcopy(stack)
return (executer,stack)
class NoiseExecutor:
def __init__(self, V,F, expr,stack):
self.V = V
self.F = F
if isinstance(expr,str):
self.tree = parse_expr(expr)
else:
self.tree = expr
self.stack = stack
seed = -1
def generate_noise(self, input_latent: torch.Tensor) -> torch.Tensor:
samples = input_latent["samples"]
vals = {v: (self.V[v].generate_noise(input_latent) if self.V[v] is not None else make_zero_like(samples)) for v in self.V}
ndim = samples.ndim
batch_dim = 0
channel_dim = -3
height_dim = -2
width_dim = -1
time_dim = None
if ndim >= 5:
time_dim = -4
frame_count = samples.shape[time_dim] if time_dim is not None else samples.shape[batch_dim]
B = getIndexTensorAlongDim(samples, batch_dim)
W = getIndexTensorAlongDim(samples, width_dim)
H = getIndexTensorAlongDim(samples, height_dim)
C = getIndexTensorAlongDim(samples, channel_dim)
variables = {
"a": vals.get("V0") if "V0" in vals else make_zero_like(samples),
"b": vals.get("V1") if "V1" in vals else make_zero_like(samples),
"c": vals.get("V2") if "V2" in vals else make_zero_like(samples),
"d": vals.get("V3") if "V3" in vals else make_zero_like(samples),
"w": self.F.get("F0", 0.0),
"x": self.F.get("F1", 0.0),
"y": self.F.get("F2", 0.0),
"z": self.F.get("F3", 0.0),
"B": B, "batch": B,
"X": W, "width": float(samples.shape[width_dim]),
"Y": H, "height": float(samples.shape[height_dim]),
"C": C, "channel": C,
"W": float(samples.shape[width_dim]), "H": float(samples.shape[height_dim]), "I": samples,
"T": float(frame_count), "N": float(samples.shape[channel_dim]),
"batch_count": float(samples.shape[batch_dim]), "channel_count": float(samples.shape[channel_dim]),
"input_latent": samples,
} | generate_dim_variables(samples) | vals | self.F
v_stacked, v_cnt = get_v_variable(vals)
if v_stacked is not None:
variables["V"] = v_stacked
variables["Vcnt"] = float(v_cnt)
variables["V_count"] = float(v_cnt)
f_stacked, f_cnt = get_f_variable(self.F)
if f_stacked is not None:
variables["F"] = f_stacked
variables["Fcnt"] = float(f_cnt)
variables["F_count"] = float(f_cnt)
if time_dim is not None:
F = getIndexTensorAlongDim(samples, time_dim)
variables.update({"frame": F, "frame_count": frame_count})
visitor = UnifiedMathVisitor(variables, samples.shape,samples.device,state_storage=self.stack)
result = visitor.visit(self.tree)
result = as_tensor(result, samples.shape)
return result