diff --git a/nodes/batch.py b/nodes/batch.py index 1d4abc1..d301a84 100644 --- a/nodes/batch.py +++ b/nodes/batch.py @@ -206,8 +206,8 @@ class BatchFloat: {"default": "Steps"}, ), "count": ("INT", {"default": 1}), - "min": ("FLOAT", {"default": 0.0}), - "max": ("FLOAT", {"default": 1.0}), + "min": ("FLOAT", {"default": 0.0, "step": 0.001}), + "max": ("FLOAT", {"default": 1.0, "step": 0.001}), "easing": ( [ "Linear", @@ -276,7 +276,7 @@ class BatchMerge: FUNCTION = "merge_batches" CATEGORY = "mtb/batch" - def merge_batches(self, fusion_mode, fill, **kwargs): + def merge_batches(self, fusion_mode: str, fill: str, **kwargs): images = kwargs.values() max_frames = max(img.shape[0] for img in images) @@ -340,9 +340,12 @@ class Batch2dTransform: FUNCTION = "transform_batch" CATEGORY = "mtb/batch" - def get_num_elements(self, param) -> int: + def get_num_elements( + self, param: None | torch.Tensor | list[torch.Tensor] | list[float] + ) -> int: if isinstance(param, torch.Tensor): return torch.numel(param) + elif isinstance(param, list): return len(param) @@ -351,13 +354,13 @@ class Batch2dTransform: def transform_batch( self, image: torch.Tensor, - border_handling, - constant_color, - x=None, - y=None, - zoom=None, - angle=None, - shear=None, + border_handling: str, + constant_color: str, + x: list[float] | None = None, + y: list[float] | None = None, + zoom: list[float] | None = None, + angle: list[float] | None = None, + shear: list[float] | None = None, ): if all( self.get_num_elements(param) <= 0 @@ -367,19 +370,25 @@ class Batch2dTransform: "At least one transform parameter must be provided" ) - keyframes = {"x": [], "y": [], "zoom": [], "angle": [], "shear": []} + keyframes: dict[str, list[float]] = { + "x": [], + "y": [], + "zoom": [], + "angle": [], + "shear": [], + } default_vals = {"x": 0, "y": 0, "zoom": 1.0, "angle": 0, "shear": 0} - if self.get_num_elements(x) > 0: + if x and self.get_num_elements(x) > 0: keyframes["x"] = x - if self.get_num_elements(y) > 0: + if y and self.get_num_elements(y) > 0: keyframes["y"] = y - if self.get_num_elements(zoom) > 0: + if zoom and self.get_num_elements(zoom) > 0: keyframes["zoom"] = zoom - if self.get_num_elements(angle) > 0: + if angle and self.get_num_elements(angle) > 0: keyframes["angle"] = angle - if self.get_num_elements(shear) > 0: + if shear and self.get_num_elements(shear) > 0: keyframes["shear"] = shear for name, values in keyframes.items(): @@ -585,10 +594,12 @@ class BatchShake: interpolant: The interpolation function, defaults to t*t*t*(t*(t*6 - 15) + 10). - Returns: + Returns + ------- A numpy array of shape shape with the generated noise. - Raises: + Raises + ------ ValueError: If shape is not a multiple of res. """ interpolant = interpolant or DEFAULT_INTERPOLANT @@ -651,11 +662,13 @@ class BatchShake: interpolant: The, interpolation function, defaults to t*t*t*(t*(t*6 - 15) + 10). - Returns: + Returns + ------- A numpy array of fractal noise and of shape shape generated by combining several octaves of perlin noise. - Raises: + Raises + ------ ValueError: If shape is not a multiple of (lacunarity**(octaves-1)*res). """