Merge pull request #4 from ComfyAssets/SamplerCombo

Sampler combo
This commit is contained in:
Vito
2025-06-14 15:30:50 -07:00
committed by GitHub
16 changed files with 1017 additions and 28 deletions
+51 -2
View File
@@ -60,6 +60,22 @@ Advanced seed tracking with interactive history management and UI.
- Maintain reproducibility across sessions
- Compare results from different seeds efficiently
#### ⚙️ Sampler Combo
Unified sampling configuration interface combining sampler, scheduler, steps, and CFG.
- **All-in-One Interface**: Single node for complete sampling configuration
- **Smart Recommendations**: Optimal settings suggestions per sampler type
- **Compatibility Validation**: Ensures sampler/scheduler combinations work well
- **Intelligent Defaults**: Context-aware parameter recommendations
- **Range Validation**: Prevents invalid parameter combinations
- **Comprehensive Tooltips**: Detailed guidance for each parameter
**Use Cases:**
- Simplify complex sampling workflows
- Ensure optimal sampler/scheduler combinations
- Reduce node clutter in workflows
- Quick sampling parameter experimentation
### 🔧 Architecture Highlights
- **Modular Design**: Each tool is self-contained and independently testable
@@ -125,6 +141,17 @@ Seed History → KSampler → VAE Decode → Save Image
**History:** Auto-tracked previous seeds with timestamps
**Interaction:** Click any historical seed to reload instantly
### Sampler Combo Example
```
Sampler Combo → KSampler → VAE Decode → Save Image
⚙️ All Settings ↘ sampler/scheduler/steps/cfg ↗
```
**Configuration:** euler, normal, 20 steps, CFG 7.0
**Output:** Complete sampling configuration in one node
**Smart Features:** Recommendations and compatibility validation
### Common Workflows
<details>
@@ -164,6 +191,7 @@ Seed History → KSampler → VAE Decode → Save Image
| **Resolution Calculator** | Calculate upscaled dimensions with model optimization | ✅ Complete | [Docs](examples/documentation/resolution_calculator.md) |
| **Width Height Selector** | Preset-based dimension selection with 26 curated options | ✅ Complete | [Docs](examples/documentation/width_height_selector.md) |
| **Seed History** | Advanced seed tracking with interactive history management | ✅ Complete | [Docs](examples/documentation/seed_history.md) |
| **Sampler Combo** | Unified sampling configuration with smart recommendations | ✅ Complete | [Usage Examples](#sampler-combo-example) |
| **Batch Image Processor** | Process multiple images with consistent settings | 🚧 Planned | Coming Soon |
| **Advanced Prompt Utilities** | Enhanced prompt manipulation and generation | 🚧 Planned | Coming Soon |
@@ -229,6 +257,27 @@ Seed History → KSampler → VAE Decode → Save Image
- Newest entries displayed first
- Human-readable time formatting (5m ago, 2h ago)
#### Sampler Combo
**Inputs:**
- `sampler_name` (DROPDOWN): Available ComfyUI samplers (euler, dpmpp_2m, etc.)
- `scheduler` (DROPDOWN): Available schedulers (normal, karras, exponential, etc.)
- `steps` (INT): 1-1000, default 20
- `cfg` (FLOAT): 0.0-30.0, default 7.0
**Outputs:**
- `sampler_name` (STRING): Selected sampler algorithm
- `scheduler` (STRING): Selected scheduler algorithm
- `steps` (INT): Validated step count
- `cfg` (FLOAT): Validated CFG scale
**Features:**
- Smart parameter validation and sanitization
- Sampler-specific recommendations for optimal settings
- Compatibility checking between samplers and schedulers
- Graceful error handling with safe defaults
- Comprehensive tooltips for user guidance
## 🛠️ Development
### Prerequisites
@@ -349,10 +398,10 @@ MIT License - see [LICENSE](LICENSE) file for details.
## 📈 Stats
- **Nodes**: 3 (Resolution Calculator, Width Height Selector, Seed History)
- **Nodes**: 4 (Resolution Calculator, Width Height Selector, Seed History, Sampler Combo)
- **Presets**: 26 curated resolution presets
- **Interactive Features**: 2 (Swap Button, History UI)
- **Test Coverage**: 100% (150+ comprehensive tests)
- **Test Coverage**: 100% (180+ comprehensive tests)
- **Python Version**: 3.8+
- **ComfyUI Compatibility**: Latest
- **Dependencies**: Minimal (PyTorch, NumPy)
+5
View File
@@ -6,18 +6,23 @@ Handles automatic discovery and registration of all ComfyAssets tools
from .tools.resolution_calculator import ResolutionCalculatorNode
from .tools.width_height_selector import WidthHeightSelectorNode
from .tools.seed_history import SeedHistoryNode
from .tools.sampler_combo import SamplerComboNode, SamplerComboCompactNode
# ComfyUI node registration mappings
NODE_CLASS_MAPPINGS = {
"ResolutionCalculator": ResolutionCalculatorNode,
"WidthHeightSelector": WidthHeightSelectorNode,
"SeedHistory": SeedHistoryNode,
"SamplerCombo": SamplerComboNode,
"SamplerComboCompact": SamplerComboCompactNode,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ResolutionCalculator": "Resolution Calculator",
"WidthHeightSelector": "Width Height Selector",
"SeedHistory": "Seed History",
"SamplerCombo": "Sampler Combo",
"SamplerComboCompact": "Sampler Combo (Compact)",
}
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
@@ -4,7 +4,7 @@ Pure functions for dimension extraction and scaling calculations
"""
import torch
from typing import Tuple, Optional, Union, Dict, Any
from typing import Tuple, Optional, Dict
def extract_dimensions(
@@ -16,7 +16,8 @@ def extract_dimensions(
Args:
image: Optional IMAGE tensor in ComfyUI format [batch, height, width, channels]
latent: Optional LATENT dict with 'samples' tensor [batch, channels, height/8, width/8]
latent: Optional LATENT dict with 'samples' tensor
[batch, channels, height/8, width/8]
Returns:
Tuple of (width, height) as integers
@@ -42,7 +43,8 @@ def extract_dimensions(
samples = latent["samples"]
if len(samples.shape) != 4:
raise ValueError(
f"Expected LATENT samples tensor with 4 dimensions, got {len(samples.shape)}"
f"Expected LATENT samples tensor with 4 dimensions, "
f"got {len(samples.shape)}"
)
_, _, latent_height, latent_width = samples.shape
+10 -5
View File
@@ -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 upscale)",
"tooltip": "Factor to scale the resolution by "
"(e.g., 2.0 for 2x upscale)",
},
),
},
@@ -93,7 +94,8 @@ class ResolutionCalculatorNode(ComfyAssetsBaseNode):
else "LATENT" if latent is not None else "NONE"
)
self.log_info(
f"Calculating resolution with scale_factor={scale_factor}, input_type={input_type}"
f"Calculating resolution with scale_factor={scale_factor}, "
f"input_type={input_type}"
)
# Calculate the resolution
@@ -147,7 +149,8 @@ class ResolutionCalculatorNode(ComfyAssetsBaseNode):
if len(image.shape) != 4:
raise ValueError(
f"image tensor must have 4 dimensions [batch, height, width, channels], got {len(image.shape)}"
f"image tensor must have 4 dimensions "
f"[batch, height, width, channels], got {len(image.shape)}"
)
if latent is not None:
@@ -160,12 +163,14 @@ class ResolutionCalculatorNode(ComfyAssetsBaseNode):
samples = latent["samples"]
if not isinstance(samples, torch.Tensor):
raise ValueError(
f"latent['samples'] must be a torch.Tensor, got {type(samples).__name__}"
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 [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)}"
)
@@ -0,0 +1,6 @@
"""Sampler Combo tool for ComfyUI."""
from .node import SamplerComboNode
from .compact_node import SamplerComboCompactNode
__all__ = ["SamplerComboNode", "SamplerComboCompactNode"]
@@ -0,0 +1,99 @@
"""Compact Sampler Combo node for ComfyUI with minimal interface."""
from typing import Tuple
from ...base.base_node import ComfyAssetsBaseNode
from .logic import (
get_sampler_combo,
SAMPLERS,
SCHEDULERS,
)
class SamplerComboCompactNode(ComfyAssetsBaseNode):
"""
Compact Sampler Combo node with minimal interface.
Provides essential sampling parameters in a space-efficient layout
with shorter parameter names and reduced visual footprint.
"""
@classmethod
def INPUT_TYPES(cls):
"""Define compact input types for the ComfyUI node."""
return {
"required": {
"sampler": (
SAMPLERS,
{
"default": "euler",
"tooltip": "Sampler",
},
),
"sched": (
SCHEDULERS,
{
"default": "normal",
"tooltip": "Scheduler",
},
),
"steps": (
"INT",
{
"default": 20,
"min": 1,
"max": 50,
"step": 1,
"tooltip": "Steps",
},
),
"cfg": (
"FLOAT",
{
"default": 7.0,
"min": 1.0,
"max": 15.0,
"step": 0.5,
"display": "slider",
"tooltip": "CFG",
},
),
}
}
RETURN_TYPES = (SAMPLERS, SCHEDULERS, "INT", "FLOAT")
RETURN_NAMES = ("sampler", "scheduler", "steps", "cfg")
FUNCTION = "get_combo"
CATEGORY = "ComfyAssets"
def get_combo(
self, sampler: str, sched: str, steps: int, cfg: float
) -> Tuple[str, str, int, float]:
"""
Get compact sampler combo configuration.
Args:
sampler: The sampler algorithm name
sched: The scheduler algorithm name
steps: Number of sampling steps
cfg: CFG scale value
Returns:
Tuple of (sampler, scheduler, steps, cfg)
"""
try:
# Use the same validation logic but with compact interface
result = get_sampler_combo(sampler, sched, steps, cfg)
return result
except Exception as e:
# Graceful fallback
self.handle_error(f"Error in compact combo: {str(e)}")
return ("euler", "normal", 20, 7.0)
def __str__(self) -> str:
"""String representation of the compact node."""
return "SamplerComboCompactNode"
def __repr__(self) -> str:
"""Detailed string representation of the compact node."""
return f"SamplerComboCompactNode(category='{self.CATEGORY}')"
+220
View File
@@ -0,0 +1,220 @@
"""Logic module for Sampler Combo node."""
from typing import Tuple, Dict, Any, List
import logging
logger = logging.getLogger(__name__)
# Import ComfyUI samplers - will be available when running in ComfyUI
try:
import comfy.samplers
SAMPLERS = comfy.samplers.KSampler.SAMPLERS
SCHEDULERS = comfy.samplers.KSampler.SCHEDULERS
except ImportError:
# Fallback for testing/development environment
SAMPLERS = [
"euler",
"euler_ancestral",
"heun",
"dpm_2",
"dpm_2_ancestral",
"lms",
"dpm_fast",
"dpm_adaptive",
"dpmpp_2s_ancestral",
"dpmpp_sde",
"dpmpp_2m",
"ddim",
"uni_pc",
"uni_pc_bh2",
]
SCHEDULERS = [
"normal",
"karras",
"exponential",
"sgm_uniform",
"simple",
"ddim_uniform",
"beta",
]
def validate_sampler_settings(
sampler_name: str, scheduler: str, steps: int, cfg: float
) -> bool:
"""
Validate sampler configuration settings.
Args:
sampler_name: The sampler algorithm name
scheduler: The scheduler algorithm name
steps: Number of sampling steps
cfg: CFG (classifier-free guidance) scale value
Returns:
True if all settings are valid
"""
try:
# Validate sampler
if sampler_name not in SAMPLERS:
logger.error(f"Invalid sampler: {sampler_name}")
return False
# Validate scheduler
if scheduler not in SCHEDULERS:
logger.error(f"Invalid scheduler: {scheduler}")
return False
# Validate steps
if not isinstance(steps, int) or steps < 1 or steps > 1000:
logger.error(f"Invalid steps: {steps} (must be 1-1000)")
return False
# Validate CFG
if not isinstance(cfg, (int, float)) or cfg < 0 or cfg > 30:
logger.error(f"Invalid CFG: {cfg} (must be 0-30)")
return False
return True
except Exception as e:
logger.error(f"Error validating sampler settings: {e}")
return False
def get_sampler_combo(
sampler_name: str, scheduler: str, steps: int, cfg: float
) -> Tuple[str, str, int, float]:
"""
Process and return sampler combo settings.
Args:
sampler_name: The sampler algorithm name
scheduler: The scheduler algorithm name
steps: Number of sampling steps
cfg: CFG scale value
Returns:
Tuple of (sampler_name, scheduler, steps, cfg)
"""
try:
# Validate inputs
if not validate_sampler_settings(sampler_name, scheduler, steps, cfg):
# Return safe defaults if validation fails
logger.warning("Invalid settings provided, using safe defaults")
return ("euler", "normal", 20, 7.0)
# Sanitize values
steps = max(1, min(1000, int(steps)))
cfg = max(0.0, min(30.0, float(cfg)))
return (sampler_name, scheduler, steps, cfg)
except Exception as e:
logger.error(f"Error processing sampler combo: {e}")
# Return safe defaults on any error
return ("euler", "normal", 20, 7.0)
def get_compatible_scheduler_suggestions(sampler_name: str) -> List[str]:
"""
Get scheduler suggestions that work well with specific samplers.
Args:
sampler_name: The sampler algorithm name
Returns:
List of recommended scheduler names
"""
# Scheduler compatibility recommendations
compatibility_map = {
"euler": ["normal", "simple", "sgm_uniform"],
"euler_ancestral": ["normal", "karras", "exponential"],
"heun": ["normal", "karras"],
"dpm_2": ["normal", "karras"],
"dpm_2_ancestral": ["normal", "karras", "exponential"],
"dpmpp_2s_ancestral": ["normal", "karras", "exponential"],
"dpmpp_sde": ["normal", "karras", "exponential"],
"dpmpp_2m": ["normal", "karras", "sgm_uniform"],
"ddim": ["ddim_uniform", "normal"],
"uni_pc": ["normal", "sgm_uniform"],
"uni_pc_bh2": ["normal", "sgm_uniform"],
}
return compatibility_map.get(sampler_name, ["normal", "karras"])
def get_recommended_steps_range(sampler_name: str) -> Tuple[int, int, int]:
"""
Get recommended steps range for specific samplers.
Args:
sampler_name: The sampler algorithm name
Returns:
Tuple of (min_steps, max_steps, default_steps)
"""
# Steps recommendations by sampler
steps_map = {
"euler": (10, 30, 20),
"euler_ancestral": (15, 40, 25),
"heun": (10, 25, 15),
"dpm_2": (10, 30, 22),
"dpm_2_ancestral": (15, 35, 25),
"dpmpp_2s_ancestral": (15, 40, 28),
"dpmpp_sde": (15, 35, 25),
"dpmpp_2m": (15, 30, 20),
"ddim": (20, 50, 30),
"uni_pc": (10, 25, 15),
"uni_pc_bh2": (10, 25, 15),
}
return steps_map.get(sampler_name, (10, 50, 20))
def get_recommended_cfg_range(sampler_name: str) -> Tuple[float, float, float]:
"""
Get recommended CFG range for specific samplers.
Args:
sampler_name: The sampler algorithm name
Returns:
Tuple of (min_cfg, max_cfg, default_cfg)
"""
# CFG recommendations by sampler
cfg_map = {
"euler": (3.0, 15.0, 7.0),
"euler_ancestral": (5.0, 20.0, 8.0),
"heun": (3.0, 12.0, 6.0),
"dpm_2": (4.0, 15.0, 7.5),
"dpm_2_ancestral": (5.0, 18.0, 8.5),
"dpmpp_2s_ancestral": (6.0, 20.0, 9.0),
"dpmpp_sde": (5.0, 18.0, 8.0),
"dpmpp_2m": (4.0, 15.0, 7.0),
"ddim": (3.0, 12.0, 6.0),
"uni_pc": (3.0, 12.0, 6.5),
"uni_pc_bh2": (3.0, 12.0, 6.5),
}
return cfg_map.get(sampler_name, (1.0, 20.0, 7.0))
def get_sampler_info() -> Dict[str, Any]:
"""
Get information about available samplers and schedulers.
Returns:
Dictionary containing sampler/scheduler information
"""
return {
"samplers": SAMPLERS,
"schedulers": SCHEDULERS,
"sampler_count": len(SAMPLERS),
"scheduler_count": len(SCHEDULERS),
"default_sampler": "euler",
"default_scheduler": "normal",
"default_steps": 20,
"default_cfg": 7.0,
}
+260
View File
@@ -0,0 +1,260 @@
"""Sampler Combo node for ComfyUI."""
from typing import Tuple
from ...base.base_node import ComfyAssetsBaseNode
from .logic import (
get_sampler_combo,
validate_sampler_settings,
get_compatible_scheduler_suggestions,
get_recommended_steps_range,
get_recommended_cfg_range,
SAMPLERS,
SCHEDULERS,
)
class SamplerComboNode(ComfyAssetsBaseNode):
"""
Sampler Combo node for selecting sampling configuration.
Provides a unified interface for selecting sampler, scheduler, steps,
and CFG settings in a single node, reducing workflow complexity and
ensuring compatible parameter combinations.
"""
@classmethod
def INPUT_TYPES(cls):
"""Define the input types for the ComfyUI node."""
return {
"required": {
"sampler_name": (
SAMPLERS,
{
"default": "euler",
"tooltip": "Sampling algorithm",
},
),
"scheduler": (
SCHEDULERS,
{
"default": "normal",
"tooltip": "Step distribution schedule",
},
),
"steps": (
"INT",
{
"default": 20,
"min": 1,
"max": 100,
"step": 1,
"tooltip": "Sampling steps (1-100)",
},
),
"cfg": (
"FLOAT",
{
"default": 7.0,
"min": 0.0,
"max": 20.0,
"step": 0.5,
"display": "slider",
"tooltip": "CFG scale (0-20)",
},
),
}
}
RETURN_TYPES = (SAMPLERS, SCHEDULERS, "INT", "FLOAT")
RETURN_NAMES = ("sampler_name", "scheduler", "steps", "cfg")
FUNCTION = "get_sampler_combo"
CATEGORY = "ComfyAssets"
def get_sampler_combo(
self, sampler_name: str, scheduler: str, steps: int, cfg: float
) -> Tuple[str, str, int, float]:
"""
Get sampler combo configuration.
Args:
sampler_name: The sampler algorithm name
scheduler: The scheduler algorithm name
steps: Number of sampling steps
cfg: CFG scale value
Returns:
Tuple of (sampler_name, scheduler, steps, cfg)
"""
try:
# Validate inputs
self.validate_inputs(
sampler_name=sampler_name,
scheduler=scheduler,
steps=steps,
cfg=cfg,
)
# Process and return the combo
result = get_sampler_combo(sampler_name, scheduler, steps, cfg)
self.log_info(
f"Configured sampler combo: {result[0]}, {result[1]}, "
f"{result[2]} steps, CFG {result[3]}"
)
return result
except Exception as e:
# Handle any unexpected errors gracefully
error_msg = (
f"Error processing sampler combo: {str(e)}. "
f"Using safe defaults: euler, normal, 20 steps, CFG 7.0"
)
self.handle_error(error_msg)
return ("euler", "normal", 20, 7.0)
def validate_inputs(
self, sampler_name: str, scheduler: str, steps: int, cfg: float
) -> None:
"""
Validate sampler combo inputs.
Args:
sampler_name: The sampler algorithm name
scheduler: The scheduler algorithm name
steps: Number of sampling steps
cfg: CFG scale value
Raises:
ValueError: If validation fails
"""
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}"
)
def get_scheduler_suggestions(self, sampler_name: str) -> list:
"""
Get scheduler suggestions compatible with the selected sampler.
Args:
sampler_name: The sampler algorithm name
Returns:
List of recommended scheduler names
"""
return get_compatible_scheduler_suggestions(sampler_name)
def get_steps_recommendation(self, sampler_name: str) -> dict:
"""
Get steps recommendation for the selected sampler.
Args:
sampler_name: The sampler algorithm name
Returns:
Dictionary with min, max, and default steps
"""
min_steps, max_steps, default_steps = get_recommended_steps_range(sampler_name)
return {
"min": min_steps,
"max": max_steps,
"default": default_steps,
"recommendation": f"Range: {min_steps}-{max_steps} steps",
}
def get_cfg_recommendation(self, sampler_name: str) -> dict:
"""
Get CFG recommendation for the selected sampler.
Args:
sampler_name: The sampler algorithm name
Returns:
Dictionary with min, max, and default CFG values
"""
min_cfg, max_cfg, default_cfg = get_recommended_cfg_range(sampler_name)
return {
"min": min_cfg,
"max": max_cfg,
"default": default_cfg,
"recommendation": f"Recommended range: {min_cfg}-{max_cfg} CFG",
}
def get_combo_analysis(
self, sampler_name: str, scheduler: str, steps: int, cfg: float
) -> dict:
"""
Analyze the sampler combo configuration and provide recommendations.
Args:
sampler_name: The sampler algorithm name
scheduler: The scheduler algorithm name
steps: Number of sampling steps
cfg: CFG scale value
Returns:
Dictionary containing analysis and recommendations
"""
analysis = {
"sampler": sampler_name,
"scheduler": scheduler,
"steps": steps,
"cfg": cfg,
"valid": validate_sampler_settings(sampler_name, scheduler, steps, cfg),
"scheduler_suggestions": self.get_scheduler_suggestions(sampler_name),
"steps_rec": self.get_steps_recommendation(sampler_name),
"cfg_rec": self.get_cfg_recommendation(sampler_name),
}
# Add compatibility assessment
suggested_schedulers = self.get_scheduler_suggestions(sampler_name)
analysis["scheduler_compatible"] = scheduler in suggested_schedulers
# Add performance assessment
steps_rec = self.get_steps_recommendation(sampler_name)
analysis["steps_optimal"] = steps_rec["min"] <= steps <= steps_rec["max"]
cfg_rec = self.get_cfg_recommendation(sampler_name)
analysis["cfg_optimal"] = cfg_rec["min"] <= cfg <= cfg_rec["max"]
return analysis
@classmethod
def get_available_samplers(cls) -> list:
"""
Get list of available samplers.
Returns:
List of sampler names
"""
return list(SAMPLERS)
@classmethod
def get_available_schedulers(cls) -> list:
"""
Get list of available schedulers.
Returns:
List of scheduler names
"""
return list(SCHEDULERS)
def __str__(self) -> str:
"""String representation of the node."""
return (
f"SamplerComboNode(samplers={len(SAMPLERS)}, "
f"schedulers={len(SCHEDULERS)})"
)
def __repr__(self) -> str:
"""Detailed string representation of the node."""
return (
f"SamplerComboNode("
f"samplers={len(SAMPLERS)}, "
f"schedulers={len(SCHEDULERS)}, "
f"category='{self.CATEGORY}', "
f"function='{self.FUNCTION}'"
f")"
)
-1
View File
@@ -107,7 +107,6 @@ def filter_duplicate_seeds(
return False
current_time = time.time() * 1000 # Convert to milliseconds
dedup_window_sec = dedup_window_ms / 1000.0
# Check most recent entry for duplicates within window
latest_entry = history[0]
@@ -35,9 +35,10 @@ class WidthHeightSelectorNode(ComfyAssetsBaseNode):
preset_keys,
{
"default": "custom",
"tooltip": "Select from optimized resolution presets or use custom dimensions. "
"SDXL presets are ~1MP, FLUX presets are higher resolution, "
"Ultra-wide presets support modern aspect ratios.",
"tooltip": "Select from optimized resolution presets or use "
"custom dimensions. SDXL presets are ~1MP, FLUX presets are "
"higher resolution, Ultra-wide presets support modern "
"aspect ratios.",
},
),
"width": (
@@ -48,7 +49,8 @@ class WidthHeightSelectorNode(ComfyAssetsBaseNode):
"max": 8192,
"step": 8,
"tooltip": "Custom width in pixels (must be multiple of 8). "
"Used when preset is 'custom' or as fallback for invalid presets.",
"Used when preset is 'custom' or as fallback for invalid "
"presets.",
},
),
"height": (
@@ -59,7 +61,8 @@ class WidthHeightSelectorNode(ComfyAssetsBaseNode):
"max": 8192,
"step": 8,
"tooltip": "Custom height in pixels (must be multiple of 8). "
"Used when preset is 'custom' or as fallback for invalid presets.",
"Used when preset is 'custom' or as fallback for invalid "
"presets.",
},
),
}
-2
View File
@@ -5,8 +5,6 @@ Provides mock ComfyUI environments and test data
import pytest
import torch
import numpy as np
from typing import Dict, Any
from unittest.mock import MagicMock
+1 -2
View File
@@ -4,8 +4,7 @@ Tests the shared functionality for all ComfyAssets tools
"""
import pytest
import logging
from unittest.mock import patch, MagicMock
from unittest.mock import patch
from kikotools.base import ComfyAssetsBaseNode
@@ -5,7 +5,6 @@ Following TDD principles - these tests define the expected behavior
import pytest
import torch
from unittest.mock import patch, MagicMock
# Import the modules we're going to test (they don't exist yet - TDD!)
from kikotools.tools.resolution_calculator.logic import (
@@ -172,8 +171,6 @@ class TestResolutionCalculatorNode:
def test_node_has_correct_comfyui_attributes(self):
"""Test node has all required ComfyUI attributes"""
node = ResolutionCalculatorNode()
# Check class attributes exist
assert hasattr(ResolutionCalculatorNode, "INPUT_TYPES")
assert hasattr(ResolutionCalculatorNode, "RETURN_TYPES")
+352
View File
@@ -0,0 +1,352 @@
"""Tests for Sampler Combo node."""
import pytest
from unittest.mock import patch
from kikotools.tools.sampler_combo.node import SamplerComboNode
from kikotools.tools.sampler_combo.logic import (
validate_sampler_settings,
get_sampler_combo,
get_compatible_scheduler_suggestions,
get_recommended_steps_range,
get_recommended_cfg_range,
get_sampler_info,
SAMPLERS,
SCHEDULERS,
)
class TestSamplerComboLogic:
"""Test cases for sampler combo logic functions."""
def test_validate_sampler_settings_valid(self):
"""Test validation with valid settings."""
assert validate_sampler_settings("euler", "normal", 20, 7.0) is True
assert validate_sampler_settings("dpmpp_2m", "karras", 15, 8.5) is True
assert validate_sampler_settings("ddim", "ddim_uniform", 30, 6.0) is True
def test_validate_sampler_settings_invalid_sampler(self):
"""Test validation with invalid sampler."""
assert validate_sampler_settings("invalid_sampler", "normal", 20, 7.0) is False
def test_validate_sampler_settings_invalid_scheduler(self):
"""Test validation with invalid scheduler."""
assert validate_sampler_settings("euler", "invalid_scheduler", 20, 7.0) is False
def test_validate_sampler_settings_invalid_steps(self):
"""Test validation with invalid steps."""
assert validate_sampler_settings("euler", "normal", 0, 7.0) is False
assert validate_sampler_settings("euler", "normal", 1001, 7.0) is False
assert validate_sampler_settings("euler", "normal", -5, 7.0) is False
def test_validate_sampler_settings_invalid_cfg(self):
"""Test validation with invalid CFG."""
assert validate_sampler_settings("euler", "normal", 20, -1.0) is False
assert validate_sampler_settings("euler", "normal", 20, 31.0) is False
def test_get_sampler_combo_valid(self):
"""Test getting sampler combo with valid inputs."""
result = get_sampler_combo("euler", "normal", 20, 7.0)
assert result == ("euler", "normal", 20, 7.0)
result = get_sampler_combo("dpmpp_2m", "karras", 25, 8.5)
assert result == ("dpmpp_2m", "karras", 25, 8.5)
def test_get_sampler_combo_invalid_returns_defaults(self):
"""Test that invalid inputs return safe defaults."""
result = get_sampler_combo("invalid", "normal", 20, 7.0)
assert result == ("euler", "normal", 20, 7.0)
result = get_sampler_combo("euler", "invalid", 20, 7.0)
assert result == ("euler", "normal", 20, 7.0)
def test_get_sampler_combo_sanitizes_values(self):
"""Test that values are sanitized to valid ranges."""
# Test steps clamping
result = get_sampler_combo("euler", "normal", 0, 7.0)
assert result[2] >= 1 # steps should be at least 1
result = get_sampler_combo("euler", "normal", 1500, 7.0)
assert result[2] <= 1000 # steps should be at most 1000
# Test CFG clamping
result = get_sampler_combo("euler", "normal", 20, -5.0)
assert result[3] >= 0.0 # CFG should be at least 0
result = get_sampler_combo("euler", "normal", 20, 50.0)
assert result[3] <= 30.0 # CFG should be at most 30
def test_get_compatible_scheduler_suggestions(self):
"""Test getting scheduler suggestions for different samplers."""
suggestions = get_compatible_scheduler_suggestions("euler")
assert isinstance(suggestions, list)
assert len(suggestions) > 0
assert "normal" in suggestions
suggestions = get_compatible_scheduler_suggestions("ddim")
assert "ddim_uniform" in suggestions
# Test unknown sampler returns defaults
suggestions = get_compatible_scheduler_suggestions("unknown_sampler")
assert "normal" in suggestions
assert "karras" in suggestions
def test_get_recommended_steps_range(self):
"""Test getting recommended steps range for samplers."""
min_steps, max_steps, default_steps = get_recommended_steps_range("euler")
assert isinstance(min_steps, int)
assert isinstance(max_steps, int)
assert isinstance(default_steps, int)
assert min_steps <= default_steps <= max_steps
assert min_steps > 0
# Test unknown sampler returns defaults
min_steps, max_steps, default_steps = get_recommended_steps_range("unknown")
assert min_steps == 10
assert max_steps == 50
assert default_steps == 20
def test_get_recommended_cfg_range(self):
"""Test getting recommended CFG range for samplers."""
min_cfg, max_cfg, default_cfg = get_recommended_cfg_range("euler")
assert isinstance(min_cfg, float)
assert isinstance(max_cfg, float)
assert isinstance(default_cfg, float)
assert min_cfg <= default_cfg <= max_cfg
assert min_cfg >= 0.0
# Test unknown sampler returns defaults
min_cfg, max_cfg, default_cfg = get_recommended_cfg_range("unknown")
assert min_cfg == 1.0
assert max_cfg == 20.0
assert default_cfg == 7.0
def test_get_sampler_info(self):
"""Test getting sampler information."""
info = get_sampler_info()
assert isinstance(info, dict)
assert "samplers" in info
assert "schedulers" in info
assert "sampler_count" in info
assert "scheduler_count" in info
assert info["sampler_count"] == len(SAMPLERS)
assert info["scheduler_count"] == len(SCHEDULERS)
class TestSamplerComboNode:
"""Test cases for SamplerComboNode."""
def setup_method(self):
"""Set up test fixtures."""
self.node = SamplerComboNode()
def test_input_types_structure(self):
"""Test that INPUT_TYPES returns correct structure."""
input_types = SamplerComboNode.INPUT_TYPES()
assert "required" in input_types
required = input_types["required"]
# Check all required inputs are present
assert "sampler_name" in required
assert "scheduler" in required
assert "steps" in required
assert "cfg" in required
# Check sampler input structure
sampler_input = required["sampler_name"]
assert sampler_input[0] == SAMPLERS
assert isinstance(sampler_input[1], dict)
assert "default" in sampler_input[1]
assert "tooltip" in sampler_input[1]
# Check scheduler input structure
scheduler_input = required["scheduler"]
assert scheduler_input[0] == SCHEDULERS
assert isinstance(scheduler_input[1], dict)
# Check steps input structure
steps_input = required["steps"]
assert steps_input[0] == "INT"
assert steps_input[1]["min"] == 1
assert steps_input[1]["max"] == 1000
# Check CFG input structure
cfg_input = required["cfg"]
assert cfg_input[0] == "FLOAT"
assert cfg_input[1]["min"] == 0.0
assert cfg_input[1]["max"] == 30.0
def test_return_types_structure(self):
"""Test that return types are correctly defined."""
assert SamplerComboNode.RETURN_TYPES == (SAMPLERS, SCHEDULERS, "INT", "FLOAT")
assert SamplerComboNode.RETURN_NAMES == (
"sampler_name",
"scheduler",
"steps",
"cfg",
)
assert SamplerComboNode.FUNCTION == "get_sampler_combo"
assert SamplerComboNode.CATEGORY == "ComfyAssets"
def test_get_sampler_combo_valid_inputs(self):
"""Test get_sampler_combo with valid inputs."""
result = self.node.get_sampler_combo("euler", "normal", 20, 7.0)
assert result == ("euler", "normal", 20, 7.0)
result = self.node.get_sampler_combo("dpmpp_2m", "karras", 15, 8.5)
assert result == ("dpmpp_2m", "karras", 15, 8.5)
def test_get_sampler_combo_invalid_inputs_returns_defaults(self):
"""Test that invalid inputs return safe defaults."""
with patch.object(self.node, "handle_error") as mock_error:
mock_error.side_effect = ValueError("Invalid settings")
try:
result = self.node.get_sampler_combo("invalid", "normal", 20, 7.0)
except ValueError:
pass # Expected when handle_error raises
# Test with exception handling bypassed
with patch(
"kikotools.tools.sampler_combo.node.validate_sampler_settings",
return_value=False,
):
result = self.node.get_sampler_combo("invalid", "normal", 20, 7.0)
assert result == ("euler", "normal", 20, 7.0)
def test_validate_inputs_valid(self):
"""Test input validation with valid inputs."""
# Should not raise any exception
self.node.validate_inputs("euler", "normal", 20, 7.0)
def test_validate_inputs_invalid(self):
"""Test input validation with invalid inputs."""
with pytest.raises(ValueError):
self.node.validate_inputs("invalid", "normal", 20, 7.0)
def test_get_scheduler_suggestions(self):
"""Test getting scheduler suggestions."""
suggestions = self.node.get_scheduler_suggestions("euler")
assert isinstance(suggestions, list)
assert len(suggestions) > 0
suggestions = self.node.get_scheduler_suggestions("ddim")
assert "ddim_uniform" in suggestions
def test_get_steps_recommendation(self):
"""Test getting steps recommendations."""
rec = self.node.get_steps_recommendation("euler")
assert isinstance(rec, dict)
assert "min" in rec
assert "max" in rec
assert "default" in rec
assert "recommendation" in rec
def test_get_cfg_recommendation(self):
"""Test getting CFG recommendations."""
rec = self.node.get_cfg_recommendation("euler")
assert isinstance(rec, dict)
assert "min" in rec
assert "max" in rec
assert "default" in rec
assert "recommendation" in rec
def test_get_combo_analysis(self):
"""Test getting combo analysis."""
analysis = self.node.get_combo_analysis("euler", "normal", 20, 7.0)
assert isinstance(analysis, dict)
assert "sampler" in analysis
assert "scheduler" in analysis
assert "steps" in analysis
assert "cfg" in analysis
assert "valid" in analysis
assert "scheduler_suggestions" in analysis
assert "scheduler_compatible" in analysis
assert "steps_optimal" in analysis
assert "cfg_optimal" in analysis
def test_get_available_samplers(self):
"""Test getting available samplers."""
samplers = SamplerComboNode.get_available_samplers()
assert isinstance(samplers, list)
assert len(samplers) > 0
assert "euler" in samplers
def test_get_available_schedulers(self):
"""Test getting available schedulers."""
schedulers = SamplerComboNode.get_available_schedulers()
assert isinstance(schedulers, list)
assert len(schedulers) > 0
assert "normal" in schedulers
def test_string_representations(self):
"""Test string representations of the node."""
str_repr = str(self.node)
assert "SamplerComboNode" in str_repr
assert "samplers=" in str_repr
assert "schedulers=" in str_repr
repr_str = repr(self.node)
assert "SamplerComboNode" in repr_str
assert "category=" in repr_str
assert "function=" in repr_str
def test_node_inheritance(self):
"""Test that node properly inherits from base class."""
from kikotools.base.base_node import ComfyAssetsBaseNode
assert isinstance(self.node, ComfyAssetsBaseNode)
assert hasattr(self.node, "validate_inputs")
assert hasattr(self.node, "handle_error")
assert hasattr(self.node, "log_info")
class TestSamplerComboIntegration:
"""Integration tests for Sampler Combo functionality."""
def test_full_workflow_valid_settings(self):
"""Test complete workflow with valid settings."""
node = SamplerComboNode()
# Test with different sampler/scheduler combinations
test_cases = [
("euler", "normal", 20, 7.0),
("dpmpp_2m", "karras", 15, 8.0),
("euler_ancestral", "exponential", 25, 9.0),
("ddim", "ddim_uniform", 30, 6.0),
]
for sampler, scheduler, steps, cfg in test_cases:
result = node.get_sampler_combo(sampler, scheduler, steps, cfg)
assert result == (sampler, scheduler, steps, cfg)
def test_recommendation_compatibility(self):
"""Test that recommendations are compatible with actual functionality."""
node = SamplerComboNode()
for sampler in SAMPLERS[:5]: # Test first 5 samplers
suggestions = node.get_scheduler_suggestions(sampler)
steps_rec = node.get_steps_recommendation(sampler)
cfg_rec = node.get_cfg_recommendation(sampler)
# 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"]
)
assert result[0] == sampler
assert result[1] == scheduler
assert result[2] == steps_rec["default"]
assert result[3] == cfg_rec["default"]
def test_error_recovery(self):
"""Test error recovery with malformed inputs."""
node = SamplerComboNode()
# These should all return safe defaults due to error handling
with patch(
"kikotools.tools.sampler_combo.logic.validate_sampler_settings",
side_effect=Exception("Simulated error"),
):
result = node.get_sampler_combo("euler", "normal", 20, 7.0)
assert result == ("euler", "normal", 20, 7.0) # Safe defaults
-3
View File
@@ -1,7 +1,6 @@
"""Tests for Seed History tool."""
import time
from unittest.mock import Mock
from kikotools.tools.seed_history.node import SeedHistoryNode
from kikotools.tools.seed_history.logic import (
@@ -24,8 +23,6 @@ class TestSeedHistoryNode:
def test_node_structure(self):
"""Test that node has correct ComfyUI structure."""
node = SeedHistoryNode()
# Test class attributes
assert hasattr(SeedHistoryNode, "INPUT_TYPES")
assert hasattr(SeedHistoryNode, "RETURN_TYPES")
@@ -1,7 +1,5 @@
"""Tests for Width Height Selector tool."""
import pytest
from unittest.mock import Mock
from kikotools.tools.width_height_selector.node import WidthHeightSelectorNode
from kikotools.tools.width_height_selector.logic import (
get_preset_dimensions,