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
This commit is contained in:
Vito Sansevero
2025-08-04 07:17:55 -07:00
parent 1a3efd3802
commit 972c487dd4
5 changed files with 432 additions and 0 deletions
@@ -0,0 +1,5 @@
"""Image Scale Down By tool for ComfyUI."""
from .node import ImageScaleDownByNode
__all__ = ["ImageScaleDownByNode"]
@@ -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)
@@ -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}")
@@ -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)
+131
View File
@@ -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...");
};
}
}
});