diff --git a/src/more_math/LatentMathNode.py b/src/more_math/LatentMathNode.py index 1dc244d..c74fb43 100644 --- a/src/more_math/LatentMathNode.py +++ b/src/more_math/LatentMathNode.py @@ -67,30 +67,6 @@ class LatentMathNode(io.ComfyNode): #OUTPUT_NODE = False #OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node - async def check_lazy_status(self,Latent,a,b='',c='',d='',w='',x='',y='',z=''): - input_stream = InputStream(Latent) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - need_load = ['a'] if a is None else [] - for v in stream.getTokens(): - if v.type == MathExprParser.VARIABLE: - var_name = v.text - if var_name == 'b' and b is None: - need_load.append('b') - elif var_name == 'c' and c is None: - need_load.append('c') - elif var_name == 'd' and d is None: - need_load.append('d') - elif var_name == 'w' and w is None: - need_load.append('w') - elif var_name == 'x' and x is None: - need_load.append('x') - elif var_name == 'y' and y is None: - need_load.append('y') - elif var_name == 'z' and z is None: - need_load.append('z') - - return need_load @classmethod def execute(cls, Latent, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0) -> io.NodeOutput: diff --git a/src/more_math/NoiseMathNode.py b/src/more_math/NoiseMathNode.py index cd4cbd2..785fe63 100644 --- a/src/more_math/NoiseMathNode.py +++ b/src/more_math/NoiseMathNode.py @@ -164,26 +164,27 @@ class NoiseExecutor(): merged_c = merge_to_tensor(c_val, merged_samples) merged_d = merge_to_tensor(d_val, merged_samples) - # evaluate once - ndim = merged_samples.ndim - batch_dim = 0 - channel_dim = -3 - height_dim = -2 - width_dim = -1 - time_dim = None - if ndim >= 5: - time_dim = -4 - - B = getIndexTensorAlongDim(merged_samples, batch_dim) - W = getIndexTensorAlongDim(merged_samples, width_dim) - H = getIndexTensorAlongDim(merged_samples, height_dim) - C = getIndexTensorAlongDim(merged_samples, channel_dim) else: merged_samples = samples merged_a = a_val merged_b = b_val merged_c = c_val merged_d = d_val + + # evaluate once + ndim = merged_samples.ndim + batch_dim = 0 + channel_dim = -3 + height_dim = -2 + width_dim = -1 + time_dim = None + if ndim >= 5: + time_dim = -4 + + B = getIndexTensorAlongDim(merged_samples, batch_dim) + W = getIndexTensorAlongDim(merged_samples, width_dim) + H = getIndexTensorAlongDim(merged_samples, height_dim) + C = getIndexTensorAlongDim(merged_samples, channel_dim) variables = { 'a': merged_a, 'b': merged_b, 'c': merged_c, 'd': merged_d, 'w': self.w, 'x': self.x, 'y': self.y, 'z': self.z, @@ -200,57 +201,8 @@ class NoiseExecutor(): visitor = TensorEvalVisitor(variables, variables['a'].shape) merged_result = visitor.visit(self.tree) - split_results = list(merged_result.split(sizes, dim=0)) - return _nested_tensor_module.NestedTensor(split_results) - return split_results - - # non-nested path - # fallback zeros if any generator missing - def to_tensor(val, ref): - if val is None: - return torch.zeros_like(ref) - # get list of tensors from NestedTensor or list/tuple - if hasattr(val, 'is_nested') and getattr(val, 'is_nested'): - lst = val.unbind() - elif isinstance(val, (list, tuple)): - lst = list(val) - else: - return val - - if len(lst) == 0: - return torch.zeros_like(ref) - if len(lst) == 1: - return lst[0] - - # try to concatenate along batch dim - try: - cat = torch.cat(lst, dim=0) - if cat.shape == ref.shape or (cat.shape[0] == ref.shape[0] and cat.shape[1:] == ref.shape[1:]): - return cat - except Exception: - pass - - # try stack - try: - stk = torch.stack(lst, dim=0) - # direct match - if stk.shape == ref.shape or (stk.shape[0] == ref.shape[0] and stk.shape[1:] == ref.shape[1:]): - return stk - # try merging first dim into batch if possible - try: - merged = stk.view(-1, *stk.shape[2:]) if stk.ndim >= 3 else stk.view(-1) - if merged.shape == ref.shape: - return merged - except Exception: - pass - except Exception: - pass - - # fallback to first element - return lst[0] - - visitor = TensorEvalVisitor(variables,variables['a'].shape) - result = visitor.visit(self.tree) - if hasattr(input_latent, 'is_nested') and getattr(val, 'is_nested'): - input_latent - return result + if hasattr(samples, 'is_nested') and getattr(samples, 'is_nested'): + split_results = list(merged_result.split(sizes, dim=0)) + return _nested_tensor_module.NestedTensor(split_results) + + return merged_result