diff --git a/.gitignore b/.gitignore index cabae2f..0bbb58e 100644 --- a/.gitignore +++ b/.gitignore @@ -305,5 +305,45 @@ pyrightconfig.json .history .ionide + +### macOS ### +# General +.DS_Store +.AppleDouble +.LSOverride + +# Icon must end with two \r +Icon + + +# Thumbnails +._* + +# Files that might appear in the root of a volume +.DocumentRevisions-V100 +.fseventsd +.Spotlight-V100 +.TemporaryItems +.Trashes +.VolumeIcon.icns +.com.apple.timemachine.donotpresent + +# Directories potentially created on remote AFP share +.AppleDB +.AppleDesktop +Network Trash Folder +Temporary Items +.apdisk + +### macOS Patch ### +# iCloud generated files +*.icloud + # End of https://www.toptal.com/developers/gitignore/api/python,visualstudiocode,pycharm + .dev.md +.dev +CLAUDE.md +AGENTS.md +.gitnexus/ +.claude/ \ No newline at end of file diff --git a/__init__.py b/__init__.py index 0c373b5..7bd2ed6 100644 --- a/__init__.py +++ b/__init__.py @@ -1,18 +1,47 @@ +"""ComfyUI-Lightx2vWrapper entrypoint. + +ComfyUI discovers custom nodes by importing this package and reading +``NODE_CLASS_MAPPINGS`` / ``NODE_DISPLAY_NAME_MAPPINGS``. The actual node +classes live under the ``nodes/`` subpackage. +""" + import os import sys from pathlib import Path -os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" -os.environ["PROFILING_DEBUG_LEVEL"] = "2" -os.environ["TOKENIZERS_PARALLELISM"] = "false" -os.environ["ENABLE_GRAPH_MODE"] = "false" -os.environ["ENABLE_PROFILING_DEBUG"] = "true" -# os.environ["SENSITIVE_LAYER_DTYPE"] = "FP32" -os.environ["DTYPE"] = "BF16" -current_path = Path(__file__).parent.absolute() -print("Current path set to:", current_path) -sys.path.insert(0, os.path.join(current_path, "lightx2v")) # Adjust the path as needed +def _setup_env() -> None: + """Set environment variables consumed by the bundled lightx2v engine. + + Done in a function (instead of bare module-level statements) so the + side effects are explicit and easy to audit. ComfyUI imports this + module exactly once at startup, which is when these need to be set. + """ + os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") + os.environ.setdefault("PROFILING_DEBUG_LEVEL", "2") + os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") + os.environ.setdefault("ENABLE_GRAPH_MODE", "false") + os.environ.setdefault("ENABLE_PROFILING_DEBUG", "true") + os.environ.setdefault("DTYPE", "BF16") + + +def _register_lightx2v_submodule() -> None: + """Expose the bundled ``lightx2v/`` git submodule on ``sys.path``. + + The submodule ships its own top-level package also named ``lightx2v``; + putting the outer directory on ``sys.path`` lets internal modules import + ``lightx2v.xxx`` directly (as they do, e.g. ``lightx2v.common.ops``). + Our own nodes import via the relative path ``..lightx2v.lightx2v.xxx`` + and do not depend on this entry, but third-party / lightx2v-internal + code does. + """ + submodule_root = Path(__file__).parent.absolute() / "lightx2v" + if str(submodule_root) not in sys.path: + sys.path.insert(0, str(submodule_root)) + + +_setup_env() +_register_lightx2v_submodule() from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS # noqa: E402 diff --git a/bridge.py b/bridge.py deleted file mode 100644 index af547de..0000000 --- a/bridge.py +++ /dev/null @@ -1,543 +0,0 @@ -import copy -import importlib.util -import json -import logging -import os -from typing import Any, Dict, List, Tuple - -import torch -from easydict import EasyDict - - -def get_gpu_capability(): - if not torch.cuda.is_available(): - return None, None - try: - return torch.cuda.get_device_capability(0) - except Exception as e: - logging.warning(f"Failed to get GPU capability: {e}") - return None, None - - -def is_fp8_supported_gpu(): - major, minor = get_gpu_capability() - if major is None: - return False - return (major == 8 and minor == 9) or (major >= 9) - - -def is_ada_architecture_gpu(): - major, minor = get_gpu_capability() - if major is None: - return False - return major == 8 and minor == 9 - - -def is_module_installed(module_name): - try: - spec = importlib.util.find_spec(module_name) - return spec is not None - except ModuleNotFoundError: - return False - - -def get_available_ops(op_mapping): - available_ops = [] - for op_name, module_name in op_mapping.items(): - is_available = is_module_installed(module_name) - available_ops.append((op_name, is_available)) - return available_ops - - -def get_available_quant_ops(): - quant_mapping = { - "sgl": "sgl_kernel", - "vllm": "vllm", - "q8f": "q8_kernels", - "torchao": "torchao", - } - - available_ops = get_available_ops(quant_mapping) - - # Prefer q8f for Ada architecture GPUs - if is_ada_architecture_gpu(): - q8f_available = next((op for op in available_ops if op[0] == "q8f" and op[1]), None) - if q8f_available: - available_ops.remove(q8f_available) - available_ops.insert(0, q8f_available) - - return available_ops - - -def get_available_attn_ops(): - attn_mapping = { - "sage_attn2": "sageattention", - "sage_attn3": "sageattn3", - "flash_attn3": "flash_attn_interface", - "flash_attn2": "flash_attn", - "torch_sdpa": "torch", - } - - return get_available_ops(attn_mapping) - - -class LightX2VDefaultConfig: - """Central default configuration for LightX2V.""" - - DEFAULT_ATTENTION_TYPE = "flash_attn3" - DEFAULT_QUANTIZATION_SCHEMES = { - "dit": "Default", - "t5": "Default", - "clip": "Default", - "adapter": "Default", - } - DEFAULT_VIDEO_PARAMS = { - "height": 480, - "width": 832, - "length": 81, - "fps": 16, - "vae_stride": [4, 8, 8], - "patch_size": [1, 2, 2], - } - - DEFAULT_CONFIG = { - # Model Configuration - "model_cls": "wan2.1", - "model_path": "", - "task": "t2v", - # Inference Parameters - "infer_steps": 40, - "seed": 42, - "sample_guide_scale": 5.0, - "sample_shift": 5, - "enable_cfg": True, - "prompt": "", - "negative_prompt": "", - # Video Parameters - "target_height": DEFAULT_VIDEO_PARAMS["height"], - "target_width": DEFAULT_VIDEO_PARAMS["width"], - "target_video_length": DEFAULT_VIDEO_PARAMS["length"], - "fps": DEFAULT_VIDEO_PARAMS["fps"], - "vae_stride": DEFAULT_VIDEO_PARAMS["vae_stride"], - "patch_size": DEFAULT_VIDEO_PARAMS["patch_size"], - # TeaCache - "feature_caching": "NoCaching", - "teacache_thresh": 0.26, - "coefficients": None, - "use_ret_steps": False, - # Quantization - "dit_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["dit"], - "t5_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["t5"], - "clip_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["clip"], - "adapter_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["adapter"], - # Memory Optimization - "rotary_chunk": False, - "rotary_chunk_size": 100, - "clean_cuda_cache": False, - "torch_compile": False, - "self_attn_1_type": DEFAULT_ATTENTION_TYPE, - "cross_attn_1_type": DEFAULT_ATTENTION_TYPE, - "cross_attn_2_type": DEFAULT_ATTENTION_TYPE, - # CPU Offloading - "cpu_offload": False, - "offload_granularity": "block", - "offload_ratio": 1.0, - "t5_cpu_offload": False, - "t5_offload_granularity": "model", - "lazy_load": False, - "unload_modules": False, - # VAE Settings - "use_tiling_vae": False, - # Other Settings - "do_mm_calib": False, - "max_area": False, - "use_prompt_enhancer": False, - "text_len": 512, - "use_31_block": True, - "parallel": False, - "seq_parallel": False, - "cfg_parallel": False, - "audio_sr": 16000, - # "return_video": True, - "talk_objects": None, - "boundary_step_index": 2, - "rope_type": "torch", - } - - -class CoefficientCalculator: - """Calculate TeaCache coefficients based on model and resolution.""" - - COEFFICIENTS = { - "t2v": { - "1.3b": { - "default": [ - [ - -5.21862437e04, - 9.23041404e03, - -5.28275948e02, - 1.36987616e01, - -4.99875664e-02, - ], - [ - 2.39676752e03, - -1.31110545e03, - 2.01331979e02, - -8.29855975e00, - 1.37887774e-01, - ], - ] - }, - "14b": { - "default": [ - [ - -3.03318725e05, - 4.90537029e04, - -2.65530556e03, - 5.87365115e01, - -3.15583525e-01, - ], - [ - -5784.54975374, - 5449.50911966, - -1811.16591783, - 256.27178429, - -13.02252404, - ], - ] - }, - }, - "i2v": { - "720p": [ - [ - 8.10705460e03, - 2.13393892e03, - -3.72934672e02, - 1.66203073e01, - -4.17769401e-02, - ], - [-114.36346466, 65.26524496, -18.82220707, 4.91518089, -0.23412683], - ], - "480p": [ - [ - 2.57151496e05, - -3.54229917e04, - 1.40286849e03, - -1.35890334e01, - 1.32517977e-01, - ], - [ - -3.02331670e02, - 2.23948934e02, - -5.25463970e01, - 5.87348440e00, - -2.01973289e-01, - ], - ], - }, - } - - @classmethod - def get_coefficients( - cls, - task: str, - model_size: str, - resolution: Tuple[int, int], - use_ret_steps: bool, - ) -> List[List[float]]: - """Get appropriate coefficients for TeaCache.""" - if task == "t2v": - coeffs = cls.COEFFICIENTS["t2v"].get(model_size, {}).get("default", None) - else: # i2v - width, height = resolution - if height >= 720 or width >= 720: - coeffs = cls.COEFFICIENTS["i2v"]["720p"] - else: - coeffs = cls.COEFFICIENTS["i2v"]["480p"] - - if coeffs: - return coeffs[0] if use_ret_steps else coeffs[1] - raise ValueError( - f"No coefficients found for task: {task}, model_size: {model_size}, resolution: {resolution}, use_ret_steps: {use_ret_steps}" - ) - - -class ModularConfigManager: - """Manages modular configuration without presets.""" - - def __init__(self): - self.base_config = copy.deepcopy(LightX2VDefaultConfig.DEFAULT_CONFIG) - self._available_attn_ops = None - self._available_quant_ops = None - - def _get_available_ops(self, ops_list: List[Tuple[str, bool]], fallback: str = None) -> List[str]: - available = [op_name for op_name, is_available in ops_list if is_available] - if fallback and fallback not in available: - available.append(fallback) - return available - - @property - def available_attention_types(self) -> List[str]: - """Get available attention types.""" - if self._available_attn_ops is None: - self._available_attn_ops = get_available_attn_ops() - return self._get_available_ops(self._available_attn_ops, "torch_sdpa") - - @property - def available_quant_schemes(self) -> List[str]: - """Get available quantization schemes.""" - if self._available_quant_ops is None: - self._available_quant_ops = get_available_quant_ops() - return self._get_available_ops(self._available_quant_ops) - - def _update_from_config(self, updates: Dict, config: Dict, mappings: Dict[str, str]) -> None: - for config_key, update_key in mappings.items(): - if config_key in config: - if config_key == "seed" and config[config_key] == -1: - continue - updates[update_key] = config[config_key] - - def apply_inference_config(self, config: Dict[str, Any]) -> Dict[str, Any]: - updates = {} - - basic_mappings = { - "model_cls": "model_cls", - "model_path": "model_path", - "task": "task", - "infer_steps": "infer_steps", - "seed": "seed", - "sample_shift": "sample_shift", - "height": "target_height", - "width": "target_width", - "video_length": "target_video_length", - "fps": "target_fps", - "video_duration": "video_duration", - "resize_mode": "resize_mode", - "denoising_step_list": "denoising_step_list", - "use_31_block": "use_31_block", - "prev_frame_length": "prev_frame_length", - "fixed_area": "fixed_area", - } - - self._update_from_config(updates, config, basic_mappings) - - if "cfg_scale" in config: - updates["sample_guide_scale"] = config["cfg_scale"] - updates["enable_cfg"] = config["cfg_scale"] != 1.0 - - if "wan2.2_moe" in config["model_cls"]: - updates["boundary"] = 0.9 - updates["sample_guide_scale"] = [config["cfg_scale"], config["cfg_scale2"]] - if "wan2.2" in config["model_cls"]: - updates["use_image_encoder"] = False - - attention_type = config.get("attention_type", LightX2VDefaultConfig.DEFAULT_ATTENTION_TYPE) - for attn_key in [ - # "attention_type", - "self_attn_1_type", - "cross_attn_1_type", - "cross_attn_2_type", - ]: - updates[attn_key] = attention_type - - if config.get("use_tiny_vae", False): - updates.update( - { - "use_tiny_vae": True, - "tiny_vae": True, - "tiny_vae_path": os.path.join(config["model_path"], "taew2_1.pth"), - } - ) - - return updates - - def apply_teacache_config(self, config: Dict[str, Any], model_info: Dict[str, Any]) -> Dict[str, Any]: - """Apply TeaCache configuration.""" - updates = {} - - if config.get("enable", False): - updates["feature_caching"] = "Tea" - updates["teacache_thresh"] = config.get("threshold", 0.26) - updates["use_ret_steps"] = config.get("use_ret_steps", False) - - task = model_info.get("task", "t2v") - model_size = "14b" if "14b" in model_info.get("model_cls", "") else "1.3b" - resolution = ( - model_info.get("target_width", 832), - model_info.get("target_height", 480), - ) - - coeffs = CoefficientCalculator.get_coefficients(task, model_size, resolution, updates["use_ret_steps"]) - updates["coefficients"] = coeffs - else: - updates["feature_caching"] = "NoCaching" - - return updates - - def apply_quantization_config(self, config: Dict[str, Any]) -> Dict[str, Any]: - """Apply quantization configuration.""" - updates = {} - defaults = LightX2VDefaultConfig.DEFAULT_QUANTIZATION_SCHEMES - - dit_scheme = config.get("dit_quant_scheme", defaults["dit"]) - t5_scheme = config.get("t5_quant_scheme", defaults["t5"]) - clip_scheme = config.get("clip_quant_scheme", defaults["clip"]) - adapter_scheme = config.get("adapter_quant_scheme", defaults["adapter"]) - - updates.update( - { - "clip_quantized": clip_scheme != "Default", - "clip_quant_scheme": clip_scheme, - "t5_quantized": t5_scheme != "Default", - "t5_quant_scheme": t5_scheme, - "dit_quantized": dit_scheme != "Default", - "dit_quant_scheme": dit_scheme, - "adapter_quantized": adapter_scheme != "Default", - "adapter_quant_scheme": adapter_scheme, - } - ) - - return updates - - def apply_memory_optimization(self, config: Dict[str, Any]) -> Dict[str, Any]: - """Apply memory optimization settings.""" - updates = {} - - direct_mappings = { - "enable_rotary_chunk": "rotary_chunk", - "clean_cuda_cache": "clean_cuda_cache", - "cpu_offload": "cpu_offload", - "t5_cpu_offload": "t5_cpu_offload", - "vae_cpu_offload": "vae_cpu_offload", - "audio_encoder_cpu_offload": "audio_encoder_cpu_offload", - "audio_adapter_cpu_offload": "audio_adapter_cpu_offload", - "lazy_load": "lazy_load", - "unload_after_inference": "unload_modules", - "use_tiling_vae": "use_tiling_vae", - } - - for config_key, update_key in direct_mappings.items(): - updates[update_key] = config.get(config_key, config.get("cpu_offload", False)) - - if updates.get("rotary_chunk"): - updates["rotary_chunk_size"] = config.get("rotary_chunk_size", 100) - - if updates.get("cpu_offload"): - updates.update( - { - "offload_granularity": config.get("offload_granularity", "phase"), - "offload_ratio": config.get("offload_ratio", 1.0), - } - ) - - if updates.get("t5_cpu_offload"): - updates["t5_offload_granularity"] = config.get("t5_offload_granularity", "model") - - return updates - - def _load_model_config(self, model_path: str) -> Dict[str, Any]: - config_path = os.path.join(model_path, "config.json") - if not os.path.exists(config_path): - return {} - - try: - with open(config_path, "r") as f: - return json.load(f) - except Exception as e: - logging.warning(f"Failed to load model config: {e}") - return {} - - def build_final_config_from_combined(self, combined_config) -> EasyDict: - """Build final configuration directly from CombinedConfig object.""" - final_config = copy.deepcopy(self.base_config) - - # Apply inference configuration - if combined_config.inference: - updates = self.apply_inference_config(combined_config.inference.to_dict()) - final_config.update(updates) - - # Apply memory optimization configuration - if combined_config.memory: - memory_updates = self.apply_memory_optimization(combined_config.memory.to_dict()) - final_config.update(memory_updates) - - # Apply TeaCache configuration - if combined_config.teacache: - teacache_updates = self.apply_teacache_config(combined_config.teacache.to_dict(), final_config) - final_config.update(teacache_updates) - - # Apply quantization configuration - if combined_config.quantization: - quant_updates = self.apply_quantization_config(combined_config.quantization.to_dict()) - final_config.update(quant_updates) - - # Handle LoRA configurations - if combined_config.lora_configs: - lora_chain = [lora.to_dict() for lora in combined_config.lora_configs] - final_config["lora_configs"] = lora_chain - - # Handle talk objects configuration - if combined_config.talk_objects: - talk_objects_dict = combined_config.talk_objects.to_dict() - final_config.update(talk_objects_dict) - - # Load model-specific configuration - model_config = self._load_model_config(final_config.get("model_path", "")) - for key, value in model_config.items(): - if key not in final_config or final_config[key] is None: - final_config[key] = value - - return EasyDict(final_config) - - def build_final_config(self, configs: Dict[str, Dict[str, Any]]) -> EasyDict: - """Build final configuration from module configs. - - This method is kept for backward compatibility. - It converts dict configs to CombinedConfig and uses the new method. - """ - from .data_models import ( - CombinedConfig, - InferenceConfig, - LoRAConfig, - MemoryOptimizationConfig, - QuantizationConfig, - TalkObject, - TalkObjectsConfig, - TeaCacheConfig, - ) - - # Create CombinedConfig from dictionary configs - combined = CombinedConfig() - - # Process inference config - if "inference" in configs: - combined.inference = InferenceConfig(**configs["inference"]) - - # Process teacache config - if "teacache" in configs: - combined.teacache = TeaCacheConfig(**configs["teacache"]) - - # Process quantization config - if "quantization" in configs: - combined.quantization = QuantizationConfig(**configs["quantization"]) - - # Process memory config - if "memory" in configs: - combined.memory = MemoryOptimizationConfig(**configs["memory"]) - - # Process lora configs - if "lora_configs" in configs: - for lora_dict in configs["lora_configs"]: - lora_config = LoRAConfig(**lora_dict) - combined.lora_configs.append(lora_config) - - # Process talk objects - if "talk_objects" in configs: - talk_objects = TalkObjectsConfig() - for obj_dict in configs["talk_objects"]: - talk_obj = TalkObject(**obj_dict) - talk_objects.add_object(talk_obj) - combined.talk_objects = talk_objects - - # Use the new method to build final config - return self.build_final_config_from_combined(combined) diff --git a/bridge/__init__.py b/bridge/__init__.py new file mode 100644 index 0000000..f8474c4 --- /dev/null +++ b/bridge/__init__.py @@ -0,0 +1,41 @@ +"""Bridge between ComfyUI widget values and lightx2v's internal config schema. + +Submodules: + - ``capability`` GPU + backend-op detection (pure functions) + - ``defaults`` ``LightX2VDefaultConfig`` — wrapper-side starting values + - ``teacache_coeffs`` ``CoefficientCalculator`` — polynomial constants + - ``translator/`` per-feature wrapper-key -> lightx2v-key translators, + plus ``ModularConfigManager`` that orchestrates them + +Public surface (re-exported here for backward compat with existing imports +``from .bridge import …``): +""" + +from .capability import ( + get_available_attn_ops, + get_available_ops, + get_available_quant_ops, + get_gpu_capability, + is_ada_architecture_gpu, + is_fp8_supported_gpu, + is_module_installed, +) +from .defaults import LightX2VDefaultConfig +from .teacache_coeffs import CoefficientCalculator +from .translator import ModularConfigManager + +__all__ = [ + # capability + "get_gpu_capability", + "is_fp8_supported_gpu", + "is_ada_architecture_gpu", + "is_module_installed", + "get_available_ops", + "get_available_quant_ops", + "get_available_attn_ops", + # defaults / coeffs + "LightX2VDefaultConfig", + "CoefficientCalculator", + # translator orchestrator + "ModularConfigManager", +] diff --git a/bridge/capability.py b/bridge/capability.py new file mode 100644 index 0000000..821b683 --- /dev/null +++ b/bridge/capability.py @@ -0,0 +1,80 @@ +"""GPU and backend-op capability detection. + +Pure functions — no state, no module-level side effects. Cheap to call +(the underlying ``torch.cuda`` / ``importlib`` probes are fast). +""" + +import importlib.util +import logging +from typing import List, Tuple + +import torch + + +def get_gpu_capability(): + if not torch.cuda.is_available(): + return None, None + try: + return torch.cuda.get_device_capability(0) + except Exception as e: + logging.warning(f"Failed to get GPU capability: {e}") + return None, None + + +def is_fp8_supported_gpu() -> bool: + major, minor = get_gpu_capability() + if major is None: + return False + return (major == 8 and minor == 9) or (major >= 9) + + +def is_ada_architecture_gpu() -> bool: + major, minor = get_gpu_capability() + if major is None: + return False + return major == 8 and minor == 9 + + +def is_module_installed(module_name: str) -> bool: + try: + spec = importlib.util.find_spec(module_name) + return spec is not None + except ModuleNotFoundError: + return False + + +def get_available_ops(op_mapping: dict) -> List[Tuple[str, bool]]: + return [(op_name, is_module_installed(module_name)) for op_name, module_name in op_mapping.items()] + + +_QUANT_OP_MAPPING = { + "sgl": "sgl_kernel", + "vllm": "vllm", + "q8f": "q8_kernels", + "torchao": "torchao", +} + +_ATTN_OP_MAPPING = { + "sage_attn2": "sageattention", + "sage_attn3": "sageattn3", + "flash_attn3": "flash_attn_interface", + "flash_attn2": "flash_attn", + "torch_sdpa": "torch", +} + + +def get_available_quant_ops() -> List[Tuple[str, bool]]: + available_ops = get_available_ops(_QUANT_OP_MAPPING) + + # Prefer q8f on Ada (sm_8.9) GPUs — best perf/precision tradeoff there. + if is_ada_architecture_gpu(): + q8f_available = next((op for op in available_ops if op[0] == "q8f" and op[1]), None) + if q8f_available: + available_ops.remove(q8f_available) + available_ops.insert(0, q8f_available) + + return available_ops + + +def get_available_attn_ops() -> List[Tuple[str, bool]]: + return get_available_ops(_ATTN_OP_MAPPING) diff --git a/bridge/defaults.py b/bridge/defaults.py new file mode 100644 index 0000000..f37ee55 --- /dev/null +++ b/bridge/defaults.py @@ -0,0 +1,80 @@ +"""Default config values that the wrapper provides to lightx2v. + +These are the wrapper's *starting point* — lightx2v's own ``set_config`` will +further merge from ``config_json`` and the model's own ``config.json`` on disk +(see ``lightx2v/utils/set_config.py:set_config``). Anything lightx2v sets +internally (``vae_stride``, ``patch_size``, etc.) should NOT be duplicated here. +""" + + +class LightX2VDefaultConfig: + """Central default configuration for LightX2V.""" + + DEFAULT_ATTENTION_TYPE = "flash_attn3" + DEFAULT_QUANTIZATION_SCHEMES = { + "dit": "Default", + "t5": "Default", + "clip": "Default", + "adapter": "Default", + } + + DEFAULT_CONFIG = { + # Model + "model_cls": "wan2.1", + "model_path": "", + "task": "t2v", + # Inference + "infer_steps": 40, + "seed": 42, + "sample_guide_scale": 5.0, + "sample_shift": 5, + "enable_cfg": True, + "prompt": "", + "negative_prompt": "", + # Video / Image output (lightx2v field names — see translator/inference.py) + "target_height": 480, + "target_width": 832, + "target_video_length": 81, + "fps": 16, + # TeaCache + "feature_caching": "NoCaching", + "teacache_thresh": 0.26, + "coefficients": None, + "use_ret_steps": False, + # Quantization + "dit_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["dit"], + "t5_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["t5"], + "clip_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["clip"], + "adapter_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["adapter"], + # Attention + "self_attn_1_type": DEFAULT_ATTENTION_TYPE, + "cross_attn_1_type": DEFAULT_ATTENTION_TYPE, + "cross_attn_2_type": DEFAULT_ATTENTION_TYPE, + # Memory / offload + "rotary_chunk": False, + "rotary_chunk_size": 100, + "clean_cuda_cache": False, + "torch_compile": False, + "cpu_offload": False, + "offload_granularity": "block", + "offload_ratio": 1.0, + "t5_cpu_offload": False, + "t5_offload_granularity": "model", + "lazy_load": False, + "unload_modules": False, + # VAE + "use_tiling_vae": False, + # Misc + "do_mm_calib": False, + "max_area": False, + "use_prompt_enhancer": False, + "text_len": 512, + "use_31_block": True, + "parallel": False, + "seq_parallel": False, + "cfg_parallel": False, + "audio_sr": 16000, + "talk_objects": None, + "boundary_step_index": 2, + "rope_type": "torch", + } diff --git a/bridge/teacache_coeffs.py b/bridge/teacache_coeffs.py new file mode 100644 index 0000000..73d7308 --- /dev/null +++ b/bridge/teacache_coeffs.py @@ -0,0 +1,64 @@ +"""TeaCache polynomial coefficients per (task, model size, resolution). + +These constants come from the upstream TeaCache calibration runs (one set per +task/resolution bucket). They are pure data; no logic here other than picking +the right bucket. +""" + +from typing import List, Tuple + + +class CoefficientCalculator: + """Pick TeaCache polynomial coefficients for a given task/model/resolution.""" + + COEFFICIENTS = { + "t2v": { + "1.3b": { + "default": [ + [-5.21862437e04, 9.23041404e03, -5.28275948e02, 1.36987616e01, -4.99875664e-02], + [2.39676752e03, -1.31110545e03, 2.01331979e02, -8.29855975e00, 1.37887774e-01], + ] + }, + "14b": { + "default": [ + [-3.03318725e05, 4.90537029e04, -2.65530556e03, 5.87365115e01, -3.15583525e-01], + [-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429, -13.02252404], + ] + }, + }, + "i2v": { + "720p": [ + [8.10705460e03, 2.13393892e03, -3.72934672e02, 1.66203073e01, -4.17769401e-02], + [-114.36346466, 65.26524496, -18.82220707, 4.91518089, -0.23412683], + ], + "480p": [ + [2.57151496e05, -3.54229917e04, 1.40286849e03, -1.35890334e01, 1.32517977e-01], + [-3.02331670e02, 2.23948934e02, -5.25463970e01, 5.87348440e00, -2.01973289e-01], + ], + }, + } + + @classmethod + def get_coefficients( + cls, + task: str, + model_size: str, + resolution: Tuple[int, int], + use_ret_steps: bool, + ) -> List[float]: + """Pick the right coefficient row for this (task, model_size, resolution). + + ``use_ret_steps`` selects between the two calibration runs (cache key + steps only vs. cache all steps). + """ + if task == "t2v": + coeffs = cls.COEFFICIENTS["t2v"].get(model_size, {}).get("default", None) + else: # i2v + width, height = resolution + coeffs = cls.COEFFICIENTS["i2v"]["720p"] if height >= 720 or width >= 720 else cls.COEFFICIENTS["i2v"]["480p"] + + if coeffs: + return coeffs[0] if use_ret_steps else coeffs[1] + raise ValueError( + f"No coefficients found for task: {task}, model_size: {model_size}, resolution: {resolution}, use_ret_steps: {use_ret_steps}" + ) diff --git a/bridge/translator/__init__.py b/bridge/translator/__init__.py new file mode 100644 index 0000000..b8563b6 --- /dev/null +++ b/bridge/translator/__init__.py @@ -0,0 +1,26 @@ +"""Field translators: ComfyUI widget value -> lightx2v config dict. + +Each module here mirrors one ComfyUI Config node in ``nodes/config.py``: + + LightX2VInferenceConfig <-> translator/inference.py + LightX2VTeaCache <-> translator/teacache.py + LightX2VQuantization <-> translator/quant.py + LightX2VMemoryOptimization <-> translator/memory.py + +``translator/pipeline.py`` orchestrates them and adds the model's own +``config.json`` (read from disk by lightx2v's ``set_config``). +""" + +from .inference import apply_inference_config +from .memory import apply_memory_optimization +from .pipeline import ModularConfigManager +from .quant import apply_quantization_config +from .teacache import apply_teacache_config + +__all__ = [ + "apply_inference_config", + "apply_teacache_config", + "apply_quantization_config", + "apply_memory_optimization", + "ModularConfigManager", +] diff --git a/bridge/translator/inference.py b/bridge/translator/inference.py new file mode 100644 index 0000000..14dc1af --- /dev/null +++ b/bridge/translator/inference.py @@ -0,0 +1,85 @@ +"""Translate ``LightX2VInferenceConfig`` widget values into lightx2v config keys. + +The wrapper's widget naming follows ComfyUI conventions (``height``, ``width``, +``video_length``, ``cfg_scale`` …). lightx2v's internal naming is different +(``target_height``, ``target_width``, ``target_video_length``, +``sample_guide_scale`` …). The single source of truth for that translation +is the ``WRAPPER_TO_LIGHTX2V_FIELDS`` table below — when adding a new field, +add a row there rather than burying the rename inside the function body. +""" + +import os +from typing import Any, Dict + +from ..defaults import LightX2VDefaultConfig + +# Direct rename map: wrapper-side key -> lightx2v-side key. +# A row of ("foo", "foo") means the name matches but we still want to forward +# the value explicitly (rather than relying on the default config). +WRAPPER_TO_LIGHTX2V_FIELDS: Dict[str, str] = { + # Model selection — names match. + "model_cls": "model_cls", + "model_path": "model_path", + "task": "task", + # Inference loop. + "infer_steps": "infer_steps", + "seed": "seed", + "sample_shift": "sample_shift", + # Output shape — wrapper uses bare names, lightx2v prefixes with target_. + "height": "target_height", + "width": "target_width", + "video_length": "target_video_length", + "fps": "target_fps", + "video_duration": "video_duration", + # Image preprocessing. + "resize_mode": "resize_mode", + "fixed_area": "fixed_area", + # Sekotalk-specific. + "prev_frame_length": "prev_frame_length", + # Distillation. + "denoising_step_list": "denoising_step_list", + "use_31_block": "use_31_block", +} + +# Attention type fans out to three internal slots in lightx2v. +_ATTN_TYPE_SLOTS = ("self_attn_1_type", "cross_attn_1_type", "cross_attn_2_type") + + +def apply_inference_config(config: Dict[str, Any]) -> Dict[str, Any]: + """Translate inference widget values to a partial lightx2v config dict.""" + updates: Dict[str, Any] = {} + + # Bulk rename via the table. + for wrapper_key, lightx2v_key in WRAPPER_TO_LIGHTX2V_FIELDS.items(): + if wrapper_key not in config: + continue + # seed=-1 means "use lightx2v's default / random"; leave it out. + if wrapper_key == "seed" and config[wrapper_key] == -1: + continue + updates[lightx2v_key] = config[wrapper_key] + + # cfg_scale -> sample_guide_scale (and toggle enable_cfg). + if "cfg_scale" in config: + updates["sample_guide_scale"] = config["cfg_scale"] + updates["enable_cfg"] = config["cfg_scale"] != 1.0 + + # Wan2.2 MoE has two CFG scales (high/low noise) and a boundary param. + model_cls = config.get("model_cls", "") + if "wan2.2_moe" in model_cls: + updates["boundary"] = 0.9 + updates["sample_guide_scale"] = [config.get("cfg_scale"), config.get("cfg_scale2")] + if "wan2.2" in model_cls: + updates["use_image_encoder"] = False + + # One widget value drives three lightx2v attention slots. + attention_type = config.get("attention_type", LightX2VDefaultConfig.DEFAULT_ATTENTION_TYPE) + for slot in _ATTN_TYPE_SLOTS: + updates[slot] = attention_type + + # TAEW2.1 lightweight VAE lives next to the model. + if config.get("use_tiny_vae", False): + updates["use_tiny_vae"] = True + updates["tiny_vae"] = True + updates["tiny_vae_path"] = os.path.join(config["model_path"], "taew2_1.pth") + + return updates diff --git a/bridge/translator/memory.py b/bridge/translator/memory.py new file mode 100644 index 0000000..8046540 --- /dev/null +++ b/bridge/translator/memory.py @@ -0,0 +1,48 @@ +"""Translate ``LightX2VMemoryOptimization`` widget values into lightx2v config keys. + +Several toggles only matter when their parent is enabled (e.g. ``offload_granularity`` +only when ``cpu_offload=True``). Those nested keys are written only on the +true-branch to keep the resulting config dict minimal. +""" + +from typing import Any, Dict + +# Direct rename: wrapper key -> lightx2v key. +WRAPPER_TO_LIGHTX2V_FIELDS: Dict[str, str] = { + "enable_rotary_chunk": "rotary_chunk", + "clean_cuda_cache": "clean_cuda_cache", + "cpu_offload": "cpu_offload", + "t5_cpu_offload": "t5_cpu_offload", + "vae_cpu_offload": "vae_cpu_offload", + "audio_encoder_cpu_offload": "audio_encoder_cpu_offload", + "audio_adapter_cpu_offload": "audio_adapter_cpu_offload", + "lazy_load": "lazy_load", + "unload_after_inference": "unload_modules", + "use_tiling_vae": "use_tiling_vae", +} + + +def apply_memory_optimization(config: Dict[str, Any]) -> Dict[str, Any]: + """Translate memory-optimization widget values.""" + updates: Dict[str, Any] = {} + + # NOTE: legacy behavior — when a specific offload key is missing, fall back + # to the global ``cpu_offload`` flag. This means if the user only sets + # ``cpu_offload=True``, every sub-offload (T5/VAE/audio…) silently follows. + # Preserved as-is for backward compat; revisit when audio_* offloads + # become widget-exposed everywhere. + global_cpu_offload = config.get("cpu_offload", False) + for wrapper_key, lightx2v_key in WRAPPER_TO_LIGHTX2V_FIELDS.items(): + updates[lightx2v_key] = config.get(wrapper_key, global_cpu_offload) + + if updates.get("rotary_chunk"): + updates["rotary_chunk_size"] = config.get("rotary_chunk_size", 100) + + if updates.get("cpu_offload"): + updates["offload_granularity"] = config.get("offload_granularity", "phase") + updates["offload_ratio"] = config.get("offload_ratio", 1.0) + + if updates.get("t5_cpu_offload"): + updates["t5_offload_granularity"] = config.get("t5_offload_granularity", "model") + + return updates diff --git a/bridge/translator/pipeline.py b/bridge/translator/pipeline.py new file mode 100644 index 0000000..e80b395 --- /dev/null +++ b/bridge/translator/pipeline.py @@ -0,0 +1,111 @@ +"""Orchestrate per-feature translators into a final lightx2v config dict. + +Flow: + 1. start from ``LightX2VDefaultConfig.DEFAULT_CONFIG`` + 2. apply inference / memory / teacache / quantization translators in order + (teacache runs after inference so it can see the resolved task/resolution) + 3. attach LoRA chain and talk_objects (no rename needed) + 4. shallow-merge the model's own ``config.json`` for keys still unset + (lightx2v's ``set_config`` will further read its own model config later) + 5. wrap as ``EasyDict`` so consumers can use attribute access + +NOTE: ``input_info`` (the per-call dataclass in ``lightx2v.utils.input_info``) +is NOT built here. The inference node constructs it dynamically from this +config plus the runtime image/audio paths, because lightx2v itself distinguishes +"persistent config" from "per-call input_info". +""" + +import copy +import json +import logging +import os +from typing import Any, Dict + +from easydict import EasyDict + +from ..capability import get_available_attn_ops, get_available_quant_ops +from ..defaults import LightX2VDefaultConfig +from .inference import apply_inference_config +from .memory import apply_memory_optimization +from .quant import apply_quantization_config +from .teacache import apply_teacache_config + + +class ModularConfigManager: + """Compose translators into a final lightx2v config.""" + + def __init__(self): + self.base_config = copy.deepcopy(LightX2VDefaultConfig.DEFAULT_CONFIG) + self._available_attn_ops = None + self._available_quant_ops = None + + @staticmethod + def _filter_available(ops_list, fallback=None): + available = [name for name, ok in ops_list if ok] + if fallback and fallback not in available: + available.append(fallback) + return available + + @property + def available_attention_types(self): + if self._available_attn_ops is None: + self._available_attn_ops = get_available_attn_ops() + return self._filter_available(self._available_attn_ops, "torch_sdpa") + + @property + def available_quant_schemes(self): + if self._available_quant_ops is None: + self._available_quant_ops = get_available_quant_ops() + return self._filter_available(self._available_quant_ops) + + # Exposed for tests/debugging; the public entrypoint is build_final_config_from_combined. + apply_inference_config = staticmethod(apply_inference_config) + apply_teacache_config = staticmethod(apply_teacache_config) + apply_quantization_config = staticmethod(apply_quantization_config) + apply_memory_optimization = staticmethod(apply_memory_optimization) + + @staticmethod + def _load_model_config(model_path: str) -> Dict[str, Any]: + config_path = os.path.join(model_path, "config.json") + if not os.path.exists(config_path): + return {} + try: + with open(config_path, "r") as f: + return json.load(f) + except Exception as e: + logging.warning(f"Failed to load model config: {e}") + return {} + + def build_final_config_from_combined(self, combined_config) -> EasyDict: + """Build the final lightx2v config from a CombinedConfig dataclass.""" + final_config = copy.deepcopy(self.base_config) + + if combined_config.inference: + final_config.update(apply_inference_config(combined_config.inference.to_dict())) + + if combined_config.memory: + final_config.update(apply_memory_optimization(combined_config.memory.to_dict())) + + # teacache reads the (already-resolved) task and resolution off final_config. + if combined_config.teacache: + final_config.update(apply_teacache_config(combined_config.teacache.to_dict(), final_config)) + + if combined_config.quantization: + final_config.update(apply_quantization_config(combined_config.quantization.to_dict())) + + if combined_config.lora_configs: + final_config["lora_configs"] = [lora.to_dict() for lora in combined_config.lora_configs] + + if combined_config.talk_objects: + final_config.update(combined_config.talk_objects.to_dict()) + + # Shallow-merge the model's own config.json for keys still unset. + # lightx2v's own set_config.auto_calc_config will do its own deeper + # merge of model_path/config.json — this just gives translators a + # chance to see model-side hints (e.g. text_len) when they run. + model_config = self._load_model_config(final_config.get("model_path", "")) + for key, value in model_config.items(): + if key not in final_config or final_config[key] is None: + final_config[key] = value + + return EasyDict(final_config) diff --git a/bridge/translator/quant.py b/bridge/translator/quant.py new file mode 100644 index 0000000..401b526 --- /dev/null +++ b/bridge/translator/quant.py @@ -0,0 +1,25 @@ +"""Translate ``LightX2VQuantization`` widget values into lightx2v config keys. + +Each of dit/t5/clip/adapter contributes two keys to lightx2v: +``{component}_quantized`` (bool) and ``{component}_quant_scheme`` (str). +``"Default"`` means "leave as-is". +""" + +from typing import Any, Dict + +from ..defaults import LightX2VDefaultConfig + +_COMPONENTS = ("dit", "t5", "clip", "adapter") + + +def apply_quantization_config(config: Dict[str, Any]) -> Dict[str, Any]: + """Translate quantization widget values.""" + updates: Dict[str, Any] = {} + defaults = LightX2VDefaultConfig.DEFAULT_QUANTIZATION_SCHEMES + + for component in _COMPONENTS: + scheme = config.get(f"{component}_quant_scheme", defaults[component]) + updates[f"{component}_quantized"] = scheme != "Default" + updates[f"{component}_quant_scheme"] = scheme + + return updates diff --git a/bridge/translator/teacache.py b/bridge/translator/teacache.py new file mode 100644 index 0000000..1a70251 --- /dev/null +++ b/bridge/translator/teacache.py @@ -0,0 +1,35 @@ +"""Translate ``LightX2VTeaCache`` widget values into lightx2v config keys. + +Wrapper-side ``enable / threshold / use_ret_steps`` -> lightx2v-side +``feature_caching / teacache_thresh / use_ret_steps`` + polynomial coefficients +picked from the calibration table. +""" + +from typing import Any, Dict + +from ..teacache_coeffs import CoefficientCalculator + + +def apply_teacache_config(config: Dict[str, Any], model_info: Dict[str, Any]) -> Dict[str, Any]: + """Translate TeaCache widget values. + + ``model_info`` is the partially-built lightx2v config so we can pick + coefficients matched to the actual task and output resolution. + """ + if not config.get("enable", False): + return {"feature_caching": "NoCaching"} + + use_ret_steps = config.get("use_ret_steps", False) + task = model_info.get("task", "t2v") + model_size = "14b" if "14b" in model_info.get("model_cls", "") else "1.3b" + resolution = ( + model_info.get("target_width", 832), + model_info.get("target_height", 480), + ) + + return { + "feature_caching": "Tea", + "teacache_thresh": config.get("threshold", 0.26), + "use_ret_steps": use_ret_steps, + "coefficients": CoefficientCalculator.get_coefficients(task, model_size, resolution, use_ret_steps), + } diff --git a/lightx2v b/lightx2v index 7216292..27e5c90 160000 --- a/lightx2v +++ b/lightx2v @@ -1 +1 @@ -Subproject commit 7216292de236447480d936e2dcaa7084dab48147 +Subproject commit 27e5c906eacfc894abc22c3ecf4cb7689d5b6db3 diff --git a/model_utils.py b/model_utils.py index 7fcc15b..bd5eb66 100644 --- a/model_utils.py +++ b/model_utils.py @@ -42,6 +42,7 @@ def support_model_cls_list() -> List[str]: "wan2.2_audio", "wan2.2_moe_distill", "qwen_image", + "seedvr2", ] diff --git a/nodes.py b/nodes.py deleted file mode 100644 index 1c641d0..0000000 --- a/nodes.py +++ /dev/null @@ -1,1664 +0,0 @@ -import gc -import io -import json -import logging -import os -import subprocess as sp -import wave - -import numpy as np -import torch -from comfy.utils import ProgressBar -from PIL import Image - -from .bridge import get_available_attn_ops, get_available_quant_ops -from .config_builder import ( - ConfigBuilder, - InferenceConfigBuilder, - LoRAChainBuilder, - TalkObjectConfigBuilder, -) -from .data_models import ( - InferenceConfig, - MemoryOptimizationConfig, - QuantizationConfig, - TalkObjectsConfig, - TeaCacheConfig, -) -from .file_handlers import ( - AudioFileHandler, - ComfyUIFileResolver, - HTTPFileDownloader, - ImageFileHandler, - TempFileManager, -) -from .lightx2v.lightx2v.infer import init_runner -from .lightx2v.lightx2v.utils.input_info import init_empty_input_info, update_input_info_from_dict -from .lightx2v.lightx2v.utils.set_config import set_config -from .model_utils import scan_loras, scan_models, support_model_cls_list - - -class LightX2VInferenceConfig: - @classmethod - def INPUT_TYPES(cls): - available_models = scan_models() - support_model_classes = support_model_cls_list() - available_attn = get_available_attn_ops() - attn_types = [] - - for op_name, is_available in available_attn: - if is_available: - attn_types.append(op_name) - - if "torch_sdpa" not in attn_types: - attn_types.append("torch_sdpa") - - return { - "required": { - "model_cls": ( - support_model_classes, - {"default": "wan2.1", "tooltip": "Model type"}, - ), - "model_name": ( - available_models, - { - "default": available_models[0], - "tooltip": "Select model from available models", - }, - ), - "task": ( - ["t2v", "i2v", "s2v", "rs2v"], - { - "default": "i2v", - "tooltip": "Task type: text-to-video or image-to-video or reference_image and audio to video (shot)", - }, - ), - "infer_steps": ( - "INT", - {"default": 4, "min": 1, "max": 100, "tooltip": "Inference steps"}, - ), - "seed": ( - "INT", - { - "default": 42, - "min": -1, - "max": 2**32 - 1, - "tooltip": "Random seed, -1 for random", - }, - ), - "cfg_scale": ( - "FLOAT", - { - "default": 5.0, - "min": 1.0, - "max": 10.0, - "step": 0.1, - "tooltip": "CFG guidance strength", - }, - ), - "cfg_scale2": ( - "FLOAT", - { - "default": 5.0, - "min": 1.0, - "max": 10.0, - "step": 0.1, - "tooltip": "CFG guidance, lower noise when model cls is Wan2.2 MoE", - }, - ), - "sample_shift": ( - "INT", - {"default": 5, "min": 0, "max": 10, "tooltip": "Sample shift"}, - ), - "height": ( - "INT", - { - "default": 1280, - "min": 64, - "max": 2048, - "step": 8, - "tooltip": "Video height", - }, - ), - "width": ( - "INT", - { - "default": 720, - "min": 64, - "max": 2048, - "step": 8, - "tooltip": "Video width", - }, - ), - "duration": ( - "FLOAT", - { - "default": 5.0, - "min": 1.0, - "max": 999, - "step": 0.1, - "tooltip": "Video duration in seconds", - }, - ), - "attention_type": ( - attn_types, - {"default": attn_types[0], "tooltip": "Attention mechanism type"}, - ), - }, - "optional": { - "denoising_steps": ( - "STRING", - { - "default": "", - "tooltip": "Custom denoising steps for distillation models (comma-separated, e.g., '999,750,500,250'). Leave empty to use model defaults.", - }, - ), - "resize_mode": ( - [ - "adaptive", - "keep_ratio_fixed_area", - "fixed_min_area", - "fixed_max_area", - "fixed_shape", - "fixed_min_side", - ], - { - "default": "adaptive", - "tooltip": "Adaptive resize input image to target aspect ratio", - }, - ), - "fixed_area": ( - "STRING", - { - "default": "720p", - "tooltip": "Fixed shape for input image, e.g., '720p', '480p', when resize_mode is 'keep_ratio_fixed_area' or 'fixed_min_side'", - }, - ), - "segment_length": ( - "INT", - { - "default": 81, - "min": 16, - "max": 256, - "tooltip": "Segment length in frames for sekotalk models (target_video_length)", - }, - ), - "prev_frame_length": ( - "INT", - { - "default": 5, - "min": 0, - "max": 16, - "tooltip": "Previous frame overlap for sekotalk models", - }, - ), - "use_tiny_vae": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Use lightweight VAE to accelerate decoding", - }, - ), - }, - } - - RETURN_TYPES = ("INFERENCE_CONFIG",) - RETURN_NAMES = ("inference_config",) - FUNCTION = "create_config" - CATEGORY = "LightX2V/Config" - - def create_config( - self, - model_cls, - model_name, - task, - infer_steps, - seed, - cfg_scale, - cfg_scale2, - sample_shift, - height, - width, - duration, - attention_type, - denoising_steps="", - resize_mode="adaptive", - fixed_area="720p", - segment_length=81, - prev_frame_length=5, - use_tiny_vae=False, - ): - """Create basic inference configuration.""" - builder = InferenceConfigBuilder() - - config = builder.build( - model_cls=model_cls, - model_name=model_name, - task=task, - infer_steps=infer_steps, - seed=seed, - cfg_scale=cfg_scale, - cfg_scale2=cfg_scale2, - sample_shift=sample_shift, - height=height, - width=width, - duration=duration, - attention_type=attention_type, - denoising_steps=denoising_steps, - resize_mode=resize_mode, - fixed_area=fixed_area, - segment_length=segment_length, - prev_frame_length=prev_frame_length, - use_tiny_vae=use_tiny_vae, - ) - - return (config.to_dict(),) - - -class LightX2VTeaCache: - """TeaCache configuration node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "enable": ( - "BOOLEAN", - {"default": False, "tooltip": "Enable TeaCache feature caching"}, - ), - "threshold": ( - "FLOAT", - { - "default": 0.26, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "tooltip": "Cache threshold, lower values provide more speedup: 0.1 ~2x speedup, 0.2 ~3x speedup", - }, - ), - "use_ret_steps": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Only cache key steps to balance quality and speed", - }, - ), - } - } - - RETURN_TYPES = ("TEACACHE_CONFIG",) - RETURN_NAMES = ("teacache_config",) - FUNCTION = "create_config" - CATEGORY = "LightX2V/Config" - - def create_config(self, enable, threshold, use_ret_steps): - """Create TeaCache configuration.""" - config = TeaCacheConfig(enable=enable, threshold=threshold, use_ret_steps=use_ret_steps) - return (config.to_dict(),) - - -class LightX2VQuantization: - @classmethod - def INPUT_TYPES(cls): - available_ops = get_available_quant_ops() - quant_backends = [] - - for op_name, is_available in available_ops: - if is_available: - quant_backends.append(op_name) - - common_schema = ["fp8", "int8"] - supported_quant_schemes = ["Default"] - for schema in common_schema: - for backend in quant_backends: - supported_quant_schemes.append(f"{schema}-{backend}") - - return { - "required": { - "dit_quant_scheme": ( - supported_quant_schemes, - { - "default": supported_quant_schemes[0], - "tooltip": "DIT model quantization precision", - }, - ), - "t5_quant_scheme": ( - supported_quant_schemes, - { - "default": supported_quant_schemes[0], - "tooltip": "T5 encoder quantization precision", - }, - ), - "clip_quant_scheme": ( - supported_quant_schemes, - { - "default": supported_quant_schemes[0], - "tooltip": "CLIP encoder quantization precision", - }, - ), - "adapter_quant_scheme": ( - supported_quant_schemes, - { - "default": supported_quant_schemes[0], - "tooltip": "Adapter quantization precision", - }, - ), - } - } - - RETURN_TYPES = ("QUANT_CONFIG",) - RETURN_NAMES = ("quantization_config",) - FUNCTION = "create_config" - CATEGORY = "LightX2V/Config" - - def create_config( - self, - dit_quant_scheme, - t5_quant_scheme, - clip_quant_scheme, - adapter_quant_scheme, - ): - """Create quantization configuration.""" - config = QuantizationConfig( - dit_quant_scheme=dit_quant_scheme, - t5_quant_scheme=t5_quant_scheme, - clip_quant_scheme=clip_quant_scheme, - adapter_quant_scheme=adapter_quant_scheme, - ) - return (config.to_dict(),) - - -class LightX2VMemoryOptimization: - """Memory optimization configuration node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "enable_rotary_chunk": ( - "BOOLEAN", - {"default": False, "tooltip": "Enable rotary encoding chunking"}, - ), - "rotary_chunk_size": ( - "INT", - {"default": 100, "min": 100, "max": 10000, "step": 100}, - ), - "clean_cuda_cache": ( - "BOOLEAN", - {"default": False, "tooltip": "Clean CUDA cache promptly"}, - ), - "cpu_offload": ( - "BOOLEAN", - {"default": True, "tooltip": "Enable CPU offloading"}, - ), - "offload_granularity": ( - ["block", "phase", "model"], - {"default": "block", "tooltip": "Offload granularity"}, - ), - "offload_ratio": ( - "FLOAT", - {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.1}, - ), - "t5_cpu_offload": ( - "BOOLEAN", - {"default": True, "tooltip": "Enable T5 CPU offloading"}, - ), - "t5_offload_granularity": ( - ["model", "block"], - {"default": "model", "tooltip": "T5 offload granularity"}, - ), - "audio_encoder_cpu_offload": ( - "BOOLEAN", - {"default": True, "tooltip": "Enable audio encoder CPU offloading"}, - ), - "audio_adapter_cpu_offload": ( - "BOOLEAN", - {"default": True, "tooltip": "Enable audio adapter CPU offloading"}, - ), - "vae_cpu_offload": ( - "BOOLEAN", - {"default": True, "tooltip": "Enable VAE CPU offloading"}, - ), - "use_tiling_vae": ( - "BOOLEAN", - {"default": True, "tooltip": "Enable VAE tiling inference"}, - ), - "lazy_load": ( - "BOOLEAN", - {"default": False, "tooltip": "Lazy load model"}, - ), - "unload_after_inference": ( - "BOOLEAN", - {"default": False, "tooltip": "Unload modules after inference"}, - ), - }, - } - - RETURN_TYPES = ("MEMORY_CONFIG",) - RETURN_NAMES = ("memory_config",) - FUNCTION = "create_config" - CATEGORY = "LightX2V/Config" - - def create_config( - self, - enable_rotary_chunk=False, - rotary_chunk_size=100, - clean_cuda_cache=False, - cpu_offload=False, - offload_granularity="phase", - offload_ratio=1.0, - t5_cpu_offload=True, - t5_offload_granularity="model", - audio_encoder_cpu_offload=False, - audio_adapter_cpu_offload=False, - vae_cpu_offload=False, - use_tiling_vae=False, - lazy_load=False, - unload_after_inference=False, - ): - """Create memory optimization configuration.""" - config = MemoryOptimizationConfig( - enable_rotary_chunk=enable_rotary_chunk, - rotary_chunk_size=rotary_chunk_size, - clean_cuda_cache=clean_cuda_cache, - cpu_offload=cpu_offload, - offload_granularity=offload_granularity, - offload_ratio=offload_ratio, - t5_cpu_offload=t5_cpu_offload, - t5_offload_granularity=t5_offload_granularity, - audio_encoder_cpu_offload=audio_encoder_cpu_offload, - audio_adapter_cpu_offload=audio_adapter_cpu_offload, - vae_cpu_offload=vae_cpu_offload, - use_tiling_vae=use_tiling_vae, - lazy_load=lazy_load, - unload_after_inference=unload_after_inference, - ) - return (config.to_dict(),) - - -class LightX2VLoRALoader: - @classmethod - def INPUT_TYPES(cls): - available_loras = scan_loras() - - return { - "required": { - "lora_name": ( - available_loras, - { - "default": available_loras[0], - "tooltip": "Select LoRA from available LoRAs", - }, - ), - "strength": ( - "FLOAT", - { - "default": 1.0, - "min": 0.0, - "max": 2.0, - "step": 0.1, - "tooltip": "LoRA strength", - }, - ), - }, - "optional": { - "lora_chain": ( - "LORA_CHAIN", - {"tooltip": "Previous LoRA chain to append to"}, - ), - }, - } - - RETURN_TYPES = ("LORA_CHAIN",) - RETURN_NAMES = ("lora_chain",) - FUNCTION = "load_lora" - CATEGORY = "LightX2V/LoRA" - - def load_lora(self, lora_name, strength, lora_chain=None): - """Load and chain LoRA configurations.""" - chain = LoRAChainBuilder.build_chain(lora_name=lora_name, strength=strength, existing_chain=lora_chain) - return (chain,) - - -class TalkObjectInput: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "name": ( - "STRING", - {"default": "person_1", "tooltip": "speaker name identifier"}, - ), - }, - "optional": { - "audio": ("AUDIO", {"tooltip": "uploaded audio file"}), - "mask": ("MASK", {"tooltip": "uploaded mask image (optional)"}), - "save_to_input": ( - "BOOLEAN", - {"default": True, "tooltip": "save to input folder"}, - ), - }, - } - - RETURN_TYPES = ("TALK_OBJECT",) - RETURN_NAMES = ("talk_object",) - FUNCTION = "create_talk_object" - CATEGORY = "LightX2V/Audio" - - def create_talk_object(self, name, audio=None, mask=None, save_to_input=True): - """Create a talk object from input data.""" - builder = TalkObjectConfigBuilder() - - talk_object = builder.build_from_input(name=name, audio=audio, mask=mask, save_to_input=save_to_input) - - if talk_object: - return (talk_object,) - return (None,) - - -class TalkObjectsCombiner: - PREDEFINED_SLOTS = 16 - - @classmethod - def INPUT_TYPES(cls): - inputs = {"required": {}, "optional": {}} - - for i in range(cls.PREDEFINED_SLOTS): - inputs["optional"][f"talk_object_{i + 1}"] = ( - "TALK_OBJECT", - {"tooltip": f"talk object {i + 1}"}, - ) - - return inputs - - RETURN_TYPES = ("TALK_OBJECTS_CONFIG",) - RETURN_NAMES = ("talk_objects_config",) - FUNCTION = "combine_talk_objects" - CATEGORY = "LightX2V/Audio" - - def combine_talk_objects(self, **kwargs): - config = TalkObjectsConfig() - - for i in range(self.PREDEFINED_SLOTS): - talk_obj = kwargs.get(f"talk_object_{i + 1}") - - if talk_obj is not None: - config.add_object(talk_obj) - - if not config.talk_objects: - return (None,) - - return (config,) - - -class TalkObjectsFromJSON: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "json_config": ( - "STRING", - { - "multiline": True, - "default": '[{"name": "person1", "audio": "/path/to/audio1.wav", "mask": "/path/to/mask1.png"}]', - "tooltip": "JSON format talk objects configuration", - }, - ), - }, - } - - RETURN_TYPES = ("TALK_OBJECTS_CONFIG",) - RETURN_NAMES = ("talk_objects_config",) - FUNCTION = "parse_json_config" - CATEGORY = "LightX2V/Audio" - - def parse_json_config(self, json_config): - builder = TalkObjectConfigBuilder() - talk_objects_config = builder.build_from_json(json_config) - return (talk_objects_config,) - - -class TalkObjectsFromFiles: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "audio_files": ( - "STRING", - { - "multiline": True, - "default": "audio1.wav\naudio2.wav", - "tooltip": "audio file list (one per line)", - }, - ), - }, - "optional": { - "mask_files": ( - "STRING", - { - "multiline": True, - "default": "mask1.png\nmask2.png", - "tooltip": "mask file list (one per line, optional)", - }, - ), - "names": ( - "STRING", - { - "multiline": True, - "default": "person1\nperson2", - "tooltip": "talk object name list (one per line, optional)", - }, - ), - }, - } - - RETURN_TYPES = ("TALK_OBJECTS_CONFIG",) - RETURN_NAMES = ("talk_objects_config",) - FUNCTION = "build_from_files" - CATEGORY = "LightX2V/Audio" - - def build_from_files(self, audio_files, mask_files="", names=""): - builder = TalkObjectConfigBuilder() - talk_objects_config = builder.build_from_files(audio_files, mask_files, names) - return (talk_objects_config,) - - -class LightX2VConfigCombiner: - def __init__(self): - self.config_builder = ConfigBuilder() - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "inference_config": ( - "INFERENCE_CONFIG", - {"tooltip": "Basic inference configuration"}, - ), - }, - "optional": { - "teacache_config": ( - "TEACACHE_CONFIG", - {"tooltip": "TeaCache configuration"}, - ), - "quantization_config": ( - "QUANT_CONFIG", - {"tooltip": "Quantization configuration"}, - ), - "memory_config": ( - "MEMORY_CONFIG", - {"tooltip": "Memory optimization configuration"}, - ), - "lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}), - }, - } - - RETURN_TYPES = ("COMBINED_CONFIG",) - RETURN_NAMES = ("combined_config",) - FUNCTION = "combine_configs" - CATEGORY = "LightX2V/Config" - - def combine_configs( - self, - inference_config, - teacache_config=None, - quantization_config=None, - memory_config=None, - lora_chain=None, - talk_objects_config=None, - ): - """Combine multiple configurations into final config.""" - # Convert dict configs back to objects if needed - - # Create objects from dicts - inf_config = InferenceConfig(**inference_config) if isinstance(inference_config, dict) else None - tea_config = TeaCacheConfig(**teacache_config) if teacache_config and isinstance(teacache_config, dict) else None - quant_config = QuantizationConfig(**quantization_config) if quantization_config and isinstance(quantization_config, dict) else None - mem_config = MemoryOptimizationConfig(**memory_config) if memory_config and isinstance(memory_config, dict) else None - - config = self.config_builder.combine_configs( - inference_config=inf_config, - teacache_config=tea_config, - quantization_config=quant_config, - memory_config=mem_config, - lora_chain=lora_chain, - talk_objects_config=talk_objects_config, - ) - - return (config,) - - -class LightX2VConfigCombinerV2: - """Config combiner that also handles data preparation (image/audio/prompts).""" - - def __init__(self): - self.config_builder = ConfigBuilder() - self.temp_manager = TempFileManager() - self.image_handler = ImageFileHandler() - self.audio_handler = AudioFileHandler() - self.resolver = ComfyUIFileResolver() - self.http_downloader = HTTPFileDownloader() - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "inference_config": ( - "INFERENCE_CONFIG", - {"tooltip": "Basic inference configuration"}, - ), - "prompt": ( - "STRING", - {"multiline": True, "default": "", "tooltip": "Generation prompt"}, - ), - "negative_prompt": ( - "STRING", - {"multiline": True, "default": "", "tooltip": "Negative prompt"}, - ), - }, - "optional": { - "teacache_config": ( - "TEACACHE_CONFIG", - {"tooltip": "TeaCache configuration"}, - ), - "quantization_config": ( - "QUANT_CONFIG", - {"tooltip": "Quantization configuration"}, - ), - "memory_config": ( - "MEMORY_CONFIG", - {"tooltip": "Memory optimization configuration"}, - ), - "lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}), - "talk_objects_config": ("TALK_OBJECTS_CONFIG", {"tooltip": "Talk objects configuration"}), - "image": ("IMAGE", {"tooltip": "Input image for i2v or s2v task"}), - "audio": ( - "AUDIO", - {"tooltip": "Input audio for audio-driven generation for s2v or rs2v task"}, - ), - }, - } - - RETURN_TYPES = ("PREPARED_CONFIG",) - RETURN_NAMES = ("prepared_config",) - FUNCTION = "prepare_config" - CATEGORY = "LightX2V/ConfigV2" - - def prepare_config( - self, - inference_config, - prompt, - negative_prompt, - teacache_config=None, - quantization_config=None, - memory_config=None, - lora_chain=None, - talk_objects_config=None, - image=None, - audio=None, - ): - """Combine configurations and prepare data for inference.""" - - # Convert dict configs back to objects if needed - inf_config = InferenceConfig(**inference_config) if isinstance(inference_config, dict) else inference_config - tea_config = TeaCacheConfig(**teacache_config) if teacache_config and isinstance(teacache_config, dict) else teacache_config - quant_config = ( - QuantizationConfig(**quantization_config) if quantization_config and isinstance(quantization_config, dict) else quantization_config - ) - mem_config = MemoryOptimizationConfig(**memory_config) if memory_config and isinstance(memory_config, dict) else memory_config - - # Build combined config - config = self.config_builder.combine_configs( - inference_config=inf_config, - teacache_config=tea_config, - quantization_config=quant_config, - memory_config=mem_config, - lora_chain=lora_chain, - talk_objects_config=talk_objects_config, - ) - - # Add prompts to config - config.prompt = prompt - config.negative_prompt = negative_prompt - - # Validate task requirements - if config.task in ["i2v", "s2v", "rs2v"] and image is None: - raise ValueError("i2v or s2v or rs2v task requires input image") - - # Handle image input - if config.task in ["i2v", "s2v", "rs2v"] and image is not None: - image_np = (image[0].cpu().numpy() * 255).astype(np.uint8) - pil_image = Image.fromarray(image_np) - - temp_path = self.temp_manager.create_temp_file(suffix=".png") - pil_image.save(temp_path) - config.image_path = temp_path - logging.info(f"Image saved to {temp_path}") - - # Handle audio input for seko models - if audio is not None and hasattr(config, "model_cls") and "seko" in config.model_cls: - temp_path = self.temp_manager.create_temp_file(suffix=".wav") - self.audio_handler.save(audio, temp_path) - config.audio_path = temp_path - logging.info(f"Audio saved to {temp_path}") - - # Handle talk objects - if hasattr(config, "talk_objects") and config.talk_objects: - talk_objects = config.talk_objects - processed_talk_objects = [] - - for talk_obj in talk_objects: - processed_obj = {} - - if "audio" in talk_obj: - processed_obj["audio"] = talk_obj["audio"] - - if "mask" in talk_obj: - processed_obj["mask"] = talk_obj["mask"] - - if "audio" in processed_obj: - processed_talk_objects.append(processed_obj) - - # Resolve paths and download URLs - for obj in processed_talk_objects: - if "audio" in obj and obj["audio"]: - audio_path = obj["audio"] - - # Check if it's a URL and download if needed - if self.http_downloader.is_url(audio_path): - try: - downloaded_path = self.http_downloader.download_if_url(audio_path, prefix="audio") - obj["audio"] = downloaded_path - logging.info(f"Downloaded audio from URL: {audio_path} -> {downloaded_path}") - except Exception as e: - logging.error(f"Failed to download audio from {audio_path}: {e}") - continue - # Handle relative paths - elif not os.path.isabs(audio_path) and not audio_path.startswith("/tmp"): - obj["audio"] = self.resolver.resolve_input_path(audio_path) - logging.info(f"Resolved audio path: {audio_path} -> {obj['audio']}") - - # Check if file exists - if not os.path.exists(obj["audio"]): - logging.warning(f"Audio file not found: {obj['audio']}") - - if "mask" in obj and obj["mask"]: - mask_path = obj["mask"] - - # Check if it's a URL and download if needed - if self.http_downloader.is_url(mask_path): - try: - downloaded_path = self.http_downloader.download_if_url(mask_path, prefix="mask") - obj["mask"] = downloaded_path - logging.info(f"Downloaded mask from URL: {mask_path} -> {downloaded_path}") - except Exception as e: - logging.error(f"Failed to download mask from {mask_path}: {e}") - # Don't skip the object if mask download fails (mask is optional) - # Handle relative paths - elif not os.path.isabs(mask_path) and not mask_path.startswith("/tmp"): - obj["mask"] = self.resolver.resolve_input_path(mask_path) - logging.info(f"Resolved mask path: {mask_path} -> {obj['mask']}") - - # Check if file exists - if not os.path.exists(obj["mask"]): - logging.warning(f"Mask file not found: {obj['mask']}") - - if processed_talk_objects: - if len(processed_talk_objects) == 1 and not processed_talk_objects[0].get("mask", "").strip(): - config.audio_path = processed_talk_objects[0]["audio"] - logging.info(f"Convert Processed 1 talk object to audio path: {config.audio_path}") - else: - temp_dir = self.temp_manager.create_temp_dir() - with open(os.path.join(temp_dir, "config.json"), "w") as f: - json.dump({"talk_objects": processed_talk_objects}, f) - config.audio_path = temp_dir - logging.info(f"Processed {len(processed_talk_objects)} talk objects") - - logging.info("lightx2v prepared config: " + json.dumps(config, indent=2, ensure_ascii=False)) - - return (config,) - - -class LightX2VConfigCombinerV3: - """Config combiner that also handles data preparation (image/audio/prompts).""" - - def __init__(self): - self.config_builder = ConfigBuilder() - self.temp_manager = TempFileManager() - self.image_handler = ImageFileHandler() - self.audio_handler = AudioFileHandler() - self.resolver = ComfyUIFileResolver() - self.http_downloader = HTTPFileDownloader() - - @staticmethod - def extend_mp3(input_path: str, output_path: str, duration: float) -> bool: - """Extend or truncate MP3 audio file. - - Extend or truncate the input audio based on its duration and target - duration: - - If input duration > duration + 0.1, raise an error - - If input duration is in [duration, duration + 0.1), truncate audio - - If input duration < duration, extend audio using silence padding - - - Args: - input_path (str): - Path to the input MP3 file. - output_path (str): - Path to the output MP3 file. - duration (float): - Target duration in seconds. - - Returns: - bool: - Returns True if the operation succeeds. - - Raises: - ValueError: - Raised when input audio duration exceeds duration + 0.1. - """ - cmd_probe = [ - "ffprobe", - "-v", - "error", - "-select_streams", - "a:0", - "-show_entries", - "stream=duration,sample_rate,bit_rate,channels", - "-of", - "json", - input_path, - ] - - try: - output = sp.check_output(cmd_probe, encoding="utf-8", errors="replace") - data = json.loads(output) - streams = data.get("streams", []) - if not streams: - raise ValueError(f"Failed to get audio stream information: {input_path}") - - stream_info = streams[0] - input_duration = float(stream_info.get("duration", 0)) - sample_rate = stream_info.get("sample_rate", "44100") - bit_rate = stream_info.get("bit_rate", "128000") - channels = stream_info.get("channels", 2) - - if input_duration > duration: - raise ValueError(f"Input audio duration ({input_duration:.2f}s) exceeds target duration + 0.1s ({duration + 0.1:.2f}s)") - else: - pad_duration = duration - input_duration - cmd = [ - "ffmpeg", - "-i", - input_path, - "-af", - f"apad=pad_dur={pad_duration}", - "-ar", - str(sample_rate), - "-b:a", - str(bit_rate), - "-ac", - str(channels), - "-c:a", - "libmp3lame", - "-y", - output_path, - ] - - sp.run(cmd, capture_output=True, text=True, check=True, encoding="utf-8", errors="replace") - return True - - except sp.CalledProcessError as e: - if e.stderr: - logging.error(f"Subprocess execution failed, stderr: {e.stderr}") - raise - except json.JSONDecodeError as e: - raise ValueError(f"Failed to parse audio information: {input_path}") - except Exception as e: - raise - - @staticmethod - def get_audio_duration(input_path: str) -> float: - """Get the duration of an audio file. - - Uses ffprobe to extract audio stream information and returns the - duration in seconds. - - - Args: - input_path (str): - Path to the audio file. - - Returns: - float: - Audio duration in seconds. - - Raises: - ValueError: - Raised when audio stream information cannot be retrieved or - parsed. - """ - cmd_probe = [ - "ffprobe", - "-v", - "error", - "-select_streams", - "a:0", - "-show_entries", - "stream=duration,sample_rate,bit_rate,channels", - "-of", - "json", - input_path, - ] - try: - output = sp.check_output(cmd_probe, encoding="utf-8", errors="replace") - data = json.loads(output) - streams = data.get("streams", []) - if not streams: - raise ValueError(f"Failed to get audio stream information: {input_path}") - - stream_info = streams[0] - input_duration = float(stream_info.get("duration", 0)) - return input_duration - - except sp.CalledProcessError as e: - if e.stderr: - logging.error(f"Subprocess execution failed, stderr: {e.stderr}") - raise e - except json.JSONDecodeError as e: - raise ValueError(f"Failed to parse audio information: {input_path}") from e - except Exception as e: - raise e - - @staticmethod - def generate_white_noise( - duration: float, framerate: int, n_channels: int = 1, rms: float = None, std_dev: float = None, seed: int = None - ) -> np.ndarray: - """Generate white noise audio. - - Generate white noise audio data with optional normalization using - RMS or standard deviation. The noise is generated using a normal - distribution and can be normalized to a target RMS value or standard - deviation. - - - Args: - duration (float): - Audio duration in seconds. - framerate (int): - Sample rate in Hz. - n_channels (int, optional): - Number of audio channels. Defaults to 1 (mono). - rms (float, optional): - Target RMS value for normalization. If provided, the noise - will be normalized to this RMS value. Defaults to None. - std_dev (float, optional): - Target standard deviation for normalization. If provided, the - noise will be normalized to this standard deviation. - Defaults to None. - seed (int, optional): - Random seed for reproducible generation. Defaults to None. - - Returns: - np.ndarray: - Generated audio data with shape (n_samples, n_channels) for - multi-channel or (n_samples,) for mono channel, where - n_samples = duration * framerate. - """ - if seed is not None: - np.random.seed(seed) - - n_samples = int(duration * framerate) - - if n_channels == 1: - noise = np.random.normal(0, 1, n_samples).astype(np.float32) - else: - noise = np.random.normal(0, 1, (n_samples, n_channels)).astype(np.float32) - - if std_dev is not None: - current_std = np.std(noise) - if current_std > 0: - noise = noise * (std_dev / current_std) - elif rms is not None: - current_rms = np.sqrt(np.mean(noise**2)) - if current_rms > 0: - noise = noise * (rms / current_rms) - return noise - - @staticmethod - def save_wav_file(audio_data: np.ndarray, output_path: str | io.BytesIO, framerate: int, sample_width: int = 2) -> None: - """Save audio data as WAV file or BytesIO object. - - Convert normalized float audio data to integer format and save as - WAV file. Supports mono and multi-channel audio with configurable - sample width. - - - Args: - audio_data (np.ndarray): - Audio data with shape (n_samples,) for mono or - (n_samples, n_channels) for multi-channel. Values should - be in the range [-1.0, 1.0]. - output_path (str | io.BytesIO): - Output file path as string or BytesIO object. - framerate (int): - Sample rate in Hz. - sample_width (int, optional): - Sample width in bytes. Supported values are 1 (8-bit), - 2 (16-bit), and 4 (32-bit). Defaults to 2. - """ - if audio_data.ndim == 1: - n_channels = 1 - audio_data = audio_data.reshape(-1, 1) - else: - n_channels = audio_data.shape[1] - - audio_data = np.clip(audio_data, -1.0, 1.0) - - if sample_width == 1: - audio_int = ((audio_data + 1.0) * 127.5).astype(np.uint8) - elif sample_width == 2: - audio_int = (audio_data * 32767).astype(np.int16) - elif sample_width == 4: - audio_int = (audio_data * 2147483647).astype(np.int32) - else: - raise ValueError(f"Unsupported sample width: {sample_width}") - - if n_channels == 1: - audio_int = audio_int.flatten() - else: - audio_int = audio_int.reshape(-1, n_channels) - - with wave.open(output_path, "wb") as wav_file: - wav_file.setnchannels(n_channels) - wav_file.setsampwidth(sample_width) - wav_file.setframerate(framerate) - wav_file.writeframes(audio_int.tobytes()) - - @staticmethod - def generate_background_mask(positive_mask_paths: list[str]) -> io.BytesIO: - """Generate background mask from positive mask images. - - Generate a background mask by finding pixels that are zero (or - below threshold) in all input positive mask images. The resulting - mask marks background regions (all masks are zero) as white (255) - and foreground regions (any mask has non-zero values) as black (0). - - - Args: - positive_mask_paths (list[str]): - List of paths to positive mask image files. All images - must have the same width and height. - - Returns: - io.BytesIO: - BytesIO object containing the background mask image in JPEG - format. The mask is a grayscale image where white (255) - represents background regions and black (0) represents - foreground regions. - - Raises: - ValueError: - Raised when mask images have different dimensions. - """ - width = None - height = None - opened_imgs = list() - for path in positive_mask_paths: - img = Image.open(path) - if width is None: - width = img.width - elif width != img.width: - raise ValueError(f"Widths of masks are not the same: {width} != {img.width}") - if height is None: - height = img.height - elif height != img.height: - raise ValueError(f"Heights of masks are not the same: {height} != {img.height}") - opened_imgs.append(img) - img_arrays = [] - for img in opened_imgs: - img_array = np.array(img) - if img_array.ndim == 2: - img_array = img_array[:, :, np.newaxis] - img_arrays.append(img_array) - - threshold = 1 - zero_masks = [] - for img_array in img_arrays: - if img_array.shape[-1] == 1: - zero_mask = img_array[:, :, 0] <= threshold - else: - zero_mask = np.all(img_array <= threshold, axis=-1) - zero_masks.append(zero_mask) - - if zero_masks: - all_zero_mask = np.logical_and.reduce(zero_masks) - bg_array = np.where(all_zero_mask, 255, 0).astype(np.uint8) - else: - bg_array = np.full((height, width), 255, dtype=np.uint8) - - bg_img = Image.fromarray(bg_array, mode="L") - img_io = io.BytesIO() - bg_img.save(img_io, format="JPEG") - img_io.seek(0) - for img in opened_imgs: - img.close() - return img_io - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "inference_config": ( - "INFERENCE_CONFIG", - {"tooltip": "Basic inference configuration"}, - ), - "prompt": ( - "STRING", - {"multiline": True, "default": "", "tooltip": "Generation prompt"}, - ), - "negative_prompt": ( - "STRING", - {"multiline": True, "default": "", "tooltip": "Negative prompt"}, - ), - }, - "optional": { - "teacache_config": ( - "TEACACHE_CONFIG", - {"tooltip": "TeaCache configuration"}, - ), - "quantization_config": ( - "QUANT_CONFIG", - {"tooltip": "Quantization configuration"}, - ), - "memory_config": ( - "MEMORY_CONFIG", - {"tooltip": "Memory optimization configuration"}, - ), - "lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}), - "talk_objects_config": ("TALK_OBJECTS_CONFIG", {"tooltip": "Talk objects configuration"}), - "image": ("IMAGE", {"tooltip": "Input image for i2v or s2v or rs2v task"}), - "audio": ( - "AUDIO", - {"tooltip": "Input audio for audio-driven generation for s2v or rs2v task"}, - ), - }, - } - - RETURN_TYPES = ("PREPARED_CONFIG",) - RETURN_NAMES = ("prepared_config",) - FUNCTION = "prepare_config" - CATEGORY = "LightX2V/ConfigV2" - - def prepare_config( - self, - inference_config, - prompt, - negative_prompt, - teacache_config=None, - quantization_config=None, - memory_config=None, - lora_chain=None, - talk_objects_config=None, - image=None, - audio=None, - ): - """Combine configurations and prepare data for inference.""" - - # Convert dict configs back to objects if needed - inf_config = InferenceConfig(**inference_config) if isinstance(inference_config, dict) else inference_config - tea_config = TeaCacheConfig(**teacache_config) if teacache_config and isinstance(teacache_config, dict) else teacache_config - quant_config = ( - QuantizationConfig(**quantization_config) if quantization_config and isinstance(quantization_config, dict) else quantization_config - ) - mem_config = MemoryOptimizationConfig(**memory_config) if memory_config and isinstance(memory_config, dict) else memory_config - - # Build combined config - config = self.config_builder.combine_configs( - inference_config=inf_config, - teacache_config=tea_config, - quantization_config=quant_config, - memory_config=mem_config, - lora_chain=lora_chain, - talk_objects_config=talk_objects_config, - ) - - # Add prompts to config - config.prompt = prompt - config.negative_prompt = negative_prompt - - # Validate task requirements - if config.task in ["i2v", "s2v", "rs2v"] and image is None: - raise ValueError("i2v or s2v or rs2v task requires input image") - - # Handle image input - if config.task in ["i2v", "s2v", "rs2v"] and image is not None: - image_np = (image[0].cpu().numpy() * 255).astype(np.uint8) - pil_image = Image.fromarray(image_np) - - temp_path = self.temp_manager.create_temp_file(suffix=".png") - pil_image.save(temp_path) - config.image_path = temp_path - logging.info(f"Image saved to {temp_path}") - - # Handle audio input for seko models - if audio is not None and hasattr(config, "model_cls") and "seko" in config.model_cls: - temp_path = self.temp_manager.create_temp_file(suffix=".wav") - self.audio_handler.save(audio, temp_path) - config.audio_path = temp_path - logging.info(f"Audio saved to {temp_path}") - - # Handle talk objects - if hasattr(config, "talk_objects") and config.talk_objects: - talk_objects = config.talk_objects - src_talk_objects = [] - - for talk_obj in talk_objects: - src_obj = {} - - if "audio" in talk_obj: - src_obj["audio"] = talk_obj["audio"] - - if "mask" in talk_obj: - src_obj["mask"] = talk_obj["mask"] - - if "audio" in src_obj: - src_talk_objects.append(src_obj) - - # Resolve paths and download URLs, - # record the max duration of the src talk objects - max_src_duration = None - for obj in src_talk_objects: - if "audio" in obj and obj["audio"]: - audio_path = obj["audio"] - - # Check if it's a URL and download if needed - if self.http_downloader.is_url(audio_path): - try: - downloaded_path = self.http_downloader.download_if_url(audio_path, prefix="audio") - obj["audio"] = downloaded_path - logging.info(f"Downloaded audio from URL: {audio_path} -> {downloaded_path}") - except Exception as e: - logging.error(f"Failed to download audio from {audio_path}: {e}") - continue - # Handle relative paths - elif not os.path.isabs(audio_path) and not audio_path.startswith("/tmp"): - obj["audio"] = self.resolver.resolve_input_path(audio_path) - logging.info(f"Resolved audio path: {audio_path} -> {obj['audio']}") - - # Check if file exists - if not os.path.exists(obj["audio"]): - logging.warning(f"Audio file not found: {obj['audio']}") - duration = self.get_audio_duration(obj["audio"]) - obj["duration"] = duration - if max_src_duration is None or duration > max_src_duration: - max_src_duration = duration - - if "mask" in obj and obj["mask"]: - mask_path = obj["mask"] - - # Check if it's a URL and download if needed - if self.http_downloader.is_url(mask_path): - try: - downloaded_path = self.http_downloader.download_if_url(mask_path, prefix="mask") - obj["mask"] = downloaded_path - logging.info(f"Downloaded mask from URL: {mask_path} -> {downloaded_path}") - except Exception as e: - logging.error(f"Failed to download mask from {mask_path}: {e}") - # Don't skip the object if mask download fails (mask is optional) - # Handle relative paths - elif not os.path.isabs(mask_path) and not mask_path.startswith("/tmp"): - obj["mask"] = self.resolver.resolve_input_path(mask_path) - logging.info(f"Resolved mask path: {mask_path} -> {obj['mask']}") - - # Check if file exists - if not os.path.exists(obj["mask"]): - logging.warning(f"Mask file not found: {obj['mask']}") - - if len(src_talk_objects) > 1: - # extend audio duration to the max duration of the src talk objects - processed_talk_objects: list[dict[str, str]] = list() - mask_img_paths = list() - extend_count = 0 - for obj in src_talk_objects: - dst_obj = dict() - src_audio_path = obj["audio"] - src_audio_duration = obj["duration"] - if max_src_duration - src_audio_duration > 0.1: - dst_audio_path = self.temp_manager.create_temp_file(suffix=".mp3") - self.extend_mp3(src_audio_path, dst_audio_path, max_src_duration) - extend_count += 1 - dst_obj["audio"] = dst_audio_path - else: - dst_obj["audio"] = src_audio_path - src_mask = obj.get("mask", None) - if src_mask: - dst_obj["mask"] = src_mask - mask_img_paths.append(src_mask) - processed_talk_objects.append(dst_obj) - logging.info(f"Extended {extend_count} audio files") - # generate background mask and audio - bg_mask_io = self.generate_background_mask(mask_img_paths) - bg_mask_path = self.temp_manager.create_temp_file(suffix=".jpg") - with open(bg_mask_path, "wb") as f: - f.write(bg_mask_io.getvalue()) - bg_noise_data = self.generate_white_noise( - duration=max_src_duration, - framerate=16000, - n_channels=1, - rms=0.00232, - std_dev=0.00232, - ) - wav_io = io.BytesIO() - self.save_wav_file(audio_data=bg_noise_data, output_path=wav_io, framerate=16000, sample_width=2) - bg_audio_path = self.temp_manager.create_temp_file(suffix=".wav") - with open(bg_audio_path, "wb") as f: - f.write(wav_io.getvalue()) - bg_obj = dict( - audio=bg_audio_path, - mask=bg_mask_path, - ) - processed_talk_objects.append(bg_obj) - logging.info(f"Generated background mask and audio: {bg_mask_path}, {bg_audio_path}") - else: - processed_talk_objects = src_talk_objects - - if processed_talk_objects: - if len(processed_talk_objects) == 1 and not processed_talk_objects[0].get("mask", "").strip(): - config.audio_path = processed_talk_objects[0]["audio"] - logging.info(f"Convert Processed 1 talk object to audio path: {config.audio_path}") - else: - temp_dir = self.temp_manager.create_temp_dir() - with open(os.path.join(temp_dir, "config.json"), "w") as f: - json.dump({"talk_objects": processed_talk_objects}, f) - config.audio_path = temp_dir - logging.info(f"Processed {len(processed_talk_objects)} talk objects") - - logging.info("lightx2v prepared config: " + json.dumps(config, indent=2, ensure_ascii=False)) - - return (config,) - - -class LightX2VModularInferenceV2: - """Pure inference node that takes prepared config and runs inference.""" - - _current_runner = None - _current_config_hash = None - - def __init__(self): - if not hasattr(self.__class__, "_current_runner"): - self.__class__._current_runner = None - if not hasattr(self.__class__, "_current_config_hash"): - self.__class__._current_config_hash = None - - self.config_builder = ConfigBuilder() - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "prepared_config": ( - "PREPARED_CONFIG", - {"tooltip": "Fully prepared configuration from ConfigCombinerV2"}, - ), - }, - } - - RETURN_TYPES = ("IMAGE", "AUDIO") - RETURN_NAMES = ("images", "audio") - FUNCTION = "generate" - CATEGORY = "LightX2V/InferenceV2" - - def _get_config_hash(self, config) -> str: - """Get hash of configuration to detect changes.""" - return self.config_builder.get_config_hash(config) - - def _build_rs2v_shot_config(self, config): - from .lightx2v.lightx2v.shot_runner.shot_base import load_clip_configs - from .lightx2v.lightx2v.utils.lockable_dict import LockableDict - - config_json = config.get("config_json") - if config_json: - main_cfg = config_json - elif config.get("clip_configs"): - main_cfg = config - else: - main_cfg = { - "lightx2v_path": "", - "clip_configs": [ - { - "name": "rs2v_clip", - "config": LockableDict(config), - } - ], - } - if "task" not in main_cfg["clip_configs"][0]["config"]: - main_cfg["clip_configs"][0]["config"]["task"] = "rs2v" - - if isinstance(main_cfg, dict) and "lightx2v_path" not in main_cfg: - main_cfg = dict(main_cfg) - main_cfg["lightx2v_path"] = "" - - return load_clip_configs(main_cfg) - # return dict( - # seed=config.get("seed", 42), - # image_path=config.get("image_path", ""), - # audio_path=config.get("audio_path", ""), - # prompt=config.get("prompt", ""), - # negative_prompt=config.get("negative_prompt", ""), - # save_result_path=config.get("save_result_path", ""), - # clip_configs=clip_configs, - # target_shape=config.get("target_shape", []), - # ) - - def generate(self, prepared_config): - """Run inference with prepared configuration.""" - config = prepared_config - - try: - config_hash = self._get_config_hash(config) - - current_runner = getattr(self.__class__, "_current_runner", None) - current_config_hash = getattr(self.__class__, "_current_config_hash", None) - - needs_reinit = current_runner is None or current_config_hash != config_hash or getattr(config, "lazy_load", False) - - logging.info(f"Needs reinit: {needs_reinit}, old config hash: {current_config_hash}, new config hash: {config_hash}") - if needs_reinit: - if current_runner is not None: - # current_runner.end_run() - del self.__class__._current_runner - torch.cuda.empty_cache() - gc.collect() - if config.get("task") == "rs2v": - from .lightx2v.lightx2v.shot_runner.rs2v_infer import ShotRS2VPipeline - - shot_cfg = self._build_rs2v_shot_config(config) - self.__class__._current_runner = ShotRS2VPipeline(shot_cfg) - else: - formatted_config = set_config(config) - self.__class__._current_runner = init_runner(formatted_config) - self.__class__._current_config_hash = config_hash - - progress = ProgressBar(100) - - def update_progress(current_step, _total): - progress.update_absolute(current_step) - - current_runner = self.__class__._current_runner - - if hasattr(current_runner, "set_progress_callback"): - current_runner.set_progress_callback(update_progress) - - config["return_result_tensor"] = True - config["save_result_path"] = "" - config["negative_prompt"] = config.get("negative_prompt", "") - if config.get("task") == "rs2v": - # rs2v 使用 shot_runner 管线 - # current_runner.set_config(config) - result_dict = current_runner.run_pipeline(config) - else: - input_data = init_empty_input_info(config.task) - update_input_info_from_dict(input_data, config) - current_runner.set_config(config) - result_dict = current_runner.run_pipeline(input_data) - - images = result_dict.get("video", None) - audio = result_dict.get("audio", None) - - if images is not None and images.numel() > 0: - images = images.cpu() - if images.dtype != torch.float32: - images = images.float() - - if getattr(config, "unload_after_inference", False): - if hasattr(self.__class__, "_current_runner"): - del self.__class__._current_runner - self.__class__._current_runner = None - self.__class__._current_config_hash = None - - torch.cuda.empty_cache() - gc.collect() - - return (images, audio) - - except Exception as e: - logging.error(f"Error during inference: {e}") - raise - - finally: - # Cleanup is handled by TempFileManager destructor - pass - - -NODE_CLASS_MAPPINGS = { - "LightX2VInferenceConfig": LightX2VInferenceConfig, - "LightX2VTeaCache": LightX2VTeaCache, - "LightX2VQuantization": LightX2VQuantization, - "LightX2VMemoryOptimization": LightX2VMemoryOptimization, - "LightX2VLoRALoader": LightX2VLoRALoader, - "LightX2VConfigCombiner": LightX2VConfigCombiner, - "LightX2VConfigCombinerV2": LightX2VConfigCombinerV2, - "LightX2VConfigCombinerV3": LightX2VConfigCombinerV3, - "LightX2VModularInferenceV2": LightX2VModularInferenceV2, - "LightX2VTalkObjectInput": TalkObjectInput, - "LightX2VTalkObjectsCombiner": TalkObjectsCombiner, - "LightX2VTalkObjectsFromJSON": TalkObjectsFromJSON, - "LightX2VTalkObjectsFromFiles": TalkObjectsFromFiles, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "LightX2VInferenceConfig": "LightX2V Inference Config", - "LightX2VTeaCache": "LightX2V TeaCache", - "LightX2VQuantization": "LightX2V Quantization", - "LightX2VMemoryOptimization": "LightX2V Memory Optimization", - "LightX2VLoRALoader": "LightX2V LoRA Loader", - "LightX2VConfigCombiner": "LightX2V Config Combiner", - "LightX2VConfigCombinerV2": "LightX2V Config Combiner V2", - "LightX2VConfigCombinerV3": "LightX2V Config Combiner V3", - "LightX2VModularInferenceV2": "LightX2V Modular Inference V2", - "LightX2VTalkObjectInput": "LightX2V Talk Object Input (Single)", - "LightX2VTalkObjectsCombiner": "LightX2V Talk Objects Combiner", - "LightX2VTalkObjectsFromFiles": "LightX2V Talk Objects From Files", - "LightX2VTalkObjectsFromJSON": "LightX2V Talk Objects From JSON (API)", -} diff --git a/nodes/__init__.py b/nodes/__init__.py new file mode 100644 index 0000000..8b79c7b --- /dev/null +++ b/nodes/__init__.py @@ -0,0 +1,64 @@ +"""ComfyUI node definitions for LightX2V. + +Each submodule groups a category of nodes: +- ``config`` : per-feature configuration nodes (inference / teacache / quant / memory) +- ``lora`` : LoRA chain loader +- ``talk`` : talk-object input/combiner nodes +- ``combiner`` : config combiners (V1/V2/V3) that aggregate the above +- ``inference`` : the modular inference runner +- ``seedvr`` : SeedVR2 super-resolution runner +""" + +from .combiner import ( + LightX2VConfigCombinerV2, + LightX2VConfigCombinerV3, +) +from .config import ( + LightX2VInferenceConfig, + LightX2VMemoryOptimization, + LightX2VQuantization, + LightX2VTeaCache, +) +from .inference import LightX2VModularInferenceV2 +from .lora import LightX2VLoRALoader +from .seedvr import LightX2VSeedVRSR +from .talk import ( + TalkObjectInput, + TalkObjectsCombiner, + TalkObjectsFromFiles, + TalkObjectsFromJSON, +) + +NODE_CLASS_MAPPINGS = { + "LightX2VInferenceConfig": LightX2VInferenceConfig, + "LightX2VTeaCache": LightX2VTeaCache, + "LightX2VQuantization": LightX2VQuantization, + "LightX2VMemoryOptimization": LightX2VMemoryOptimization, + "LightX2VLoRALoader": LightX2VLoRALoader, + "LightX2VConfigCombinerV2": LightX2VConfigCombinerV2, + "LightX2VConfigCombinerV3": LightX2VConfigCombinerV3, + "LightX2VModularInferenceV2": LightX2VModularInferenceV2, + "LightX2VSeedVRSR": LightX2VSeedVRSR, + "LightX2VTalkObjectInput": TalkObjectInput, + "LightX2VTalkObjectsCombiner": TalkObjectsCombiner, + "LightX2VTalkObjectsFromJSON": TalkObjectsFromJSON, + "LightX2VTalkObjectsFromFiles": TalkObjectsFromFiles, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LightX2VInferenceConfig": "LightX2V Inference Config", + "LightX2VTeaCache": "LightX2V TeaCache", + "LightX2VQuantization": "LightX2V Quantization", + "LightX2VMemoryOptimization": "LightX2V Memory Optimization", + "LightX2VLoRALoader": "LightX2V LoRA Loader", + "LightX2VConfigCombinerV2": "LightX2V Config Combiner V2", + "LightX2VConfigCombinerV3": "LightX2V Config Combiner V3", + "LightX2VModularInferenceV2": "LightX2V Modular Inference V2", + "LightX2VSeedVRSR": "LightX2V SeedVR2 Super-Resolution", + "LightX2VTalkObjectInput": "LightX2V Talk Object Input (Single)", + "LightX2VTalkObjectsCombiner": "LightX2V Talk Objects Combiner", + "LightX2VTalkObjectsFromFiles": "LightX2V Talk Objects From Files", + "LightX2VTalkObjectsFromJSON": "LightX2V Talk Objects From JSON (API)", +} + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/nodes/combiner.py b/nodes/combiner.py new file mode 100644 index 0000000..48d9cc5 --- /dev/null +++ b/nodes/combiner.py @@ -0,0 +1,639 @@ +"""Config combiner nodes. + +- V2 ``LightX2VConfigCombinerV2`` : config aggregation + data prep (image/audio/talk_objects), + emits ``PREPARED_CONFIG``. +- V3 ``LightX2VConfigCombinerV3`` : V2 + equal-duration audio padding and background-mask + synthesis for multi-talker setups. +""" + +import io +import json +import logging +import os +import subprocess as sp +import wave + +import numpy as np +from PIL import Image + +from ..config_builder import ConfigBuilder +from ..data_models import ( + InferenceConfig, + MemoryOptimizationConfig, + QuantizationConfig, + TeaCacheConfig, +) +from ..file_handlers import ( + AudioFileHandler, + ComfyUIFileResolver, + HTTPFileDownloader, + ImageFileHandler, + TempFileManager, +) + + +class LightX2VConfigCombinerV2: + """Config combiner that also handles data preparation (image/audio/prompts).""" + + def __init__(self): + self.config_builder = ConfigBuilder() + self.temp_manager = TempFileManager() + self.image_handler = ImageFileHandler() + self.audio_handler = AudioFileHandler() + self.resolver = ComfyUIFileResolver() + self.http_downloader = HTTPFileDownloader() + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "inference_config": ( + "INFERENCE_CONFIG", + {"tooltip": "Basic inference configuration"}, + ), + "prompt": ( + "STRING", + {"multiline": True, "default": "", "tooltip": "Generation prompt"}, + ), + "negative_prompt": ( + "STRING", + {"multiline": True, "default": "", "tooltip": "Negative prompt"}, + ), + }, + "optional": { + "teacache_config": ( + "TEACACHE_CONFIG", + {"tooltip": "TeaCache configuration"}, + ), + "quantization_config": ( + "QUANT_CONFIG", + {"tooltip": "Quantization configuration"}, + ), + "memory_config": ( + "MEMORY_CONFIG", + {"tooltip": "Memory optimization configuration"}, + ), + "lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}), + "talk_objects_config": ("TALK_OBJECTS_CONFIG", {"tooltip": "Talk objects configuration"}), + "image": ("IMAGE", {"tooltip": "Input image for i2v or s2v task"}), + "audio": ( + "AUDIO", + {"tooltip": "Input audio for audio-driven generation for s2v or rs2v task"}, + ), + }, + } + + RETURN_TYPES = ("PREPARED_CONFIG",) + RETURN_NAMES = ("prepared_config",) + FUNCTION = "prepare_config" + CATEGORY = "LightX2V/ConfigV2" + + def prepare_config( + self, + inference_config, + prompt, + negative_prompt, + teacache_config=None, + quantization_config=None, + memory_config=None, + lora_chain=None, + talk_objects_config=None, + image=None, + audio=None, + ): + """Combine configurations and prepare data for inference.""" + + inf_config = InferenceConfig(**inference_config) if isinstance(inference_config, dict) else inference_config + tea_config = TeaCacheConfig(**teacache_config) if teacache_config and isinstance(teacache_config, dict) else teacache_config + quant_config = ( + QuantizationConfig(**quantization_config) if quantization_config and isinstance(quantization_config, dict) else quantization_config + ) + mem_config = MemoryOptimizationConfig(**memory_config) if memory_config and isinstance(memory_config, dict) else memory_config + + config = self.config_builder.combine_configs( + inference_config=inf_config, + teacache_config=tea_config, + quantization_config=quant_config, + memory_config=mem_config, + lora_chain=lora_chain, + talk_objects_config=talk_objects_config, + ) + + config.prompt = prompt + config.negative_prompt = negative_prompt + + if config.task in ["i2v", "s2v", "rs2v"] and image is None: + raise ValueError("i2v or s2v or rs2v task requires input image") + + if config.task in ["i2v", "s2v", "rs2v"] and image is not None: + image_np = (image[0].cpu().numpy() * 255).astype(np.uint8) + pil_image = Image.fromarray(image_np) + + temp_path = self.temp_manager.create_temp_file(suffix=".png") + pil_image.save(temp_path) + config.image_path = temp_path + logging.info(f"Image saved to {temp_path}") + + if audio is not None and hasattr(config, "model_cls") and "seko" in config.model_cls: + temp_path = self.temp_manager.create_temp_file(suffix=".wav") + self.audio_handler.save(audio, temp_path) + config.audio_path = temp_path + logging.info(f"Audio saved to {temp_path}") + + if hasattr(config, "talk_objects") and config.talk_objects: + talk_objects = config.talk_objects + processed_talk_objects = [] + + for talk_obj in talk_objects: + processed_obj = {} + + if "audio" in talk_obj: + processed_obj["audio"] = talk_obj["audio"] + + if "mask" in talk_obj: + processed_obj["mask"] = talk_obj["mask"] + + if "audio" in processed_obj: + processed_talk_objects.append(processed_obj) + + for obj in processed_talk_objects: + if "audio" in obj and obj["audio"]: + audio_path = obj["audio"] + + if self.http_downloader.is_url(audio_path): + try: + downloaded_path = self.http_downloader.download_if_url(audio_path, prefix="audio") + obj["audio"] = downloaded_path + logging.info(f"Downloaded audio from URL: {audio_path} -> {downloaded_path}") + except Exception as e: + logging.error(f"Failed to download audio from {audio_path}: {e}") + continue + elif not os.path.isabs(audio_path) and not audio_path.startswith("/tmp"): + obj["audio"] = self.resolver.resolve_input_path(audio_path) + logging.info(f"Resolved audio path: {audio_path} -> {obj['audio']}") + + if not os.path.exists(obj["audio"]): + logging.warning(f"Audio file not found: {obj['audio']}") + + if "mask" in obj and obj["mask"]: + mask_path = obj["mask"] + + if self.http_downloader.is_url(mask_path): + try: + downloaded_path = self.http_downloader.download_if_url(mask_path, prefix="mask") + obj["mask"] = downloaded_path + logging.info(f"Downloaded mask from URL: {mask_path} -> {downloaded_path}") + except Exception as e: + logging.error(f"Failed to download mask from {mask_path}: {e}") + elif not os.path.isabs(mask_path) and not mask_path.startswith("/tmp"): + obj["mask"] = self.resolver.resolve_input_path(mask_path) + logging.info(f"Resolved mask path: {mask_path} -> {obj['mask']}") + + if not os.path.exists(obj["mask"]): + logging.warning(f"Mask file not found: {obj['mask']}") + + if processed_talk_objects: + if len(processed_talk_objects) == 1 and not processed_talk_objects[0].get("mask", "").strip(): + config.audio_path = processed_talk_objects[0]["audio"] + logging.info(f"Convert Processed 1 talk object to audio path: {config.audio_path}") + else: + temp_dir = self.temp_manager.create_temp_dir() + with open(os.path.join(temp_dir, "config.json"), "w") as f: + json.dump({"talk_objects": processed_talk_objects}, f) + config.audio_path = temp_dir + logging.info(f"Processed {len(processed_talk_objects)} talk objects") + + logging.info("lightx2v prepared config: " + json.dumps(config, indent=2, ensure_ascii=False)) + + return (config,) + + +class LightX2VConfigCombinerV3: + """V2 + equal-duration audio padding and background-mask synthesis for multi-talker setups.""" + + def __init__(self): + self.config_builder = ConfigBuilder() + self.temp_manager = TempFileManager() + self.image_handler = ImageFileHandler() + self.audio_handler = AudioFileHandler() + self.resolver = ComfyUIFileResolver() + self.http_downloader = HTTPFileDownloader() + + @staticmethod + def extend_mp3(input_path: str, output_path: str, duration: float) -> bool: + """Extend or truncate MP3 audio file. + + - If input duration > duration + 0.1, raise an error + - If input duration is in [duration, duration + 0.1), truncate audio + - If input duration < duration, extend audio using silence padding + """ + cmd_probe = [ + "ffprobe", + "-v", + "error", + "-select_streams", + "a:0", + "-show_entries", + "stream=duration,sample_rate,bit_rate,channels", + "-of", + "json", + input_path, + ] + + try: + output = sp.check_output(cmd_probe, encoding="utf-8", errors="replace") + data = json.loads(output) + streams = data.get("streams", []) + if not streams: + raise ValueError(f"Failed to get audio stream information: {input_path}") + + stream_info = streams[0] + input_duration = float(stream_info.get("duration", 0)) + sample_rate = stream_info.get("sample_rate", "44100") + bit_rate = stream_info.get("bit_rate", "128000") + channels = stream_info.get("channels", 2) + + if input_duration > duration: + raise ValueError(f"Input audio duration ({input_duration:.2f}s) exceeds target duration + 0.1s ({duration + 0.1:.2f}s)") + else: + pad_duration = duration - input_duration + cmd = [ + "ffmpeg", + "-i", + input_path, + "-af", + f"apad=pad_dur={pad_duration}", + "-ar", + str(sample_rate), + "-b:a", + str(bit_rate), + "-ac", + str(channels), + "-c:a", + "libmp3lame", + "-y", + output_path, + ] + + sp.run(cmd, capture_output=True, text=True, check=True, encoding="utf-8", errors="replace") + return True + + except sp.CalledProcessError as e: + if e.stderr: + logging.error(f"Subprocess execution failed, stderr: {e.stderr}") + raise + except json.JSONDecodeError: + raise ValueError(f"Failed to parse audio information: {input_path}") + except Exception: + raise + + @staticmethod + def get_audio_duration(input_path: str) -> float: + """Get the duration of an audio file in seconds via ffprobe.""" + cmd_probe = [ + "ffprobe", + "-v", + "error", + "-select_streams", + "a:0", + "-show_entries", + "stream=duration,sample_rate,bit_rate,channels", + "-of", + "json", + input_path, + ] + try: + output = sp.check_output(cmd_probe, encoding="utf-8", errors="replace") + data = json.loads(output) + streams = data.get("streams", []) + if not streams: + raise ValueError(f"Failed to get audio stream information: {input_path}") + + stream_info = streams[0] + return float(stream_info.get("duration", 0)) + + except sp.CalledProcessError as e: + if e.stderr: + logging.error(f"Subprocess execution failed, stderr: {e.stderr}") + raise e + except json.JSONDecodeError as e: + raise ValueError(f"Failed to parse audio information: {input_path}") from e + except Exception as e: + raise e + + @staticmethod + def generate_white_noise( + duration: float, framerate: int, n_channels: int = 1, rms: float = None, std_dev: float = None, seed: int = None + ) -> np.ndarray: + """Generate white noise audio with optional RMS/std-dev normalization.""" + if seed is not None: + np.random.seed(seed) + + n_samples = int(duration * framerate) + + if n_channels == 1: + noise = np.random.normal(0, 1, n_samples).astype(np.float32) + else: + noise = np.random.normal(0, 1, (n_samples, n_channels)).astype(np.float32) + + if std_dev is not None: + current_std = np.std(noise) + if current_std > 0: + noise = noise * (std_dev / current_std) + elif rms is not None: + current_rms = np.sqrt(np.mean(noise**2)) + if current_rms > 0: + noise = noise * (rms / current_rms) + return noise + + @staticmethod + def save_wav_file(audio_data: np.ndarray, output_path, framerate: int, sample_width: int = 2) -> None: + """Save audio data as WAV file or BytesIO object.""" + if audio_data.ndim == 1: + n_channels = 1 + audio_data = audio_data.reshape(-1, 1) + else: + n_channels = audio_data.shape[1] + + audio_data = np.clip(audio_data, -1.0, 1.0) + + if sample_width == 1: + audio_int = ((audio_data + 1.0) * 127.5).astype(np.uint8) + elif sample_width == 2: + audio_int = (audio_data * 32767).astype(np.int16) + elif sample_width == 4: + audio_int = (audio_data * 2147483647).astype(np.int32) + else: + raise ValueError(f"Unsupported sample width: {sample_width}") + + if n_channels == 1: + audio_int = audio_int.flatten() + else: + audio_int = audio_int.reshape(-1, n_channels) + + with wave.open(output_path, "wb") as wav_file: + wav_file.setnchannels(n_channels) + wav_file.setsampwidth(sample_width) + wav_file.setframerate(framerate) + wav_file.writeframes(audio_int.tobytes()) + + @staticmethod + def generate_background_mask(positive_mask_paths): + """Generate a background mask: white where all positive masks are ~zero, else black.""" + width = None + height = None + opened_imgs = [] + for path in positive_mask_paths: + img = Image.open(path) + if width is None: + width = img.width + elif width != img.width: + raise ValueError(f"Widths of masks are not the same: {width} != {img.width}") + if height is None: + height = img.height + elif height != img.height: + raise ValueError(f"Heights of masks are not the same: {height} != {img.height}") + opened_imgs.append(img) + img_arrays = [] + for img in opened_imgs: + img_array = np.array(img) + if img_array.ndim == 2: + img_array = img_array[:, :, np.newaxis] + img_arrays.append(img_array) + + threshold = 1 + zero_masks = [] + for img_array in img_arrays: + if img_array.shape[-1] == 1: + zero_mask = img_array[:, :, 0] <= threshold + else: + zero_mask = np.all(img_array <= threshold, axis=-1) + zero_masks.append(zero_mask) + + if zero_masks: + all_zero_mask = np.logical_and.reduce(zero_masks) + bg_array = np.where(all_zero_mask, 255, 0).astype(np.uint8) + else: + bg_array = np.full((height, width), 255, dtype=np.uint8) + + bg_img = Image.fromarray(bg_array, mode="L") + img_io = io.BytesIO() + bg_img.save(img_io, format="JPEG") + img_io.seek(0) + for img in opened_imgs: + img.close() + return img_io + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "inference_config": ( + "INFERENCE_CONFIG", + {"tooltip": "Basic inference configuration"}, + ), + "prompt": ( + "STRING", + {"multiline": True, "default": "", "tooltip": "Generation prompt"}, + ), + "negative_prompt": ( + "STRING", + {"multiline": True, "default": "", "tooltip": "Negative prompt"}, + ), + }, + "optional": { + "teacache_config": ( + "TEACACHE_CONFIG", + {"tooltip": "TeaCache configuration"}, + ), + "quantization_config": ( + "QUANT_CONFIG", + {"tooltip": "Quantization configuration"}, + ), + "memory_config": ( + "MEMORY_CONFIG", + {"tooltip": "Memory optimization configuration"}, + ), + "lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}), + "talk_objects_config": ("TALK_OBJECTS_CONFIG", {"tooltip": "Talk objects configuration"}), + "image": ("IMAGE", {"tooltip": "Input image for i2v or s2v or rs2v task"}), + "audio": ( + "AUDIO", + {"tooltip": "Input audio for audio-driven generation for s2v or rs2v task"}, + ), + }, + } + + RETURN_TYPES = ("PREPARED_CONFIG",) + RETURN_NAMES = ("prepared_config",) + FUNCTION = "prepare_config" + CATEGORY = "LightX2V/ConfigV2" + + def prepare_config( + self, + inference_config, + prompt, + negative_prompt, + teacache_config=None, + quantization_config=None, + memory_config=None, + lora_chain=None, + talk_objects_config=None, + image=None, + audio=None, + ): + """Combine configurations and prepare data for inference.""" + + inf_config = InferenceConfig(**inference_config) if isinstance(inference_config, dict) else inference_config + tea_config = TeaCacheConfig(**teacache_config) if teacache_config and isinstance(teacache_config, dict) else teacache_config + quant_config = ( + QuantizationConfig(**quantization_config) if quantization_config and isinstance(quantization_config, dict) else quantization_config + ) + mem_config = MemoryOptimizationConfig(**memory_config) if memory_config and isinstance(memory_config, dict) else memory_config + + config = self.config_builder.combine_configs( + inference_config=inf_config, + teacache_config=tea_config, + quantization_config=quant_config, + memory_config=mem_config, + lora_chain=lora_chain, + talk_objects_config=talk_objects_config, + ) + + config.prompt = prompt + config.negative_prompt = negative_prompt + + if config.task in ["i2v", "s2v", "rs2v"] and image is None: + raise ValueError("i2v or s2v or rs2v task requires input image") + + if config.task in ["i2v", "s2v", "rs2v"] and image is not None: + image_np = (image[0].cpu().numpy() * 255).astype(np.uint8) + pil_image = Image.fromarray(image_np) + + temp_path = self.temp_manager.create_temp_file(suffix=".png") + pil_image.save(temp_path) + config.image_path = temp_path + logging.info(f"Image saved to {temp_path}") + + if audio is not None and hasattr(config, "model_cls") and "seko" in config.model_cls: + temp_path = self.temp_manager.create_temp_file(suffix=".wav") + self.audio_handler.save(audio, temp_path) + config.audio_path = temp_path + logging.info(f"Audio saved to {temp_path}") + + if hasattr(config, "talk_objects") and config.talk_objects: + talk_objects = config.talk_objects + src_talk_objects = [] + + for talk_obj in talk_objects: + src_obj = {} + + if "audio" in talk_obj: + src_obj["audio"] = talk_obj["audio"] + + if "mask" in talk_obj: + src_obj["mask"] = talk_obj["mask"] + + if "audio" in src_obj: + src_talk_objects.append(src_obj) + + # Resolve paths / download URLs, and record max source duration. + max_src_duration = None + for obj in src_talk_objects: + if "audio" in obj and obj["audio"]: + audio_path = obj["audio"] + + if self.http_downloader.is_url(audio_path): + try: + downloaded_path = self.http_downloader.download_if_url(audio_path, prefix="audio") + obj["audio"] = downloaded_path + logging.info(f"Downloaded audio from URL: {audio_path} -> {downloaded_path}") + except Exception as e: + logging.error(f"Failed to download audio from {audio_path}: {e}") + continue + elif not os.path.isabs(audio_path) and not audio_path.startswith("/tmp"): + obj["audio"] = self.resolver.resolve_input_path(audio_path) + logging.info(f"Resolved audio path: {audio_path} -> {obj['audio']}") + + if not os.path.exists(obj["audio"]): + logging.warning(f"Audio file not found: {obj['audio']}") + duration = self.get_audio_duration(obj["audio"]) + obj["duration"] = duration + if max_src_duration is None or duration > max_src_duration: + max_src_duration = duration + + if "mask" in obj and obj["mask"]: + mask_path = obj["mask"] + + if self.http_downloader.is_url(mask_path): + try: + downloaded_path = self.http_downloader.download_if_url(mask_path, prefix="mask") + obj["mask"] = downloaded_path + logging.info(f"Downloaded mask from URL: {mask_path} -> {downloaded_path}") + except Exception as e: + logging.error(f"Failed to download mask from {mask_path}: {e}") + elif not os.path.isabs(mask_path) and not mask_path.startswith("/tmp"): + obj["mask"] = self.resolver.resolve_input_path(mask_path) + logging.info(f"Resolved mask path: {mask_path} -> {obj['mask']}") + + if not os.path.exists(obj["mask"]): + logging.warning(f"Mask file not found: {obj['mask']}") + + if len(src_talk_objects) > 1: + # Extend each talker's audio to max duration, then synthesize a background track. + processed_talk_objects = [] + mask_img_paths = [] + extend_count = 0 + for obj in src_talk_objects: + dst_obj = {} + src_audio_path = obj["audio"] + src_audio_duration = obj["duration"] + if max_src_duration - src_audio_duration > 0.1: + dst_audio_path = self.temp_manager.create_temp_file(suffix=".mp3") + self.extend_mp3(src_audio_path, dst_audio_path, max_src_duration) + extend_count += 1 + dst_obj["audio"] = dst_audio_path + else: + dst_obj["audio"] = src_audio_path + src_mask = obj.get("mask", None) + if src_mask: + dst_obj["mask"] = src_mask + mask_img_paths.append(src_mask) + processed_talk_objects.append(dst_obj) + logging.info(f"Extended {extend_count} audio files") + + bg_mask_io = self.generate_background_mask(mask_img_paths) + bg_mask_path = self.temp_manager.create_temp_file(suffix=".jpg") + with open(bg_mask_path, "wb") as f: + f.write(bg_mask_io.getvalue()) + bg_noise_data = self.generate_white_noise( + duration=max_src_duration, + framerate=16000, + n_channels=1, + rms=0.00232, + std_dev=0.00232, + ) + wav_io = io.BytesIO() + self.save_wav_file(audio_data=bg_noise_data, output_path=wav_io, framerate=16000, sample_width=2) + bg_audio_path = self.temp_manager.create_temp_file(suffix=".wav") + with open(bg_audio_path, "wb") as f: + f.write(wav_io.getvalue()) + processed_talk_objects.append({"audio": bg_audio_path, "mask": bg_mask_path}) + logging.info(f"Generated background mask and audio: {bg_mask_path}, {bg_audio_path}") + else: + processed_talk_objects = src_talk_objects + + if processed_talk_objects: + if len(processed_talk_objects) == 1 and not processed_talk_objects[0].get("mask", "").strip(): + config.audio_path = processed_talk_objects[0]["audio"] + logging.info(f"Convert Processed 1 talk object to audio path: {config.audio_path}") + else: + temp_dir = self.temp_manager.create_temp_dir() + with open(os.path.join(temp_dir, "config.json"), "w") as f: + json.dump({"talk_objects": processed_talk_objects}, f) + config.audio_path = temp_dir + logging.info(f"Processed {len(processed_talk_objects)} talk objects") + + logging.info("lightx2v prepared config: " + json.dumps(config, indent=2, ensure_ascii=False)) + + return (config,) diff --git a/nodes/config.py b/nodes/config.py new file mode 100644 index 0000000..4d85ef3 --- /dev/null +++ b/nodes/config.py @@ -0,0 +1,448 @@ +"""Per-feature configuration nodes: inference / teacache / quantization / memory.""" + +from ..bridge import get_available_attn_ops, get_available_quant_ops +from ..config_builder import InferenceConfigBuilder +from ..data_models import ( + MemoryOptimizationConfig, + QuantizationConfig, + TeaCacheConfig, +) +from ..model_utils import scan_models, support_model_cls_list + + +class LightX2VInferenceConfig: + @classmethod + def INPUT_TYPES(cls): + available_models = scan_models() + support_model_classes = support_model_cls_list() + available_attn = get_available_attn_ops() + attn_types = [] + + for op_name, is_available in available_attn: + if is_available: + attn_types.append(op_name) + + if "torch_sdpa" not in attn_types: + attn_types.append("torch_sdpa") + + return { + "required": { + "model_cls": ( + support_model_classes, + {"default": "wan2.1", "tooltip": "Model type"}, + ), + "model_name": ( + available_models, + { + "default": available_models[0], + "tooltip": "Select model from available models", + }, + ), + "task": ( + ["t2v", "i2v", "s2v", "rs2v"], + { + "default": "i2v", + "tooltip": "Task type: text-to-video or image-to-video or reference_image and audio to video (shot)", + }, + ), + "infer_steps": ( + "INT", + {"default": 4, "min": 1, "max": 100, "tooltip": "Inference steps"}, + ), + "seed": ( + "INT", + { + "default": 42, + "min": -1, + "max": 2**32 - 1, + "tooltip": "Random seed, -1 for random", + }, + ), + "cfg_scale": ( + "FLOAT", + { + "default": 5.0, + "min": 1.0, + "max": 10.0, + "step": 0.1, + "tooltip": "CFG guidance strength", + }, + ), + "cfg_scale2": ( + "FLOAT", + { + "default": 5.0, + "min": 1.0, + "max": 10.0, + "step": 0.1, + "tooltip": "CFG guidance, lower noise when model cls is Wan2.2 MoE", + }, + ), + "sample_shift": ( + "INT", + {"default": 5, "min": 0, "max": 10, "tooltip": "Sample shift"}, + ), + "height": ( + "INT", + { + "default": 1280, + "min": 64, + "max": 2048, + "step": 8, + "tooltip": "Video height", + }, + ), + "width": ( + "INT", + { + "default": 720, + "min": 64, + "max": 2048, + "step": 8, + "tooltip": "Video width", + }, + ), + "duration": ( + "FLOAT", + { + "default": 5.0, + "min": 1.0, + "max": 999, + "step": 0.1, + "tooltip": "Video duration in seconds", + }, + ), + "attention_type": ( + attn_types, + {"default": attn_types[0], "tooltip": "Attention mechanism type"}, + ), + }, + "optional": { + "denoising_steps": ( + "STRING", + { + "default": "", + "tooltip": "Custom denoising steps for distillation models (comma-separated, e.g., '999,750,500,250'). Leave empty to use model defaults.", + }, + ), + "resize_mode": ( + [ + "adaptive", + "keep_ratio_fixed_area", + "fixed_min_area", + "fixed_max_area", + "fixed_shape", + "fixed_min_side", + ], + { + "default": "adaptive", + "tooltip": "Adaptive resize input image to target aspect ratio", + }, + ), + "fixed_area": ( + "STRING", + { + "default": "720p", + "tooltip": "Fixed shape for input image, e.g., '720p', '480p', when resize_mode is 'keep_ratio_fixed_area' or 'fixed_min_side'", + }, + ), + "segment_length": ( + "INT", + { + "default": 81, + "min": 16, + "max": 256, + "tooltip": "Segment length in frames for sekotalk models (target_video_length)", + }, + ), + "prev_frame_length": ( + "INT", + { + "default": 5, + "min": 0, + "max": 16, + "tooltip": "Previous frame overlap for sekotalk models", + }, + ), + "use_tiny_vae": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Use lightweight VAE to accelerate decoding", + }, + ), + }, + } + + RETURN_TYPES = ("INFERENCE_CONFIG",) + RETURN_NAMES = ("inference_config",) + FUNCTION = "create_config" + CATEGORY = "LightX2V/Config" + + def create_config( + self, + model_cls, + model_name, + task, + infer_steps, + seed, + cfg_scale, + cfg_scale2, + sample_shift, + height, + width, + duration, + attention_type, + denoising_steps="", + resize_mode="adaptive", + fixed_area="720p", + segment_length=81, + prev_frame_length=5, + use_tiny_vae=False, + ): + """Create basic inference configuration.""" + builder = InferenceConfigBuilder() + + config = builder.build( + model_cls=model_cls, + model_name=model_name, + task=task, + infer_steps=infer_steps, + seed=seed, + cfg_scale=cfg_scale, + cfg_scale2=cfg_scale2, + sample_shift=sample_shift, + height=height, + width=width, + duration=duration, + attention_type=attention_type, + denoising_steps=denoising_steps, + resize_mode=resize_mode, + fixed_area=fixed_area, + segment_length=segment_length, + prev_frame_length=prev_frame_length, + use_tiny_vae=use_tiny_vae, + ) + + return (config.to_dict(),) + + +class LightX2VTeaCache: + """TeaCache configuration node.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "enable": ( + "BOOLEAN", + {"default": False, "tooltip": "Enable TeaCache feature caching"}, + ), + "threshold": ( + "FLOAT", + { + "default": 0.26, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "Cache threshold, lower values provide more speedup: 0.1 ~2x speedup, 0.2 ~3x speedup", + }, + ), + "use_ret_steps": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Only cache key steps to balance quality and speed", + }, + ), + } + } + + RETURN_TYPES = ("TEACACHE_CONFIG",) + RETURN_NAMES = ("teacache_config",) + FUNCTION = "create_config" + CATEGORY = "LightX2V/Config" + + def create_config(self, enable, threshold, use_ret_steps): + """Create TeaCache configuration.""" + config = TeaCacheConfig(enable=enable, threshold=threshold, use_ret_steps=use_ret_steps) + return (config.to_dict(),) + + +class LightX2VQuantization: + @classmethod + def INPUT_TYPES(cls): + available_ops = get_available_quant_ops() + quant_backends = [] + + for op_name, is_available in available_ops: + if is_available: + quant_backends.append(op_name) + + common_schema = ["fp8", "int8"] + supported_quant_schemes = ["Default"] + for schema in common_schema: + for backend in quant_backends: + supported_quant_schemes.append(f"{schema}-{backend}") + + return { + "required": { + "dit_quant_scheme": ( + supported_quant_schemes, + { + "default": supported_quant_schemes[0], + "tooltip": "DIT model quantization precision", + }, + ), + "t5_quant_scheme": ( + supported_quant_schemes, + { + "default": supported_quant_schemes[0], + "tooltip": "T5 encoder quantization precision", + }, + ), + "clip_quant_scheme": ( + supported_quant_schemes, + { + "default": supported_quant_schemes[0], + "tooltip": "CLIP encoder quantization precision", + }, + ), + "adapter_quant_scheme": ( + supported_quant_schemes, + { + "default": supported_quant_schemes[0], + "tooltip": "Adapter quantization precision", + }, + ), + } + } + + RETURN_TYPES = ("QUANT_CONFIG",) + RETURN_NAMES = ("quantization_config",) + FUNCTION = "create_config" + CATEGORY = "LightX2V/Config" + + def create_config( + self, + dit_quant_scheme, + t5_quant_scheme, + clip_quant_scheme, + adapter_quant_scheme, + ): + """Create quantization configuration.""" + config = QuantizationConfig( + dit_quant_scheme=dit_quant_scheme, + t5_quant_scheme=t5_quant_scheme, + clip_quant_scheme=clip_quant_scheme, + adapter_quant_scheme=adapter_quant_scheme, + ) + return (config.to_dict(),) + + +class LightX2VMemoryOptimization: + """Memory optimization configuration node.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "enable_rotary_chunk": ( + "BOOLEAN", + {"default": False, "tooltip": "Enable rotary encoding chunking"}, + ), + "rotary_chunk_size": ( + "INT", + {"default": 100, "min": 100, "max": 10000, "step": 100}, + ), + "clean_cuda_cache": ( + "BOOLEAN", + {"default": False, "tooltip": "Clean CUDA cache promptly"}, + ), + "cpu_offload": ( + "BOOLEAN", + {"default": True, "tooltip": "Enable CPU offloading"}, + ), + "offload_granularity": ( + ["block", "phase", "model"], + {"default": "block", "tooltip": "Offload granularity"}, + ), + "offload_ratio": ( + "FLOAT", + {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.1}, + ), + "t5_cpu_offload": ( + "BOOLEAN", + {"default": True, "tooltip": "Enable T5 CPU offloading"}, + ), + "t5_offload_granularity": ( + ["model", "block"], + {"default": "model", "tooltip": "T5 offload granularity"}, + ), + "audio_encoder_cpu_offload": ( + "BOOLEAN", + {"default": True, "tooltip": "Enable audio encoder CPU offloading"}, + ), + "audio_adapter_cpu_offload": ( + "BOOLEAN", + {"default": True, "tooltip": "Enable audio adapter CPU offloading"}, + ), + "vae_cpu_offload": ( + "BOOLEAN", + {"default": True, "tooltip": "Enable VAE CPU offloading"}, + ), + "use_tiling_vae": ( + "BOOLEAN", + {"default": True, "tooltip": "Enable VAE tiling inference"}, + ), + "lazy_load": ( + "BOOLEAN", + {"default": False, "tooltip": "Lazy load model"}, + ), + "unload_after_inference": ( + "BOOLEAN", + {"default": False, "tooltip": "Unload modules after inference"}, + ), + }, + } + + RETURN_TYPES = ("MEMORY_CONFIG",) + RETURN_NAMES = ("memory_config",) + FUNCTION = "create_config" + CATEGORY = "LightX2V/Config" + + def create_config( + self, + enable_rotary_chunk=False, + rotary_chunk_size=100, + clean_cuda_cache=False, + cpu_offload=False, + offload_granularity="phase", + offload_ratio=1.0, + t5_cpu_offload=True, + t5_offload_granularity="model", + audio_encoder_cpu_offload=False, + audio_adapter_cpu_offload=False, + vae_cpu_offload=False, + use_tiling_vae=False, + lazy_load=False, + unload_after_inference=False, + ): + """Create memory optimization configuration.""" + config = MemoryOptimizationConfig( + enable_rotary_chunk=enable_rotary_chunk, + rotary_chunk_size=rotary_chunk_size, + clean_cuda_cache=clean_cuda_cache, + cpu_offload=cpu_offload, + offload_granularity=offload_granularity, + offload_ratio=offload_ratio, + t5_cpu_offload=t5_cpu_offload, + t5_offload_granularity=t5_offload_granularity, + audio_encoder_cpu_offload=audio_encoder_cpu_offload, + audio_adapter_cpu_offload=audio_adapter_cpu_offload, + vae_cpu_offload=vae_cpu_offload, + use_tiling_vae=use_tiling_vae, + lazy_load=lazy_load, + unload_after_inference=unload_after_inference, + ) + return (config.to_dict(),) diff --git a/nodes/inference.py b/nodes/inference.py new file mode 100644 index 0000000..2ca9267 --- /dev/null +++ b/nodes/inference.py @@ -0,0 +1,148 @@ +"""Modular inference runner that consumes a PREPARED_CONFIG.""" + +import gc +import logging + +import torch +from comfy.utils import ProgressBar + +from ..config_builder import ConfigBuilder +from ..lightx2v.lightx2v.infer import init_runner +from ..lightx2v.lightx2v.utils.input_info import init_empty_input_info, update_input_info_from_dict +from ..lightx2v.lightx2v.utils.set_config import set_config + + +class LightX2VModularInferenceV2: + """Pure inference node that takes prepared config and runs inference.""" + + _current_runner = None + _current_config_hash = None + + def __init__(self): + if not hasattr(self.__class__, "_current_runner"): + self.__class__._current_runner = None + if not hasattr(self.__class__, "_current_config_hash"): + self.__class__._current_config_hash = None + + self.config_builder = ConfigBuilder() + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "prepared_config": ( + "PREPARED_CONFIG", + {"tooltip": "Fully prepared configuration from ConfigCombinerV2"}, + ), + }, + } + + RETURN_TYPES = ("IMAGE", "AUDIO") + RETURN_NAMES = ("images", "audio") + FUNCTION = "generate" + CATEGORY = "LightX2V/InferenceV2" + + def _get_config_hash(self, config) -> str: + """Get hash of configuration to detect changes.""" + return self.config_builder.get_config_hash(config) + + def _build_rs2v_shot_config(self, config): + from ..lightx2v.lightx2v.shot_runner.shot_base import load_clip_configs + from ..lightx2v.lightx2v.utils.lockable_dict import LockableDict + + config_json = config.get("config_json") + if config_json: + main_cfg = config_json + elif config.get("clip_configs"): + main_cfg = config + else: + main_cfg = { + "lightx2v_path": "", + "clip_configs": [ + { + "name": "rs2v_clip", + "config": LockableDict(config), + } + ], + } + if "task" not in main_cfg["clip_configs"][0]["config"]: + main_cfg["clip_configs"][0]["config"]["task"] = "rs2v" + + if isinstance(main_cfg, dict) and "lightx2v_path" not in main_cfg: + main_cfg = dict(main_cfg) + main_cfg["lightx2v_path"] = "" + + return load_clip_configs(main_cfg) + + def generate(self, prepared_config): + """Run inference with prepared configuration.""" + + config = prepared_config + + try: + config_hash = self._get_config_hash(config) + + current_runner = getattr(self.__class__, "_current_runner", None) + current_config_hash = getattr(self.__class__, "_current_config_hash", None) + + needs_reinit = current_runner is None or current_config_hash != config_hash or getattr(config, "lazy_load", False) + + logging.info(f"Needs reinit: {needs_reinit}, old config hash: {current_config_hash}, new config hash: {config_hash}") + if needs_reinit: + if current_runner is not None: + del self.__class__._current_runner + torch.cuda.empty_cache() + gc.collect() + if config.get("task") == "rs2v": + from ..lightx2v.lightx2v.shot_runner.rs2v_infer import ShotRS2VPipeline + + shot_cfg = self._build_rs2v_shot_config(config) + self.__class__._current_runner = ShotRS2VPipeline(shot_cfg) + else: + formatted_config = set_config(config) + self.__class__._current_runner = init_runner(formatted_config) + self.__class__._current_config_hash = config_hash + + progress = ProgressBar(100) + + def update_progress(current_step, _total): + progress.update_absolute(current_step) + + current_runner = self.__class__._current_runner + + if hasattr(current_runner, "set_progress_callback"): + current_runner.set_progress_callback(update_progress) + + config["return_result_tensor"] = True + config["save_result_path"] = "" + config["negative_prompt"] = config.get("negative_prompt", "") + if config.get("task") == "rs2v": + result_dict = current_runner.run_pipeline(config) + else: + input_data = init_empty_input_info(config.task) + update_input_info_from_dict(input_data, config) + current_runner.set_config(config) + result_dict = current_runner.run_pipeline(input_data) + + images = result_dict.get("video", None) + audio = result_dict.get("audio", None) + + if images is not None and images.numel() > 0: + images = images.cpu() + if images.dtype != torch.float32: + images = images.float() + + if getattr(config, "unload_after_inference", False): + if hasattr(self.__class__, "_current_runner"): + del self.__class__._current_runner + self.__class__._current_runner = None + self.__class__._current_config_hash = None + + torch.cuda.empty_cache() + gc.collect() + + return (images, audio) + + except Exception as e: + logging.error(f"Error during inference: {e}") + raise diff --git a/nodes/lora.py b/nodes/lora.py new file mode 100644 index 0000000..3a6f4fd --- /dev/null +++ b/nodes/lora.py @@ -0,0 +1,48 @@ +"""LoRA chain loader node.""" + +from ..config_builder import LoRAChainBuilder +from ..model_utils import scan_loras + + +class LightX2VLoRALoader: + @classmethod + def INPUT_TYPES(cls): + available_loras = scan_loras() + + return { + "required": { + "lora_name": ( + available_loras, + { + "default": available_loras[0], + "tooltip": "Select LoRA from available LoRAs", + }, + ), + "strength": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 2.0, + "step": 0.1, + "tooltip": "LoRA strength", + }, + ), + }, + "optional": { + "lora_chain": ( + "LORA_CHAIN", + {"tooltip": "Previous LoRA chain to append to"}, + ), + }, + } + + RETURN_TYPES = ("LORA_CHAIN",) + RETURN_NAMES = ("lora_chain",) + FUNCTION = "load_lora" + CATEGORY = "LightX2V/LoRA" + + def load_lora(self, lora_name, strength, lora_chain=None): + """Load and chain LoRA configurations.""" + chain = LoRAChainBuilder.build_chain(lora_name=lora_name, strength=strength, existing_chain=lora_chain) + return (chain,) diff --git a/nodes/seedvr.py b/nodes/seedvr.py new file mode 100644 index 0000000..52b6593 --- /dev/null +++ b/nodes/seedvr.py @@ -0,0 +1,339 @@ +"""SeedVR2 super-resolution node.""" + +import gc +import hashlib +import logging + +import torch +from comfy.utils import ProgressBar + +from ..model_utils import get_model_base_path, get_model_full_path, scan_models + + +class LightX2VSeedVRSR: + """SeedVR2 video/image super-resolution node for ComfyUI. + + Wraps the SeedVR2-3B model via LightX2V's SeedVRRunner to perform + single-pass diffusion super-resolution on a video (mp4) or a single image. + """ + + _current_runner = None + _current_config_hash = None + + @classmethod + def INPUT_TYPES(cls): + available_models = scan_models() + return { + "required": { + "model_name": ( + available_models, + { + "default": available_models[0] if available_models else "None", + "tooltip": "SeedVR2 model directory under models/lightx2v/", + }, + ), + "input_type": ( + ["video", "image"], + {"default": "video", "tooltip": "Whether to SR a video file or a single image"}, + ), + "input_path": ( + "STRING", + { + "default": "", + "tooltip": "Absolute path to input .mp4 (for video) or .png/.jpg (for image). For video, also accepts a directory of frames.", + }, + ), + "sr_ratio": ( + "FLOAT", + { + "default": 2.0, + "min": 1.0, + "max": 8.0, + "step": 0.5, + "tooltip": "Super-resolution ratio (e.g. 2.0 = 2x, 4.0 = 4x)", + }, + ), + "target_height": ( + "INT", + { + "default": 720, + "min": 64, + "max": 4096, + "step": 8, + "tooltip": "Output frame height (SeedVR NaDiT processes at native resolution)", + }, + ), + "target_width": ( + "INT", + { + "default": 1280, + "min": 64, + "max": 4096, + "step": 8, + "tooltip": "Output frame width (must be divisible by 16 for VAE)", + }, + ), + "fps": ( + "FLOAT", + { + "default": 16.0, + "min": 1.0, + "max": 60.0, + "step": 0.5, + "tooltip": "Output FPS for video SR (input video FPS is preserved if available)", + }, + ), + "segment_length": ( + "INT", + { + "default": 81, + "min": 16, + "max": 256, + "step": 1, + "tooltip": "Frames per segment for long video SR (1-step diffusion per segment)", + }, + ), + "segment_overlap": ( + "INT", + { + "default": 1, + "min": 0, + "max": 32, + "step": 1, + "tooltip": "Overlap frames between segments to prevent seams", + }, + ), + "seed": ( + "INT", + { + "default": 42, + "min": -1, + "max": 2**32 - 1, + "tooltip": "Random seed, -1 for random", + }, + ), + "prompt": ( + "STRING", + { + "default": "", + "multiline": True, + "tooltip": "Optional text prompt to guide detail synthesis (SeedVR uses pre-computed embeddings; prompt mostly affects style)", + }, + ), + "negative_prompt": ( + "STRING", + { + "default": "", + "multiline": True, + "tooltip": "Negative prompt for guidance", + }, + ), + "save_output": ( + "BOOLEAN", + { + "default": False, + "tooltip": "If True, also write the SR result to disk in addition to returning IMAGE tensor", + }, + ), + "output_path": ( + "STRING", + { + "default": "", + "tooltip": "Where to save (only used if save_output=True). Leave empty for auto-generated name next to input.", + }, + ), + "unload_after_inference": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Unload SeedVR runner from VRAM after inference (frees ~6GB+ for other nodes)", + }, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "run_seedvr" + CATEGORY = "LightX2V/SeedVR" + + def _config_hash( + self, + model_name, + input_type, + input_path, + sr_ratio, + target_height, + target_width, + fps, + segment_length, + segment_overlap, + seed, + prompt, + negative_prompt, + save_output, + output_path, + ): + """Hash of all parameters that should trigger runner reinit.""" + raw = ( + f"{model_name}|{input_type}|{input_path}|{sr_ratio}|" + f"{target_height}|{target_width}|{fps}|" + f"{segment_length}|{segment_overlap}|{seed}|" + f"{prompt}|{negative_prompt}|{save_output}|{output_path}" + ) + return hashlib.md5(raw.encode("utf-8")).hexdigest() + + def run_seedvr( + self, + model_name, + input_type, + input_path, + sr_ratio, + target_height, + target_width, + fps, + segment_length, + segment_overlap, + seed, + prompt, + negative_prompt, + save_output, + output_path, + unload_after_inference, + ): + """Run SeedVR2 super-resolution and return IMAGE tensor.""" + from ..lightx2v.lightx2v.infer import init_runner + from ..lightx2v.lightx2v.utils.input_info import ( + init_empty_input_info, + update_input_info_from_dict, + ) + from ..lightx2v.lightx2v.utils.set_config import set_config + + if not model_name or model_name == "None": + raise ValueError("model_name is required — select a SeedVR2 model directory under models/lightx2v/") + if not input_path: + raise ValueError("input_path is required — provide an absolute path to a video or image file") + + model_full_path = get_model_full_path(model_name) + if not model_full_path: + raise FileNotFoundError( + f"Model '{model_name}' not found under models/lightx2v/. Expected directory: {get_model_base_path() / model_name}" + ) + + cfg_hash = self._config_hash( + model_name, + input_type, + input_path, + sr_ratio, + target_height, + target_width, + fps, + segment_length, + segment_overlap, + seed, + prompt, + negative_prompt, + save_output, + output_path, + ) + + try: + needs_reinit = ( + getattr(self.__class__, "_current_runner", None) is None or getattr(self.__class__, "_current_config_hash", None) != cfg_hash + ) + + if needs_reinit: + if getattr(self.__class__, "_current_runner", None) is not None: + del self.__class__._current_runner + torch.cuda.empty_cache() + gc.collect() + + config = { + "model_cls": "seedvr2", + "task": "sr", + "model_path": model_full_path, + "sr_ratio": float(sr_ratio), + "target_height": int(target_height), + "target_width": int(target_width), + "target_video_length": int(segment_length), + "sr_segment_length": int(segment_length), + "sr_overlap": int(segment_overlap), + "fps": float(fps), + "infer_steps": 1, + "seed": int(seed), + "prompt": prompt, + "negative_prompt": negative_prompt, + } + + formatted_config = set_config(config) + self.__class__._current_runner = init_runner(formatted_config) + self.__class__._current_config_hash = cfg_hash + + runner = self.__class__._current_runner + + progress = ProgressBar(100) + + def _update_progress(current_step, _total): + progress.update_absolute(current_step) + + if hasattr(runner, "set_progress_callback"): + runner.set_progress_callback(_update_progress) + + input_info = init_empty_input_info("sr") + update_input_info_from_dict( + input_info, + { + "video_path": input_path if input_type == "video" else "", + "image_path": input_path if input_type == "image" else "", + "prompt": prompt, + "negative_prompt": negative_prompt, + "seed": int(seed), + "save_result_path": output_path if (save_output and output_path) else "", + "return_result_tensor": True, + }, + ) + + runner.set_config( + { + "video_path": input_path if input_type == "video" else "", + "image_path": input_path if input_type == "image" else "", + "prompt": prompt, + "negative_prompt": negative_prompt, + "seed": int(seed), + "save_result_path": output_path if (save_output and output_path) else "", + "return_result_tensor": True, + } + ) + + result_dict = runner.run_pipeline(input_info) + images = result_dict.get("video", None) + + if images is None or images.numel() == 0: + raise RuntimeError("SeedVR returned empty result") + + images = images.cpu() + if images.dtype != torch.float32: + images = images.float() + + if images.dim() == 4 and images.shape[0] > 0: + images = images[0] + + if unload_after_inference: + if hasattr(self.__class__, "_current_runner"): + del self.__class__._current_runner + self.__class__._current_runner = None + self.__class__._current_config_hash = None + + torch.cuda.empty_cache() + gc.collect() + + return (images,) + + except Exception as e: + logging.error(f"SeedVR SR failed: {e}") + if unload_after_inference: + if hasattr(self.__class__, "_current_runner"): + del self.__class__._current_runner + self.__class__._current_runner = None + self.__class__._current_config_hash = None + raise diff --git a/nodes/talk.py b/nodes/talk.py new file mode 100644 index 0000000..6f2ee4c --- /dev/null +++ b/nodes/talk.py @@ -0,0 +1,147 @@ +"""Talk-object input and combiner nodes (for multi-speaker audio-driven generation).""" + +from ..config_builder import TalkObjectConfigBuilder +from ..data_models import TalkObjectsConfig + + +class TalkObjectInput: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "name": ( + "STRING", + {"default": "person_1", "tooltip": "speaker name identifier"}, + ), + }, + "optional": { + "audio": ("AUDIO", {"tooltip": "uploaded audio file"}), + "mask": ("MASK", {"tooltip": "uploaded mask image (optional)"}), + "save_to_input": ( + "BOOLEAN", + {"default": True, "tooltip": "save to input folder"}, + ), + }, + } + + RETURN_TYPES = ("TALK_OBJECT",) + RETURN_NAMES = ("talk_object",) + FUNCTION = "create_talk_object" + CATEGORY = "LightX2V/Audio" + + def create_talk_object(self, name, audio=None, mask=None, save_to_input=True): + """Create a talk object from input data.""" + builder = TalkObjectConfigBuilder() + + talk_object = builder.build_from_input(name=name, audio=audio, mask=mask, save_to_input=save_to_input) + + if talk_object: + return (talk_object,) + return (None,) + + +class TalkObjectsCombiner: + PREDEFINED_SLOTS = 16 + + @classmethod + def INPUT_TYPES(cls): + inputs = {"required": {}, "optional": {}} + + for i in range(cls.PREDEFINED_SLOTS): + inputs["optional"][f"talk_object_{i + 1}"] = ( + "TALK_OBJECT", + {"tooltip": f"talk object {i + 1}"}, + ) + + return inputs + + RETURN_TYPES = ("TALK_OBJECTS_CONFIG",) + RETURN_NAMES = ("talk_objects_config",) + FUNCTION = "combine_talk_objects" + CATEGORY = "LightX2V/Audio" + + def combine_talk_objects(self, **kwargs): + config = TalkObjectsConfig() + + for i in range(self.PREDEFINED_SLOTS): + talk_obj = kwargs.get(f"talk_object_{i + 1}") + + if talk_obj is not None: + config.add_object(talk_obj) + + if not config.talk_objects: + return (None,) + + return (config,) + + +class TalkObjectsFromJSON: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "json_config": ( + "STRING", + { + "multiline": True, + "default": '[{"name": "person1", "audio": "/path/to/audio1.wav", "mask": "/path/to/mask1.png"}]', + "tooltip": "JSON format talk objects configuration", + }, + ), + }, + } + + RETURN_TYPES = ("TALK_OBJECTS_CONFIG",) + RETURN_NAMES = ("talk_objects_config",) + FUNCTION = "parse_json_config" + CATEGORY = "LightX2V/Audio" + + def parse_json_config(self, json_config): + builder = TalkObjectConfigBuilder() + talk_objects_config = builder.build_from_json(json_config) + return (talk_objects_config,) + + +class TalkObjectsFromFiles: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "audio_files": ( + "STRING", + { + "multiline": True, + "default": "audio1.wav\naudio2.wav", + "tooltip": "audio file list (one per line)", + }, + ), + }, + "optional": { + "mask_files": ( + "STRING", + { + "multiline": True, + "default": "mask1.png\nmask2.png", + "tooltip": "mask file list (one per line, optional)", + }, + ), + "names": ( + "STRING", + { + "multiline": True, + "default": "person1\nperson2", + "tooltip": "talk object name list (one per line, optional)", + }, + ), + }, + } + + RETURN_TYPES = ("TALK_OBJECTS_CONFIG",) + RETURN_NAMES = ("talk_objects_config",) + FUNCTION = "build_from_files" + CATEGORY = "LightX2V/Audio" + + def build_from_files(self, audio_files, mask_files="", names=""): + builder = TalkObjectConfigBuilder() + talk_objects_config = builder.build_from_files(audio_files, mask_files, names) + return (talk_objects_config,)