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"), 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",stack={}): return checkLazyNew(Expression,V,F) @classmethod def execute(cls, V, F, Expression,batching, length_mismatch="tile",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 = 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": ae.shape[width_dim], "width": ae.shape[width_dim], "H": ae.shape[height_dim], "height": ae.shape[height_dim], "T": frame_count, "batch_count": ae.shape[batch_dim], "N": ae.shape[channel_dim], "channel_count": 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