AI: fix NestedTensor phase2
revert splitting of input so no-op is easy / possible
This commit is contained in:
+17
-19
@@ -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
|
||||
|
||||
@@ -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
@@ -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))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user