From f9d2ebf91d09fc214fecf7501a5490b33c30aca2 Mon Sep 17 00:00:00 2001 From: Mel Massadian Date: Tue, 7 May 2024 23:42:02 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20=E2=9C=A8=20add=20BatchFloatMath?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Simple math operations on FLOATS (list of floats) --- nodes/batch.py | 84 ++++++++++++++++++++++++++++++++++++++++++---- web/mtb_widgets.js | 1 + 2 files changed, 78 insertions(+), 7 deletions(-) diff --git a/nodes/batch.py b/nodes/batch.py index 7a539be..d9e252f 100644 --- a/nodes/batch.py +++ b/nodes/batch.py @@ -10,7 +10,7 @@ from ..utils import EASINGS, apply_easing, pil2tensor from .transform import MTB_TransformImage -def hex_to_rgb(hex_color, bgr=False): +def hex_to_rgb(hex_color: str, bgr: bool = False): hex_color = hex_color.lstrip("#") if bgr: return tuple(int(hex_color[i : i + 2], 16) for i in (4, 2, 0)) @@ -18,6 +18,68 @@ def hex_to_rgb(hex_color, bgr=False): return tuple(int(hex_color[i : i + 2], 16) for i in (0, 2, 4)) +class MTB_BatchFloatMath: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "reverse": ("BOOLEAN", {"default": False}), + "operation": ( + ["add", "sub", "mul", "div", "pow", "abs"], + {"default": "add"}, + ), + } + } + + RETURN_TYPES = ("FLOATS",) + CATEGORY = "mtb/utils" + FUNCTION = "execute" + + def execute(self, reverse: bool, operation: str, **kwargs: list[float]): + res: list[float] = [] + vals = list(kwargs.values()) + + if reverse: + vals = vals[::-1] + + ref_count = len(vals[0]) + for v in vals: + if len(v) != ref_count: + raise ValueError( + f"All values must have the same length (current: {len(v)}, ref: {ref_count}" + ) + + match operation: + case "add": + for i in range(ref_count): + result = sum(v[i] for v in vals) + res.append(result) + case "sub": + for i in range(ref_count): + result = vals[0][i] - sum(v[i] for v in vals[1:]) + res.append(result) + case "mul": + for i in range(ref_count): + result = vals[0][i] * vals[1][i] + res.append(result) + case "div": + for i in range(ref_count): + result = vals[0][i] / vals[1][i] + res.append(result) + case "pow": + for i in range(ref_count): + result: float = vals[0][i] ** vals[1][i] + res.append(result) + case "abs": + for i in range(ref_count): + result = abs(vals[0][i]) + res.append(result) + case _: + log.info(f"For now this mode ({operation}) is not implemented") + + return (res,) + + class MTB_BatchFloatNormalize: """Normalize the values in the list of floats""" @@ -281,18 +343,21 @@ class MTB_BatchFloatAssemble: def INPUT_TYPES(cls): return {"required": {"reverse": ("BOOLEAN", {"default": False})}} - FUNCTION = "assemble_floats" RETURN_TYPES = ("FLOATS",) CATEGORY = "mtb/batch" + FUNCTION = "assemble_floats" + + def assemble_floats(self, reverse: bool, **kwargs: list[float]): + res: list[float] = [] - def assemble_floats(self, reverse, **kwargs): - res = [] if reverse: for x in reversed(kwargs.values()): - res += x + if x: + res += x else: for x in kwargs.values(): - res += x + if x: + res += x return (res,) @@ -308,7 +373,7 @@ class MTB_BatchFloat: ["Single", "Steps"], {"default": "Steps"}, ), - "count": ("INT", {"default": 1}), + "count": ("INT", {"default": 2}), "min": ("FLOAT", {"default": 0.0, "step": 0.001}), "max": ("FLOAT", {"default": 1.0, "step": 0.001}), "easing": ( @@ -346,6 +411,10 @@ class MTB_BatchFloat: CATEGORY = "mtb/batch" def set_floats(self, mode, count, min, max, easing): + if mode == "Steps" and count == 1: + raise ValueError( + "Steps mode requires at least a count of 2 values" + ) keyframes = [] if mode == "Single": keyframes = [min] * count @@ -969,4 +1038,5 @@ __nodes__ = [ MTB_PlotBatchFloat, MTB_BatchTimeWrap, MTB_BatchFloatFit, + MTB_BatchFloatMath, ] diff --git a/web/mtb_widgets.js b/web/mtb_widgets.js index 2b78152..971544b 100644 --- a/web/mtb_widgets.js +++ b/web/mtb_widgets.js @@ -1137,6 +1137,7 @@ const mtb_widgets = { break } case 'Batch Float Assemble (mtb)': + case 'Batch Float Math (mtb)': case 'Plot Batch Float (mtb)': { shared.setupDynamicConnections(nodeType, 'floats', 'FLOATS') break