AI: fix NestedTensor phase2

revert splitting of input so no-op is easy / possible
This commit is contained in:
mcDandy
2026-08-08 21:02:11 +02:00
parent 70d6b75853
commit c5d81146df
3 changed files with 73 additions and 40 deletions
+17 -19
View File
@@ -84,7 +84,10 @@ class LatentMathNode(io.ComfyNode):
raise ValueError("At least one input is required.")
stack = stack if remember_stack else (copy.deepcopy(stack) if stack is not None else {})
# Identify all present tensors and their keys, expanding NestedTensors into component variables
# Identify all present tensors and their keys.
# NestedTensor inputs keep their original latent dict under the base key (V0, V1, ...)
# so downstream nodes receive the original NestedTensor, while individual components
# are also exposed as V0_0, V0_1, ... for math expressions.
tensor_keys = []
V_norm_samples = {}
nested_component_keys = {} # maps original key -> list of expanded component variable names
@@ -99,20 +102,20 @@ class LatentMathNode(io.ComfyNode):
V_norm_samples[comp_key] = t
component_names.append(comp_key)
nested_component_keys[k] = component_names
# Preserve the original key as a list of components for the V variable
# Keep the original NestedTensor under the base key so V0 returns it unchanged
tensor_keys.append(k)
V_norm_samples[k] = [V_norm_samples[comp] for comp in component_names]
V_norm_samples[k] = samples
else:
tensor_keys.append(k)
V_norm_samples[k] = samples
# Normalize all together (only regular tensors; NestedTensor originals as lists are kept separate)
at_list = [V_norm_samples[k] for k in tensor_keys if not isinstance(V_norm_samples[k], list)]
# Normalize all together; NestedTensor base entries are skipped but other non-tensor values are still supported
at_list = [V_norm_samples[k] for k in tensor_keys if torch.is_tensor(V_norm_samples[k])]
if at_list:
normalized_samples = normalize_to_common_shape(*at_list, mode=length_mismatch)
norm_iter = iter(normalized_samples)
for key in tensor_keys:
if isinstance(V_norm_samples[key], list):
if not torch.is_tensor(V_norm_samples[key]):
continue
V_norm_samples[key] = next(norm_iter)
@@ -120,19 +123,16 @@ class LatentMathNode(io.ComfyNode):
tensor_keys.extend([c for components in nested_component_keys.values() for c in components])
def _resolve_alias(base):
value = V_norm_samples.get(base)
if isinstance(value, list):
return value[0] if value else None
if value is not None:
return value
# Aliases a/b/c/d should point to the first tensor component of a NestedTensor input
components = nested_component_keys.get(base)
if components:
return V_norm_samples[components[0]]
return None
return V_norm_samples.get(base)
first_sample = next(iter(V_norm_samples.values()))
if isinstance(first_sample, list):
first_sample = first_sample[0]
if not torch.is_tensor(first_sample):
# If the first value is a NestedTensor, grab its first component for metadata
first_sample = first_sample.tensors[0]
ae_res = _resolve_alias("V0")
ae = ae_res if ae_res is not None else make_zero_like(first_sample)
be_res = _resolve_alias("V1")
@@ -150,11 +150,9 @@ class LatentMathNode(io.ComfyNode):
sample = V_norm_samples.get(name)
if sample is None:
continue
if isinstance(sample, list):
for comp in sample:
if comp is not None and comp.shape[0] != ae.shape[0]:
raise ValueError(f"Input '{name}' component has shape {comp.shape[0]}, expected {ae.shape[0]} to match input.")
elif sample.shape[0] != ae.shape[0]:
if not torch.is_tensor(sample):
continue
if sample.shape[0] != ae.shape[0]:
raise ValueError(f"Input '{name}' has shape {sample.shape[0]}, expected {ae.shape[0]} to match input.")
# parse expression once
+36 -5
View File
@@ -97,6 +97,12 @@ class UnifiedMathVisitor(MathExprVisitor):
def _is_tensor(self, val):
return isinstance(val, torch.Tensor) or getattr(val, "is_nested", False)
def _is_plain_tensor(self, val):
return isinstance(val, torch.Tensor)
def _is_nested_tensor(self, val):
return getattr(val, "is_nested", False)
def _is_list(self, val):
return isinstance(val, (list, tuple))
@@ -116,9 +122,9 @@ class UnifiedMathVisitor(MathExprVisitor):
Generic binary operation handler.
"""
try:
if self._is_tensor(a) and a.numel() == 1:
if self._is_plain_tensor(a) and a.numel() == 1:
a = float(a.flatten()[0].item())
if self._is_tensor(b) and b.numel() == 1:
if self._is_plain_tensor(b) and b.numel() == 1:
b = float(b.flatten()[0].item())
# one of them is a list and one is tensor
@@ -153,6 +159,27 @@ class UnifiedMathVisitor(MathExprVisitor):
return [self._bin_op(a, x, torch_op, scalar_op, ctx) for x in b]
# Handle tensor operations
if self._is_nested_tensor(a) or self._is_nested_tensor(b):
# Treat NestedTensors as lists of components for arithmetic
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)]
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]
# 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]
if self._is_tensor(a) or self._is_tensor(b):
orig_dtype = None
@@ -190,7 +217,7 @@ class UnifiedMathVisitor(MathExprVisitor):
def _unary_op(self, a, torch_op, scalar_op):
if self._is_tensor(a) and a.numel() == 1:
if self._is_plain_tensor(a) and a.numel() == 1:
a = float(a.flatten()[0].item())
if self._is_list(a):
@@ -213,11 +240,13 @@ class UnifiedMathVisitor(MathExprVisitor):
return scalar_op(a)
def _reduction_op(self, val, torch_op, list_op):
if self._is_tensor(val):
if self._is_plain_tensor(val):
res = torch_op(val)
if self._is_tensor(res) and res.numel() == 1:
if self._is_plain_tensor(res) and res.numel() == 1:
return float(res.item())
return res
if self._is_nested_tensor(val):
return list_op(val.tensors)
if self._is_list(val):
return list_op(val)
return val
@@ -350,6 +379,8 @@ class UnifiedMathVisitor(MathExprVisitor):
idx_tuple = tuple(indices)
result = val[idx_tuple]
if self._is_tensor(result):
if getattr(result, "is_nested", False):
return result
if result.numel() == 1:
return result.item()
return result.contiguous()
+20 -16
View File
@@ -418,8 +418,15 @@ def test_nested_tensor_support():
nt_in = NestedTensor([t1, t2])
l_in = {"samples": nt_in}
# NestedTensor inputs are expanded into component variables (V0_0, V0_1, ...)
# so users can address each component explicitly instead of auto-stacking them.
# The base V0 variable returns the original NestedTensor unchanged.
result_list, stack = node.execute(Expression="V0", V={"V0": l_in}, F={}, batching=0)
assert len(result_list) == 1
samples = result_list[0]["samples"]
assert getattr(samples, "is_nested", False), "V0 should return the original NestedTensor"
assert torch.allclose(samples.tensors[0], torch.full((1, 4, 32, 32), 1.0))
assert torch.allclose(samples.tensors[1], torch.full((2, 4, 32, 32), 2.0))
# Component access uses V0_0, V0_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)}"
@@ -442,6 +449,14 @@ def test_nested_tensor_support_position_v1():
base_latent = {"samples": torch.zeros((1, 4, 32, 32))}
# V1 returns the original NestedTensor unchanged for downstream compatibility
result_list, _ = node.execute(Expression="V1", 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), 10.0))
assert torch.allclose(samples.tensors[1], torch.full((2, 4, 32, 32), 20.0))
# 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))
@@ -450,26 +465,15 @@ def test_nested_tensor_support_position_v1():
result_list, _ = node.execute(Expression="b * 2.0", V={"V0": base_latent, "V1": l_in}, F={}, batching=0)
assert torch.allclose(result_list[0]["samples"], torch.full((1, 4, 32, 32), 20.0))
# Directly using V1 in an expression (without indexing) resolves to the list of components.
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))
# V list variable should contain the original V inputs; NestedTensor V1 is a list of
# components, so V[1][0] accesses the first component.
# V list variable contains the original V inputs; NestedTensor V1 supports indexing its components
result_list, _ = node.execute(Expression="V[1][0] + 5.0", V={"V0": base_latent, "V1": l_in}, F={}, batching=0)
assert torch.allclose(result_list[0]["samples"], torch.full((1, 4, 32, 32), 15.0))
# V0 is a regular latent (not NestedTensor) so V[0] is the tensor itself
# V0 is a regular latent (not NestedTensor)
result_list, _ = node.execute(Expression="V[0] + 5.0", V={"V0": base_latent, "V1": l_in}, F={}, batching=0)
assert torch.allclose(result_list[0]["samples"], torch.full((1, 4, 32, 32), 5.0))
# Using just V1 returns all components as separate latents
result_list, _ = node.execute(Expression="V1", V={"V0": base_latent, "V1": l_in}, F={}, batching=0)
assert len(result_list) == 2
assert torch.allclose(result_list[0]["samples"], torch.full((1, 4, 32, 32), 10.0))
assert torch.allclose(result_list[1]["samples"], torch.full((2, 4, 32, 32), 20.0))
# V1_0 and V1_1 are direct component variables
# V1_1 is a direct component variable
result_list, _ = node.execute(Expression="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), 20.0))