fix: 🐛 allow smaller values in BatchTransform

Using a smaller step to avoid 0 division on low values.
Also added some typing
This commit is contained in:
Mel Massadian
2024-04-01 13:57:01 +02:00
parent eeac8c002a
commit 9a4b27d2e0
+34 -21
View File
@@ -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).
"""