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

223 lines
9.0 KiB
Python

from inspect import cleandoc
from comfy_api.latest import io
from .helper_functions import (
generate_dim_variables,
getIndexTensorAlongDim,
parse_expr,
as_tensor,
normalize_to_common_shape,
make_zero_like,
get_v_variable,
get_f_variable,
checkLazyNew
)
from .Parser.UnifiedMathVisitor import UnifiedMathVisitor
import torch
from comfy.nested_tensor import NestedTensor
from .Stack import MrmthStack
from .ParseTree import MrmthParseTree
import copy
class LatentMathNode(io.ComfyNode):
"""
This node enables the use of math expressions on Latents using Autogrow inputs.
"""
def __init__(self):
pass
@classmethod
def define_schema(cls) -> io.Schema:
""" """
return io.Schema(
node_id="mrmth_ag_LatentMathNode",
display_name="Latent math",
category="More math",
inputs=[
io.Autogrow.Input(id="V",template=io.Autogrow.TemplatePrefix(io.Latent.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("Expression", default="I0*(1-F0)+I1*F0", multiline=False),
types=[io.String,MrmthParseTree],
tooltip="Expression to apply on input latents",
),
io.Combo.Input(
id="length_mismatch",
options=["do nothing","error","tile", "pad"],
display_name="on size mismatch",
default="error",
tooltip="How to handle mismatched latent batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
),
io.Int.Input(id="batching"),
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.Latent.Output(is_output_list=True),
MrmthStack.Output(),
],
)
tooltip = cleandoc(__doc__)
@classmethod
def check_lazy_status(cls, Expression, V, F,batching, length_mismatch="tile",remember_stack=False,stack={}):
return checkLazyNew(Expression,V,F)
@classmethod
def execute(cls, V, F, Expression,batching, length_mismatch="tile",remember_stack=False,stack={}) -> io.NodeOutput:
# Determine reference latent
ref_latent = None
for lat in V.values():
if lat is not None:
ref_latent = lat
break
if ref_latent is None:
raise ValueError("At least one input is required.")
stack = stack if remember_stack else (copy.deepcopy(stack) if stack is not None else {})
# Identify if any input is a NestedTensor and track original sizes for restoration
stacked = False
orig_split_sizes = None
# Check all present inputs for nested tensors
for item in V.values():
if item is not None:
samples = item.get("samples")
if getattr(samples, "is_nested", False):
stacked = True
# Store original split sizes (batch dimension) - assume all nested inputs share structure if mixed?
# Or just take from the first one found.
orig_split_sizes = [t.shape[0] for t in samples.tensors]
break
# Flatten nested tensors in V
if stacked:
for k, val in V.items():
if val is not None and getattr(val.get("samples"), "is_nested", False):
new_val = val.copy()
new_val["samples"] = torch.cat(new_val["samples"].tensors, dim=0)
V[k] = new_val
# Identify all present tensors and their keys
tensor_keys = [k for k, v in V.items() if v is not None]
at_list = [V[k]["samples"] for k in tensor_keys]
# Normalize all together
normalized_samples = normalize_to_common_shape(*at_list, mode=length_mismatch)
V_norm_samples = dict(zip(tensor_keys, normalized_samples))
ae = V_norm_samples.get("V0", make_zero_like(normalized_samples[0]))
be = V_norm_samples.get("V1", make_zero_like(ae))
ce = V_norm_samples.get("V2", make_zero_like(ae))
de = V_norm_samples.get("V3", make_zero_like(ae))
# Ensure legacy are normalized
ae, be, ce, de = normalize_to_common_shape(ae, be, ce, de, mode=length_mismatch)
if(length_mismatch == "error"):
for name in tensor_keys:
if V[name]["samples"].shape[0] != ae.shape[0]:
raise ValueError(f"Input '{name}' has shape {V[name]['samples'].shape[0]}, expected {ae.shape[0]} to match input.")
# parse expression once
tree = None
if isinstance(Expression,str):
tree = parse_expr(Expression)
else:
tree = Expression
ndim = ae.ndim
batch_dim = 0
channel_dim = -3
height_dim = -2
width_dim = -1
time_dim = None
if ndim >= 5:
time_dim = -4
frame_count = ae.shape[time_dim] if time_dim is not None else ae.shape[batch_dim]
variables = {
"a": ae, "b": be, "c": ce, "d": de,
"w": F.get("F0", 0.0) if F.get("F0") is not None else 0.0,
"x": F.get("F1", 0.0) if F.get("F1") is not None else 0.0,
"y": F.get("F2", 0.0) if F.get("F2") is not None else 0.0,
"z": F.get("F3", 0.0) if F.get("F3") is not None else 0.0,
"X": getIndexTensorAlongDim(ae, width_dim),
"Y": getIndexTensorAlongDim(ae, height_dim),
"B": getIndexTensorAlongDim(ae, batch_dim),
"batch": getIndexTensorAlongDim(ae, batch_dim),
"C": getIndexTensorAlongDim(ae, channel_dim),
"channel": getIndexTensorAlongDim(ae, channel_dim),
"W": float(ae.shape[width_dim]),
"width": float(ae.shape[width_dim]),
"H": float(ae.shape[height_dim]),
"height": float(ae.shape[height_dim]),
"T": float(frame_count),
"batch_count": float(ae.shape[batch_dim]),
"N": float(ae.shape[channel_dim]),
"channel_count": float(ae.shape[channel_dim]),
} | generate_dim_variables(ae)
if time_dim is not None:
F_idx = getIndexTensorAlongDim(ae, time_dim)
variables.update({"frame_idx": F_idx, "frame": F_idx, "frame_count": frame_count})
# Add all dynamic inputs
variables.update(V_norm_samples)
v_stacked, v_cnt = get_v_variable(V_norm_samples, length_mismatch=length_mismatch)
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(F)
if f_stacked is not None:
variables["F"] = f_stacked
variables["Fcnt"] = float(f_cnt)
variables["F_count"] = float(f_cnt)
for k, v in F.items():
variables[k] = v if v is not None else 0.0
visitor = UnifiedMathVisitor(variables, ae.shape,ae.device,state_storage=stack)
result_t = as_tensor(visitor.visit(tree), ae.shape)
result_latent = ref_latent.copy()
if(batching>0):
res = torch.split(result_t,batching)
results=[]
results1=[]
for i in range(len(res)):
result_tensor = res[i] if i<len(res) else torch.zeros([1])
results.append(result_tensor)
for result_t in results:
rl = result_latent.copy()
if stacked and orig_split_sizes is not None:
# Restore original split sizes
try:
rl["samples"] = NestedTensor(torch.split(result_t, orig_split_sizes, dim=0))
except Exception:
# Fallback if split fails (e.g. result shape changed)
rl["samples"] = result_t
else:
rl["samples"] = result_t
results1.append(rl)
stack = stack if remember_stack else copy.deepcopy(stack)
return (results1,stack)
rl = result_latent.copy()
rl["samples"] = result_t
stack = stack if remember_stack else copy.deepcopy(stack)
return ([rl],stack)