From 972c487dd4e6f664adb6b11dd02dda55d1c0e20c Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Mon, 4 Aug 2025 07:17:55 -0700 Subject: [PATCH] Add image_scale_down_by tool and display_any.js - Add new image_scale_down_by tool for downscaling images/latents - Add display_any.js web component for node display - Include comprehensive unit tests for the new tool --- .../tools/image_scale_down_by/__init__.py | 5 + kikotools/tools/image_scale_down_by/logic.py | 40 +++++ kikotools/tools/image_scale_down_by/node.py | 86 +++++++++ tests/unit/tools/test_image_scale_down_by.py | 170 ++++++++++++++++++ web/display_any.js | 131 ++++++++++++++ 5 files changed, 432 insertions(+) create mode 100644 kikotools/tools/image_scale_down_by/__init__.py create mode 100644 kikotools/tools/image_scale_down_by/logic.py create mode 100644 kikotools/tools/image_scale_down_by/node.py create mode 100644 tests/unit/tools/test_image_scale_down_by.py create mode 100644 web/display_any.js diff --git a/kikotools/tools/image_scale_down_by/__init__.py b/kikotools/tools/image_scale_down_by/__init__.py new file mode 100644 index 0000000..589accd --- /dev/null +++ b/kikotools/tools/image_scale_down_by/__init__.py @@ -0,0 +1,5 @@ +"""Image Scale Down By tool for ComfyUI.""" + +from .node import ImageScaleDownByNode + +__all__ = ["ImageScaleDownByNode"] diff --git a/kikotools/tools/image_scale_down_by/logic.py b/kikotools/tools/image_scale_down_by/logic.py new file mode 100644 index 0000000..d003697 --- /dev/null +++ b/kikotools/tools/image_scale_down_by/logic.py @@ -0,0 +1,40 @@ +"""Core logic for ImageScaleDownBy tool.""" + +import torch.nn.functional as F +from torch import Tensor + + +def scale_down_image(image: Tensor, scale_by: float) -> Tensor: + """Scale down an image by a given factor. + + Args: + image: Input image tensor of shape (batch, height, width, channels) + scale_by: Scale factor between 0.01 and 1.0 + + Returns: + Scaled down image tensor + """ + batch, height, width, channels = image.shape + + # Calculate new dimensions + new_height = int(height * scale_by) + new_width = int(width * scale_by) + + # Ensure minimum size of 1x1 + new_height = max(1, new_height) + new_width = max(1, new_width) + + # Convert from BHWC to BCHW for interpolation + image_chw = image.permute(0, 3, 1, 2) + + # Scale down the image using bilinear interpolation + scaled = F.interpolate( + image_chw, + size=(new_height, new_width), + mode="bilinear", + align_corners=False, + antialias=True, + ) + + # Convert back to BHWC + return scaled.permute(0, 2, 3, 1) diff --git a/kikotools/tools/image_scale_down_by/node.py b/kikotools/tools/image_scale_down_by/node.py new file mode 100644 index 0000000..74af327 --- /dev/null +++ b/kikotools/tools/image_scale_down_by/node.py @@ -0,0 +1,86 @@ +"""ComfyUI node implementation for ImageScaleDownBy.""" + +from typing import Dict, Any, Tuple + +from torch import Tensor + +from ...base import ComfyAssetsBaseNode +from .logic import scale_down_image + + +class ImageScaleDownByNode(ComfyAssetsBaseNode): + """ + Scales down images by a specified factor. + + Reduces image dimensions proportionally using bilinear interpolation + with antialiasing for smooth downscaling. + """ + + @classmethod + def INPUT_TYPES(cls) -> Dict[str, Any]: + return { + "required": { + "images": ("IMAGE",), + "scale_by": ( + "FLOAT", + { + "default": 0.5, + "min": 0.01, + "max": 1.0, + "step": 0.01, + "display": "number", + }, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "scale_down" + + def scale_down(self, images: Tensor, scale_by: float) -> Tuple[Tensor]: + """ + Scale down images by the specified factor. + + Args: + images: Input image tensor + scale_by: Scale factor between 0.01 and 1.0 + + Returns: + Tuple containing scaled down image tensor + """ + try: + self.validate_inputs(images=images, scale_by=scale_by) + + # Scale down the images + scaled_images = scale_down_image(images, scale_by) + + _, new_height, new_width, _ = scaled_images.shape + _, orig_height, orig_width, _ = images.shape + + self.log_info( + f"Scaled down images from {orig_height}x{orig_width} " + f"to {new_height}x{new_width} (scale factor: {scale_by})" + ) + + return (scaled_images,) + + except Exception as e: + self.handle_error(f"Failed to scale down images: {str(e)}", e) + + def validate_inputs(self, **kwargs) -> None: + """Validate inputs for ImageScaleDownBy node.""" + images = kwargs.get("images") + scale_by = kwargs.get("scale_by") + + if images is None: + raise ValueError("Images input is required") + + if not isinstance(images, Tensor) or len(images.shape) != 4: + raise ValueError( + f"Expected image tensor with shape (batch, height, width, channels), " + f"got shape {images.shape if isinstance(images, Tensor) else 'non-tensor'}" + ) + + if scale_by <= 0 or scale_by > 1.0: + raise ValueError(f"scale_by must be between 0.01 and 1.0, got {scale_by}") diff --git a/tests/unit/tools/test_image_scale_down_by.py b/tests/unit/tools/test_image_scale_down_by.py new file mode 100644 index 0000000..921b379 --- /dev/null +++ b/tests/unit/tools/test_image_scale_down_by.py @@ -0,0 +1,170 @@ +"""Unit tests for ImageScaleDownBy tool.""" + +import pytest +import torch +from kikotools.tools.image_scale_down_by.logic import scale_down_image +from kikotools.tools.image_scale_down_by.node import ImageScaleDownByNode + + +class TestImageScaleDownByLogic: + """Test the core logic for scaling down images.""" + + def test_scale_down_by_half(self): + """Test scaling down an image by 0.5.""" + # Create a test image (batch=1, height=512, width=512, channels=3) + image = torch.randn(1, 512, 512, 3) + scale_by = 0.5 + + result = scale_down_image(image, scale_by) + + assert result.shape == (1, 256, 256, 3) + + def test_scale_down_by_quarter(self): + """Test scaling down an image by 0.25.""" + image = torch.randn(1, 1024, 768, 3) + scale_by = 0.25 + + result = scale_down_image(image, scale_by) + + assert result.shape == (1, 256, 192, 3) + + def test_scale_down_by_custom_factor(self): + """Test scaling down by a custom factor.""" + image = torch.randn(1, 800, 600, 3) + scale_by = 0.75 + + result = scale_down_image(image, scale_by) + + assert result.shape == (1, 600, 450, 3) + + def test_scale_down_maintains_batch_size(self): + """Test that batch size is maintained.""" + # Test with batch size > 1 + image = torch.randn(4, 512, 512, 3) + scale_by = 0.5 + + result = scale_down_image(image, scale_by) + + assert result.shape == (4, 256, 256, 3) + + def test_scale_by_one_returns_same_size(self): + """Test that scale_by=1.0 returns the same size.""" + image = torch.randn(1, 512, 512, 3) + scale_by = 1.0 + + result = scale_down_image(image, scale_by) + + assert result.shape == image.shape + + def test_non_square_image(self): + """Test scaling non-square images.""" + image = torch.randn(1, 720, 1280, 3) + scale_by = 0.5 + + result = scale_down_image(image, scale_by) + + assert result.shape == (1, 360, 640, 3) + + def test_small_scale_factor(self): + """Test with very small scale factor.""" + image = torch.randn(1, 1000, 1000, 3) + scale_by = 0.01 + + result = scale_down_image(image, scale_by) + + assert result.shape == (1, 10, 10, 3) + + +class TestImageScaleDownByNode: + """Test the ComfyUI node implementation.""" + + @pytest.fixture + def node(self): + """Create a node instance.""" + return ImageScaleDownByNode() + + def test_input_types(self): + """Test that INPUT_TYPES is properly defined.""" + input_types = ImageScaleDownByNode.INPUT_TYPES() + + assert "required" in input_types + assert "images" in input_types["required"] + assert input_types["required"]["images"] == ("IMAGE",) + assert "scale_by" in input_types["required"] + + # Check scale_by configuration + scale_config = input_types["required"]["scale_by"] + assert scale_config[0] == "FLOAT" + assert scale_config[1]["default"] == 0.5 + assert scale_config[1]["min"] == 0.01 + assert scale_config[1]["max"] == 1.0 + assert scale_config[1]["step"] == 0.01 + + def test_return_types(self): + """Test that return types are properly defined.""" + assert ImageScaleDownByNode.RETURN_TYPES == ("IMAGE",) + assert ImageScaleDownByNode.RETURN_NAMES == ("images",) + assert ImageScaleDownByNode.FUNCTION == "scale_down" + + def test_scale_down_execution(self, node): + """Test the scale_down method.""" + images = torch.randn(1, 512, 512, 3) + scale_by = 0.5 + + result = node.scale_down(images, scale_by) + + assert isinstance(result, tuple) + assert len(result) == 1 + assert result[0].shape == (1, 256, 256, 3) + + def test_input_validation_no_images(self, node): + """Test validation with missing images.""" + with pytest.raises(ValueError, match="Images input is required"): + node.validate_inputs(images=None, scale_by=0.5) + + def test_input_validation_invalid_tensor_shape(self, node): + """Test validation with invalid tensor shape.""" + invalid_image = torch.randn(512, 512, 3) # Missing batch dimension + + with pytest.raises(ValueError, match="Expected image tensor with shape"): + node.validate_inputs(images=invalid_image, scale_by=0.5) + + def test_input_validation_scale_too_small(self, node): + """Test validation with scale_by too small.""" + images = torch.randn(1, 512, 512, 3) + + with pytest.raises(ValueError, match="scale_by must be between"): + node.validate_inputs(images=images, scale_by=0.0) + + def test_input_validation_scale_too_large(self, node): + """Test validation with scale_by too large.""" + images = torch.randn(1, 512, 512, 3) + + with pytest.raises(ValueError, match="scale_by must be between"): + node.validate_inputs(images=images, scale_by=1.5) + + def test_category_is_comfyassets(self): + """Test that the node is in the ComfyAssets category.""" + assert ImageScaleDownByNode.CATEGORY == "ComfyAssets" + + def test_scale_down_with_batch(self, node): + """Test scaling down with batch of images.""" + images = torch.randn(3, 640, 480, 3) + scale_by = 0.25 + + result = node.scale_down(images, scale_by) + + assert result[0].shape == (3, 160, 120, 3) + + def test_error_handling(self, node, mocker): + """Test that errors are properly handled.""" + # Mock the scale_down_image function to raise an exception + mocker.patch( + "kikotools.tools.image_scale_down_by.node.scale_down_image", + side_effect=RuntimeError("Test error"), + ) + + images = torch.randn(1, 512, 512, 3) + + with pytest.raises(ValueError, match="Failed to scale down images"): + node.scale_down(images, 0.5) diff --git a/web/display_any.js b/web/display_any.js new file mode 100644 index 0000000..3a1cd6f --- /dev/null +++ b/web/display_any.js @@ -0,0 +1,131 @@ +import { app } from "../../../scripts/app.js"; + +app.registerExtension({ + name: "ComfyAssets.DisplayAny", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if (nodeData.name === "DisplayAny") { + const onExecuted = nodeType.prototype.onExecuted; + + nodeType.prototype.onExecuted = function(message) { + onExecuted?.apply(this, arguments); + + if (message?.text && message.text.length > 0) { + const displayText = message.text[0]; + // Update the display widget with the value + this.updateDisplay(displayText); + + // Also show a condensed version in the title + const condensed = displayText.length > 20 + ? displayText.substring(0, 20) + "..." + : displayText; + this.title = `DisplayAny: ${condensed}`; + } + }; + + nodeType.prototype.updateDisplay = function(text) { + // Remove existing display widget if any + const existingWidget = this.widgets?.find(w => w.name === "display_value"); + if (existingWidget) { + const index = this.widgets.indexOf(existingWidget); + this.widgets.splice(index, 1); + } + + // Create display widget + const widget = { + type: "custom_display", + name: "display_value", + size: [this.size[0] - 20, 80], + displayText: text, + + draw: function(ctx, node, widget_width, y, H) { + const margin = 10; + const padding = 10; + const lineHeight = 16; + const minHeight = 60; + + // Calculate needed height based on text + ctx.font = "12px monospace"; + const lines = this.displayText ? this.displayText.split('\n') : [""]; + const textHeight = Math.max(minHeight, lines.length * lineHeight + padding * 2); + + // Draw background + ctx.fillStyle = "#2a2a2a"; + ctx.fillRect(margin, y, widget_width - margin * 2, textHeight); + + // Draw border + ctx.strokeStyle = "#444"; + ctx.strokeRect(margin, y, widget_width - margin * 2, textHeight); + + // Draw text area background + ctx.fillStyle = "#1e1e1e"; + ctx.fillRect(margin + 1, y + 1, widget_width - margin * 2 - 2, textHeight - 2); + + // Prepare text + ctx.fillStyle = "#ddd"; + ctx.textAlign = "left"; + ctx.textBaseline = "top"; + + // Draw each line + const maxWidth = widget_width - margin * 2 - padding * 2; + let currentY = y + padding; + + for (let i = 0; i < lines.length && i < 3; i++) { // Show max 3 lines + let line = lines[i]; + const metrics = ctx.measureText(line); + + if (metrics.width > maxWidth) { + // Truncate line to fit + while (ctx.measureText(line + "...").width > maxWidth && line.length > 0) { + line = line.slice(0, -1); + } + line = line + "..."; + } + + ctx.fillText(line, margin + padding, currentY); + currentY += lineHeight; + } + + if (lines.length > 3) { + ctx.fillStyle = "#888"; + ctx.fillText("...", margin + padding, currentY); + } + + return textHeight; + }, + + computeSize: function(width) { + const lines = this.displayText ? this.displayText.split('\n') : [""]; + const lineHeight = 16; + const padding = 10; + const minHeight = 60; + const textHeight = Math.max(minHeight, Math.min(lines.length, 3) * lineHeight + padding * 2); + return [width, textHeight]; + } + }; + + // Add the widget + if (!this.widgets) { + this.widgets = []; + } + this.widgets.push(widget); + + // Adjust node size + this.computeSize(); + this.setDirtyCanvas(true); + }; + + // Initialize on node creation + const onNodeCreated = nodeType.prototype.onNodeCreated; + nodeType.prototype.onNodeCreated = function() { + onNodeCreated?.apply(this, arguments); + + // Set minimum size + this.size[0] = Math.max(this.size[0], 250); + this.size[1] = Math.max(this.size[1], 150); + + // Add placeholder text + this.updateDisplay("Value will appear here..."); + }; + } + } +}); \ No newline at end of file