From c207f97e7c3b13136c5a7470720079ec42d43807 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Sat, 8 Aug 2026 21:19:16 +0200 Subject: [PATCH] AI: fix nested tensor phase3 Fix math actually mathing --- more_math/LatentMathNode.py | 13 +++++++++++-- more_math/Parser/UnifiedMathVisitor.py | 17 ++++++++++------- more_math/helper_functions.py | 12 ++++++++++-- tests/test_more_math.py | 16 ++++++++++++++++ 4 files changed, 47 insertions(+), 11 deletions(-) diff --git a/more_math/LatentMathNode.py b/more_math/LatentMathNode.py index ccad5cc..7d92735 100644 --- a/more_math/LatentMathNode.py +++ b/more_math/LatentMathNode.py @@ -13,6 +13,7 @@ from .helper_functions import ( ) from .Parser.UnifiedMathVisitor import UnifiedMathVisitor import torch +from comfy.nested_tensor import NestedTensor from .Stack import MrmthStack from .ParseTree import MrmthParseTree import copy @@ -222,8 +223,16 @@ class LatentMathNode(io.ComfyNode): result_t = as_tensor(raw_result, ae.shape) result_latent = ref_latent.copy() - # If result is a list of tensors (e.g. user wrote just V1 for a NestedTensor), - # return each component as a separate latent output. + + # If the visitor produced a NestedTensor (whole-tensor arithmetic on a NestedTensor + # input), propagate that through as a single downstream-compatible latent. + if getattr(result_t, "is_nested", False): + rl = result_latent.copy() + rl["samples"] = result_t + stack = stack if remember_stack else copy.deepcopy(stack) + return ([rl], stack) + + # If the visitor produced a list/tuple, emit each element as a separate latent. if isinstance(result_t, (list, tuple)) and result_t and isinstance(result_t[0], torch.Tensor): results = [] for comp in result_t: diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index 38594eb..d6b6b15 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -4,6 +4,7 @@ import torch import math import inspect import torch.nn.functional as F +from comfy.nested_tensor import NestedTensor from . import optical_flow_utils as ofu from .Func import LambdaFunction from antlr4 import TerminalNode @@ -160,25 +161,25 @@ class UnifiedMathVisitor(MathExprVisitor): # Handle tensor operations if self._is_nested_tensor(a) or self._is_nested_tensor(b): - # Treat NestedTensors as lists of components for arithmetic + # NestedTensor arithmetic is performed component-wise and re-wrapped as NestedTensor a_comps = a.tensors if self._is_nested_tensor(a) else a b_comps = b.tensors if self._is_nested_tensor(b) else b if self._is_nested_tensor(a) and self._is_nested_tensor(b): if len(a_comps) != len(b_comps): raise ValueError("NestedTensor component counts must match") - return [self._bin_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(a_comps, b_comps)] + return NestedTensor([self._bin_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(a_comps, b_comps)]) if self._is_nested_tensor(a): if self._is_list(b_comps): if len(a_comps) != len(b_comps): raise ValueError("NestedTensor and list length mismatch") - return [self._bin_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(a_comps, b_comps)] - return [self._bin_op(x, b_comps, torch_op, scalar_op, ctx) for x in a_comps] + return NestedTensor([self._bin_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(a_comps, b_comps)]) + return NestedTensor([self._bin_op(x, b_comps, torch_op, scalar_op, ctx) for x in a_comps]) # b is nested if self._is_list(a_comps): if len(a_comps) != len(b_comps): raise ValueError("List and NestedTensor length mismatch") - return [self._bin_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(a_comps, b_comps)] - return [self._bin_op(a_comps, x, torch_op, scalar_op, ctx) for x in b_comps] + return NestedTensor([self._bin_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(a_comps, b_comps)]) + return NestedTensor([self._bin_op(a_comps, x, torch_op, scalar_op, ctx) for x in b_comps]) if self._is_tensor(a) or self._is_tensor(b): orig_dtype = None @@ -220,6 +221,8 @@ class UnifiedMathVisitor(MathExprVisitor): if self._is_plain_tensor(a) and a.numel() == 1: a = float(a.flatten()[0].item()) + if self._is_nested_tensor(a): + return NestedTensor([self._unary_op(x, torch_op, scalar_op) for x in a.tensors]) if self._is_list(a): return [self._unary_op(x, torch_op, scalar_op) for x in a] if self._is_tensor(a): @@ -246,7 +249,7 @@ class UnifiedMathVisitor(MathExprVisitor): return float(res.item()) return res if self._is_nested_tensor(val): - return list_op(val.tensors) + return NestedTensor([torch_op(t) for t in val.tensors]) if self._is_list(val): return list_op(val) return val diff --git a/more_math/helper_functions.py b/more_math/helper_functions.py index d1b0bd0..f5827ab 100644 --- a/more_math/helper_functions.py +++ b/more_math/helper_functions.py @@ -25,7 +25,7 @@ def as_tensor(value, shape): value = (value,) return torch.broadcast_to(torch.Tensor(value).to(dtype=torch.float32), shape).contiguous() if isinstance(value, (list, tuple)) and value and isinstance(value[0], torch.Tensor): - # Keep lists of tensors as-is (e.g. returning a NestedTensor as a list of components) + # Keep lists of tensors as-is (e.g. intentionally returning a list of latents). return value return torch.cat(value) @@ -203,12 +203,20 @@ def get_v_variable(v_norm_dict, length_mismatch="error"): """ Collects V0, V1, ... from the dict (ignoring NestedTensor component keys like V0_0, V1_1), and returns them as a list and count. + NestedTensor entries are represented as a list of their components so that + V[idx][comp] indexing works as expected. """ base_keys = sorted( [k for k in v_norm_dict.keys() if re.fullmatch(r"V\d+", k)], key=lambda x: int(x[1:]) ) - listt = [v_norm_dict[k] for k in base_keys] + listt = [] + for k in base_keys: + val = v_norm_dict[k] + if getattr(val, "is_nested", False): + listt.append(list(val.tensors)) + else: + listt.append(val) return listt, len(listt) def get_f_variable(f_dict): diff --git a/tests/test_more_math.py b/tests/test_more_math.py index cf92659..3b99b8b 100644 --- a/tests/test_more_math.py +++ b/tests/test_more_math.py @@ -427,6 +427,14 @@ def test_nested_tensor_support(): assert torch.allclose(samples.tensors[1], torch.full((2, 4, 32, 32), 2.0)) # Component access uses V0_0, V0_1, ... + # Arithmetic on the whole NestedTensor returns a single NestedTensor again + result_list, stack = node.execute(Expression="V0 * 0.5 + 0.1", V={"V0": l_in}, F={}, batching=0) + assert len(result_list) == 1 + samples = result_list[0]["samples"] + assert getattr(samples, "is_nested", False), "Arithmetic on a NestedTensor should return a NestedTensor" + assert torch.allclose(samples.tensors[0], torch.full((1, 4, 32, 32), 0.6)) + assert torch.allclose(samples.tensors[1], torch.full((2, 4, 32, 32), 1.1)) + result_list, stack = node.execute(Expression="V0_0 + 1.0", V={"V0": l_in}, F={}, batching=0) res_0 = result_list[0]["samples"] assert isinstance(res_0, torch.Tensor), f"Expected torch.Tensor, got {type(res_0)}" @@ -457,6 +465,14 @@ def test_nested_tensor_support_position_v1(): assert torch.allclose(samples.tensors[0], torch.full((1, 4, 32, 32), 10.0)) assert torch.allclose(samples.tensors[1], torch.full((2, 4, 32, 32), 20.0)) + # Arithmetic over the whole NestedTensor returns a NestedTensor again + result_list, _ = node.execute(Expression="V1 * 0.5 + 0.1", V={"V0": base_latent, "V1": l_in}, F={}, batching=0) + assert len(result_list) == 1 + samples = result_list[0]["samples"] + assert getattr(samples, "is_nested", False) + assert torch.allclose(samples.tensors[0], torch.full((1, 4, 32, 32), 5.1)) + assert torch.allclose(samples.tensors[1], torch.full((2, 4, 32, 32), 10.1)) + # Component access for a NestedTensor at position V1 result_list, _ = node.execute(Expression="V1_0 + V1_1", V={"V0": base_latent, "V1": l_in}, F={}, batching=0) assert torch.allclose(result_list[0]["samples"], torch.full((2, 4, 32, 32), 30.0))