AI: fix nested tensor phase3
Fix math actually mathing
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user