AI: fix nested tensor phase3

Fix math actually mathing
This commit is contained in:
mcDandy
2026-08-08 21:19:16 +02:00
parent c5d81146df
commit c207f97e7c
4 changed files with 47 additions and 11 deletions
+11 -2
View File
@@ -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:
+10 -7
View File
@@ -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
+10 -2
View File
@@ -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):
+16
View File
@@ -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))