diff --git a/comfy_extras/nodes_images.py b/comfy_extras/nodes_images.py index d42b319e6..6ebf1dbd8 100644 --- a/comfy_extras/nodes_images.py +++ b/comfy_extras/nodes_images.py @@ -241,11 +241,11 @@ class ImageStitch: "image1": ("IMAGE",), "direction": (["right", "down", "left", "up"], {"default": "right"}), "match_image_size": ("BOOLEAN", {"default": True}), - "border_width": ( + "spacing_width": ( "INT", - {"default": 8, "min": 0, "max": 1024, "step": 2}, + {"default": 0, "min": 0, "max": 1024, "step": 2}, ), - "border_color": ( + "spacing_color": ( ["white", "black", "red", "green", "blue"], {"default": "white"}, ), @@ -261,6 +261,7 @@ class ImageStitch: DESCRIPTION = """ Stitches image2 to image1 in the specified direction. If image2 is not provided, returns image1 unchanged. +Optional spacing can be added between images. """ def stitch( @@ -268,8 +269,8 @@ If image2 is not provided, returns image1 unchanged. image1, direction, match_image_size, - border_width, - border_color, + spacing_width, + spacing_color, image2=None, ): if image2 is None: @@ -360,9 +361,9 @@ If image2 is not provided, returns image1 unchanged. dim=-1, ) - # Add border if specified - if border_width > 0: - border_width = border_width + (border_width % 2) # Ensure even + # Add spacing if specified + if spacing_width > 0: + spacing_width = spacing_width + (spacing_width % 2) # Ensure even color_map = { "white": 1.0, @@ -371,39 +372,39 @@ If image2 is not provided, returns image1 unchanged. "green": (0.0, 1.0, 0.0), "blue": (0.0, 0.0, 1.0), } - color_val = color_map[border_color] + color_val = color_map[spacing_color] if direction in ["left", "right"]: - border_shape = ( + spacing_shape = ( image1.shape[0], max(image1.shape[1], image2.shape[1]), - border_width, + spacing_width, image1.shape[-1], ) else: - border_shape = ( + spacing_shape = ( image1.shape[0], - border_width, + spacing_width, max(image1.shape[2], image2.shape[2]), image1.shape[-1], ) - border = torch.full(border_shape, 0.0, device=image1.device) + spacing = torch.full(spacing_shape, 0.0, device=image1.device) if isinstance(color_val, tuple): for i, c in enumerate(color_val): - if i < border.shape[-1]: - border[..., i] = c - if border.shape[-1] == 4: # Add alpha - border[..., 3] = 1.0 + if i < spacing.shape[-1]: + spacing[..., i] = c + if spacing.shape[-1] == 4: # Add alpha + spacing[..., 3] = 1.0 else: - border[..., : min(3, border.shape[-1])] = color_val - if border.shape[-1] == 4: - border[..., 3] = 1.0 + spacing[..., : min(3, spacing.shape[-1])] = color_val + if spacing.shape[-1] == 4: + spacing[..., 3] = 1.0 # Concatenate images images = [image2, image1] if direction in ["left", "up"] else [image1, image2] - if border_width > 0: - images.insert(1, border) + if spacing_width > 0: + images.insert(1, spacing) concat_dim = 2 if direction in ["left", "right"] else 1 return (torch.cat(images, dim=concat_dim),) diff --git a/tests-unit/comfy_extras_test/image_stitch_test.py b/tests-unit/comfy_extras_test/image_stitch_test.py index dcad0adec..fbaef756c 100644 --- a/tests-unit/comfy_extras_test/image_stitch_test.py +++ b/tests-unit/comfy_extras_test/image_stitch_test.py @@ -111,54 +111,54 @@ class TestImageStitch: # Both images should be padded to width 64 assert result[0].shape == (1, 56, 64, 3) # 32 + 24 height, max(64,48) width - def test_border_horizontal(self): - """Test border addition in horizontal concatenation""" + def test_spacing_horizontal(self): + """Test spacing addition in horizontal concatenation""" node = ImageStitch() image1 = self.create_test_image(height=32, width=32) image2 = self.create_test_image(height=32, width=24) - border_width = 16 + spacing_width = 16 - result = node.stitch(image1, "right", False, border_width, "white", image2) + result = node.stitch(image1, "right", False, spacing_width, "white", image2) - # Expected width: 32 + 16 (border) + 24 = 72 + # Expected width: 32 + 16 (spacing) + 24 = 72 assert result[0].shape == (1, 32, 72, 3) - def test_border_vertical(self): - """Test border addition in vertical concatenation""" + def test_spacing_vertical(self): + """Test spacing addition in vertical concatenation""" node = ImageStitch() image1 = self.create_test_image(height=32, width=32) image2 = self.create_test_image(height=24, width=32) - border_width = 16 + spacing_width = 16 - result = node.stitch(image1, "down", False, border_width, "white", image2) + result = node.stitch(image1, "down", False, spacing_width, "white", image2) - # Expected height: 32 + 16 (border) + 24 = 72 + # Expected height: 32 + 16 (spacing) + 24 = 72 assert result[0].shape == (1, 72, 32, 3) - def test_border_color_values(self): - """Test that border colors are applied correctly""" + def test_spacing_color_values(self): + """Test that spacing colors are applied correctly""" node = ImageStitch() image1 = self.create_test_image(height=32, width=32) image2 = self.create_test_image(height=32, width=32) - # Test white border + # Test white spacing result_white = node.stitch(image1, "right", False, 16, "white", image2) - # Check that border region contains white values (close to 1.0) - border_region = result_white[0][:, :, 32:48, :] # Middle 16 pixels - assert torch.all(border_region >= 0.9) # Should be close to white + # Check that spacing region contains white values (close to 1.0) + spacing_region = result_white[0][:, :, 32:48, :] # Middle 16 pixels + assert torch.all(spacing_region >= 0.9) # Should be close to white - # Test black border + # Test black spacing result_black = node.stitch(image1, "right", False, 16, "black", image2) - border_region = result_black[0][:, :, 32:48, :] - assert torch.all(border_region <= 0.1) # Should be close to black + spacing_region = result_black[0][:, :, 32:48, :] + assert torch.all(spacing_region <= 0.1) # Should be close to black - def test_odd_border_width_made_even(self): - """Test that odd border widths are made even""" + def test_odd_spacing_width_made_even(self): + """Test that odd spacing widths are made even""" node = ImageStitch() image1 = self.create_test_image(height=32, width=32) image2 = self.create_test_image(height=32, width=32) - # Use odd border width + # Use odd spacing width result = node.stitch(image1, "right", False, 15, "white", image2) # Should be made even (16), so total width = 32 + 16 + 32 = 80 @@ -221,19 +221,19 @@ class TestImageStitch: result = node.stitch(image1, direction, False, 0, "white", image2) assert result[0].shape == (1, 32, 64, 3) if direction in ["right", "left"] else (1, 64, 32, 3) - def test_batch_size_channel_border_integration(self): - """Test integration of batch matching, channel matching, size matching, and borders""" + def test_batch_size_channel_spacing_integration(self): + """Test integration of batch matching, channel matching, size matching, and spacings""" node = ImageStitch() image1 = self.create_test_image(batch_size=2, height=64, width=48, channels=3) image2 = self.create_test_image(batch_size=1, height=32, width=32, channels=4) result = node.stitch(image1, "right", True, 8, "red", image2) - # Should handle: batch matching, size matching, channel matching, border + # Should handle: batch matching, size matching, channel matching, spacing assert result[0].shape[0] == 2 # Batch size matched assert result[0].shape[-1] == 4 # Channels matched to max assert result[0].shape[1] == 64 # Height from image1 (size matching) - # Width should be: 48 + 8 (border) + resized_image2_width + # Width should be: 48 + 8 (spacing) + resized_image2_width expected_image2_width = int(64 * (32/32)) # Resized to height 64 expected_total_width = 48 + 8 + expected_image2_width assert result[0].shape[2] == expected_total_width