Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3ecab5ac08 | ||
|
|
70592114f9 | ||
|
|
407fc4ca7b | ||
|
|
cb7d5246f9 | ||
|
|
9829fc001d | ||
|
|
e84ec6721c | ||
|
|
80fac8e544 | ||
|
|
90c1aa402d | ||
|
|
f9540bd984 |
+10
-6
@@ -226,9 +226,11 @@ class HFHubLoraLoader:
|
||||
|
||||
lora_path = hf_hub_download(
|
||||
repo_id=repo_id.strip(),
|
||||
subfolder=None
|
||||
if subfolder is None or subfolder.strip() == ""
|
||||
else subfolder.strip(),
|
||||
subfolder=(
|
||||
None
|
||||
if subfolder is None or subfolder.strip() == ""
|
||||
else subfolder.strip()
|
||||
),
|
||||
filename=filename.strip(),
|
||||
cache_dir=find_or_create_cache(),
|
||||
)
|
||||
@@ -281,9 +283,11 @@ class HFHubEmbeddingLoader:
|
||||
):
|
||||
hf_hub_download(
|
||||
repo_id=repo_id.strip(),
|
||||
subfolder=None
|
||||
if subfolder is None or subfolder.strip() == ""
|
||||
else subfolder.strip(),
|
||||
subfolder=(
|
||||
None
|
||||
if subfolder is None or subfolder.strip() == ""
|
||||
else subfolder.strip()
|
||||
),
|
||||
filename=filename.strip(),
|
||||
local_dir=get_folder_paths("embeddings")[0],
|
||||
)
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
# Display Any
|
||||
|
||||
The Display Any node is a debugging and inspection tool that can display any type of input value in ComfyUI. It's particularly useful for understanding data structures and tensor shapes during workflow development.
|
||||
|
||||
## Features
|
||||
|
||||
- **Universal Input**: Accepts any type of input data (tensors, strings, numbers, lists, dictionaries, etc.)
|
||||
- **Two Display Modes**:
|
||||
- **Raw Value**: Shows the string representation of the input
|
||||
- **Tensor Shape**: Extracts and displays the shapes of any tensors found in the input
|
||||
- **Nested Structure Support**: Can find tensors within nested dictionaries and lists
|
||||
- **UI Output**: Displays results directly in the ComfyUI interface
|
||||
|
||||
## Inputs
|
||||
|
||||
- **input** (*): Any value you want to display or inspect
|
||||
- **mode** (DROPDOWN): Display mode selection
|
||||
- `raw value`: Shows the complete string representation of the input
|
||||
- `tensor shape`: Extracts and shows shapes of any tensors in the input
|
||||
|
||||
## Outputs
|
||||
|
||||
- **display_text** (STRING): The formatted display text
|
||||
|
||||
## Usage Examples
|
||||
|
||||
### 1. Display Simple Values
|
||||
|
||||
Connect any output to see its raw value:
|
||||
```
|
||||
String Input: "Hello, ComfyUI!"
|
||||
Mode: raw value
|
||||
Output: "Hello, ComfyUI!"
|
||||
```
|
||||
|
||||
### 2. Inspect Tensor Shapes
|
||||
|
||||
Great for debugging image processing pipelines:
|
||||
```
|
||||
Image Tensor: [1, 3, 512, 512]
|
||||
Mode: tensor shape
|
||||
Output: "[[1, 3, 512, 512]]"
|
||||
```
|
||||
|
||||
### 3. Debug Complex Data Structures
|
||||
|
||||
View nested data structures with multiple tensors:
|
||||
```python
|
||||
Input: {
|
||||
"images": tensor([1, 3, 256, 256]),
|
||||
"masks": [tensor([256, 256]), tensor([256, 256, 1])],
|
||||
"config": {"steps": 20}
|
||||
}
|
||||
Mode: tensor shape
|
||||
Output: "[[1, 3, 256, 256], [256, 256], [256, 256, 1]]"
|
||||
```
|
||||
|
||||
### 4. Workflow Debugging
|
||||
|
||||
Use Display Any nodes at various points in your workflow to understand data flow:
|
||||
- After loading images to verify dimensions
|
||||
- Before/after processing nodes to track shape changes
|
||||
- To inspect conditioning or latent data structures
|
||||
- To view metadata or configuration dictionaries
|
||||
|
||||
## Use Cases
|
||||
|
||||
### Image Pipeline Debugging
|
||||
Place Display Any nodes after image loading and processing nodes to track dimension changes:
|
||||
```
|
||||
Load Image → Display Any (tensor shape) → Resize → Display Any (tensor shape)
|
||||
```
|
||||
|
||||
### Latent Space Inspection
|
||||
Understand latent dimensions in your workflows:
|
||||
```
|
||||
VAE Encode → Display Any (tensor shape) → KSampler → Display Any (raw value)
|
||||
```
|
||||
|
||||
### Configuration Verification
|
||||
Display complex configuration objects to ensure correct settings:
|
||||
```
|
||||
Config Node → Display Any (raw value) → Processing Node
|
||||
```
|
||||
|
||||
## Tips
|
||||
|
||||
1. **Multiple Display Nodes**: You can use multiple Display Any nodes in a single workflow to track data at different stages
|
||||
|
||||
2. **Tensor Shape Mode**: Particularly useful when working with:
|
||||
- Image batches to verify batch size
|
||||
- Latent tensors to understand dimensions
|
||||
- Mask arrays to check compatibility
|
||||
|
||||
3. **Raw Value Mode**: Best for:
|
||||
- String prompts and text
|
||||
- Configuration dictionaries
|
||||
- Debugging node outputs
|
||||
- Understanding data structure
|
||||
|
||||
4. **No Tensors Found**: If you see "No tensors found in input" in tensor shape mode, the input doesn't contain any tensor-like objects (numpy arrays, torch tensors, etc.)
|
||||
|
||||
## Technical Notes
|
||||
|
||||
- The node uses `str()` for raw value display, providing Python's string representation
|
||||
- Tensor shape detection works with any object that has a `shape` attribute
|
||||
- Nested structure traversal supports dictionaries, lists, and tuples
|
||||
- The output is both displayed in the UI and available as a string output for further processing
|
||||
|
||||
## Example Workflow Integration
|
||||
|
||||
```
|
||||
[Load Image] → [Image Processing] → [Display Any (tensor shape)]
|
||||
↓
|
||||
"[[1, 3, 512, 512]]"
|
||||
↓
|
||||
[Text Multiline] ← [Concatenate] ← "Image dimensions: "
|
||||
```
|
||||
|
||||
This creates a text output showing the current image dimensions that can be used elsewhere in your workflow.
|
||||
@@ -11,6 +11,7 @@ from .tools.empty_latent_batch import EmptyLatentBatchNode
|
||||
from .tools.kiko_save_image import KikoSaveImageNode
|
||||
from .tools.image_to_multiple_of import ImageToMultipleOfNode
|
||||
from .tools.gemini_prompt import GeminiPromptNode
|
||||
from .tools.display_any import DisplayAnyNode
|
||||
|
||||
# ComfyUI node registration mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -23,6 +24,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"KikoSaveImage": KikoSaveImageNode,
|
||||
"ImageToMultipleOf": ImageToMultipleOfNode,
|
||||
"GeminiPrompt": GeminiPromptNode,
|
||||
"DisplayAny": DisplayAnyNode,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -35,6 +37,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"KikoSaveImage": "Kiko Save Image",
|
||||
"ImageToMultipleOf": "Image to Multiple of",
|
||||
"GeminiPrompt": "Gemini Prompt Engineer",
|
||||
"DisplayAny": "Display Any",
|
||||
}
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
|
||||
@@ -35,7 +35,9 @@ class ComfyAssetsBaseNode:
|
||||
"""
|
||||
pass
|
||||
|
||||
def handle_error(self, error_msg: str, exception: Optional[Exception] = None) -> None:
|
||||
def handle_error(
|
||||
self, error_msg: str, exception: Optional[Exception] = None
|
||||
) -> None:
|
||||
"""
|
||||
Standardized error handling with logging
|
||||
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""DisplayAny tool for ComfyUI."""
|
||||
|
||||
from .node import DisplayAnyNode
|
||||
|
||||
__all__ = ["DisplayAnyNode"]
|
||||
@@ -0,0 +1,64 @@
|
||||
"""Logic for DisplayAny node - displays any input value or tensor shape."""
|
||||
|
||||
from typing import Any, List, Union
|
||||
|
||||
|
||||
def get_tensor_shapes(input_value: Any) -> List[List[int]]:
|
||||
"""Extract tensor shapes from nested structures.
|
||||
|
||||
Args:
|
||||
input_value: Any input value that may contain tensors
|
||||
|
||||
Returns:
|
||||
List of tensor shapes found in the input
|
||||
"""
|
||||
shapes = []
|
||||
|
||||
def extract_shapes(value: Any) -> None:
|
||||
"""Recursively extract shapes from nested structures."""
|
||||
if isinstance(value, dict):
|
||||
for v in value.values():
|
||||
extract_shapes(v)
|
||||
elif isinstance(value, (list, tuple)):
|
||||
for item in value:
|
||||
extract_shapes(item)
|
||||
elif hasattr(value, "shape"):
|
||||
# Handle tensors (numpy arrays, torch tensors, etc.)
|
||||
shapes.append(list(value.shape))
|
||||
|
||||
extract_shapes(input_value)
|
||||
return shapes
|
||||
|
||||
|
||||
def format_display_value(input_value: Any, mode: str = "raw value") -> str:
|
||||
"""Format input value for display based on selected mode.
|
||||
|
||||
Args:
|
||||
input_value: Any input value to display
|
||||
mode: Display mode - "raw value" or "tensor shape"
|
||||
|
||||
Returns:
|
||||
Formatted string representation of the input
|
||||
"""
|
||||
if mode == "tensor shape":
|
||||
shapes = get_tensor_shapes(input_value)
|
||||
if shapes:
|
||||
return str(shapes)
|
||||
else:
|
||||
return "No tensors found in input"
|
||||
|
||||
# Default to raw value display
|
||||
return str(input_value)
|
||||
|
||||
|
||||
def validate_display_mode(mode: str) -> bool:
|
||||
"""Validate if the display mode is supported.
|
||||
|
||||
Args:
|
||||
mode: Display mode to validate
|
||||
|
||||
Returns:
|
||||
True if mode is valid, False otherwise
|
||||
"""
|
||||
valid_modes = ["raw value", "tensor shape"]
|
||||
return mode in valid_modes
|
||||
@@ -0,0 +1,66 @@
|
||||
"""DisplayAny node for ComfyUI - displays any input value or tensor information."""
|
||||
|
||||
from typing import Any, Dict, Tuple
|
||||
|
||||
from ...base import ComfyAssetsBaseNode
|
||||
from .logic import format_display_value, validate_display_mode
|
||||
|
||||
|
||||
# Define AnyType for wildcard input matching
|
||||
class AnyType(str):
|
||||
"""A special type that matches any input type in ComfyUI."""
|
||||
|
||||
def __ne__(self, other):
|
||||
return False
|
||||
|
||||
|
||||
class DisplayAnyNode(ComfyAssetsBaseNode):
|
||||
"""Display any input value or tensor shape information.
|
||||
|
||||
This node can display any type of input in two modes:
|
||||
- Raw value: Shows the string representation of the input
|
||||
- Tensor shape: Extracts and displays shapes of any tensors in the input
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
"""Define input types for the node."""
|
||||
return {
|
||||
"required": {
|
||||
"input": (AnyType("*"), {}), # Accept any type of input
|
||||
"mode": (["raw value", "tensor shape"],),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, **kwargs) -> bool:
|
||||
"""Validate inputs - always returns True as we accept any input."""
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("display_text",)
|
||||
FUNCTION = "display"
|
||||
OUTPUT_NODE = True # This node displays output in the UI
|
||||
|
||||
def display(self, input: Any, mode: str = "raw value") -> Dict[str, Any]:
|
||||
"""Display the input value according to the selected mode.
|
||||
|
||||
Args:
|
||||
input: Any input value to display
|
||||
mode: Display mode - "raw value" or "tensor shape"
|
||||
|
||||
Returns:
|
||||
Dictionary with UI display and result
|
||||
"""
|
||||
# Validate mode
|
||||
if not validate_display_mode(mode):
|
||||
mode = "raw value" # Default to raw value if invalid
|
||||
|
||||
# Format the display text
|
||||
display_text = format_display_value(input, mode)
|
||||
|
||||
# Return both UI display and result
|
||||
return {
|
||||
"ui": {"text": display_text},
|
||||
"result": (display_text,),
|
||||
}
|
||||
@@ -4,7 +4,9 @@ import torch
|
||||
from typing import Dict, Tuple
|
||||
|
||||
|
||||
def create_empty_latent_batch(width: int, height: int, batch_size: int = 1) -> Dict[str, torch.Tensor]:
|
||||
def create_empty_latent_batch(
|
||||
width: int, height: int, batch_size: int = 1
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
"""
|
||||
Create empty latent tensor with batch support.
|
||||
|
||||
@@ -28,7 +30,9 @@ def create_empty_latent_batch(width: int, height: int, batch_size: int = 1) -> D
|
||||
|
||||
# Ensure dimensions are divisible by 8 (VAE requirement)
|
||||
if width % 8 != 0 or height % 8 != 0:
|
||||
raise ValueError(f"Width and height must be divisible by 8, got {width}x{height}")
|
||||
raise ValueError(
|
||||
f"Width and height must be divisible by 8, got {width}x{height}"
|
||||
)
|
||||
|
||||
# Convert pixel dimensions to latent space (divide by 8)
|
||||
latent_width = width // 8
|
||||
|
||||
@@ -36,7 +36,8 @@ class EmptyLatentBatchNode(ComfyAssetsBaseNode):
|
||||
metadata = PRESET_METADATA.get(preset_name)
|
||||
if metadata:
|
||||
formatted_option = (
|
||||
f"{preset_name} - {metadata.aspect_ratio} " f"({metadata.megapixels:.1f}MP) - {metadata.model_group}"
|
||||
f"{preset_name} - {metadata.aspect_ratio} "
|
||||
f"({metadata.megapixels:.1f}MP) - {metadata.model_group}"
|
||||
)
|
||||
preset_options.append(formatted_option)
|
||||
else:
|
||||
@@ -85,7 +86,8 @@ class EmptyLatentBatchNode(ComfyAssetsBaseNode):
|
||||
"min": 1,
|
||||
"max": 64,
|
||||
"step": 1,
|
||||
"tooltip": "Number of empty latents to create in the batch. " "Useful for batch processing workflows.",
|
||||
"tooltip": "Number of empty latents to create in the batch. "
|
||||
"Useful for batch processing workflows.",
|
||||
},
|
||||
),
|
||||
}
|
||||
@@ -116,7 +118,9 @@ class EmptyLatentBatchNode(ComfyAssetsBaseNode):
|
||||
original_preset = self._extract_preset_name(preset)
|
||||
|
||||
# Get base dimensions from preset or custom input
|
||||
base_width, base_height = get_preset_dimensions(original_preset, width, height)
|
||||
base_width, base_height = get_preset_dimensions(
|
||||
original_preset, width, height
|
||||
)
|
||||
|
||||
# Sanitize dimensions to ensure they meet requirements
|
||||
final_width, final_height = sanitize_dimensions(base_width, base_height)
|
||||
@@ -130,17 +134,23 @@ class EmptyLatentBatchNode(ComfyAssetsBaseNode):
|
||||
|
||||
# Validate final dimensions
|
||||
if not validate_dimensions(final_width, final_height):
|
||||
self.handle_error(f"Invalid dimensions after sanitization: {final_width}×{final_height}")
|
||||
self.handle_error(
|
||||
f"Invalid dimensions after sanitization: {final_width}×{final_height}"
|
||||
)
|
||||
|
||||
# Validate batch size
|
||||
if batch_size <= 0:
|
||||
self.handle_error(f"Batch size must be positive, got {batch_size}")
|
||||
|
||||
if batch_size > 64:
|
||||
self.log_info(f"Large batch size ({batch_size}) may use significant memory")
|
||||
self.log_info(
|
||||
f"Large batch size ({batch_size}) may use significant memory"
|
||||
)
|
||||
|
||||
# Create the empty latent batch
|
||||
latent_dict = create_empty_latent_batch(final_width, final_height, batch_size)
|
||||
latent_dict = create_empty_latent_batch(
|
||||
final_width, final_height, batch_size
|
||||
)
|
||||
|
||||
# Log the operation
|
||||
latent_height = final_height // 8
|
||||
@@ -188,7 +198,9 @@ class EmptyLatentBatchNode(ComfyAssetsBaseNode):
|
||||
# Default to "custom" if we can't parse it
|
||||
return "custom"
|
||||
|
||||
def validate_inputs(self, preset: str, width: int, height: int, batch_size: int) -> bool:
|
||||
def validate_inputs(
|
||||
self, preset: str, width: int, height: int, batch_size: int
|
||||
) -> bool:
|
||||
"""
|
||||
Validate node inputs.
|
||||
|
||||
@@ -279,7 +291,12 @@ class EmptyLatentBatchNode(ComfyAssetsBaseNode):
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""Detailed string representation of the node."""
|
||||
return f"EmptyLatentBatchNode(" f"category='{self.CATEGORY}', " f"function='{self.FUNCTION}'" f")"
|
||||
return (
|
||||
f"EmptyLatentBatchNode("
|
||||
f"category='{self.CATEGORY}', "
|
||||
f"function='{self.FUNCTION}'"
|
||||
f")"
|
||||
)
|
||||
|
||||
|
||||
# Node class mappings for ComfyUI registration
|
||||
|
||||
@@ -48,7 +48,9 @@ def get_save_image_path(
|
||||
prefix_name = os.path.basename(filename_prefix)
|
||||
|
||||
# Sanitize only the filename part (not the directory path)
|
||||
safe_prefix = prefix_name.replace(":", "_") # Only sanitize problematic chars for filenames
|
||||
safe_prefix = prefix_name.replace(
|
||||
":", "_"
|
||||
) # Only sanitize problematic chars for filenames
|
||||
safe_prefix = "".join(c for c in safe_prefix if c.isalnum() or c in "._-")
|
||||
|
||||
# Create unique filename with timestamp to avoid conflicts
|
||||
@@ -114,7 +116,9 @@ def convert_tensor_to_pil(image_tensor: torch.Tensor) -> Image.Image:
|
||||
return img
|
||||
|
||||
|
||||
def create_png_metadata(prompt: Optional[Dict] = None, extra_pnginfo: Optional[Dict] = None) -> Optional[PngInfo]:
|
||||
def create_png_metadata(
|
||||
prompt: Optional[Dict] = None, extra_pnginfo: Optional[Dict] = None
|
||||
) -> Optional[PngInfo]:
|
||||
"""
|
||||
Create PNG metadata with workflow information
|
||||
|
||||
@@ -242,7 +246,10 @@ def process_image_batch(
|
||||
format_extensions = {"PNG": ".png", "JPEG": ".jpg", "WEBP": ".webp"}
|
||||
|
||||
if format_type not in format_extensions:
|
||||
raise ValueError(f"Unsupported format: {format_type}. " f"Supported: {list(format_extensions.keys())}")
|
||||
raise ValueError(
|
||||
f"Unsupported format: {format_type}. "
|
||||
f"Supported: {list(format_extensions.keys())}"
|
||||
)
|
||||
|
||||
format_ext = format_extensions[format_type]
|
||||
|
||||
@@ -308,7 +315,9 @@ def process_image_batch(
|
||||
return results, enhanced_data
|
||||
|
||||
|
||||
def validate_save_inputs(images: torch.Tensor, format_type: str, quality: int, png_compress_level: int) -> None:
|
||||
def validate_save_inputs(
|
||||
images: torch.Tensor, format_type: str, quality: int, png_compress_level: int
|
||||
) -> None:
|
||||
"""
|
||||
Validate inputs for image saving
|
||||
|
||||
@@ -326,19 +335,31 @@ def validate_save_inputs(images: torch.Tensor, format_type: str, quality: int, p
|
||||
raise ValueError(f"images must be a torch.Tensor, got {type(images).__name__}")
|
||||
|
||||
if len(images.shape) != 4:
|
||||
raise ValueError(f"images tensor must have 4 dimensions [batch, height, width, channels], " f"got {len(images.shape)}")
|
||||
raise ValueError(
|
||||
f"images tensor must have 4 dimensions [batch, height, width, channels], "
|
||||
f"got {len(images.shape)}"
|
||||
)
|
||||
|
||||
# Validate format
|
||||
supported_formats = ["PNG", "JPEG", "WEBP"]
|
||||
if format_type not in supported_formats:
|
||||
raise ValueError(f"format must be one of {supported_formats}, got {format_type}")
|
||||
raise ValueError(
|
||||
f"format must be one of {supported_formats}, got {format_type}"
|
||||
)
|
||||
|
||||
# Validate quality (for JPEG/WebP)
|
||||
if format_type in ["JPEG", "WEBP"]:
|
||||
if not isinstance(quality, int) or not (1 <= quality <= 100):
|
||||
raise ValueError(f"quality must be an integer between 1 and 100, got {quality}")
|
||||
raise ValueError(
|
||||
f"quality must be an integer between 1 and 100, got {quality}"
|
||||
)
|
||||
|
||||
# Validate PNG compression level
|
||||
if format_type == "PNG":
|
||||
if not isinstance(png_compress_level, int) or not (0 <= png_compress_level <= 9):
|
||||
raise ValueError(f"png_compress_level must be an integer between 0 and 9, " f"got {png_compress_level}")
|
||||
if not isinstance(png_compress_level, int) or not (
|
||||
0 <= png_compress_level <= 9
|
||||
):
|
||||
raise ValueError(
|
||||
f"png_compress_level must be an integer between 0 and 9, "
|
||||
f"got {png_compress_level}"
|
||||
)
|
||||
|
||||
@@ -79,7 +79,8 @@ class KikoSaveImageNode(ComfyAssetsBaseNode):
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Use lossless WebP compression " "(ignores quality setting)",
|
||||
"tooltip": "Use lossless WebP compression "
|
||||
"(ignores quality setting)",
|
||||
},
|
||||
),
|
||||
"popup": (
|
||||
@@ -162,7 +163,10 @@ class KikoSaveImageNode(ComfyAssetsBaseNode):
|
||||
|
||||
# Log results
|
||||
total_size = sum(data["file_size"] for data in enhanced_data)
|
||||
self.log_info(f"Successfully saved {len(results)} images " f"(total size: {total_size / 1024:.1f} KB)")
|
||||
self.log_info(
|
||||
f"Successfully saved {len(results)} images "
|
||||
f"(total size: {total_size / 1024:.1f} KB)"
|
||||
)
|
||||
|
||||
# Return UI data for ComfyUI preview (clean) + enhanced data for our JS
|
||||
return {
|
||||
@@ -204,7 +208,9 @@ class KikoSaveImageNode(ComfyAssetsBaseNode):
|
||||
|
||||
# Additional node-specific validation
|
||||
if not isinstance(webp_lossless, bool):
|
||||
raise ValueError(f"webp_lossless must be a boolean, got {type(webp_lossless).__name__}")
|
||||
raise ValueError(
|
||||
f"webp_lossless must be a boolean, got {type(webp_lossless).__name__}"
|
||||
)
|
||||
|
||||
if not isinstance(popup, bool):
|
||||
raise ValueError(f"popup must be a boolean, got {type(popup).__name__}")
|
||||
|
||||
@@ -28,7 +28,9 @@ def extract_dimensions(
|
||||
if image is not None:
|
||||
# IMAGE tensor format: [batch, height, width, channels]
|
||||
if len(image.shape) != 4:
|
||||
raise ValueError(f"Expected IMAGE tensor with 4 dimensions, got {len(image.shape)}")
|
||||
raise ValueError(
|
||||
f"Expected IMAGE tensor with 4 dimensions, got {len(image.shape)}"
|
||||
)
|
||||
|
||||
_, height, width, _ = image.shape
|
||||
return int(width), int(height)
|
||||
@@ -40,7 +42,10 @@ def extract_dimensions(
|
||||
|
||||
samples = latent["samples"]
|
||||
if len(samples.shape) != 4:
|
||||
raise ValueError(f"Expected LATENT samples tensor with 4 dimensions, " f"got {len(samples.shape)}")
|
||||
raise ValueError(
|
||||
f"Expected LATENT samples tensor with 4 dimensions, "
|
||||
f"got {len(samples.shape)}"
|
||||
)
|
||||
|
||||
_, _, latent_height, latent_width = samples.shape
|
||||
|
||||
@@ -74,7 +79,9 @@ def ensure_divisible_by_8(width: int, height: int) -> Tuple[int, int]:
|
||||
return int(new_width), int(new_height)
|
||||
|
||||
|
||||
def calculate_scaled_dimensions(width: int, height: int, scale_factor: float) -> Tuple[int, int]:
|
||||
def calculate_scaled_dimensions(
|
||||
width: int, height: int, scale_factor: float
|
||||
) -> Tuple[int, int]:
|
||||
"""
|
||||
Calculate new dimensions with scale factor and ensure divisible by 8
|
||||
|
||||
@@ -97,7 +104,9 @@ def calculate_scaled_dimensions(width: int, height: int, scale_factor: float) ->
|
||||
return ensure_divisible_by_8(new_width, new_height)
|
||||
|
||||
|
||||
def validate_scale_factor(scale_factor: float, min_scale: float = 0.1, max_scale: float = 8.0) -> None:
|
||||
def validate_scale_factor(
|
||||
scale_factor: float, min_scale: float = 0.1, max_scale: float = 8.0
|
||||
) -> None:
|
||||
"""
|
||||
Validate scale factor is within reasonable bounds
|
||||
|
||||
@@ -110,13 +119,19 @@ def validate_scale_factor(scale_factor: float, min_scale: float = 0.1, max_scale
|
||||
ValueError: If scale factor is out of bounds
|
||||
"""
|
||||
if not isinstance(scale_factor, (int, float)):
|
||||
raise ValueError(f"Scale factor must be a number, got {type(scale_factor).__name__}")
|
||||
raise ValueError(
|
||||
f"Scale factor must be a number, got {type(scale_factor).__name__}"
|
||||
)
|
||||
|
||||
if scale_factor < min_scale:
|
||||
raise ValueError(f"Scale factor {scale_factor} is too small (minimum: {min_scale})")
|
||||
raise ValueError(
|
||||
f"Scale factor {scale_factor} is too small (minimum: {min_scale})"
|
||||
)
|
||||
|
||||
if scale_factor > max_scale:
|
||||
raise ValueError(f"Scale factor {scale_factor} is too large (maximum: {max_scale})")
|
||||
raise ValueError(
|
||||
f"Scale factor {scale_factor} is too large (maximum: {max_scale})"
|
||||
)
|
||||
|
||||
|
||||
def calculate_resolution_from_input(
|
||||
@@ -146,6 +161,8 @@ def calculate_resolution_from_input(
|
||||
original_width, original_height = extract_dimensions(image=image, latent=latent)
|
||||
|
||||
# Calculate scaled dimensions
|
||||
new_width, new_height = calculate_scaled_dimensions(original_width, original_height, scale_factor)
|
||||
new_width, new_height = calculate_scaled_dimensions(
|
||||
original_width, original_height, scale_factor
|
||||
)
|
||||
|
||||
return new_width, new_height
|
||||
|
||||
@@ -42,7 +42,8 @@ class ResolutionCalculatorNode(ComfyAssetsBaseNode):
|
||||
"max": 8.0,
|
||||
"step": 0.1,
|
||||
"display": "slider",
|
||||
"tooltip": "Factor to scale the resolution by " "(e.g., 2.0 for 2x, 0.5 for half scale)",
|
||||
"tooltip": "Factor to scale the resolution by "
|
||||
"(e.g., 2.0 for 2x, 0.5 for half scale)",
|
||||
},
|
||||
),
|
||||
},
|
||||
@@ -87,11 +88,20 @@ class ResolutionCalculatorNode(ComfyAssetsBaseNode):
|
||||
self.validate_inputs(scale_factor=scale_factor, image=image, latent=latent)
|
||||
|
||||
# Log the operation
|
||||
input_type = "IMAGE" if image is not None else "LATENT" if latent is not None else "NONE"
|
||||
self.log_info(f"Calculating resolution with scale_factor={scale_factor}, " f"input_type={input_type}")
|
||||
input_type = (
|
||||
"IMAGE"
|
||||
if image is not None
|
||||
else "LATENT" if latent is not None else "NONE"
|
||||
)
|
||||
self.log_info(
|
||||
f"Calculating resolution with scale_factor={scale_factor}, "
|
||||
f"input_type={input_type}"
|
||||
)
|
||||
|
||||
# Calculate the resolution
|
||||
width, height = calculate_resolution_from_input(scale_factor=scale_factor, image=image, latent=latent)
|
||||
width, height = calculate_resolution_from_input(
|
||||
scale_factor=scale_factor, image=image, latent=latent
|
||||
)
|
||||
|
||||
# Log the result
|
||||
self.log_info(f"Calculated resolution: {width}x{height}")
|
||||
@@ -126,7 +136,9 @@ class ResolutionCalculatorNode(ComfyAssetsBaseNode):
|
||||
|
||||
# Validate scale factor type
|
||||
if not isinstance(scale_factor, (int, float)):
|
||||
raise ValueError(f"scale_factor must be a number, got {type(scale_factor).__name__}")
|
||||
raise ValueError(
|
||||
f"scale_factor must be a number, got {type(scale_factor).__name__}"
|
||||
)
|
||||
|
||||
# Validate tensors using helper methods
|
||||
if image is not None:
|
||||
@@ -138,11 +150,14 @@ class ResolutionCalculatorNode(ComfyAssetsBaseNode):
|
||||
def _validate_image_tensor(self, image: torch.Tensor) -> None:
|
||||
"""Validate image tensor format"""
|
||||
if not isinstance(image, torch.Tensor):
|
||||
raise ValueError(f"image must be a torch.Tensor, got {type(image).__name__}")
|
||||
raise ValueError(
|
||||
f"image must be a torch.Tensor, got {type(image).__name__}"
|
||||
)
|
||||
|
||||
if len(image.shape) != 4:
|
||||
raise ValueError(
|
||||
f"image tensor must have 4 dimensions " f"[batch, height, width, channels], got {len(image.shape)}"
|
||||
f"image tensor must have 4 dimensions "
|
||||
f"[batch, height, width, channels], got {len(image.shape)}"
|
||||
)
|
||||
|
||||
def _validate_latent_dict(self, latent: Dict[str, torch.Tensor]) -> None:
|
||||
@@ -155,11 +170,15 @@ class ResolutionCalculatorNode(ComfyAssetsBaseNode):
|
||||
|
||||
samples = latent["samples"]
|
||||
if not isinstance(samples, torch.Tensor):
|
||||
raise ValueError(f"latent['samples'] must be a torch.Tensor, " f"got {type(samples).__name__}")
|
||||
raise ValueError(
|
||||
f"latent['samples'] must be a torch.Tensor, "
|
||||
f"got {type(samples).__name__}"
|
||||
)
|
||||
|
||||
if len(samples.shape) != 4:
|
||||
raise ValueError(
|
||||
f"latent samples tensor must have 4 dimensions " f"[batch, channels, height, width], got {len(samples.shape)}"
|
||||
f"latent samples tensor must have 4 dimensions "
|
||||
f"[batch, channels, height, width], got {len(samples.shape)}"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -65,7 +65,9 @@ class SamplerComboCompactNode(ComfyAssetsBaseNode):
|
||||
FUNCTION = "get_combo"
|
||||
CATEGORY = "ComfyAssets"
|
||||
|
||||
def get_combo(self, sampler: str, sched: str, steps: int, cfg: float) -> Tuple[object, str, int, float]:
|
||||
def get_combo(
|
||||
self, sampler: str, sched: str, steps: int, cfg: float
|
||||
) -> Tuple[object, str, int, float]:
|
||||
"""
|
||||
Get compact sampler combo configuration.
|
||||
|
||||
|
||||
@@ -40,7 +40,9 @@ except ImportError:
|
||||
]
|
||||
|
||||
|
||||
def validate_sampler_settings(sampler_name: str, scheduler: str, steps: int, cfg: float) -> bool:
|
||||
def validate_sampler_settings(
|
||||
sampler_name: str, scheduler: str, steps: int, cfg: float
|
||||
) -> bool:
|
||||
"""
|
||||
Validate sampler configuration settings.
|
||||
|
||||
@@ -81,7 +83,9 @@ def validate_sampler_settings(sampler_name: str, scheduler: str, steps: int, cfg
|
||||
return False
|
||||
|
||||
|
||||
def get_sampler_combo(sampler_name: str, scheduler: str, steps: int, cfg: float) -> Tuple[str, str, int, float]:
|
||||
def get_sampler_combo(
|
||||
sampler_name: str, scheduler: str, steps: int, cfg: float
|
||||
) -> Tuple[str, str, int, float]:
|
||||
"""
|
||||
Process and return sampler combo settings.
|
||||
|
||||
|
||||
@@ -70,7 +70,9 @@ class SamplerComboNode(ComfyAssetsBaseNode):
|
||||
FUNCTION = "get_sampler_combo"
|
||||
CATEGORY = "ComfyAssets"
|
||||
|
||||
def get_sampler_combo(self, sampler_name: str, scheduler: str, steps: int, cfg: float) -> Tuple[object, str, int, float]:
|
||||
def get_sampler_combo(
|
||||
self, sampler_name: str, scheduler: str, steps: int, cfg: float
|
||||
) -> Tuple[object, str, int, float]:
|
||||
"""
|
||||
Get sampler combo configuration.
|
||||
|
||||
@@ -117,7 +119,10 @@ class SamplerComboNode(ComfyAssetsBaseNode):
|
||||
# Return sampler name for testing
|
||||
sampler = result[0]
|
||||
|
||||
self.log_info(f"Configured sampler combo: {result[0]}, {result[1]}, " f"{result[2]} steps, CFG {result[3]}")
|
||||
self.log_info(
|
||||
f"Configured sampler combo: {result[0]}, {result[1]}, "
|
||||
f"{result[2]} steps, CFG {result[3]}"
|
||||
)
|
||||
|
||||
return (sampler, result[1], result[2], result[3])
|
||||
|
||||
@@ -139,7 +144,9 @@ class SamplerComboNode(ComfyAssetsBaseNode):
|
||||
sampler = "euler"
|
||||
return (sampler, "normal", 20, 7.0)
|
||||
|
||||
def validate_inputs(self, sampler_name: str, scheduler: str, steps: int, cfg: float) -> None:
|
||||
def validate_inputs(
|
||||
self, sampler_name: str, scheduler: str, steps: int, cfg: float
|
||||
) -> None:
|
||||
"""
|
||||
Validate sampler combo inputs.
|
||||
|
||||
@@ -154,7 +161,8 @@ class SamplerComboNode(ComfyAssetsBaseNode):
|
||||
"""
|
||||
if not validate_sampler_settings(sampler_name, scheduler, steps, cfg):
|
||||
self.handle_error(
|
||||
f"Invalid sampler settings: sampler={sampler_name}, " f"scheduler={scheduler}, steps={steps}, cfg={cfg}"
|
||||
f"Invalid sampler settings: sampler={sampler_name}, "
|
||||
f"scheduler={scheduler}, steps={steps}, cfg={cfg}"
|
||||
)
|
||||
|
||||
def get_scheduler_suggestions(self, sampler_name: str) -> list:
|
||||
@@ -205,7 +213,9 @@ class SamplerComboNode(ComfyAssetsBaseNode):
|
||||
"recommendation": f"Recommended range: {min_cfg}-{max_cfg} CFG",
|
||||
}
|
||||
|
||||
def get_combo_analysis(self, sampler_name: str, scheduler: str, steps: int, cfg: float) -> dict:
|
||||
def get_combo_analysis(
|
||||
self, sampler_name: str, scheduler: str, steps: int, cfg: float
|
||||
) -> dict:
|
||||
"""
|
||||
Analyze the sampler combo configuration and provide recommendations.
|
||||
|
||||
@@ -264,7 +274,10 @@ class SamplerComboNode(ComfyAssetsBaseNode):
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""String representation of the node."""
|
||||
return f"SamplerComboNode(samplers={len(SAMPLERS)}, " f"schedulers={len(SCHEDULERS)})"
|
||||
return (
|
||||
f"SamplerComboNode(samplers={len(SAMPLERS)}, "
|
||||
f"schedulers={len(SCHEDULERS)})"
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""Detailed string representation of the node."""
|
||||
|
||||
@@ -66,7 +66,9 @@ def sanitize_seed_value(seed: Any) -> int:
|
||||
raise ValueError(f"Invalid seed value: {seed}") from e
|
||||
|
||||
|
||||
def create_history_entry(seed: int, timestamp: Optional[float] = None) -> Dict[str, Any]:
|
||||
def create_history_entry(
|
||||
seed: int, timestamp: Optional[float] = None
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create a standardized history entry for a seed.
|
||||
|
||||
@@ -87,7 +89,9 @@ def create_history_entry(seed: int, timestamp: Optional[float] = None) -> Dict[s
|
||||
}
|
||||
|
||||
|
||||
def filter_duplicate_seeds(history: List[Dict[str, Any]], new_seed: int, dedup_window_ms: int = 500) -> bool:
|
||||
def filter_duplicate_seeds(
|
||||
history: List[Dict[str, Any]], new_seed: int, dedup_window_ms: int = 500
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a seed should be filtered as a duplicate.
|
||||
|
||||
@@ -189,7 +193,9 @@ def format_time_ago(timestamp: float) -> str:
|
||||
return f"{seconds}s ago"
|
||||
|
||||
|
||||
def search_history_by_seed(history: List[Dict[str, Any]], seed: int) -> Optional[Dict[str, Any]]:
|
||||
def search_history_by_seed(
|
||||
history: List[Dict[str, Any]], seed: int
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Search history for a specific seed value.
|
||||
|
||||
|
||||
@@ -28,7 +28,8 @@ class SeedHistoryNode(ComfyAssetsBaseNode):
|
||||
"default": 12345,
|
||||
"min": 0,
|
||||
"max": 0xFFFFFFFFFFFFFFFF,
|
||||
"tooltip": "Seed value for generation processes. " "History UI tracks all changes automatically.",
|
||||
"tooltip": "Seed value for generation processes. "
|
||||
"History UI tracks all changes automatically.",
|
||||
},
|
||||
),
|
||||
}
|
||||
@@ -56,7 +57,10 @@ class SeedHistoryNode(ComfyAssetsBaseNode):
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.error(f"{self.__class__.__name__}: Invalid seed value: {seed}. " f"Using fallback seed 12345.")
|
||||
logger.error(
|
||||
f"{self.__class__.__name__}: Invalid seed value: {seed}. "
|
||||
f"Using fallback seed 12345."
|
||||
)
|
||||
return (12345,)
|
||||
|
||||
clean_seed = sanitize_seed_value(seed)
|
||||
@@ -68,7 +72,10 @@ class SeedHistoryNode(ComfyAssetsBaseNode):
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.error(f"{self.__class__.__name__}: Error processing seed: {str(e)}. " f"Using fallback seed 12345.")
|
||||
logger.error(
|
||||
f"{self.__class__.__name__}: Error processing seed: {str(e)}. "
|
||||
f"Using fallback seed 12345."
|
||||
)
|
||||
return (12345,)
|
||||
|
||||
def generate_new_seed(self) -> int:
|
||||
|
||||
@@ -5,7 +5,9 @@ from math import gcd
|
||||
from .presets import PRESET_OPTIONS
|
||||
|
||||
|
||||
def get_preset_dimensions(preset: str, custom_width: int, custom_height: int) -> Tuple[int, int]:
|
||||
def get_preset_dimensions(
|
||||
preset: str, custom_width: int, custom_height: int
|
||||
) -> Tuple[int, int]:
|
||||
"""
|
||||
Get dimensions from preset name or use custom dimensions.
|
||||
|
||||
@@ -112,7 +114,9 @@ def sanitize_dimensions(width: int, height: int) -> Tuple[int, int]:
|
||||
return width, height
|
||||
|
||||
|
||||
def get_dimension_info(preset: str, width: int, height: int, swap_enabled: bool) -> dict:
|
||||
def get_dimension_info(
|
||||
preset: str, width: int, height: int, swap_enabled: bool
|
||||
) -> dict:
|
||||
"""
|
||||
Get comprehensive dimension information including metadata.
|
||||
|
||||
@@ -201,7 +205,9 @@ def parse_dimension_string(dimension_str: str) -> Tuple[int, int]:
|
||||
raise ValueError(f"Could not parse dimensions from {dimension_str}: {e}")
|
||||
|
||||
|
||||
def get_optimal_scale_factor(current_width: int, current_height: int, target_width: int, target_height: int) -> float:
|
||||
def get_optimal_scale_factor(
|
||||
current_width: int, current_height: int, target_width: int, target_height: int
|
||||
) -> float:
|
||||
"""
|
||||
Calculate optimal scale factor to get from current to target dimensions.
|
||||
|
||||
|
||||
@@ -36,7 +36,8 @@ class WidthHeightSelectorNode(ComfyAssetsBaseNode):
|
||||
metadata = PRESET_METADATA.get(preset_name)
|
||||
if metadata:
|
||||
formatted_option = (
|
||||
f"{preset_name} - {metadata.aspect_ratio} " f"({metadata.megapixels:.1f}MP) - {metadata.model_group}"
|
||||
f"{preset_name} - {metadata.aspect_ratio} "
|
||||
f"({metadata.megapixels:.1f}MP) - {metadata.model_group}"
|
||||
)
|
||||
preset_options.append(formatted_option)
|
||||
else:
|
||||
@@ -103,7 +104,9 @@ class WidthHeightSelectorNode(ComfyAssetsBaseNode):
|
||||
original_preset = self._extract_preset_name(preset)
|
||||
|
||||
# Get base dimensions from preset or custom input
|
||||
final_width, final_height = get_preset_dimensions(original_preset, width, height)
|
||||
final_width, final_height = get_preset_dimensions(
|
||||
original_preset, width, height
|
||||
)
|
||||
|
||||
# Sanitize dimensions to ensure they meet ComfyUI requirements
|
||||
final_width, final_height = sanitize_dimensions(final_width, final_height)
|
||||
@@ -112,7 +115,8 @@ class WidthHeightSelectorNode(ComfyAssetsBaseNode):
|
||||
if not validate_dimensions(final_width, final_height):
|
||||
# This should not happen after sanitization, but handle gracefully
|
||||
self.handle_error(
|
||||
f"Generated invalid dimensions: {final_width}×{final_height}. " f"Using fallback dimensions 1024×1024."
|
||||
f"Generated invalid dimensions: {final_width}×{final_height}. "
|
||||
f"Using fallback dimensions 1024×1024."
|
||||
)
|
||||
final_width, final_height = 1024, 1024
|
||||
|
||||
@@ -120,7 +124,9 @@ class WidthHeightSelectorNode(ComfyAssetsBaseNode):
|
||||
|
||||
except Exception as e:
|
||||
# Handle any unexpected errors gracefully
|
||||
error_msg = f"Error processing dimensions: {str(e)}. Using fallback 1024×1024."
|
||||
error_msg = (
|
||||
f"Error processing dimensions: {str(e)}. Using fallback 1024×1024."
|
||||
)
|
||||
self.handle_error(error_msg)
|
||||
return (1024, 1024)
|
||||
|
||||
@@ -170,7 +176,10 @@ class WidthHeightSelectorNode(ComfyAssetsBaseNode):
|
||||
|
||||
metadata = get_preset_metadata(preset)
|
||||
if metadata.width > 0: # Valid metadata
|
||||
return f"{preset} - {metadata.aspect_ratio} ({metadata.megapixels:.1f}MP) - " f"{metadata.description}"
|
||||
return (
|
||||
f"{preset} - {metadata.aspect_ratio} ({metadata.megapixels:.1f}MP) - "
|
||||
f"{metadata.description}"
|
||||
)
|
||||
|
||||
return f"Unknown preset: {preset}"
|
||||
|
||||
|
||||
@@ -287,15 +287,21 @@ PRESET_METADATA: Dict[str, PresetMetadata] = {
|
||||
|
||||
# Legacy compatibility - maintain old preset dictionaries
|
||||
SDXL_PRESETS: Dict[str, Tuple[int, int]] = {
|
||||
k: (v.width, v.height) for k, v in PRESET_METADATA.items() if v.model_group == "SDXL"
|
||||
k: (v.width, v.height)
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "SDXL"
|
||||
}
|
||||
|
||||
FLUX_PRESETS: Dict[str, Tuple[int, int]] = {
|
||||
k: (v.width, v.height) for k, v in PRESET_METADATA.items() if v.model_group == "FLUX"
|
||||
k: (v.width, v.height)
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "FLUX"
|
||||
}
|
||||
|
||||
ULTRA_WIDE_PRESETS: Dict[str, Tuple[int, int]] = {
|
||||
k: (v.width, v.height) for k, v in PRESET_METADATA.items() if v.model_group == "Ultra-Wide"
|
||||
k: (v.width, v.height)
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "Ultra-Wide"
|
||||
}
|
||||
|
||||
# Combined preset options for ComfyUI dropdown
|
||||
@@ -308,28 +314,78 @@ PRESET_OPTIONS: Dict[str, Tuple[int, int]] = {
|
||||
PRESET_CATEGORIES = {
|
||||
"Custom": ["custom"],
|
||||
# SDXL Categories
|
||||
"SDXL Square": [k for k, v in PRESET_METADATA.items() if v.model_group == "SDXL" and v.category == "Square"],
|
||||
"SDXL Portrait": [k for k, v in PRESET_METADATA.items() if v.model_group == "SDXL" and v.category == "Portrait"],
|
||||
"SDXL Landscape": [k for k, v in PRESET_METADATA.items() if v.model_group == "SDXL" and v.category == "Landscape"],
|
||||
"SDXL Square": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "SDXL" and v.category == "Square"
|
||||
],
|
||||
"SDXL Portrait": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "SDXL" and v.category == "Portrait"
|
||||
],
|
||||
"SDXL Landscape": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "SDXL" and v.category == "Landscape"
|
||||
],
|
||||
# FLUX Categories
|
||||
"FLUX Square": [k for k, v in PRESET_METADATA.items() if v.model_group == "FLUX" and v.category == "Square"],
|
||||
"FLUX Portrait": [k for k, v in PRESET_METADATA.items() if v.model_group == "FLUX" and v.category == "Portrait"],
|
||||
"FLUX Cinematic": [k for k, v in PRESET_METADATA.items() if v.model_group == "FLUX" and v.category == "Cinematic"],
|
||||
"FLUX Classic": [k for k, v in PRESET_METADATA.items() if v.model_group == "FLUX" and v.category == "Classic"],
|
||||
"FLUX Photography": [k for k, v in PRESET_METADATA.items() if v.model_group == "FLUX" and v.category == "Photography"],
|
||||
"FLUX Square": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "FLUX" and v.category == "Square"
|
||||
],
|
||||
"FLUX Portrait": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "FLUX" and v.category == "Portrait"
|
||||
],
|
||||
"FLUX Cinematic": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "FLUX" and v.category == "Cinematic"
|
||||
],
|
||||
"FLUX Classic": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "FLUX" and v.category == "Classic"
|
||||
],
|
||||
"FLUX Photography": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "FLUX" and v.category == "Photography"
|
||||
],
|
||||
# Ultra-Wide Categories
|
||||
"Ultra-Wide Gaming": [k for k, v in PRESET_METADATA.items() if v.model_group == "Ultra-Wide" and v.category == "Gaming"],
|
||||
"Ultra-Wide Gaming": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "Ultra-Wide" and v.category == "Gaming"
|
||||
],
|
||||
"Ultra-Wide Cinematic": [
|
||||
k for k, v in PRESET_METADATA.items() if v.model_group == "Ultra-Wide" and v.category == "Cinematic"
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "Ultra-Wide" and v.category == "Cinematic"
|
||||
],
|
||||
"Ultra-Wide Panoramic": [
|
||||
k for k, v in PRESET_METADATA.items() if v.model_group == "Ultra-Wide" and v.category == "Panoramic"
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "Ultra-Wide" and v.category == "Panoramic"
|
||||
],
|
||||
"Ultra-Wide Mobile": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "Ultra-Wide" and v.category == "Mobile"
|
||||
],
|
||||
"Ultra-Wide Mobile": [k for k, v in PRESET_METADATA.items() if v.model_group == "Ultra-Wide" and v.category == "Mobile"],
|
||||
"Ultra-Wide Vertical": [
|
||||
k for k, v in PRESET_METADATA.items() if v.model_group == "Ultra-Wide" and v.category == "Vertical"
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "Ultra-Wide" and v.category == "Vertical"
|
||||
],
|
||||
"Ultra-Wide Banner": [
|
||||
k
|
||||
for k, v in PRESET_METADATA.items()
|
||||
if v.model_group == "Ultra-Wide" and v.category == "Banner"
|
||||
],
|
||||
"Ultra-Wide Banner": [k for k, v in PRESET_METADATA.items() if v.model_group == "Ultra-Wide" and v.category == "Banner"],
|
||||
}
|
||||
|
||||
# Legacy compatibility - preset descriptions
|
||||
@@ -339,7 +395,9 @@ PRESET_DESCRIPTIONS = {k: v.description for k, v in PRESET_METADATA.items()}
|
||||
MODEL_RECOMMENDATIONS = {
|
||||
"SDXL": [k for k, v in PRESET_METADATA.items() if v.model_group == "SDXL"],
|
||||
"FLUX": [k for k, v in PRESET_METADATA.items() if v.model_group == "FLUX"],
|
||||
"Ultra-Wide": [k for k, v in PRESET_METADATA.items() if v.model_group == "Ultra-Wide"],
|
||||
"Ultra-Wide": [
|
||||
k for k, v in PRESET_METADATA.items() if v.model_group == "Ultra-Wide"
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@@ -393,7 +451,9 @@ def validate_preset_dimensions() -> bool:
|
||||
|
||||
# Check divisible by 8
|
||||
if width % 8 != 0 or height % 8 != 0:
|
||||
print(f"ERROR: {preset_name} dimensions not divisible by 8: {width}×{height}")
|
||||
print(
|
||||
f"ERROR: {preset_name} dimensions not divisible by 8: {width}×{height}"
|
||||
)
|
||||
return False
|
||||
|
||||
# Check reasonable bounds
|
||||
@@ -409,7 +469,9 @@ def validate_metadata_consistency() -> bool:
|
||||
"""Validate metadata consistency and completeness."""
|
||||
for preset_name, metadata in PRESET_METADATA.items():
|
||||
# Verify aspect ratio calculation
|
||||
expected_ratio, expected_decimal = calculate_aspect_ratio(metadata.width, metadata.height)
|
||||
expected_ratio, expected_decimal = calculate_aspect_ratio(
|
||||
metadata.width, metadata.height
|
||||
)
|
||||
if abs(metadata.aspect_decimal - expected_decimal) > 0.001:
|
||||
print(
|
||||
f"ERROR: {preset_name} aspect ratio mismatch: "
|
||||
@@ -420,7 +482,10 @@ def validate_metadata_consistency() -> bool:
|
||||
# Verify megapixel calculation
|
||||
expected_mp = (metadata.width * metadata.height) / 1_000_000
|
||||
if abs(metadata.megapixels - expected_mp) > 0.1:
|
||||
print(f"ERROR: {preset_name} megapixel mismatch: " f"expected {expected_mp:.2f}, got {metadata.megapixels}")
|
||||
print(
|
||||
f"ERROR: {preset_name} megapixel mismatch: "
|
||||
f"expected {expected_mp:.2f}, got {metadata.megapixels}"
|
||||
)
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
+1
-1
@@ -42,7 +42,7 @@ Icon = "https://avatars.githubusercontent.com/u/213204677?s=200"
|
||||
includes = []
|
||||
|
||||
[tool.black]
|
||||
line-length = 127
|
||||
line-length = 88
|
||||
target-version = ['py310']
|
||||
include = '\.pyi?$'
|
||||
extend-exclude = '''
|
||||
|
||||
+9
-3
@@ -102,12 +102,18 @@ def assert_divisible_by_8(width: int, height: int) -> None:
|
||||
assert height % 8 == 0, f"Height {height} must be divisible by 8"
|
||||
|
||||
|
||||
def assert_reasonable_dimensions(width: int, height: int, min_size: int = 64, max_size: int = 8192) -> None:
|
||||
def assert_reasonable_dimensions(
|
||||
width: int, height: int, min_size: int = 64, max_size: int = 8192
|
||||
) -> None:
|
||||
"""
|
||||
Helper function to assert dimensions are within reasonable bounds
|
||||
"""
|
||||
assert min_size <= width <= max_size, f"Width {width} out of reasonable range [{min_size}, {max_size}]"
|
||||
assert min_size <= height <= max_size, f"Height {height} out of reasonable range [{min_size}, {max_size}]"
|
||||
assert (
|
||||
min_size <= width <= max_size
|
||||
), f"Width {width} out of reasonable range [{min_size}, {max_size}]"
|
||||
assert (
|
||||
min_size <= height <= max_size
|
||||
), f"Height {height} out of reasonable range [{min_size}, {max_size}]"
|
||||
|
||||
|
||||
# Make helper functions available as pytest fixtures
|
||||
|
||||
@@ -32,7 +32,10 @@ class TestComfyAssetsBaseNode:
|
||||
node.handle_error("Test error message")
|
||||
|
||||
mock_logger.error.assert_called_once()
|
||||
assert "ComfyAssetsBaseNode: Test error message" in mock_logger.error.call_args[0][0]
|
||||
assert (
|
||||
"ComfyAssetsBaseNode: Test error message"
|
||||
in mock_logger.error.call_args[0][0]
|
||||
)
|
||||
|
||||
def test_handle_error_with_exception_logs_exception(self):
|
||||
"""Test error handling with original exception logs both messages"""
|
||||
@@ -55,7 +58,10 @@ class TestComfyAssetsBaseNode:
|
||||
node.log_info("Test information")
|
||||
|
||||
mock_logger.info.assert_called_once()
|
||||
assert "ComfyAssetsBaseNode: Test information" in mock_logger.info.call_args[0][0]
|
||||
assert (
|
||||
"ComfyAssetsBaseNode: Test information"
|
||||
in mock_logger.info.call_args[0][0]
|
||||
)
|
||||
|
||||
def test_get_node_info_returns_metadata(self):
|
||||
"""Test get_node_info returns correct metadata"""
|
||||
|
||||
@@ -0,0 +1,287 @@
|
||||
"""Unit tests for DisplayAny node."""
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from kikotools.tools.display_any import DisplayAnyNode
|
||||
from kikotools.tools.display_any.logic import (
|
||||
format_display_value,
|
||||
get_tensor_shapes,
|
||||
validate_display_mode,
|
||||
)
|
||||
from kikotools.tools.display_any.node import AnyType
|
||||
|
||||
|
||||
class TestAnyType:
|
||||
"""Test cases for AnyType class."""
|
||||
|
||||
def test_anytype_not_equal(self):
|
||||
"""Test that AnyType is never equal to other types."""
|
||||
any_type = AnyType("*")
|
||||
|
||||
# Should not be equal to any other type
|
||||
assert not (any_type != "STRING")
|
||||
assert not (any_type != "IMAGE")
|
||||
assert not (any_type != "LATENT")
|
||||
assert not (any_type != 123)
|
||||
assert not (any_type != None)
|
||||
assert not (any_type != ["LIST"])
|
||||
|
||||
def test_anytype_string_representation(self):
|
||||
"""Test string representation of AnyType."""
|
||||
any_type = AnyType("*")
|
||||
assert str(any_type) == "*"
|
||||
|
||||
|
||||
class TestDisplayAnyNode:
|
||||
"""Test cases for DisplayAnyNode."""
|
||||
|
||||
def test_node_properties(self):
|
||||
"""Test node has correct properties."""
|
||||
assert DisplayAnyNode.CATEGORY == "ComfyAssets"
|
||||
assert DisplayAnyNode.FUNCTION == "display"
|
||||
assert DisplayAnyNode.RETURN_TYPES == ("STRING",)
|
||||
assert DisplayAnyNode.RETURN_NAMES == ("display_text",)
|
||||
assert DisplayAnyNode.OUTPUT_NODE is True
|
||||
|
||||
def test_input_types(self):
|
||||
"""Test INPUT_TYPES configuration."""
|
||||
input_types = DisplayAnyNode.INPUT_TYPES()
|
||||
|
||||
# Check required inputs
|
||||
assert "required" in input_types
|
||||
assert "input" in input_types["required"]
|
||||
# Check that input is AnyType with wildcard
|
||||
input_type = input_types["required"]["input"]
|
||||
assert len(input_type) == 2
|
||||
assert isinstance(input_type[0], AnyType)
|
||||
assert str(input_type[0]) == "*"
|
||||
assert input_type[1] == {}
|
||||
assert "mode" in input_types["required"]
|
||||
assert input_types["required"]["mode"] == (["raw value", "tensor shape"],)
|
||||
|
||||
def test_validate_inputs(self):
|
||||
"""Test VALIDATE_INPUTS always returns True."""
|
||||
assert DisplayAnyNode.VALIDATE_INPUTS() is True
|
||||
assert DisplayAnyNode.VALIDATE_INPUTS(input="test") is True
|
||||
assert DisplayAnyNode.VALIDATE_INPUTS(input=123, mode="raw value") is True
|
||||
|
||||
def test_display_raw_value_string(self):
|
||||
"""Test displaying raw string value."""
|
||||
node = DisplayAnyNode()
|
||||
result = node.display("Hello, World!", "raw value")
|
||||
|
||||
assert "ui" in result
|
||||
assert "text" in result["ui"]
|
||||
assert result["ui"]["text"] == "Hello, World!"
|
||||
assert "result" in result
|
||||
assert result["result"] == ("Hello, World!",)
|
||||
|
||||
def test_display_raw_value_number(self):
|
||||
"""Test displaying raw number value."""
|
||||
node = DisplayAnyNode()
|
||||
result = node.display(42, "raw value")
|
||||
|
||||
assert result["ui"]["text"] == "42"
|
||||
assert result["result"] == ("42",)
|
||||
|
||||
def test_display_raw_value_list(self):
|
||||
"""Test displaying raw list value."""
|
||||
node = DisplayAnyNode()
|
||||
test_list = [1, 2, 3, "test"]
|
||||
result = node.display(test_list, "raw value")
|
||||
|
||||
assert result["ui"]["text"] == str(test_list)
|
||||
assert result["result"] == (str(test_list),)
|
||||
|
||||
def test_display_raw_value_dict(self):
|
||||
"""Test displaying raw dictionary value."""
|
||||
node = DisplayAnyNode()
|
||||
test_dict = {"key": "value", "number": 123}
|
||||
result = node.display(test_dict, "raw value")
|
||||
|
||||
assert result["ui"]["text"] == str(test_dict)
|
||||
assert result["result"] == (str(test_dict),)
|
||||
|
||||
def test_display_tensor_shape_numpy(self):
|
||||
"""Test displaying numpy tensor shape."""
|
||||
node = DisplayAnyNode()
|
||||
tensor = np.random.rand(4, 3, 224, 224)
|
||||
result = node.display(tensor, "tensor shape")
|
||||
|
||||
assert result["ui"]["text"] == "[[4, 3, 224, 224]]"
|
||||
assert result["result"] == ("[[4, 3, 224, 224]]",)
|
||||
|
||||
@pytest.mark.skipif(not torch, reason="PyTorch not installed")
|
||||
def test_display_tensor_shape_torch(self):
|
||||
"""Test displaying PyTorch tensor shape."""
|
||||
node = DisplayAnyNode()
|
||||
tensor = torch.randn(2, 10, 512, 512)
|
||||
result = node.display(tensor, "tensor shape")
|
||||
|
||||
assert result["ui"]["text"] == "[[2, 10, 512, 512]]"
|
||||
assert result["result"] == ("[[2, 10, 512, 512]]",)
|
||||
|
||||
def test_display_nested_tensors(self):
|
||||
"""Test displaying shapes from nested structure with tensors."""
|
||||
node = DisplayAnyNode()
|
||||
nested_data = {
|
||||
"images": np.random.rand(1, 3, 256, 256),
|
||||
"masks": [
|
||||
np.random.rand(256, 256),
|
||||
np.random.rand(256, 256, 1),
|
||||
],
|
||||
"metadata": {"info": "test", "tensor": np.random.rand(10)},
|
||||
}
|
||||
result = node.display(nested_data, "tensor shape")
|
||||
|
||||
expected = "[[1, 3, 256, 256], [256, 256], [256, 256, 1], [10]]"
|
||||
assert result["ui"]["text"] == expected
|
||||
assert result["result"] == (expected,)
|
||||
|
||||
def test_display_no_tensors(self):
|
||||
"""Test displaying when no tensors are present."""
|
||||
node = DisplayAnyNode()
|
||||
data = {"text": "hello", "number": 42, "list": [1, 2, 3]}
|
||||
result = node.display(data, "tensor shape")
|
||||
|
||||
assert result["ui"]["text"] == "No tensors found in input"
|
||||
assert result["result"] == ("No tensors found in input",)
|
||||
|
||||
def test_invalid_mode_defaults_to_raw(self):
|
||||
"""Test that invalid mode defaults to raw value."""
|
||||
node = DisplayAnyNode()
|
||||
result = node.display("test", "invalid_mode")
|
||||
|
||||
assert result["ui"]["text"] == "test"
|
||||
assert result["result"] == ("test",)
|
||||
|
||||
|
||||
class TestDisplayAnyLogic:
|
||||
"""Test cases for DisplayAny logic functions."""
|
||||
|
||||
def test_get_tensor_shapes_single(self):
|
||||
"""Test getting shape from single tensor."""
|
||||
tensor = np.random.rand(3, 224, 224)
|
||||
shapes = get_tensor_shapes(tensor)
|
||||
|
||||
assert len(shapes) == 1
|
||||
assert shapes[0] == [3, 224, 224]
|
||||
|
||||
def test_get_tensor_shapes_nested_dict(self):
|
||||
"""Test getting shapes from nested dictionary."""
|
||||
data = {
|
||||
"level1": {
|
||||
"tensor1": np.random.rand(10, 20),
|
||||
"level2": {"tensor2": np.random.rand(5, 5, 5)},
|
||||
}
|
||||
}
|
||||
shapes = get_tensor_shapes(data)
|
||||
|
||||
assert len(shapes) == 2
|
||||
assert [10, 20] in shapes
|
||||
assert [5, 5, 5] in shapes
|
||||
|
||||
def test_get_tensor_shapes_nested_list(self):
|
||||
"""Test getting shapes from nested list."""
|
||||
data = [
|
||||
np.random.rand(1, 2, 3),
|
||||
[np.random.rand(4, 5), np.random.rand(6, 7, 8)],
|
||||
"not a tensor",
|
||||
]
|
||||
shapes = get_tensor_shapes(data)
|
||||
|
||||
assert len(shapes) == 3
|
||||
assert [1, 2, 3] in shapes
|
||||
assert [4, 5] in shapes
|
||||
assert [6, 7, 8] in shapes
|
||||
|
||||
def test_get_tensor_shapes_tuple(self):
|
||||
"""Test getting shapes from tuple."""
|
||||
data = (np.random.rand(2, 2), np.random.rand(3, 3))
|
||||
shapes = get_tensor_shapes(data)
|
||||
|
||||
assert len(shapes) == 2
|
||||
assert [2, 2] in shapes
|
||||
assert [3, 3] in shapes
|
||||
|
||||
def test_format_display_value_raw(self):
|
||||
"""Test formatting for raw value display."""
|
||||
result = format_display_value({"key": "value"}, "raw value")
|
||||
assert result == "{'key': 'value'}"
|
||||
|
||||
def test_format_display_value_tensor_shape(self):
|
||||
"""Test formatting for tensor shape display."""
|
||||
tensor = np.random.rand(10, 10)
|
||||
result = format_display_value(tensor, "tensor shape")
|
||||
assert result == "[[10, 10]]"
|
||||
|
||||
def test_format_display_value_no_tensors(self):
|
||||
"""Test formatting when no tensors present."""
|
||||
result = format_display_value("just a string", "tensor shape")
|
||||
assert result == "No tensors found in input"
|
||||
|
||||
def test_validate_display_mode(self):
|
||||
"""Test display mode validation."""
|
||||
assert validate_display_mode("raw value") is True
|
||||
assert validate_display_mode("tensor shape") is True
|
||||
assert validate_display_mode("invalid") is False
|
||||
assert validate_display_mode("") is False
|
||||
assert validate_display_mode(None) is False
|
||||
|
||||
|
||||
class TestDisplayAnyEdgeCases:
|
||||
"""Test edge cases for DisplayAny."""
|
||||
|
||||
def test_display_none(self):
|
||||
"""Test displaying None value."""
|
||||
node = DisplayAnyNode()
|
||||
result = node.display(None, "raw value")
|
||||
assert result["ui"]["text"] == "None"
|
||||
|
||||
def test_display_empty_list(self):
|
||||
"""Test displaying empty list."""
|
||||
node = DisplayAnyNode()
|
||||
result = node.display([], "raw value")
|
||||
assert result["ui"]["text"] == "[]"
|
||||
|
||||
def test_display_empty_dict(self):
|
||||
"""Test displaying empty dictionary."""
|
||||
node = DisplayAnyNode()
|
||||
result = node.display({}, "raw value")
|
||||
assert result["ui"]["text"] == "{}"
|
||||
|
||||
def test_display_complex_nested_structure(self):
|
||||
"""Test displaying complex nested structure."""
|
||||
node = DisplayAnyNode()
|
||||
complex_data = {
|
||||
"images": [np.random.rand(1, 3, 64, 64) for _ in range(3)],
|
||||
"config": {
|
||||
"steps": 20,
|
||||
"cfg": 7.5,
|
||||
"sampler": "euler",
|
||||
"latents": np.random.rand(1, 4, 32, 32),
|
||||
},
|
||||
"prompts": ["test1", "test2"],
|
||||
}
|
||||
result = node.display(complex_data, "tensor shape")
|
||||
|
||||
# Should find 4 tensors total (3 images + 1 latent)
|
||||
shapes_text = result["ui"]["text"]
|
||||
assert "[1, 3, 64, 64]" in shapes_text
|
||||
assert "[1, 4, 32, 32]" in shapes_text
|
||||
|
||||
def test_display_very_long_string(self):
|
||||
"""Test displaying very long string."""
|
||||
node = DisplayAnyNode()
|
||||
long_string = "x" * 10000
|
||||
result = node.display(long_string, "raw value")
|
||||
assert result["ui"]["text"] == long_string
|
||||
|
||||
def test_display_unicode(self):
|
||||
"""Test displaying unicode characters."""
|
||||
node = DisplayAnyNode()
|
||||
unicode_text = "Hello 世界 🌍"
|
||||
result = node.display(unicode_text, "raw value")
|
||||
assert result["ui"]["text"] == unicode_text
|
||||
@@ -52,7 +52,9 @@ class TestKikoSaveImageLogic:
|
||||
"""Test save path generation"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
# Test basic path generation
|
||||
full_path, filename = get_save_image_path("test_prefix", 0, ".png", temp_dir)
|
||||
full_path, filename = get_save_image_path(
|
||||
"test_prefix", 0, ".png", temp_dir
|
||||
)
|
||||
|
||||
assert full_path.startswith(temp_dir)
|
||||
assert filename.startswith("test_prefix_")
|
||||
@@ -211,20 +213,28 @@ class TestKikoSaveImageLogic:
|
||||
images = torch.rand(1, 32, 32, 3)
|
||||
|
||||
# Quality out of range
|
||||
with pytest.raises(ValueError, match="quality must be an integer between 1 and 100"):
|
||||
with pytest.raises(
|
||||
ValueError, match="quality must be an integer between 1 and 100"
|
||||
):
|
||||
validate_save_inputs(images, "JPEG", 0, 4)
|
||||
|
||||
with pytest.raises(ValueError, match="quality must be an integer between 1 and 100"):
|
||||
with pytest.raises(
|
||||
ValueError, match="quality must be an integer between 1 and 100"
|
||||
):
|
||||
validate_save_inputs(images, "JPEG", 101, 4)
|
||||
|
||||
def test_validate_save_inputs_invalid_compress_level(self):
|
||||
"""Test validation with invalid PNG compression level"""
|
||||
images = torch.rand(1, 32, 32, 3)
|
||||
|
||||
with pytest.raises(ValueError, match="png_compress_level must be an integer between 0 and 9"):
|
||||
with pytest.raises(
|
||||
ValueError, match="png_compress_level must be an integer between 0 and 9"
|
||||
):
|
||||
validate_save_inputs(images, "PNG", 90, -1)
|
||||
|
||||
with pytest.raises(ValueError, match="png_compress_level must be an integer between 0 and 9"):
|
||||
with pytest.raises(
|
||||
ValueError, match="png_compress_level must be an integer between 0 and 9"
|
||||
):
|
||||
validate_save_inputs(images, "PNG", 90, 10)
|
||||
|
||||
def test_save_image_with_format_png(self):
|
||||
|
||||
@@ -56,9 +56,13 @@ class TestDimensionExtraction:
|
||||
with pytest.raises(ValueError, match="Either image or latent must be provided"):
|
||||
extract_dimensions()
|
||||
|
||||
def test_extract_dimensions_both_inputs_prefers_image(self, mock_image_tensor, mock_latent_tensor):
|
||||
def test_extract_dimensions_both_inputs_prefers_image(
|
||||
self, mock_image_tensor, mock_latent_tensor
|
||||
):
|
||||
"""Test that when both inputs provided, image takes precedence"""
|
||||
width, height = extract_dimensions(image=mock_image_tensor, latent=mock_latent_tensor)
|
||||
width, height = extract_dimensions(
|
||||
image=mock_image_tensor, latent=mock_latent_tensor
|
||||
)
|
||||
|
||||
# Should return image dimensions, not latent
|
||||
assert width == 832
|
||||
@@ -89,7 +93,9 @@ class TestScaledDimensionsCalculation:
|
||||
original_width, original_height = 832, 1216
|
||||
scale_factor = 1.5
|
||||
|
||||
new_width, new_height = calculate_scaled_dimensions(original_width, original_height, scale_factor)
|
||||
new_width, new_height = calculate_scaled_dimensions(
|
||||
original_width, original_height, scale_factor
|
||||
)
|
||||
|
||||
# Check aspect ratio is preserved (within floating point precision)
|
||||
original_ratio = original_width / original_height
|
||||
@@ -101,7 +107,9 @@ class TestScaledDimensionsCalculation:
|
||||
base_width, base_height = 1024, 1024
|
||||
|
||||
for scale_factor in sample_scale_factors:
|
||||
width, height = calculate_scaled_dimensions(base_width, base_height, scale_factor)
|
||||
width, height = calculate_scaled_dimensions(
|
||||
base_width, base_height, scale_factor
|
||||
)
|
||||
|
||||
expected_width = int(base_width * scale_factor)
|
||||
expected_height = int(base_height * scale_factor)
|
||||
@@ -205,7 +213,9 @@ class TestResolutionCalculatorNode:
|
||||
"""Test node calculation with IMAGE input"""
|
||||
node = ResolutionCalculatorNode()
|
||||
|
||||
width, height = node.calculate_resolution(scale_factor=2.0, image=mock_image_tensor)
|
||||
width, height = node.calculate_resolution(
|
||||
scale_factor=2.0, image=mock_image_tensor
|
||||
)
|
||||
|
||||
# Original: 832x1216, 2x scale = 1664x2432
|
||||
assert isinstance(width, int)
|
||||
@@ -220,7 +230,9 @@ class TestResolutionCalculatorNode:
|
||||
"""Test node calculation with LATENT input"""
|
||||
node = ResolutionCalculatorNode()
|
||||
|
||||
width, height = node.calculate_resolution(scale_factor=1.5, latent=mock_latent_tensor)
|
||||
width, height = node.calculate_resolution(
|
||||
scale_factor=1.5, latent=mock_latent_tensor
|
||||
)
|
||||
|
||||
# Original: 832x1216, 1.5x scale = 1248x1824
|
||||
assert isinstance(width, int)
|
||||
@@ -238,12 +250,16 @@ class TestResolutionCalculatorNode:
|
||||
with pytest.raises(ValueError):
|
||||
node.calculate_resolution(scale_factor=2.0)
|
||||
|
||||
def test_calculate_resolution_with_various_scale_factors(self, mock_image_tensor_square, sample_scale_factors):
|
||||
def test_calculate_resolution_with_various_scale_factors(
|
||||
self, mock_image_tensor_square, sample_scale_factors
|
||||
):
|
||||
"""Test calculation with various scale factors"""
|
||||
node = ResolutionCalculatorNode()
|
||||
|
||||
for scale_factor in sample_scale_factors:
|
||||
width, height = node.calculate_resolution(scale_factor=scale_factor, image=mock_image_tensor_square)
|
||||
width, height = node.calculate_resolution(
|
||||
scale_factor=scale_factor, image=mock_image_tensor_square
|
||||
)
|
||||
|
||||
# All results should be integers divisible by 8
|
||||
assert isinstance(width, int)
|
||||
|
||||
@@ -331,7 +331,9 @@ class TestSamplerComboIntegration:
|
||||
|
||||
# Test that recommendations work with the node
|
||||
for scheduler in suggestions[:2]: # Test first 2 suggestions
|
||||
result = node.get_sampler_combo(sampler, scheduler, steps_rec["default"], cfg_rec["default"])
|
||||
result = node.get_sampler_combo(
|
||||
sampler, scheduler, steps_rec["default"], cfg_rec["default"]
|
||||
)
|
||||
assert result[0] == sampler
|
||||
assert result[1] == scheduler
|
||||
assert result[2] == steps_rec["default"]
|
||||
|
||||
@@ -70,7 +70,9 @@ class TestWidthHeightSelectorNode:
|
||||
|
||||
# Test formatted preset if available
|
||||
formatted_preset = "832×1216 - 13:19 (1.0MP) - SDXL"
|
||||
result = self.node.get_dimensions(preset=formatted_preset, width=512, height=512)
|
||||
result = self.node.get_dimensions(
|
||||
preset=formatted_preset, width=512, height=512
|
||||
)
|
||||
assert result == (832, 1216)
|
||||
|
||||
def test_sdxl_landscape_preset(self):
|
||||
@@ -81,7 +83,9 @@ class TestWidthHeightSelectorNode:
|
||||
|
||||
# Test formatted preset if available
|
||||
formatted_preset = "1216×832 - 19:13 (1.0MP) - SDXL"
|
||||
result = self.node.get_dimensions(preset=formatted_preset, width=512, height=512)
|
||||
result = self.node.get_dimensions(
|
||||
preset=formatted_preset, width=512, height=512
|
||||
)
|
||||
assert result == (1216, 832)
|
||||
|
||||
def test_flux_preset(self):
|
||||
@@ -92,7 +96,9 @@ class TestWidthHeightSelectorNode:
|
||||
|
||||
# Test formatted preset
|
||||
formatted_preset = "1920×1080 - 16:9 (2.1MP) - FLUX"
|
||||
result = self.node.get_dimensions(preset=formatted_preset, width=512, height=512)
|
||||
result = self.node.get_dimensions(
|
||||
preset=formatted_preset, width=512, height=512
|
||||
)
|
||||
assert result == (1920, 1080)
|
||||
|
||||
def test_ultra_wide_preset(self):
|
||||
@@ -103,7 +109,9 @@ class TestWidthHeightSelectorNode:
|
||||
|
||||
# Test formatted preset if available
|
||||
formatted_preset = "2560×1080 - 64:27 (2.8MP) - Ultra-Wide"
|
||||
result = self.node.get_dimensions(preset=formatted_preset, width=512, height=512)
|
||||
result = self.node.get_dimensions(
|
||||
preset=formatted_preset, width=512, height=512
|
||||
)
|
||||
assert result == (2560, 1080)
|
||||
|
||||
def test_all_presets_available(self):
|
||||
@@ -135,7 +143,9 @@ class TestWidthHeightSelectorNode:
|
||||
def test_invalid_preset_fallback(self):
|
||||
"""Test handling of invalid preset."""
|
||||
# Should fall back to custom dimensions
|
||||
result = self.node.get_dimensions(preset="invalid_preset", width=800, height=600)
|
||||
result = self.node.get_dimensions(
|
||||
preset="invalid_preset", width=800, height=600
|
||||
)
|
||||
assert result == (800, 600)
|
||||
|
||||
|
||||
@@ -258,14 +268,18 @@ class TestPresetDefinitions:
|
||||
for preset_dict in [SDXL_PRESETS, FLUX_PRESETS, ULTRA_WIDE_PRESETS]:
|
||||
for preset_name, (width, height) in preset_dict.items():
|
||||
assert width % 8 == 0, f"{preset_name} width {width} not divisible by 8"
|
||||
assert height % 8 == 0, f"{preset_name} height {height} not divisible by 8"
|
||||
assert (
|
||||
height % 8 == 0
|
||||
), f"{preset_name} height {height} not divisible by 8"
|
||||
|
||||
def test_preset_dimensions_within_limits(self):
|
||||
"""Test that all preset dimensions are within acceptable limits."""
|
||||
for preset_dict in [SDXL_PRESETS, FLUX_PRESETS, ULTRA_WIDE_PRESETS]:
|
||||
for preset_name, (width, height) in preset_dict.items():
|
||||
assert 64 <= width <= 8192, f"{preset_name} width {width} out of range"
|
||||
assert 64 <= height <= 8192, f"{preset_name} height {height} out of range"
|
||||
assert (
|
||||
64 <= height <= 8192
|
||||
), f"{preset_name} height {height} out of range"
|
||||
|
||||
|
||||
class TestEdgeCases:
|
||||
@@ -467,7 +481,9 @@ class TestFormattedPresets:
|
||||
|
||||
for formatted_preset, expected in test_cases:
|
||||
result = self.node._extract_preset_name(formatted_preset)
|
||||
assert result == expected, f"Expected {expected}, got {result} for input {formatted_preset}"
|
||||
assert (
|
||||
result == expected
|
||||
), f"Expected {expected}, got {result} for input {formatted_preset}"
|
||||
|
||||
def test_formatted_preset_dimensions(self):
|
||||
"""Test that formatted presets return correct dimensions."""
|
||||
@@ -509,22 +525,34 @@ class TestFormattedPresets:
|
||||
def test_formatted_preset_metadata_accuracy(self):
|
||||
"""Test that formatted presets contain accurate metadata."""
|
||||
input_types = self.node.INPUT_TYPES()
|
||||
formatted_presets = [opt for opt in input_types["required"]["preset"][0] if " - " in opt]
|
||||
formatted_presets = [
|
||||
opt for opt in input_types["required"]["preset"][0] if " - " in opt
|
||||
]
|
||||
|
||||
for formatted_preset in formatted_presets:
|
||||
# Extract components
|
||||
parts = formatted_preset.split(" - ")
|
||||
assert len(parts) == 3, f"Formatted preset should have 3 parts: {formatted_preset}"
|
||||
assert (
|
||||
len(parts) == 3
|
||||
), f"Formatted preset should have 3 parts: {formatted_preset}"
|
||||
|
||||
resolution = parts[0]
|
||||
aspect_and_mp = parts[1]
|
||||
model_group = parts[2]
|
||||
|
||||
# Verify resolution exists in metadata
|
||||
assert resolution in PRESET_METADATA, f"Resolution {resolution} not in metadata"
|
||||
assert (
|
||||
resolution in PRESET_METADATA
|
||||
), f"Resolution {resolution} not in metadata"
|
||||
|
||||
# Verify metadata matches format
|
||||
metadata = PRESET_METADATA[resolution]
|
||||
assert metadata.model_group == model_group, f"Model group mismatch for {resolution}"
|
||||
assert metadata.aspect_ratio in aspect_and_mp, f"Aspect ratio not in {aspect_and_mp}"
|
||||
assert f"{metadata.megapixels:.1f}MP" in aspect_and_mp, f"Megapixels not in {aspect_and_mp}"
|
||||
assert (
|
||||
metadata.model_group == model_group
|
||||
), f"Model group mismatch for {resolution}"
|
||||
assert (
|
||||
metadata.aspect_ratio in aspect_and_mp
|
||||
), f"Aspect ratio not in {aspect_and_mp}"
|
||||
assert (
|
||||
f"{metadata.megapixels:.1f}MP" in aspect_and_mp
|
||||
), f"Megapixels not in {aspect_and_mp}"
|
||||
|
||||
Reference in New Issue
Block a user