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:
@@ -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)
|
||||
@@ -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...");
|
||||
};
|
||||
}
|
||||
}
|
||||
});
|
||||
Reference in New Issue
Block a user