From c5d81146dfeeb9f12a3f6bbb8cb977f0600fd20c Mon Sep 17 00:00:00 2001 From: mcDandy Date: Sat, 8 Aug 2026 21:02:11 +0200 Subject: [PATCH] AI: fix NestedTensor phase2 revert splitting of input so no-op is easy / possible --- more_math/LatentMathNode.py | 36 +++++++++++----------- more_math/Parser/UnifiedMathVisitor.py | 41 ++++++++++++++++++++++---- tests/test_more_math.py | 36 ++++++++++++---------- 3 files changed, 73 insertions(+), 40 deletions(-) diff --git a/more_math/LatentMathNode.py b/more_math/LatentMathNode.py index d31292b..ccad5cc 100644 --- a/more_math/LatentMathNode.py +++ b/more_math/LatentMathNode.py @@ -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 diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index 85f1973..38594eb 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -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() diff --git a/tests/test_more_math.py b/tests/test_more_math.py index 9f0f1e2..cf92659 100644 --- a/tests/test_more_math.py +++ b/tests/test_more_math.py @@ -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))