Files
AEmotionStudio-ComfyUI-Shad…/__init__.py
T

222 lines
7.2 KiB
Python

"""
ComfyUI-ShaderNoiseKSampler
A custom KSampler node that uses shader-based noise patterns
for creative image generation.
"""
# Import node classes from nodes package
from .nodes import (
ShaderNoiseKSampler,
DirectShaderNoiseKSampler,
AdvancedImageComparer,
VideoComparer,
)
from .shader_noise_walk import ShaderNoiseWalk
from .shader_noise_source import ShaderNoiseSource
from .shader_to_tensor import ShaderToTensor
# Import shader registry
from .shaders.registry import (
ShaderRegistry,
register_shader,
get_shader,
list_shaders,
)
# Importing the generator modules runs their @shader_generator decorators,
# which is what registers every shader type. Only the aliases below are
# registered here.
from .shaders.domain_warp import (
DomainWarpGenerator,
generate_domain_warp_tensor,
)
from .shaders.tensor_field import (
TensorFieldGenerator,
generate_tensor_field_tensor,
)
from .shaders.curl_noise import (
CurlNoiseGenerator,
generate_curl_noise_tensor,
)
from .shaders.temporal_coherent_noise import (
TemporalCoherentNoiseGenerator,
generate_temporal_coherent_noise_tensor,
)
from .shaders.spectral import (
SpectralNoiseGenerator,
generate_spectral_tensor,
)
from .shaders.gaussian import (
GaussianNoiseGenerator,
generate_gaussian_tensor,
)
register_shader("curl", CurlNoiseGenerator, {
"description": "Curl/fluid noise patterns (alias)",
"supports_temporal": True,
})
register_shader("temporal_coherent_noise", TemporalCoherentNoiseGenerator, {
"description": "Temporally coherent noise (alias)",
"supports_temporal": True,
})
# Register API routes for server-side parameter saving
try:
from server import PromptServer
from .api_routes import setup_routes
setup_routes(PromptServer.instance)
except ImportError:
# PromptServer not available (e.g., running tests without ComfyUI)
pass
except Exception as e:
print(f"[ShaderNoiseKSampler] Warning: Could not register API routes: {e}")
# Legacy SHADER_GENERATORS dict for backward compatibility
# Maps shader type names to generator functions
SHADER_GENERATORS = {
"domain_warp": generate_domain_warp_tensor,
"tensor_field": generate_tensor_field_tensor,
"curl": generate_curl_noise_tensor,
"curl_noise": generate_curl_noise_tensor,
"temporal_coherent": generate_temporal_coherent_noise_tensor,
"temporal_coherent_noise": generate_temporal_coherent_noise_tensor,
}
def _wrap_legacy_generator(legacy_func):
"""
Wrap a legacy generator function to accept the new 'params' keyword argument.
Legacy functions expect 'shader_params' as a dict, but the new convention uses
'params' which may be a ShaderParams instance. This wrapper translates between
the two conventions and converts ShaderParams to dict.
Args:
legacy_func: Legacy generator function expecting shader_params as dict
Returns:
Wrapped function accepting params (ShaderParams or dict)
"""
def wrapper(**kwargs):
# If 'params' is provided but not 'shader_params', translate it
if 'params' in kwargs and 'shader_params' not in kwargs:
params = kwargs.pop('params')
# Convert ShaderParams to dict if needed for legacy function
if hasattr(params, 'to_dict'):
shader_params = params.to_dict()
elif hasattr(params, '__iter__'):
shader_params = dict(params)
else:
shader_params = params
kwargs['shader_params'] = shader_params
return legacy_func(**kwargs)
return wrapper
def get_shader_generator(shader_type: str):
"""
Get the appropriate shader generator function based on shader type.
This function provides backward compatibility with the old API
while using the new registry system internally. The returned function
accepts both 'params' (new convention) and 'shader_params' (legacy convention).
Args:
shader_type: Name of the shader type
Returns:
Generator function for the shader type. Falls back to generate_noise_tensor
if not found (consistent with shader_noise_ksampler.py behavior).
"""
# Import here to avoid circular imports
from .shader_params_reader import generate_noise_tensor
# First try the legacy dict for backward compatibility
# Wrap legacy functions to accept 'params' keyword argument
if shader_type in SHADER_GENERATORS:
return _wrap_legacy_generator(SHADER_GENERATORS[shader_type])
# Fall back to registry - return the static generate method
generator_class = get_shader(shader_type)
if generator_class is not None:
# Return the static generate method directly (consistent with shader_noise_ksampler.py)
return generator_class.generate
# Fallback: wrap generate_noise_tensor to translate params -> shader_params
# This matches the behavior in shader_noise_ksampler.py
def fallback_wrapper(params, height, width, batch_size, device, seed, target_channels, **kwargs):
# Convert ShaderParams to dict if needed for legacy function
if hasattr(params, 'to_dict'):
shader_params = params.to_dict()
elif hasattr(params, '__iter__'):
shader_params = dict(params)
else:
shader_params = {}
return generate_noise_tensor(
shader_params=shader_params,
height=height,
width=width,
batch_size=batch_size,
device=device,
seed=seed,
target_channels=target_channels,
**kwargs
)
return fallback_wrapper
def register_shader_generator(shader_type: str, generator_function):
"""
Register a shader generator function.
This function provides backward compatibility with the old API.
Registers to both the legacy SHADER_GENERATORS dict and the new registry.
Args:
shader_type: Name of the shader type
generator_function: Generator function or class to register
"""
# Add to legacy dict for backward compatibility
SHADER_GENERATORS[shader_type] = generator_function
# Also register to the new registry so shader_noise_ksampler.py can find it
register_shader(shader_type, generator_function)
# Node class mappings
NODE_CLASS_MAPPINGS = {
"ShaderNoiseKSampler": ShaderNoiseKSampler,
"ShaderNoiseKSamplerDirect": DirectShaderNoiseKSampler,
"ShaderNoiseWalk": ShaderNoiseWalk,
"ShaderNoiseSource": ShaderNoiseSource,
"AdvancedImageComparer": AdvancedImageComparer,
"Video Comparer": VideoComparer,
}
# Display name mappings
NODE_DISPLAY_NAME_MAPPINGS = {
"ShaderNoiseKSampler": "Shader Noise KSampler",
"ShaderNoiseKSamplerDirect": "Shader Noise KSampler (Direct)",
"ShaderNoiseWalk": "Shader Noise Walk",
"ShaderNoiseSource": "Shader Noise Source",
"AdvancedImageComparer": "Advanced Image Comparer",
"Video Comparer": "Video Comparer",
}
# Add web directory for UI components
WEB_DIRECTORY = "./web"
# List of exported elements
__all__ = [
"NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS",
"WEB_DIRECTORY",
"SHADER_GENERATORS",
"get_shader_generator",
"register_shader_generator",
"ShaderRegistry",
"register_shader",
"get_shader",
"list_shaders",
]