From 04170f4ae741549e52bbfab4ae8bc7a32d18b1cd Mon Sep 17 00:00:00 2001 From: gaclove Date: Tue, 15 Jul 2025 20:30:34 +0800 Subject: [PATCH] refactor: streamline LightX2V module by removing unused files and consolidating configuration management into bridge.py for improved maintainability and clarity --- .ruff.toml | 4 +- __init__.py | 10 +- bridge.py | 450 +++++++++++++ lightx2v_nodes/__init__.py | 44 -- lightx2v_nodes/config.py | 232 ------- lightx2v_nodes/factory.py | 237 ------- lightx2v_nodes/models.py | 179 ----- lightx2v_nodes/nodes.py | 855 ----------------------- lightx2v_nodes/universal_bridge.py | 1008 ---------------------------- nodes.py | 431 +++++++++++- 10 files changed, 876 insertions(+), 2574 deletions(-) create mode 100644 bridge.py delete mode 100644 lightx2v_nodes/__init__.py delete mode 100644 lightx2v_nodes/config.py delete mode 100644 lightx2v_nodes/factory.py delete mode 100644 lightx2v_nodes/models.py delete mode 100644 lightx2v_nodes/nodes.py delete mode 100644 lightx2v_nodes/universal_bridge.py diff --git a/.ruff.toml b/.ruff.toml index 3612370..2e4c07b 100644 --- a/.ruff.toml +++ b/.ruff.toml @@ -1,2 +1,4 @@ line-length = 150 -indent-width = 4 \ No newline at end of file +indent-width = 4 + +extend-select = ["I"] diff --git a/__init__.py b/__init__.py index 3eae80c..39a8c6b 100644 --- a/__init__.py +++ b/__init__.py @@ -1,11 +1,3 @@ -import sys -import os -from pathlib import Path - -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 - -from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS # noqa: E402 +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/bridge.py b/bridge.py new file mode 100644 index 0000000..b859af5 --- /dev/null +++ b/bridge.py @@ -0,0 +1,450 @@ +"""Modular configuration system for LightX2V ComfyUI integration.""" + +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 is_fp8_supported_gpu(): + if not torch.cuda.is_available(): + return False + compute_capability = torch.cuda.get_device_capability(0) + major, minor = compute_capability + return (major == 8 and minor == 9) or (major >= 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_quant_ops(): + available_ops = [] + + vllm_installed = is_module_installed("vllm") + if vllm_installed: + available_ops.append(("vllm", True)) + else: + available_ops.append(("vllm", False)) + + sgl_installed = is_module_installed("sgl_kernel") + if sgl_installed: + available_ops.append(("sgl", True)) + else: + available_ops.append(("sgl", False)) + + q8f_installed = is_module_installed("q8_kernels") + if q8f_installed: + available_ops.append(("q8f", True)) + else: + available_ops.append(("q8f", False)) + + return available_ops + + +def get_available_attn_ops(): + available_ops = [] + + vllm_installed = is_module_installed("flash_attn") + if vllm_installed: + available_ops.append(("flash_attn2", True)) + else: + available_ops.append(("flash_attn2", False)) + + sgl_installed = is_module_installed("flash_attn_interface") + if sgl_installed: + available_ops.append(("flash_attn3", True)) + else: + available_ops.append(("flash_attn3", False)) + + q8f_installed = is_module_installed("sageattention") + if q8f_installed: + available_ops.append(("sage_attn2", True)) + else: + available_ops.append(("sage_attn2", False)) + + torch_installed = is_module_installed("torch") + if torch_installed: + available_ops.append(("torch_sdpa", True)) + else: + available_ops.append(("torch_sdpa", False)) + + return available_ops + + +class LightX2VDefaultConfig: + """Central default configuration for LightX2V.""" + + DEFAULT_CONFIG = { + # ========== Model Configuration ========== + "model_cls": "wan2.1", + "model_path": "", + "task": "t2v", + "mode": "infer", + # ========== 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": 480, + "target_width": 832, + "target_video_length": 81, + "fps": 16, + "vae_stride": [4, 8, 8], + "patch_size": [1, 2, 2], + # ========== Feature Caching (TeaCache) ========== + "feature_caching": "NoCaching", + "enable_teacache": False, + "teacache_thresh": 0.26, + "coefficients": None, # Auto-calculated + "use_ret_steps": False, + # ========== Quantization ========== + "dit_quant_scheme": "bf16", + "t5_quant_scheme": "bf16", + "clip_quant_scheme": "fp16", + "quant_op": "vllm", + "precision_mode": "fp32", + "dit_quantized_ckpt": None, + "t5_quantized_ckpt": None, + "clip_quantized_ckpt": None, + "mm_config": {"mm_type": "Default"}, + # ========== GPU Memory Optimization ========== + "rotary_chunk": False, + "rotary_chunk_size": 100, + "clean_cuda_cache": False, + "torch_compile": False, + "attention_type": "flash_attn3", + "self_attn_1_type": "flash_attn3", + "cross_attn_1_type": "flash_attn3", + "cross_attn_2_type": "flash_attn3", + # ========== Async Offloading ========== + "cpu_offload": False, + "offload_granularity": "phase", + "offload_ratio": 1.0, + "t5_cpu_offload": False, + "t5_offload_granularity": "model", + "lazy_load": False, + "unload_modules": False, + # ========== Lightweight VAE ========== + "use_tiny_vae": False, + "tiny_vae": False, + "tiny_vae_path": None, + "use_tiling_vae": False, + # ========== Other Settings ========== + "lora_path": None, + "strength_model": 1.0, + "do_mm_calib": False, + "parallel_attn_type": None, + "parallel_vae": False, + "max_area": False, + "use_prompt_enhancer": False, + "text_len": 512, + } + + +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] + return None + + +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 + + @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() + + available = [] + for op_name, is_available in self._available_attn_ops: + if is_available: + available.append(op_name) + + # Always include fallback + if "torch_sdpa" not in available: + available.append("torch_sdpa") + + return available + + @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() + + available = [] + for op_name, is_available in self._available_quant_ops: + if is_available: + available.append(op_name) + + return available + + def apply_inference_config(self, config: Dict[str, Any]) -> Dict[str, Any]: + """Apply basic inference configuration.""" + updates = {} + + # Model settings + if "model_cls" in config: + updates["model_cls"] = config["model_cls"] + if "model_path" in config: + updates["model_path"] = config["model_path"] + if "task" in config: + updates["task"] = config["task"] + + # Inference parameters + if "infer_steps" in config: + updates["infer_steps"] = config["infer_steps"] + if "seed" in config and config["seed"] != -1: + updates["seed"] = config["seed"] + if "cfg_scale" in config: + updates["sample_guide_scale"] = config["cfg_scale"] + updates["enable_cfg"] = config["cfg_scale"] != 1.0 + if "sample_shift" in config: + updates["sample_shift"] = config["sample_shift"] + + # Video parameters + if "height" in config: + updates["target_height"] = config["height"] + if "width" in config: + updates["target_width"] = config["width"] + if "video_length" in config: + updates["target_video_length"] = config["video_length"] + if "fps" in config: + updates["fps"] = config["fps"] + + 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["enable_teacache"] = True + updates["teacache_thresh"] = config.get("threshold", 0.26) + updates["use_ret_steps"] = config.get("cache_key_steps_only", False) + + # Auto-calculate coefficients + 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"]) + if coeffs: + updates["coefficients"] = coeffs + else: + updates["feature_caching"] = "NoCaching" + updates["enable_teacache"] = False + + return updates + + def apply_quantization_config(self, config: Dict[str, Any], model_path: str) -> Dict[str, Any]: + """Apply quantization configuration.""" + updates = {} + + # DIT quantization + dit_scheme = config.get("dit_precision", "bf16") + updates["dit_quant_scheme"] = dit_scheme + if dit_scheme != "bf16": + updates["dit_quantized_ckpt"] = os.path.join(model_path, dit_scheme) + + # T5 quantization + t5_scheme = config.get("t5_precision", "bf16") + updates["t5_quant_scheme"] = t5_scheme + updates["t5_quantized"] = t5_scheme != "bf16" + if t5_scheme != "bf16": + t5_path = os.path.join(model_path, t5_scheme) + updates["t5_quantized_ckpt"] = os.path.join(t5_path, f"models_t5_umt5-xxl-enc-{t5_scheme}.pth") + + # CLIP quantization + clip_scheme = config.get("clip_precision", "fp16") + updates["clip_quant_scheme"] = clip_scheme + updates["clip_quantized"] = clip_scheme != "fp16" + if clip_scheme != "fp16": + clip_path = os.path.join(model_path, clip_scheme) + updates["clip_quantized_ckpt"] = os.path.join(clip_path, f"clip-{clip_scheme}.pth") + + # Quantization backend + quant_backend = config.get("quant_backend", "vllm") + updates["quant_op"] = quant_backend + + # Determine mm_type based on quantization settings + if dit_scheme != "bf16": + if quant_backend == "vllm": + mm_type = f"W-{dit_scheme}-channel-sym-A-{dit_scheme}-channel-sym-dynamic-Vllm" + elif quant_backend == "sgl": + if dit_scheme == "int8": + mm_type = f"W-{dit_scheme}-channel-sym-A-{dit_scheme}-channel-sym-dynamic-Sgl-ActVllm" + else: + mm_type = f"W-{dit_scheme}-channel-sym-A-{dit_scheme}-channel-sym-dynamic-Sgl" + elif quant_backend == "q8f": + mm_type = f"W-{dit_scheme}-channel-sym-A-{dit_scheme}-channel-sym-dynamic-Q8F" + else: + mm_type = "Default" + + updates["mm_config"] = {"mm_type": mm_type} + else: + updates["mm_config"] = {"mm_type": "Default"} + + # Precision mode + updates["precision_mode"] = config.get("sensitive_layers_precision", "fp32") + + return updates + + def apply_memory_optimization(self, config: Dict[str, Any]) -> Dict[str, Any]: + """Apply memory optimization settings.""" + updates = {} + + level = config.get("optimization_level", "none") + + # GPU optimization + if config.get("enable_rotary_chunk", False) or level in ["high", "extreme"]: + updates["rotary_chunk"] = True + updates["rotary_chunk_size"] = config.get("rotary_chunk_size", 100) + + if config.get("clean_cuda_cache", False) or level == "extreme": + updates["clean_cuda_cache"] = True + + # CPU offloading + if config.get("enable_cpu_offload", False) or level in ["medium", "high", "extreme"]: + updates["cpu_offload"] = True + updates["offload_granularity"] = config.get("offload_granularity", "phase") + updates["offload_ratio"] = config.get("offload_ratio", 1.0) + + # T5 offloading + if level in ["high", "extreme"]: + updates["t5_cpu_offload"] = True + updates["t5_offload_granularity"] = "block" if level == "extreme" else "model" + + # Module management + if config.get("lazy_load", False) or level == "extreme": + updates["lazy_load"] = True + + if config.get("unload_after_inference", False) or level == "extreme": + updates["unload_modules"] = True + + # Attention type + attention_type = config.get("attention_type", "flash_attn3") + updates["attention_type"] = attention_type + updates["self_attn_1_type"] = attention_type + updates["cross_attn_1_type"] = attention_type + updates["cross_attn_2_type"] = attention_type + + return updates + + def apply_vae_config(self, config: Dict[str, Any], model_path: str) -> Dict[str, Any]: + """Apply VAE configuration.""" + updates = {} + + if config.get("use_tiny_vae", False): + updates["use_tiny_vae"] = True + updates["tiny_vae"] = True + updates["tiny_vae_path"] = os.path.join(model_path, "taew2_1.pth") + + if config.get("use_tiling_vae", False): + updates["use_tiling_vae"] = True + + return updates + + def build_final_config(self, configs: Dict[str, Dict[str, Any]]) -> EasyDict: + """Build final configuration from module configs.""" + final_config = copy.deepcopy(self.base_config) + + # Apply configurations in order + if "inference" in configs: + final_config.update(self.apply_inference_config(configs["inference"])) + + if "teacache" in configs: + teacache_updates = self.apply_teacache_config( + configs["teacache"], + final_config, # Pass current config for coefficient calculation + ) + final_config.update(teacache_updates) + + if "quantization" in configs: + model_path = final_config.get("model_path", "") + quant_updates = self.apply_quantization_config(configs["quantization"], model_path) + final_config.update(quant_updates) + + if "memory" in configs: + final_config.update(self.apply_memory_optimization(configs["memory"])) + + if "vae" in configs: + model_path = final_config.get("model_path", "") + final_config.update(self.apply_vae_config(configs["vae"], model_path)) + + # Load model config if exists + model_config_path = os.path.join(final_config["model_path"], "config.json") + if os.path.exists(model_config_path): + try: + with open(model_config_path, "r") as f: + model_config = json.load(f) + # Model config has lower priority than user configs + for key, value in model_config.items(): + if key not in final_config or final_config[key] is None: + final_config[key] = value + except Exception as e: + logging.warning(f"Failed to load model config: {e}") + + return EasyDict(final_config) diff --git a/lightx2v_nodes/__init__.py b/lightx2v_nodes/__init__.py deleted file mode 100644 index fca95a2..0000000 --- a/lightx2v_nodes/__init__.py +++ /dev/null @@ -1,44 +0,0 @@ -# Refactored LightX2V module -# from ..lightx2v.lightx2v.common.ops import * # noqa: F401, F403 for import global register - -from .config import LightX2VConfig -from .factory import LightX2VFactory -from .models import ( - LightX2VT5Encoder, - LightX2VClipVisionEncoder, - LightX2VVae, - LightX2VModel, -) -from .nodes import ( - Lightx2vWanVideoModelDir, - Lightx2vWanVideoT5EncoderLoader, - Lightx2vWanVideoT5Encoder, - Lightx2vWanVideoClipVisionEncoderLoader, - Lightx2vWanVideoVaeLoader, - Lightx2vWanVideoVaeDecoder, - Lightx2vWanVideoImageEncoder, - Lightx2vWanVideoEmptyEmbeds, - Lightx2vWanVideoModelLoader, - Lightx2vWanVideoSampler, - WanVideoTeaCache, -) - -__all__ = [ - "LightX2VConfig", - "LightX2VFactory", - "LightX2VT5Encoder", - "LightX2VClipVisionEncoder", - "LightX2VVae", - "LightX2VModel", - "Lightx2vWanVideoModelDir", - "Lightx2vWanVideoT5EncoderLoader", - "Lightx2vWanVideoT5Encoder", - "Lightx2vWanVideoClipVisionEncoderLoader", - "Lightx2vWanVideoVaeLoader", - "Lightx2vWanVideoVaeDecoder", - "Lightx2vWanVideoImageEncoder", - "Lightx2vWanVideoEmptyEmbeds", - "Lightx2vWanVideoModelLoader", - "Lightx2vWanVideoSampler", - "WanVideoTeaCache", -] diff --git a/lightx2v_nodes/config.py b/lightx2v_nodes/config.py deleted file mode 100644 index c3fe654..0000000 --- a/lightx2v_nodes/config.py +++ /dev/null @@ -1,232 +0,0 @@ -"""Configuration management for LightX2V.""" - -from dataclasses import dataclass, field -from typing import Optional, Dict, Any, List, Tuple -from pathlib import Path -import torch -from easydict import EasyDict -import importlib.util - - -@dataclass -class TeaCacheConfig: - """Configuration for TeaCache optimization.""" - - rel_l1_thresh: float = 0.26 - start_percent: float = 0.1 - end_percent: float = 1.0 - cache_device: torch.device = field(default_factory=lambda: torch.device("cpu")) - coefficients: List[List[float]] = field(default_factory=list) - use_ret_steps: bool = False - mode: str = "e" - - -@dataclass -class VideoConfig: - """Video generation configuration.""" - - target_width: int = 832 - target_height: int = 480 - target_video_length: int = 81 - vae_stride: Tuple[int, int, int] = (4, 8, 8) - patch_size: Tuple[int, int, int] = (1, 2, 2) - - @property - def max_area(self) -> int: - return self.target_height * self.target_width - - -@dataclass -class ModelConfig: - """Model loading and inference configuration.""" - - model_path: Path - model_type: str = "i2v" # "t2v" or "i2v" - precision: str = "bf16" # "bf16", "fp16", "fp32" - device: str = "cuda" - attention_type: str = "flash_attn3" - cpu_offload: bool = False - offload_granularity: str = "phase" # "block" or "phase" - - # Optional configurations - lora_path: Optional[Path] = None - lora_strength: float = 1.0 - mm_config: Dict[str, Any] = field(default_factory=dict) - - # Inference settings - steps: int = 20 - shift: float = 5.0 - cfg_scale: float = 5.0 - seed: int = 42 - feature_caching: str = "NoCaching" - - def to_dtype(self) -> torch.dtype: - """Convert precision string to torch dtype.""" - dtype_map = { - "bf16": torch.bfloat16, - "fp16": torch.float16, - "fp32": torch.float32, - } - return dtype_map[self.precision] - - def to_device(self) -> torch.device: - """Get torch device.""" - if self.device == "cuda": - return torch.device("cuda") - return torch.device("cpu") - - -@dataclass -class EncoderConfig: - """Encoder configuration.""" - - model_path: Path - dtype: torch.dtype - device: torch.device - - # T5 specific - text_len: int = 512 - tokenizer_path: Optional[Path] = None - cpu_offload: bool = False - - # CLIP specific - clip_quantized: bool = False - clip_quantized_ckpt: Optional[Path] = None - quant_scheme: Optional[str] = None - - # VAE specific - z_dim: int = 16 - parallel: bool = False - - -@dataclass -class LightX2VConfig: - """Main configuration container for LightX2V.""" - - model: ModelConfig - video: VideoConfig - teacache: Optional[TeaCacheConfig] = None - - @classmethod - def from_dict(cls, config_dict: Dict[str, Any]) -> "LightX2VConfig": - """Create configuration from dictionary.""" - model_config = ModelConfig(**config_dict.get("model", {})) - video_config = VideoConfig(**config_dict.get("video", {})) - - teacache_config = None - if "teacache" in config_dict: - teacache_config = TeaCacheConfig(**config_dict["teacache"]) - - return cls(model=model_config, video=video_config, teacache=teacache_config) - - def to_easydict(self) -> EasyDict: - """Convert to EasyDict for legacy compatibility.""" - config_dict = { - "model_path": str(self.model.model_path), - "task": self.model.model_type, - "dtype": self.model.to_dtype(), - "device": self.model.to_device(), - "attention_type": self.model.attention_type, - "cpu_offload": self.model.cpu_offload, - "offload_granularity": self.model.offload_granularity, - "target_height": self.video.target_height, - "target_width": self.video.target_width, - "target_video_length": self.video.target_video_length, - "vae_stride": self.video.vae_stride, - "patch_size": self.video.patch_size, - "infer_steps": self.model.steps, - "sample_shift": self.model.shift, - "sample_guide_scale": self.model.cfg_scale, - "seed": self.model.seed, - "enable_cfg": self.model.cfg_scale != 1.0, - "mm_config": self.model.mm_config, - "feature_caching": self.model.feature_caching, - } - - if self.teacache: - config_dict.update( - { - "feature_caching": "Tea", - "teacache_thresh": self.teacache.rel_l1_thresh, - "use_ret_steps": self.teacache.use_ret_steps, - "coefficients": self.teacache.coefficients, - } - ) - else: - config_dict["feature_caching"] = "NoCaching" - - if self.model.lora_path: - config_dict["lora_path"] = str(self.model.lora_path) - config_dict["strength_model"] = self.model.lora_strength - - return EasyDict(config_dict) - - -def is_fp8_supported_gpu(): - if not torch.cuda.is_available(): - return False - compute_capability = torch.cuda.get_device_capability(0) - major, minor = compute_capability - return (major == 8 and minor == 9) or (major >= 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_quant_ops(): - available_ops = [] - - vllm_installed = is_module_installed("vllm") - if vllm_installed: - available_ops.append(("vllm", True)) - else: - available_ops.append(("vllm", False)) - - sgl_installed = is_module_installed("sgl_kernel") - if sgl_installed: - available_ops.append(("sgl", True)) - else: - available_ops.append(("sgl", False)) - - q8f_installed = is_module_installed("q8_kernels") - if q8f_installed: - available_ops.append(("q8f", True)) - else: - available_ops.append(("q8f", False)) - - return available_ops - - -def get_available_attn_ops(): - available_ops = [] - - vllm_installed = is_module_installed("flash_attn") - if vllm_installed: - available_ops.append(("flash_attn2", True)) - else: - available_ops.append(("flash_attn2", False)) - - sgl_installed = is_module_installed("flash_attn_interface") - if sgl_installed: - available_ops.append(("flash_attn3", True)) - else: - available_ops.append(("flash_attn3", False)) - - q8f_installed = is_module_installed("sageattention") - if q8f_installed: - available_ops.append(("sage_attn2", True)) - else: - available_ops.append(("sage_attn2", False)) - - torch_installed = is_module_installed("torch") - if torch_installed: - available_ops.append(("torch_sdpa", True)) - else: - available_ops.append(("torch_sdpa", False)) - - return available_ops diff --git a/lightx2v_nodes/factory.py b/lightx2v_nodes/factory.py deleted file mode 100644 index e061d92..0000000 --- a/lightx2v_nodes/factory.py +++ /dev/null @@ -1,237 +0,0 @@ -"""Factory pattern for creating LightX2V components.""" - -from pathlib import Path -from typing import Optional, Dict, Any, Union -import torch -import logging - -from .config import EncoderConfig, ModelConfig, VideoConfig -from .models import ( - LightX2VT5Encoder, - LightX2VClipVisionEncoder, - LightX2VVae, - LightX2VModel, -) - -# Import original LightX2V modules -from ..lightx2v.lightx2v.models.input_encoders.hf.t5.model import T5EncoderModel -from ..lightx2v.lightx2v.models.input_encoders.hf.xlm_roberta.model import CLIPModel as ClipVisionModel -from ..lightx2v.lightx2v.models.video_encoders.hf.wan.vae import WanVAE -from ..lightx2v.lightx2v.models.networks.wan.model import WanModel -from ..lightx2v.lightx2v.models.networks.wan.lora_adapter import WanLoraWrapper - - -class LightX2VFactory: - """Factory for creating LightX2V components with proper configuration.""" - - @staticmethod - def create_t5_encoder( - model_path: Union[str, Path], - tokenizer_path: Optional[Union[str, Path]] = None, - dtype: Optional[torch.dtype] = None, - device: Optional[torch.device] = None, - cpu_offload: bool = False, - ) -> LightX2VT5Encoder: - """Create a T5 encoder with configuration.""" - model_path = Path(model_path) - - # Auto-detect tokenizer path if not provided - if tokenizer_path is None: - tokenizer_path = model_path.parent / "google" / "umt5-xxl" - if not tokenizer_path.exists(): - raise ValueError(f"Tokenizer not found at {tokenizer_path}") - - config = EncoderConfig( - model_path=model_path, - tokenizer_path=Path(tokenizer_path), - dtype=dtype or torch.bfloat16, - device=device or torch.device("cuda"), - cpu_offload=cpu_offload, - ) - - # Create underlying T5 model - t5_model = T5EncoderModel( - text_len=config.text_len, - dtype=config.dtype, - device=config.device, - checkpoint_path=str(config.model_path), - tokenizer_path=str(config.tokenizer_path), - shard_fn=None, - cpu_offload=config.cpu_offload, - ) - - return LightX2VT5Encoder(t5_model, config) - - @staticmethod - def create_clip_vision_encoder( - model_path: Union[str, Path], - dtype: Optional[torch.dtype] = None, - device: Optional[torch.device] = None, - clip_quantized: bool = False, - clip_quantized_ckpt: Optional[Union[str, Path]] = None, - quant_scheme: Optional[str] = None, - ) -> LightX2VClipVisionEncoder: - """Create a CLIP vision encoder with configuration.""" - config = EncoderConfig( - model_path=Path(model_path), - dtype=dtype or torch.float16, - device=device or torch.device("cuda"), - clip_quantized=clip_quantized, - clip_quantized_ckpt=Path(clip_quantized_ckpt) if clip_quantized_ckpt else None, - quant_scheme=quant_scheme, - ) - - # Create underlying CLIP model - clip_model = ClipVisionModel( - dtype=config.dtype, - device=config.device, - checkpoint_path=str(config.model_path), - clip_quantized=config.clip_quantized, - clip_quantized_ckpt=str(config.clip_quantized_ckpt) if config.clip_quantized_ckpt else None, - quant_scheme=config.quant_scheme, - ) - - return LightX2VClipVisionEncoder(clip_model, config) - - @staticmethod - def create_vae( - model_path: Union[str, Path], - dtype: Optional[torch.dtype] = None, - device: Optional[torch.device] = None, - parallel: bool = False, - z_dim: int = 16, - ) -> LightX2VVae: - """Create a VAE with configuration.""" - config = EncoderConfig( - model_path=Path(model_path), - dtype=dtype or torch.float16, - device=device or torch.device("cuda"), - parallel=parallel, - z_dim=z_dim, - ) - - # Create underlying VAE model - vae_model = WanVAE( - z_dim=config.z_dim, - vae_pth=str(config.model_path), - dtype=config.dtype, - device=config.device, - parallel=config.parallel, - ) - - return LightX2VVae(vae_model, config) - - @staticmethod - def create_model( - config: ModelConfig, - video_config: Optional[VideoConfig] = None, - ) -> LightX2VModel: - """Create a complete LightX2V model with configuration.""" - # Load config.json if it exists - config_json_path = config.model_path / "config.json" - config_json = {} - - if config_json_path.exists(): - import json - - with open(config_json_path, "r") as f: - config_json = json.load(f) - else: - logging.warning(f"Config file not found at {config_json_path}") - - # Create model configuration dict - model_config_dict = { - "model_path": str(config.model_path), - "task": config.model_type, - "dtype": config.to_dtype(), - "device": config.to_device(), - "attention_type": config.attention_type, - "cpu_offload": config.cpu_offload, - "offload_granularity": config.offload_granularity, - "mm_config": config.mm_config, - "model_cls": "wan2.1", - "do_mm_calib": False, - "parallel_attn_type": None, - "parallel_vae": False, - "use_bfloat16": config.to_dtype() == torch.bfloat16, - "feature_caching": config.feature_caching, - "self_attn_1_type": "flash_attn3", - "cross_attn_1_type": "flash_attn3", - "cross_attn_2_type": "flash_attn3", - } - - # Add video config if provided - if video_config: - model_config_dict.update( - { - "target_height": video_config.target_height, - "target_width": video_config.target_width, - "target_video_length": video_config.target_video_length, - "vae_stride": video_config.vae_stride, - "patch_size": video_config.patch_size, - "max_area": video_config.max_area, - } - ) - - # Merge with config.json - model_config_dict.update(config_json) - - # Create EasyDict for compatibility - from easydict import EasyDict - - easydict_config = EasyDict(model_config_dict) - - # Create underlying model - wan_model = WanModel(str(config.model_path), easydict_config, config.to_device()) - - # Apply LoRA if specified - if config.lora_path and config.lora_path.exists(): - logging.info(f"Applying LoRA from {config.lora_path}") - lora_wrapper = WanLoraWrapper(wan_model) - lora_name = lora_wrapper.load_lora(str(config.lora_path)) - lora_wrapper.apply_lora(lora_name, config.lora_strength) - logging.info(f"LoRA {lora_name} applied successfully") - - return LightX2VModel(wan_model, config, easydict_config) - - @staticmethod - def create_from_paths( - model_dir: Union[str, Path], model_name: str, model_type: str = "i2v", precision: str = "bf16", device: str = "cuda", **kwargs - ) -> Dict[str, Any]: - """Convenience method to create all components from a model directory.""" - model_dir = Path(model_dir) - - # Create model config - model_config = ModelConfig(model_path=model_dir / model_name, model_type=model_type, precision=precision, device=device, **kwargs) - - # Create components - components = { - "model": LightX2VFactory.create_model(model_config), - } - - # Try to create encoders if paths exist - t5_path = model_dir / "models_t5_umt5-xxl-enc-bf16.pth" - if t5_path.exists(): - components["t5_encoder"] = LightX2VFactory.create_t5_encoder( - t5_path, - dtype=model_config.to_dtype(), - device=model_config.to_device(), - ) - - clip_path = model_dir / "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" - if clip_path.exists(): - components["clip_encoder"] = LightX2VFactory.create_clip_vision_encoder( - clip_path, - dtype=torch.float16, # CLIP typically uses fp16 - device=model_config.to_device(), - ) - - vae_path = model_dir / "Wan2.1_VAE.pth" - if vae_path.exists(): - components["vae"] = LightX2VFactory.create_vae( - vae_path, - dtype=torch.float16, # VAE typically uses fp16 - device=model_config.to_device(), - ) - - return components diff --git a/lightx2v_nodes/models.py b/lightx2v_nodes/models.py deleted file mode 100644 index 6ba67d4..0000000 --- a/lightx2v_nodes/models.py +++ /dev/null @@ -1,179 +0,0 @@ -"""Model wrappers for LightX2V components.""" - -from typing import Any, Dict, List, Optional, Union -import torch -from abc import ABC, abstractmethod - -from .config import EncoderConfig, ModelConfig, VideoConfig, TeaCacheConfig - - -class BaseModel(ABC): - """Base class for all LightX2V models.""" - - def __init__(self, config: Union[EncoderConfig, ModelConfig]): - self.config = config - - @abstractmethod - def to(self, device: torch.device) -> "BaseModel": - """Move model to device.""" - pass - - -class LightX2VT5Encoder(BaseModel): - """Wrapper for T5 text encoder.""" - - def __init__(self, t5_model: Any, config: EncoderConfig): - super().__init__(config) - self._model = t5_model - - def encode(self, prompts: List[str]) -> Dict[str, torch.Tensor]: - """Encode text prompts.""" - context = self._model.infer(prompts) - return {"context": context} - - def encode_with_negative(self, prompt: str, negative_prompt: Optional[str] = None) -> Dict[str, torch.Tensor]: - """Encode prompt with negative prompt.""" - context = self._model.infer([prompt]) - context_null = self._model.infer([negative_prompt if negative_prompt else ""]) - return {"context": context, "context_null": context_null} - - def to(self, device: torch.device) -> "LightX2VT5Encoder": - """Move encoder to device.""" - # T5 model handles device internally - return self - - @property - def device(self) -> torch.device: - """Get current device.""" - return self.config.device - - -class LightX2VClipVisionEncoder(BaseModel): - """Wrapper for CLIP vision encoder.""" - - def __init__(self, clip_model: Any, config: EncoderConfig): - super().__init__(config) - self._model = clip_model - - def encode(self, images: torch.Tensor, video_config: Optional[VideoConfig] = None) -> torch.Tensor: - """Encode images with CLIP.""" - if video_config: - # Convert VideoConfig to dict format expected by CLIP - config_dict = { - "target_height": video_config.target_height, - "target_width": video_config.target_width, - "target_video_length": video_config.target_video_length, - "vae_stride": video_config.vae_stride, - "patch_size": video_config.patch_size, - } - else: - config_dict = {} - - # Ensure images are in correct format [B, C, T, H, W] - if images.dim() == 3: # [C, H, W] - images = images.unsqueeze(0).unsqueeze(2) # [1, C, 1, H, W] - elif images.dim() == 4: # [B, C, H, W] - images = images.unsqueeze(2) # [B, C, 1, H, W] - - return self._model.visual(images, config_dict) - - def to(self, device: torch.device) -> "LightX2VClipVisionEncoder": - """Move encoder to device.""" - # CLIP model handles device internally - return self - - @property - def device(self) -> torch.device: - """Get current device.""" - return self.config.device - - -class LightX2VVae(BaseModel): - """Wrapper for VAE.""" - - def __init__(self, vae_model: Any, config: EncoderConfig): - super().__init__(config) - self._model = vae_model - - def encode(self, videos: List[torch.Tensor], video_config: Optional[VideoConfig] = None, cpu_offload: bool = False) -> List[torch.Tensor]: - """Encode videos to latent space.""" - config_dict = {"cpu_offload": cpu_offload} - - if video_config: - config_dict.update( - { - "target_height": video_config.target_height, - "target_width": video_config.target_width, - "target_video_length": video_config.target_video_length, - "vae_stride": video_config.vae_stride, - "patch_size": video_config.patch_size, - } - ) - - from easydict import EasyDict - - return self._model.encode(videos, EasyDict(config_dict)) - - def decode(self, latents: torch.Tensor, generator: Optional[torch.Generator] = None, cpu_offload: bool = False) -> torch.Tensor: - """Decode latents to video.""" - from easydict import EasyDict - - config = EasyDict({"cpu_offload": cpu_offload}) - - return self._model.decode(latents, generator=generator, config=config) - - def to(self, device: torch.device) -> "LightX2VVae": - """Move VAE to device.""" - # VAE model handles device internally - return self - - @property - def device(self) -> torch.device: - """Get current device.""" - return self.config.device - - -class LightX2VModel(BaseModel): - """Wrapper for main LightX2V model.""" - - def __init__(self, wan_model: Any, config: ModelConfig, easydict_config: Any): - super().__init__(config) - self._model = wan_model - self._easydict_config = easydict_config - self._scheduler = None - - def set_scheduler(self, scheduler: Any): - """Set the scheduler for the model.""" - self._scheduler = scheduler - self._model.set_scheduler(scheduler) - - def infer(self, inputs: Dict[str, Any]): - """Run inference.""" - return self._model.infer(inputs) - - def prepare_inputs( - self, - text_embeddings: Dict[str, torch.Tensor], - image_embeddings: Optional[Dict[str, Any]] = None, - ) -> Dict[str, Any]: - """Prepare inputs for inference.""" - inputs = { - "text_encoder_output": text_embeddings, - "image_encoder_output": image_embeddings or {}, - } - return inputs - - def to(self, device: torch.device) -> "LightX2VModel": - """Move model to device.""" - # Model handles device internally - return self - - @property - def device(self) -> torch.device: - """Get current device.""" - return self.config.to_device() - - @property - def easydict_config(self) -> Any: - """Get EasyDict config for compatibility.""" - return self._easydict_config diff --git a/lightx2v_nodes/nodes.py b/lightx2v_nodes/nodes.py deleted file mode 100644 index bcb82b6..0000000 --- a/lightx2v_nodes/nodes.py +++ /dev/null @@ -1,855 +0,0 @@ -"""Refactored ComfyUI nodes for LightX2V.""" - -import os -import torch -import gc -import logging -import json -import numpy as np -from pathlib import Path -from typing import Any, Dict, Optional, Tuple, cast -from easydict import EasyDict -from tqdm import tqdm - -import comfy.model_management as comfy_mm -from comfy.utils import ProgressBar - -from .config import LightX2VConfig, ModelConfig, VideoConfig, TeaCacheConfig -from .factory import LightX2VFactory -from .models import ( - LightX2VT5Encoder, - LightX2VClipVisionEncoder, - LightX2VVae, - LightX2VModel, -) - -# Import original LightX2V modules -from ..lightx2v.lightx2v.utils.profiler import ProfilingContext -from ..lightx2v.lightx2v.models.schedulers.wan.scheduler import WanScheduler -from ..lightx2v.lightx2v.models.schedulers.wan.feature_caching.scheduler import ( - WanSchedulerTeaCaching, -) - - -# Coefficient values for TeaCache -TEACACHE_COEFFICIENTS = { - "i2v-14B-480p": [ - [2.57151496e05, -3.54229917e04, 1.40286849e03, -1.35890334e01, 1.32517977e-01], - [-3.02331670e02, 2.23948934e02, -5.25463970e01, 5.87348440e00, -2.01973289e-01], - ], - "i2v-14B-720p": [ - [8.10705460e03, 2.13393892e03, -3.72934672e02, 1.66203073e01, -4.17769401e-02], - [-114.36346466, 65.26524496, -18.82220707, 4.91518089, -0.23412683], - ], - "t2v-1.3B": [ - [-5.21862437e04, 9.23041404e03, -5.28275948e02, 1.36987616e01, -4.99875664e-02], - [2.39676752e03, -1.31110545e03, 2.01331979e02, -8.29855975e00, 1.37887774e-01], - ], - "t2v-14B": [ - [-3.03318725e05, 4.90537029e04, -2.65530556e03, 5.87365115e01, -3.15583525e-01], - [-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429, -13.02252404], - ], -} - - -class BaseNode: - """Base class for ComfyUI nodes with common functionality.""" - - CATEGORY = "LightX2V" - - @classmethod - def get_device(cls, device_str: str) -> torch.device: - """Convert device string to torch device.""" - if device_str == "cuda": - return comfy_mm.get_torch_device() - return torch.device("cpu") - - @classmethod - def get_dtype(cls, precision_str: str) -> torch.dtype: - """Convert precision string to torch dtype.""" - dtype_map = { - "bf16": torch.bfloat16, - "fp16": torch.float16, - "fp32": torch.float32, - } - return dtype_map[precision_str] - - -class WanVideoTeaCache(BaseNode): - """TeaCache configuration node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "rel_l1_thresh": ( - "FLOAT", - { - "default": 0.26, - "min": 0.0, - "max": 10.0, - "step": 0.001, - "tooltip": "Threshold for cache application", - }, - ), - "start_percent": ( - "FLOAT", - { - "default": 0.1, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "tooltip": "Start percentage for TeaCache", - }, - ), - "end_percent": ( - "FLOAT", - { - "default": 1.0, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "tooltip": "End percentage for TeaCache", - }, - ), - "cache_device": ( - ["main_device", "offload_device"], - {"default": "offload_device", "tooltip": "Device to cache to"}, - ), - "coefficients": ( - list(TEACACHE_COEFFICIENTS.keys()), - { - "default": "i2v-14B-720p", - "tooltip": "Coefficient preset for TeaCache", - }, - ), - "use_ret_steps": ("BOOLEAN", {"default": False}), - }, - "optional": { - "mode": ( - ["e", "e0"], - { - "default": "e", - "tooltip": "Time embedding mode", - }, - ), - }, - } - - RETURN_TYPES = ("LIGHT_TEACACHEARGS",) - RETURN_NAMES = ("teacache_args",) - FUNCTION = "process" - EXPERIMENTAL = True - - def process( - self, - rel_l1_thresh: float, - start_percent: float, - end_percent: float, - cache_device: str, - coefficients: str, - use_ret_steps: bool, - mode: str = "e", - ) -> Tuple[TeaCacheConfig]: - """Create TeaCache configuration.""" - device = comfy_mm.get_torch_device() if cache_device == "main_device" else comfy_mm.unet_offload_device() - - config = TeaCacheConfig( - rel_l1_thresh=rel_l1_thresh, - start_percent=start_percent, - end_percent=end_percent, - cache_device=device, - coefficients=TEACACHE_COEFFICIENTS[coefficients], - use_ret_steps=use_ret_steps, - mode=mode, - ) - - return (config,) - - -class Lightx2vWanVideoModelDir(BaseNode): - """Model directory specification node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_dir": ( - "STRING", - {"default": "/mnt/aigc/users/lijiaqi2/wan_model/Wan2.1-I2V-14B-480P"}, - ) - } - } - - RETURN_TYPES = ("STRING",) - RETURN_NAMES = ("model_dir",) - FUNCTION = "process" - - def process(self, model_dir: str) -> Tuple[str]: - """Validate and return model directory.""" - path = Path(model_dir) - if not path.exists(): - raise ValueError(f"Model directory {model_dir} does not exist.") - return (model_dir,) - - -class Lightx2vWanVideoT5EncoderLoader(BaseNode): - """T5 encoder loader node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_name": ( - "STRING", - {"default": "models_t5_umt5-xxl-enc-bf16.pth"}, - ), - "precision": (["bf16", "fp16", "fp32"], {"default": "bf16"}), - "device": (["cuda", "cpu"], {"default": "cuda"}), - }, - "optional": { - "model_dir": ("STRING", {"default": None}), - }, - } - - RETURN_TYPES = ("LIGHT_T5_ENCODER",) - RETURN_NAMES = ("t5_encoder",) - FUNCTION = "load_t5_encoder" - - def load_t5_encoder( - self, - model_name: str, - precision: str, - device: str, - model_dir: Optional[str] = None, - ) -> Tuple[LightX2VT5Encoder]: - """Load T5 encoder.""" - dtype = self.get_dtype(precision) - device_obj = self.get_device(device) - - if model_dir: - model_path = Path(model_dir) / model_name - else: - model_path = Path(model_name) - if not model_path.exists(): - raise ValueError(f"T5 model path {model_path} does not exist.") - - encoder = LightX2VFactory.create_t5_encoder( - model_path=model_path, - dtype=dtype, - device=device_obj, - cpu_offload=(device == "cpu"), - ) - - return (encoder,) - - -class Lightx2vWanVideoT5Encoder(BaseNode): - """T5 text encoding node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "t5_encoder": ("LIGHT_T5_ENCODER",), - "prompt": ( - "STRING", - { - "multiline": True, - "default": "Summer beach vacation style...", - }, - ), - "negative_prompt": ( - "STRING", - { - "multiline": True, - "default": "", - }, - ), - } - } - - RETURN_TYPES = ("LIGHT_TEXT_EMBEDDINGS",) - RETURN_NAMES = ("text_embeddings",) - FUNCTION = "encode_text" - - def encode_text( - self, - t5_encoder: LightX2VT5Encoder, - prompt: str, - negative_prompt: str = "", - ) -> Tuple[Dict[str, torch.Tensor]]: - """Encode text with T5.""" - embeddings = t5_encoder.encode_with_negative(prompt, negative_prompt) - return (embeddings,) - - -class Lightx2vWanVideoClipVisionEncoderLoader(BaseNode): - """CLIP vision encoder loader node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_name": ( - "STRING", - {"default": "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"}, - ), - "tokenizer_path": ( - "STRING", - {"default": "xlm-roberta-large"}, - ), - "precision": (["fp16", "fp32"], {"default": "fp16"}), - "device": (["cuda", "cpu"], {"default": "cuda"}), - }, - "optional": { - "model_dir": ("STRING", {"default": None}), - }, - } - - RETURN_TYPES = ("LIGHT_CLIP_VISION_ENCODER",) - RETURN_NAMES = ("clip_vision_encoder",) - FUNCTION = "load_clip_vision_encoder" - - def load_clip_vision_encoder( - self, - model_name: str, - tokenizer_path: str, - precision: str, - device: str, - model_dir: Optional[str] = None, - ) -> Tuple[LightX2VClipVisionEncoder]: - """Load CLIP vision encoder.""" - dtype = self.get_dtype(precision) - device_obj = self.get_device(device) - - if model_dir: - model_path = Path(model_dir) / model_name - else: - model_path = Path(model_name) - if not model_path.exists(): - raise ValueError(f"CLIP model path {model_path} does not exist.") - - encoder = LightX2VFactory.create_clip_vision_encoder( - model_path=model_path, - dtype=dtype, - device=device_obj, - ) - - return (encoder,) - - -class Lightx2vWanVideoVaeLoader(BaseNode): - """VAE loader node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_name": ( - "STRING", - {"default": "Wan2.1_VAE.pth"}, - ), - "precision": (["bf16", "fp16", "fp32"], {"default": "fp16"}), - "device": (["cuda", "cpu"], {"default": "cuda"}), - "parallel": ("BOOLEAN", {"default": False}), - }, - "optional": { - "model_dir": ("STRING", {"default": None}), - }, - } - - RETURN_TYPES = ("LIGHT_WAN_VAE",) - RETURN_NAMES = ("wan_vae",) - FUNCTION = "load_vae" - - def load_vae( - self, - model_name: str, - precision: str, - device: str, - parallel: bool, - model_dir: Optional[str] = None, - ) -> Tuple[Dict[str, Any]]: - """Load VAE.""" - dtype = self.get_dtype(precision) - device_obj = self.get_device(device) - - if model_dir: - model_path = Path(model_dir) / model_name - else: - model_path = Path(model_name) - if not model_path.exists(): - raise ValueError(f"VAE model path {model_path} does not exist.") - - vae = LightX2VFactory.create_vae( - model_path=model_path, - dtype=dtype, - device=device_obj, - parallel=parallel, - ) - - # Return in legacy format for compatibility - return ({"vae_cls": vae, "device": device},) - - -class Lightx2vWanVideoVaeDecoder(BaseNode): - """VAE decoder node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "wan_vae": ("LIGHT_WAN_VAE",), - "latent": ("LIGHT_LATENT",), - } - } - - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("images",) - FUNCTION = "decode_latent" - - def decode_latent( - self, - wan_vae: Dict[str, Any], - latent: Dict[str, Any], - ) -> Tuple[torch.Tensor]: - """Decode latents to images.""" - vae_instance = wan_vae["vae_cls"] - if isinstance(vae_instance, LightX2VVae): - vae_model = vae_instance - else: - # Legacy compatibility - vae_model = vae_instance - - latents = latent["samples"] - generator = latent["generator"] - cpu_offload = wan_vae["device"] == "cpu" - - with torch.no_grad(): - with ProfilingContext("*decoded images*"): - if isinstance(vae_model, LightX2VVae): - decoded_images = vae_model.decode(latents, generator, cpu_offload) - else: - # Legacy compatibility - config = EasyDict({"cpu_offload": cpu_offload}) - decoded_images = vae_model.decode(latents, generator=generator, config=config) - - # Normalize from [-1, 1] to [0, 1] - images = (decoded_images + 1) / 2 - - # Rearrange dimensions for ComfyUI [T, H, W, C] - images = images.squeeze(0).permute(1, 2, 3, 0).cpu() - images = torch.clamp(images, 0, 1) - - # Cleanup - torch.cuda.empty_cache() - gc.collect() - - return (images,) - - -class Lightx2vWanVideoImageEncoder(BaseNode): - """Image encoding node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "vae": ("LIGHT_WAN_VAE",), - "clip_vision_encoder": ("LIGHT_CLIP_VISION_ENCODER",), - "image": ("IMAGE",), - "width": ( - "INT", - { - "default": 832, - "min": 64, - "max": 2048, - "step": 8, - "tooltip": "Width of the image", - }, - ), - "height": ( - "INT", - { - "default": 480, - "min": 64, - "max": 2048, - "step": 8, - "tooltip": "Height of the image", - }, - ), - "num_frames": ( - "INT", - { - "default": 81, - "min": 1, - "max": 10000, - "step": 4, - "tooltip": "Number of frames", - }, - ), - } - } - - RETURN_TYPES = ("LIGHT_IMAGE_EMBEDDINGS",) - RETURN_NAMES = ("image_embeddings",) - FUNCTION = "encode_image" - - def encode_image( - self, - vae: Dict[str, Any], - clip_vision_encoder: LightX2VClipVisionEncoder, - image: torch.Tensor, - width: int, - height: int, - num_frames: int, - ) -> Tuple[Dict[str, Any]]: - """Encode image with CLIP and VAE.""" - vae_instance = vae["vae_cls"] - - # Create video configuration - video_config = VideoConfig( - target_width=width, - target_height=height, - target_video_length=num_frames, - ) - - # Convert image format - device = comfy_mm.get_torch_device() - img = image[0].permute(2, 0, 1).to(device) - img = img.sub_(0.5).div_(0.5) # Normalize to [-1, 1] - - # Encode with CLIP - with ProfilingContext("*clip encoder*"): - if isinstance(clip_vision_encoder, LightX2VClipVisionEncoder): - clip_out = clip_vision_encoder.encode(img, video_config) - clip_out = clip_out.squeeze(0).to(torch.bfloat16) - else: - # Legacy compatibility - config_dict = video_config.__dict__ if hasattr(video_config, "__dict__") else video_config - clip_out = clip_vision_encoder.visual([img[:, None, :, :]], config_dict) - clip_out = clip_out.squeeze(0).to(torch.bfloat16) - - # Calculate dimensions - h, w = img.shape[1:] - aspect_ratio = h / w - max_area = video_config.max_area - lat_h = round(np.sqrt(max_area * aspect_ratio) // video_config.vae_stride[1] // video_config.patch_size[1] * video_config.patch_size[1]) - lat_w = round(np.sqrt(max_area / aspect_ratio) // video_config.vae_stride[2] // video_config.patch_size[2] * video_config.patch_size[2]) - - # Update config - config_dict = video_config.to_easydict() if hasattr(video_config, "to_easydict") else EasyDict(video_config.__dict__) - config_dict.lat_h = lat_h - config_dict.lat_w = lat_w - - h = lat_h * video_config.vae_stride[1] - w = lat_w * video_config.vae_stride[2] - - # Create mask - msk = torch.ones(1, num_frames, lat_h, lat_w, device=device) - msk[:, 1:] = 0 - msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1) - msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w) - msk = msk.transpose(1, 2)[0] - - # Encode with VAE - with ProfilingContext("*vae encoder*"): - video_tensor = torch.concat( - [ - torch.nn.functional.interpolate(img[None].cpu(), size=(h, w), mode="bicubic").transpose(0, 1), - torch.zeros(3, num_frames - 1, h, w), - ], - dim=1, - ).cuda() - - if isinstance(vae_instance, LightX2VVae): - vae_out = vae_instance.encode([video_tensor], video_config)[0] - else: - # Legacy compatibility - vae_out = vae_instance.encode([video_tensor], config_dict)[0] - - vae_out = torch.concat([msk, vae_out]).to(torch.bfloat16) - - image_embeddings = { - "clip_encoder_out": clip_out, - "vae_encode_out": vae_out, - "config": config_dict, - } - - logging.info(f"Image encoder outputs - CLIP: {clip_out.shape}, VAE: {vae_out.shape}") - - return (image_embeddings,) - - -class Lightx2vWanVideoEmptyEmbeds(BaseNode): - """Empty embeddings for T2V.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8}), - "height": ("INT", {"default": 480, "min": 64, "max": 2048, "step": 8}), - "num_frames": ( - "INT", - {"default": 81, "min": 1, "max": 10000, "step": 4}, - ), - } - } - - RETURN_TYPES = ("LIGHT_IMAGE_EMBEDDINGS",) - RETURN_NAMES = ("image_embeddings",) - FUNCTION = "process" - - def process(self, num_frames: int, width: int, height: int) -> Tuple[Dict[str, Any]]: - """Create empty image embeddings for T2V.""" - video_config = VideoConfig( - target_width=width, - target_height=height, - target_video_length=num_frames, - ) - - return ({"config": video_config.to_easydict()},) - - -class Lightx2vWanVideoModelLoader(BaseNode): - """Main model loader node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_name": ("STRING", {"default": ""}), - "model_type": (["t2v", "i2v"], {"default": "i2v"}), - "precision": (["bf16", "fp16", "fp32"], {"default": "bf16"}), - "device": (["cuda", "cpu"], {"default": "cuda"}), - "attention_type": ( - ["sdpa", "flash_attn2", "flash_attn3"], - {"default": "flash_attn3"}, - ), - "cpu_offload": ("BOOLEAN", {"default": False}), - "offload_granularity": (["block", "phase"], {"default": "phase"}), - }, - "optional": { - "mm_type": ("STRING", {"default": None}), - "teacache_args": ("LIGHT_TEACACHEARGS", {"default": None}), - "lora_path": ("STRING", {"default": None}), - "lora_strength": ( - "FLOAT", - {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01}, - ), - "model_dir": ( - "STRING", - {"default": "/mnt/aigc/users/lijiaqi2/wan_model/Wan2.1-I2V-14B-480P"}, - ), - }, - } - - RETURN_TYPES = ("LIGHT_WAN_MODEL",) - RETURN_NAMES = ("wan_model",) - FUNCTION = "load_model" - - def load_model( - self, - model_name: str, - model_type: str, - precision: str, - device: str, - attention_type: str, - offload_granularity: str, - mm_type: Optional[str] = None, - lora_path: Optional[str] = None, - lora_strength: float = 1.0, - cpu_offload: bool = False, - teacache_args: Optional[TeaCacheConfig] = None, - model_dir: Optional[str] = None, - ) -> Tuple[Dict[str, Any]]: - """Load the main model.""" - if model_dir: - model_path = Path(model_dir) / model_name - else: - model_path = Path(model_name) - if not model_path.exists(): - raise ValueError(f"Model path {model_path} does not exist.") - - # Parse mm_config - mm_config = {} - if mm_type: - try: - mm_config = json.loads(mm_type) - except Exception as e: - logging.error(f"Invalid mm_type config: {e}") - - # Create model configuration - model_config = ModelConfig( - model_path=model_path, - model_type=model_type, - precision=precision, - device=device, - attention_type=attention_type, - cpu_offload=cpu_offload, - offload_granularity=offload_granularity, - lora_path=Path(lora_path) if lora_path and lora_path.strip() else None, - lora_strength=lora_strength, - mm_config=mm_config, - feature_caching="Tea" if teacache_args else "NoCaching", - ) - - # Create model - model = LightX2VFactory.create_model(model_config) - - # Add TeaCache config if provided - easydict_config = model.easydict_config - if teacache_args: - easydict_config.teacache_thresh = teacache_args.rel_l1_thresh - easydict_config.use_ret_steps = teacache_args.use_ret_steps - easydict_config.coefficients = teacache_args.coefficients - easydict_config.teacache_start_percent = teacache_args.start_percent - easydict_config.teacache_end_percent = teacache_args.end_percent - easydict_config.teacache_device = teacache_args.cache_device - easydict_config.teacache_mode = teacache_args.mode - - logging.info(f"Loaded model from {model_path} with type {model_type}") - - return ({"wan_model": model._model, "config": easydict_config},) - - -class Lightx2vWanVideoSampler(BaseNode): - """Video sampling node.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model": ("LIGHT_WAN_MODEL",), - "text_embeddings": ("LIGHT_TEXT_EMBEDDINGS",), - "image_embeddings": ("LIGHT_IMAGE_EMBEDDINGS",), - "steps": ("INT", {"default": 20, "min": 1, "max": 100, "step": 1}), - "shift": ("FLOAT", {"default": 5.0}), - "cfg_scale": ( - "FLOAT", - {"default": 5, "min": 1, "max": 20.0, "step": 0.1}, - ), - "seed": ( - "INT", - {"default": 42, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "step": 1}, - ), - } - } - - RETURN_TYPES = ("LIGHT_LATENT",) - RETURN_NAMES = ("latent",) - FUNCTION = "sample" - - def sample( - self, - model: Dict[str, Any], - text_embeddings: Dict[str, Any], - steps: int, - shift: float, - cfg_scale: float, - seed: int, - image_embeddings: Dict[str, Any], - ) -> Tuple[Dict[str, Any]]: - """Sample video latents.""" - model_config = cast(EasyDict, model.get("config")) - model_config.update(image_embeddings.get("config", {})) - - wan_model = model.get("wan_model") - clip_encoder_out = image_embeddings.get("clip_encoder_out", None) - vae_encode_out = image_embeddings.get("vae_encode_out", None) - - if model_config.task == "i2v" and (clip_encoder_out is None or vae_encode_out is None): - raise ValueError("Image embeddings required for i2v task") - - # Update config - model_config.infer_steps = steps - model_config.sample_shift = shift - model_config.sample_guide_scale = cfg_scale - model_config.seed = seed - model_config.enable_cfg = cfg_scale != 1.0 - - # Set target shape - num_channels_latents = model_config.get("num_channels_latents", 16) - - if model_config.task == "i2v": - model_config.target_shape = ( - num_channels_latents, - (model_config.target_video_length - 1) // model_config.vae_stride[0] + 1, - model_config.lat_h, - model_config.lat_w, - ) - else: # t2v - model_config.target_shape = ( - 16, - (model_config.target_video_length - 1) // 4 + 1, - int(model_config.target_height) // model_config.vae_stride[1], - int(model_config.target_width) // model_config.vae_stride[2], - ) - - # Create scheduler - if model_config.feature_caching == "NoCaching": - scheduler = WanScheduler(model_config) - elif model_config.feature_caching == "Tea": - scheduler = WanSchedulerTeaCaching(model_config) - else: - raise NotImplementedError(f"Unsupported caching: {model_config.feature_caching}") - - wan_model.set_scheduler(scheduler) - - # Prepare inputs - inputs = { - "text_encoder_output": text_embeddings, - "image_encoder_output": image_embeddings, - } - - scheduler.prepare(inputs.get("image_encoder_output")) - - # Run sampling - progress = ProgressBar(steps) - for step_index in tqdm(range(scheduler.infer_steps), desc="Sampling"): - scheduler.step_pre(step_index=step_index) - with ProfilingContext("model.infer"): - wan_model.infer(inputs) - scheduler.step_post() - progress.update(1) - - latents, generator = scheduler.latents, scheduler.generator - scheduler.clear() - - # Cleanup - del inputs, scheduler, text_embeddings, image_embeddings - torch.cuda.empty_cache() - - return ({"samples": latents, "generator": generator},) - - -# Node mappings -NODE_CLASS_MAPPINGS = { - "Lightx2vWanVideoModelDir": Lightx2vWanVideoModelDir, - "Lightx2vWanVideoT5EncoderLoader": Lightx2vWanVideoT5EncoderLoader, - "Lightx2vWanVideoT5Encoder": Lightx2vWanVideoT5Encoder, - "Lightx2vWanVideoClipVisionEncoderLoader": Lightx2vWanVideoClipVisionEncoderLoader, - "Lightx2vWanVideoVaeLoader": Lightx2vWanVideoVaeLoader, - "Lightx2vTeaCache": WanVideoTeaCache, - "Lightx2vWanVideoEmptyEmbeds": Lightx2vWanVideoEmptyEmbeds, - "Lightx2vWanVideoImageEncoder": Lightx2vWanVideoImageEncoder, - "Lightx2vWanVideoVaeDecoder": Lightx2vWanVideoVaeDecoder, - "Lightx2vWanVideoModelLoader": Lightx2vWanVideoModelLoader, - "Lightx2vWanVideoSampler": Lightx2vWanVideoSampler, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "Lightx2vWanVideoModelDir": "LightX2V WAN Model Directory", - "Lightx2vWanVideoT5EncoderLoader": "LightX2V WAN T5 Encoder Loader", - "Lightx2vWanVideoT5Encoder": "LightX2V WAN T5 Encoder", - "Lightx2vWanVideoClipVisionEncoderLoader": "LightX2V WAN CLIP Vision Encoder Loader", - "Lightx2vWanVideoVaeLoader": "LightX2V WAN VAE Loader", - "Lightx2vWanVideoImageEncoder": "LightX2V WAN Image Encoder", - "Lightx2vWanVideoVaeDecoder": "LightX2V WAN VAE Decoder", - "Lightx2vWanVideoModelLoader": "LightX2V WAN Model Loader", - "Lightx2vWanVideoSampler": "LightX2V WAN Video Sampler", - "Lightx2vTeaCache": "LightX2V WAN Tea Cache", - "Lightx2vWanVideoEmptyEmbeds": "LightX2V WAN Video Empty Embeds", -} diff --git a/lightx2v_nodes/universal_bridge.py b/lightx2v_nodes/universal_bridge.py deleted file mode 100644 index 158079c..0000000 --- a/lightx2v_nodes/universal_bridge.py +++ /dev/null @@ -1,1008 +0,0 @@ -"""Universal bridge between LightX2V and ComfyUI for automatic adaptation.""" - -import json -import os -import torch -import numpy as np -from pathlib import Path -from typing import Any, Dict, List, Union, Optional -from PIL import Image -import tempfile -from easydict import EasyDict -import asyncio -import gc -from comfy.utils import ProgressBar -import logging - -# from .config import get_available_attn_ops, get_available_quant_ops - -from ..lightx2v.lightx2v.utils.set_config import get_default_config -from ..lightx2v.lightx2v.infer import init_runner - - -class LightX2VBridge: - """Universal bridge for LightX2V and ComfyUI integration.""" - - def __init__(self): - self._model_registry = None - self._config_registry = None - self._quantization_registry = None - self._current_runner = None # Only store one runner at a time - self._current_runner_key = None # Track current runner identity - - @property - def model_registry(self): - """Lazy load model registry.""" - if self._model_registry is None: - self._model_registry = self._discover_models() - return self._model_registry - - @property - def config_registry(self): - """Lazy load config registry.""" - if self._config_registry is None: - self._config_registry = self._discover_configs() - return self._config_registry - - @property - def quantization_registry(self): - """Lazy load quantization registry.""" - if self._quantization_registry is None: - self._quantization_registry = self._discover_quantization_configs() - return self._quantization_registry - - def _discover_models(self) -> List[str]: - """Discover available model classes from LightX2V.""" - return ["wan2.1", "hunyuan", "wan2.1_audio", "wan2.1_distill"] - - def _discover_configs(self) -> Dict[str, Dict]: - """Discover all available config files for single-card inference.""" - configs = {} - base_dir = Path(__file__).parent.parent / "lightx2v" / "configs" - - if base_dir.exists(): - for config_file in base_dir.rglob("*.json"): - # Create a descriptive key - relative_path = config_file.relative_to(base_dir) - path_str = str(relative_path) - - # 排除分布式推理和部署相关的配置 - if any(exclude in path_str for exclude in ["deploy", "causvid", "skyreels", "cogvideox"]): - continue - - # 包含蒸馏模型配置 - if "distill" in path_str: - config_info = { - "path": config_file, - "model_type": "wan2.1_distill", - "task": self._extract_task_from_path(path_str), - "is_quantized": False, - "is_distilled": True, - "description": self._generate_config_description(relative_path, is_distilled=True), - } - key = str(relative_path).replace("/", "_").replace(".json", "") - configs[key] = config_info - continue - - # 排除量化配置(单独处理) - if "quantization" in path_str: - continue - - # 创建配置信息 - config_info = { - "path": config_file, - "model_type": relative_path.parts[0], # wan, hunyuan, etc. - "task": self._extract_task_from_path(path_str), - "is_quantized": False, - "description": self._generate_config_description(relative_path), - } - - key = str(relative_path).replace("/", "_").replace(".json", "") - configs[key] = config_info - - return configs - - def _discover_quantization_configs(self) -> Dict[str, Dict]: - """Discover quantization configurations.""" - quant_configs = {} - base_dir = Path(__file__).parent.parent / "lightx2v" / "configs" / "quantization" - - if base_dir.exists(): - for config_file in base_dir.rglob("*.json"): - relative_path = config_file.relative_to(base_dir.parent) - path_str = str(relative_path) - - config_info = { - "path": config_file, - "model_type": relative_path.parts[1], # wan, hunyuan, etc. - "task": self._extract_task_from_path(path_str), - "is_quantized": True, - "description": self._generate_config_description(relative_path, is_quantized=True), - } - - key = f"quant_{str(relative_path).replace('/', '_').replace('.json', '')}" - quant_configs[key] = config_info - - return quant_configs - - def _extract_task_from_path(self, path_str: str) -> str: - """Extract task type from config path.""" - if "i2v" in path_str: - return "i2v" - elif "t2v" in path_str: - return "t2v" - else: - return "unknown" - - def _generate_config_description( - self, - relative_path: Path, - is_quantized: bool = False, - is_distilled: bool = False, - ) -> str: - """Generate human-readable description for config.""" - parts = relative_path.parts - model_type = parts[0] if not is_quantized else parts[1] - task = self._extract_task_from_path(str(relative_path)) - - desc = f"{model_type.upper()} {task.upper()}" - if is_quantized: - desc += " (Quantized)" - if is_distilled: - desc += " (Distilled 4-step)" - - return desc - - def get_quantization_model_path(self, model_type: str, mm_type: str) -> str: - """Get quantization model path based on model type and mm_type.""" - # 定义量化模型的固定路径结构 - base_path = Path(__file__).parent.parent / "lightx2v" / "models" / "quantized" - - # 根据mm_type确定具体的模型文件名 - mm_type_to_filename = { - "W-int8-channel-sym-A-int8-channel-sym-dynamic-Vllm": "int8_dynamic_vllm.safetensors", - "W-int8-channel-sym-A-fp16-dynamic-Vllm": "int8_fp16_dynamic_vllm.safetensors", - "W-fp16-A-fp16-dynamic-Vllm": "fp16_dynamic_vllm.safetensors", - } - - filename = mm_type_to_filename.get(mm_type, "default_quantized.safetensors") - model_path = base_path / model_type / filename - - return str(model_path) - - def load_config(self, config_path: Union[str, Path]) -> Dict: - """Load a config file.""" - with open(config_path, "r") as f: - return json.load(f) - - def get_runner(self, config: EasyDict): - """Get or create a runner for the given config.""" - # Create a simple key based on essential parameters only - # Only model_cls and model_path are essential for runner identity - key = f"{config.model_cls}_{config.model_path}" # type:ignore - - logging.info(f"[lightx2v] get_runner: {key}") - - # If we have a different runner, clear the old one first - if self._current_runner_key != key: - self.clear_cache() - self._current_runner = init_runner(config) - self._current_runner_key = key - - return self._current_runner - - def clear_cache(self): - """Clear cached runner to free memory.""" - if self._current_runner is not None: - # Clean up the runner if it has cleanup methods - self._current_runner = None - self._current_runner_key = None - torch.cuda.empty_cache() - - -class InputOutputConverter: - """Convert between ComfyUI and LightX2V formats.""" - - @staticmethod - def comfy_image_to_pil(comfy_image: torch.Tensor) -> Image.Image: - """Convert ComfyUI image tensor to PIL Image. - - ComfyUI format: [B, H, W, C] float32 0-1 - """ - # Take first image from batch - image_np = (comfy_image[0].cpu().numpy() * 255).astype(np.uint8) - return Image.fromarray(image_np) - - @staticmethod - def comfy_audio_to_path(comfy_audio: Dict) -> str: - """Convert ComfyUI audio to file path.""" - # ComfyUI audio format: {"waveform": tensor, "sample_rate": int} - waveform = comfy_audio["waveform"] - sample_rate = comfy_audio["sample_rate"] - - # Save to temporary file - import torchaudio - - with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: - torchaudio.save(tmp.name, waveform, sample_rate) - return tmp.name - - @staticmethod - def video_to_latent(video_path: str) -> Dict[str, Any]: - """Convert video path to ComfyUI latent format.""" - # For now, return the path wrapped in latent format - # This allows downstream nodes to handle the video - return {"samples": video_path, "type": "video", "format": "path"} - - @staticmethod - def tensor_to_latent(video_tensor: torch.Tensor) -> Dict[str, Any]: - """Convert video tensor to ComfyUI latent format.""" - return {"samples": video_tensor, "type": "video", "format": "tensor"} - - -class ConfigManager: - """Manage configuration merging and parameter mapping.""" - - # Mapping from ComfyUI parameter names to LightX2V config keys - PARAM_MAPPING = { - "steps": "infer_steps", - "cfg_scale": "sample_guide_scale", - "seed": "seed", - "height": "target_height", - "width": "target_width", - "video_length": "target_video_length", - } - - # 用户可配置的参数(在ComfyUI中显示) - USER_CONFIGURABLE_PARAMS = { - "steps": { - "type": "INT", - "default": 20, - "min": 1, - "max": 200, - "tooltip": "推理步数 (-1使用默认值)", - }, - "cfg_scale": { - "type": "FLOAT", - "default": 7, - "min": 0.1, - "max": 30, - "step": 0.1, - "tooltip": "CFG引导强度 (-1使用默认值)", - }, - "seed": { - "type": "INT", - "default": 42, - "min": -1, - "max": 2**32 - 1, - "tooltip": "随机种子 (-1随机)", - }, - "height": { - "type": "INT", - "default": 640, - "min": 64, - "max": 2048, - "step": 8, - "tooltip": "视频高度 (-1使用默认值)", - }, - "width": { - "type": "INT", - "default": 640, - "min": 64, - "max": 2048, - "step": 8, - "tooltip": "视频宽度 (-1使用默认值)", - }, - "video_length": { - "type": "INT", - "default": 1, - "min": 1, - "max": 300, - "tooltip": "视频帧数 (-1使用默认值)", - }, - "sample_shift": { - "type": "INT", - "default": 5, - "min": 0, - "max": 20, - "tooltip": "采样偏移 (-1使用默认值)", - }, - } - - @classmethod - def create_config( - cls, - model_cls: str, - model_path: str, - task: str, - base_config: Dict, - overrides: Dict, - quantization_config: Optional[Dict] = None, - ) -> EasyDict: - """Create a complete config for LightX2V runner.""" - # Start with default config - config = get_default_config() - - # Add required fields - config.update( - { - "model_cls": model_cls, - "model_path": model_path, - "task": task, - "mode": "infer", - } - ) - - # Apply base config from file - config.update(base_config) - - # Apply quantization config if provided - if quantization_config: - config.update(quantization_config) - # 自动设置量化模型路径 - if "mm_config" in quantization_config and "mm_type" in quantization_config["mm_config"]: - mm_type = quantization_config["mm_config"]["mm_type"] - config["dit_quantized_ckpt"] = cls._get_quantization_model_path(model_cls, mm_type) - - # Apply ComfyUI overrides - for comfy_key, value in overrides.items(): - if value is not None and value != -1: # -1 means use default - if comfy_key in cls.PARAM_MAPPING: - lightx2v_key = cls.PARAM_MAPPING[comfy_key] - config[lightx2v_key] = value - else: - # Direct mapping for unknown parameters - config[comfy_key] = value - - config = cls._load_model_config(config) - - return EasyDict(config) - - @classmethod - def _load_model_config(cls, config: Dict) -> Dict: - """Load and merge config.json from model path if it exists.""" - model_path = config.get("model_path", "") - if not model_path: - return config - - model_config_path = os.path.join(model_path, "config.json") - print(f"Loading model config from: {model_config_path}", flush=True) - if os.path.exists(model_config_path): - try: - with open(model_config_path, "r") as f: - model_config = json.load(f) - config.update(model_config) - print(f"Loaded model config from: {model_config_path}") - except Exception as e: - print(f"Failed to load model config from {model_config_path}: {e}") - - return config - - @classmethod - def _get_quantization_model_path(cls, model_cls: str, mm_type: str) -> str: - """Get quantization model path based on model class and mm_type.""" - bridge = LightX2VBridge() - return bridge.get_quantization_model_path(model_cls, mm_type) - - @classmethod - def parse_custom_config(cls, custom_config_str: str) -> Dict: - """Parse custom config string.""" - if not custom_config_str.strip(): - return {} - - try: - return json.loads(custom_config_str) - except json.JSONDecodeError as e: - print(f"Failed to parse custom config: {e}") - return {} - - -class LightX2VConfigBuilder: - """Build LightX2V configuration from various sources.""" - - def __init__(self): - self.bridge = LightX2VBridge() - - @classmethod - def INPUT_TYPES(cls): - """Define inputs for the config builder.""" - bridge = LightX2VBridge() - - # Get available models and configs - model_choices = bridge.model_registry - - config_choices = ["custom"] - config_descriptions = {} - - for key, config_info in bridge.config_registry.items(): - config_choices.append(key) - config_descriptions[key] = config_info["description"] - - required_inputs = { - "model_cls": (model_choices, {"tooltip": "选择模型类型"}), - "model_path": ( - "STRING", - {"default": "", "tooltip": "模型权重路径(留空使用默认路径)"}, - ), - "task": ( - ["t2v", "i2v"], - { - "default": "t2v", - "tooltip": "任务类型: 文本到视频或图像到视频", - }, - ), - "config_preset": ( - config_choices, - { - "default": "custom", - "tooltip": "配置预设或'custom'自定义", - }, - ), - } - - for param_name, param_config in ConfigManager.USER_CONFIGURABLE_PARAMS.items(): - required_inputs[param_name] = ( - param_config["type"], - { - "default": param_config["default"], - "min": param_config["min"], - "max": param_config["max"], - "tooltip": param_config["tooltip"], - }, - ) - if "step" in param_config: - required_inputs[param_name][1]["step"] = param_config["step"] - - optional_inputs = { - "custom_config": ( - "STRING", - { - "multiline": True, - "default": "{}", - "tooltip": "自定义配置(JSON格式)", - }, - ), - "quantization_config": ( - "LIGHTX2V_CONFIG", - {"tooltip": "量化配置"}, - ), - "attention_config": ( - "LIGHTX2V_CONFIG", - {"tooltip": "注意力机制配置"}, - ), - "caching_config": ( - "LIGHTX2V_CONFIG", - {"tooltip": "缓存配置"}, - ), - "distill_config": ( - "LIGHTX2V_CONFIG", - {"tooltip": "蒸馏配置"}, - ), - "lora_path": ( - "STRING", - {"default": "", "tooltip": "LoRA权重路径"}, - ), - "strength_model": ( - "FLOAT", - { - "default": 1.0, - "min": 0.0, - "max": 2.0, - "step": 0.1, - "tooltip": "LoRA模型强度", - }, - ), - } - - return { - "required": required_inputs, - "optional": optional_inputs, - } - - RETURN_TYPES = ("LIGHTX2V_CONFIG",) - RETURN_NAMES = ("config",) - FUNCTION = "build_config" - CATEGORY = "LightX2V/Config" - - def build_config( - self, - model_cls, - model_path, - task, - config_preset, - steps, - cfg_scale, - seed, - height, - width, - video_length, - sample_shift, - custom_config="{}", - quantization_config=None, - attention_config=None, - caching_config=None, - distill_config=None, - lora_path="", - strength_model=1.0, - **kwargs, - ): - """Build configuration for LightX2V inference.""" - - if config_preset == "custom": - base_config = {} - else: - config_info = self.bridge.config_registry.get(config_preset) - if config_info: - base_config = self.bridge.load_config(config_info["path"]) - else: - raise ValueError(f"配置预设 '{config_preset}' 未找到") - - custom_cfg = ConfigManager.parse_custom_config(custom_config) - base_config.update(custom_cfg) - - if quantization_config: - base_config.update(quantization_config) - if attention_config: - base_config.update(attention_config) - if caching_config: - base_config.update(caching_config) - if distill_config: - base_config.update(distill_config) - - overrides = { - "seed": seed if seed != -1 else None, - "steps": steps, - "cfg_scale": cfg_scale, - "height": height, - "width": width, - "video_length": video_length, - "sample_shift": sample_shift, - "lora_path": lora_path if lora_path else None, - "strength_model": strength_model, - } - - config = ConfigManager.create_config(model_cls, model_path, task, base_config, overrides) - - return (config,) - - -class LightX2VQuantizationConfig: - """Configuration node for quantization settings.""" - - @classmethod - def INPUT_TYPES(cls): - """Define inputs for quantization config.""" - bridge = LightX2VBridge() - - quant_choices = ["custom"] - for key, config_info in bridge.quantization_registry.items(): - quant_choices.append(key) - - return { - "required": { - "quantization_preset": ( - quant_choices, - { - "default": "custom", - "tooltip": "量化配置预设", - }, - ), - }, - "optional": { - "mm_type": ( - [ - "W-int8-channel-sym-A-int8-channel-sym-dynamic-Vllm", - "W-int8-channel-sym-A-fp16-dynamic-Vllm", - "W-fp16-A-fp16-dynamic-Vllm", - ], - { - "default": "W-int8-channel-sym-A-int8-channel-sym-dynamic-Vllm", - "tooltip": "量化类型", - }, - ), - "dit_quantized_ckpt": ( - "STRING", - {"default": "", "tooltip": "量化模型路径"}, - ), - }, - } - - RETURN_TYPES = ("LIGHTX2V_CONFIG",) - RETURN_NAMES = ("quantization_config",) - FUNCTION = "build_quantization_config" - CATEGORY = "LightX2V/Config" - - def build_quantization_config( - self, - quantization_preset, - mm_type="W-int8-channel-sym-A-int8-channel-sym-dynamic-Vllm", - dit_quantized_ckpt="", - ): - """Build quantization configuration.""" - bridge = LightX2VBridge() - - if quantization_preset == "custom": - config: dict[str, Any] = { - "mm_config": {"mm_type": mm_type}, - } - if dit_quantized_ckpt: - config["dit_quantized_ckpt"] = dit_quantized_ckpt - else: - config_info = bridge.quantization_registry.get(quantization_preset) - if config_info: - config = bridge.load_config(config_info["path"]) - else: - raise ValueError(f"量化配置 '{quantization_preset}' 未找到") - - return (config,) - - -class LightX2VAttentionConfig: - """Configuration node for attention mechanisms.""" - - @classmethod - def INPUT_TYPES(cls): - """Define inputs for attention config.""" - return { - "required": { - "self_attn_1_type": ( - ["flash_attn3", "sage_attn2", "radial_attn", "sparge_attn"], - { - "default": "flash_attn3", - "tooltip": "自注意力类型", - }, - ), - "cross_attn_1_type": ( - ["flash_attn3", "sage_attn2", "radial_attn", "sparge_attn"], - { - "default": "flash_attn3", - "tooltip": "交叉注意力1类型", - }, - ), - "cross_attn_2_type": ( - ["flash_attn3", "sage_attn2", "radial_attn", "sparge_attn"], - { - "default": "flash_attn3", - "tooltip": "交叉注意力2类型", - }, - ), - }, - } - - RETURN_TYPES = ("LIGHTX2V_CONFIG",) - RETURN_NAMES = ("attention_config",) - FUNCTION = "build_attention_config" - CATEGORY = "LightX2V/Config" - - def build_attention_config( - self, - self_attn_1_type, - cross_attn_1_type, - cross_attn_2_type, - ): - """Build attention configuration.""" - config = { - "self_attn_1_type": self_attn_1_type, - "cross_attn_1_type": cross_attn_1_type, - "cross_attn_2_type": cross_attn_2_type, - } - - return (config,) - - -class LightX2VCachingConfig: - """Configuration node for caching mechanisms.""" - - @classmethod - def INPUT_TYPES(cls): - """Define inputs for caching config.""" - return { - "required": { - "feature_caching": ( - ["NoCaching", "Tea", "TaylorSeer", "Ada", "Custom"], - { - "default": "NoCaching", - "tooltip": "特征缓存类型", - }, - ), - }, - "optional": { - "teacache_thresh": ( - "FLOAT", - { - "default": 0.26, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "tooltip": "TeaCache阈值", - }, - ), - "use_ret_steps": ( - "BOOLEAN", - { - "default": True, - "tooltip": "使用返回步骤", - }, - ), - "coefficients": ( - "STRING", - { - "multiline": True, - "default": "", - "tooltip": "系数配置(JSON格式)", - }, - ), - }, - } - - RETURN_TYPES = ("LIGHTX2V_CONFIG",) - RETURN_NAMES = ("caching_config",) - FUNCTION = "build_caching_config" - CATEGORY = "LightX2V/Config" - - def build_caching_config( - self, - feature_caching, - teacache_thresh=0.26, - use_ret_steps=True, - coefficients="", - ): - """Build caching configuration.""" - config = {} - - if feature_caching != "NoCaching": - config["feature_caching"] = feature_caching - - if feature_caching == "Tea": - config["teacache_thresh"] = teacache_thresh - config["use_ret_steps"] = use_ret_steps - - if coefficients: - try: - config["coefficients"] = json.loads(coefficients) - except json.JSONDecodeError: - print(f"Failed to parse coefficients: {coefficients}") - else: - config["feature_caching"] = "NoCaching" - - return (config,) - - -class LightX2VInference: - """LightX2V inference node that generates images directly.""" - - def __init__(self): - self.bridge = LightX2VBridge() - self.converter = InputOutputConverter() - - @classmethod - def INPUT_TYPES(cls): - """Define inputs for the inference node.""" - return { - "required": { - "config": ("LIGHTX2V_CONFIG", {"tooltip": "LightX2V配置"}), - "prompt": ( - "STRING", - { - "multiline": True, - "default": "", - "tooltip": "生成提示词", - }, - ), - "negative_prompt": ( - "STRING", - {"multiline": True, "default": "", "tooltip": "负面提示词"}, - ), - }, - "optional": { - "image": ("IMAGE", {"tooltip": "i2v任务的输入图像"}), - "audio": ( - "AUDIO", - {"tooltip": "音频驱动生成的输入音频"}, - ), - }, - } - - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("images",) - FUNCTION = "generate" - CATEGORY = "LightX2V/Inference" - - def generate( - self, - config, - prompt, - negative_prompt, - image=None, - audio=None, - **kwargs, - ): - config.prompt = prompt - config.negative_prompt = negative_prompt - config.mode = "infer" - - os.environ["TOKENIZERS_PARALLELISM"] = "false" - if "DTYPE" not in os.environ: - os.environ["DTYPE"] = "BF16" - if "ENABLE_GRAPH_MODE" not in os.environ: - os.environ["ENABLE_GRAPH_MODE"] = "false" - if "ENABLE_PROFILING_DEBUG" not in os.environ: - os.environ["ENABLE_PROFILING_DEBUG"] = "true" - - temp_files = [] - - try: - if config.task == "i2v" and image is not None: - pil_image = self.converter.comfy_image_to_pil(image) - with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp: - pil_image.save(tmp.name) - config.image_path = tmp.name - temp_files.append(tmp.name) - elif config.task == "i2v" and image is None: - raise ValueError("i2v task requires input image") - - if audio is not None and "audio" in config.model_cls: - audio_path = self.converter.comfy_audio_to_path(audio) - config.audio_path = audio_path - temp_files.append(audio_path) - - runner = self.bridge.get_runner(config) - - if runner is None: - raise RuntimeError("Failed to initialize runner") - - try: - total_steps = runner.config.get("infer_steps", 20) - progress = ProgressBar(total_steps) - - def update_progress(current_step, total): - progress.update_absolute(current_step) - - runner.set_progress_callback(update_progress) - images = asyncio.run(runner.run_pipeline(save_video=False)) - torch.cuda.empty_cache() - gc.collect() - - return self._decode_latents_to_images(images) - - except Exception as e: - print(f"Error during pipeline execution: {e}") - raise - - except Exception as e: - print(f"Error in LightX2V generation: {e}") - raise - # finally: #TODO: refactor lightx2v input - # for temp_file in temp_files: - # if os.path.exists(temp_file): - # try: - # os.unlink(temp_file) - # except Exception: - # pass - - def _decode_latents_to_images(self, decoded_images): - """Decode latents to images using VAE.""" - - images = (decoded_images + 1) / 2 - images = images.squeeze(0).permute(1, 2, 3, 0).cpu() - images = torch.clamp(images, 0, 1) - - return (images,) - - -class LightX2VDistillConfig: - """Configuration node for distillation settings.""" - - @classmethod - def INPUT_TYPES(cls): - """Define inputs for distillation config.""" - return { - "required": { - "enable_distill": ( - "BOOLEAN", - { - "default": True, - "tooltip": "启用蒸馏模型", - }, - ), - "infer_steps": ( - "INT", - { - "default": 4, - "min": 1, - "max": 50, - "tooltip": "推理步数(蒸馏模型通常使用4步)", - }, - ), - "denoising_steps": ( - "STRING", - { - "default": "[999, 750, 500, 250]", - "tooltip": "去噪步骤列表(JSON格式)", - }, - ), - }, - "optional": { - "enable_cfg": ( - "BOOLEAN", - { - "default": False, - "tooltip": "启用CFG(分类器自由引导)", - }, - ), - "enable_dynamic_cfg": ( - "BOOLEAN", - { - "default": False, - "tooltip": "启用动态CFG", - }, - ), - "cfg_scale": ( - "FLOAT", - { - "default": 4.0, - "min": 0.1, - "max": 30.0, - "step": 0.1, - "tooltip": "动态CFG缩放比例", - }, - ), - }, - } - - RETURN_TYPES = ("LIGHTX2V_CONFIG",) - RETURN_NAMES = ("distill_config",) - FUNCTION = "build_distill_config" - CATEGORY = "LightX2V/Config" - - def build_distill_config( - self, - enable_distill, - infer_steps=4, - denoising_steps="[999, 750, 500, 250]", - enable_cfg=False, - enable_dynamic_cfg=False, - cfg_scale=4.0, - ): - """Build distillation configuration.""" - config = {} - - if enable_distill: - config["infer_steps"] = infer_steps - config["enable_cfg"] = enable_cfg - config["enable_dynamic_cfg"] = enable_dynamic_cfg - - if enable_dynamic_cfg: - config["cfg_scale"] = cfg_scale - - try: - config["denoising_step_list"] = json.loads(denoising_steps) - if len(config["denoising_step_list"]) != infer_steps: - print(f"Warning: denoising_step_list length ({len(config['denoising_step_list'])}) doesn't match infer_steps ({infer_steps})") - except json.JSONDecodeError: - print(f"Failed to parse denoising steps, using default: {denoising_steps}") - config["denoising_step_list"] = [999, 750, 500, 250] - - return (config,) - - -# Node class mapping -NODE_CLASS_MAPPINGS = { - # Modular nodes - "LightX2VConfigBuilder": LightX2VConfigBuilder, - "LightX2VQuantizationConfig": LightX2VQuantizationConfig, - "LightX2VAttentionConfig": LightX2VAttentionConfig, - "LightX2VCachingConfig": LightX2VCachingConfig, - "LightX2VDistillConfig": LightX2VDistillConfig, - "LightX2VInference": LightX2VInference, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - # Modular nodes - "LightX2VConfigBuilder": "LightX2V Config Builder", - "LightX2VQuantizationConfig": "LightX2V Quantization Config", - "LightX2VAttentionConfig": "LightX2V Attention Config", - "LightX2VCachingConfig": "LightX2V Caching Config", - "LightX2VDistillConfig": "LightX2V Distill Config", - "LightX2VInference": "LightX2V Inference", -} diff --git a/nodes.py b/nodes.py index 593adfe..26eb51b 100644 --- a/nodes.py +++ b/nodes.py @@ -1,13 +1,426 @@ -# # Import refactored modules -# from .lightx2v_nodes.nodes import ( -# NODE_CLASS_MAPPINGS, -# NODE_DISPLAY_NAME_MAPPINGS, -# ) +"""Modular ComfyUI nodes for LightX2V without presets.""" -# # Export the mappings -# __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] +import asyncio +import gc +import logging +import os +import tempfile +from typing import Any, Dict + +import numpy as np +import torch +from comfy.utils import ProgressBar +from PIL import Image + +from .bridge import ModularConfigManager, get_available_attn_ops, get_available_quant_ops +from .lightx2v.lightx2v.infer import init_runner -from .lightx2v_nodes.universal_bridge import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +class LightX2VInferenceConfig: + """Basic inference configuration node.""" -__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model_cls": (["wan2.1", "hunyuan"], {"default": "wan2.1", "tooltip": "模型类型"}), + "model_path": ("STRING", {"default": "", "tooltip": "模型路径"}), + "task": (["t2v", "i2v"], {"default": "t2v", "tooltip": "任务类型:文本到视频或图像到视频"}), + "infer_steps": ("INT", {"default": 40, "min": 1, "max": 100, "tooltip": "推理步数"}), + "seed": ("INT", {"default": 42, "min": -1, "max": 2**32 - 1, "tooltip": "随机种子,-1为随机"}), + "cfg_scale": ("FLOAT", {"default": 5.0, "min": 1.0, "max": 10.0, "step": 0.1, "tooltip": "CFG引导强度"}), + "sample_shift": ("INT", {"default": 5, "min": 0, "max": 10, "tooltip": "采样偏移"}), + "height": ("INT", {"default": 480, "min": 64, "max": 2048, "step": 8, "tooltip": "视频高度"}), + "width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "视频宽度"}), + "video_length": ("INT", {"default": 81, "min": 16, "max": 120, "tooltip": "视频帧数"}), + "fps": ("INT", {"default": 16, "min": 8, "max": 30, "tooltip": "每秒帧数"}), + } + } + + RETURN_TYPES = ("INFERENCE_CONFIG",) + RETURN_NAMES = ("inference_config",) + FUNCTION = "create_config" + CATEGORY = "LightX2V/Config" + + def create_config(self, model_cls, model_path, task, infer_steps, seed, cfg_scale, sample_shift, height, width, video_length, fps): + """Create basic inference configuration.""" + config = { + "model_cls": model_cls, + "model_path": model_path, + "task": task, + "infer_steps": infer_steps, + "seed": seed if seed != -1 else np.random.randint(0, 2**32 - 1), + "cfg_scale": cfg_scale, + "sample_shift": sample_shift, + "height": height, + "width": width, + "video_length": video_length, + "fps": fps, + } + return (config,) + + +class LightX2VTeaCache: + """TeaCache configuration node.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "enable": ("BOOLEAN", {"default": False, "tooltip": "启用TeaCache特征缓存"}), + "threshold": ( + "FLOAT", + {"default": 0.26, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "缓存阈值,越低加速越多:0.1约2倍加速,0.2约3倍加速"}, + ), + "cache_key_steps_only": ("BOOLEAN", {"default": False, "tooltip": "只缓存关键步骤以平衡质量和速度"}), + } + } + + RETURN_TYPES = ("TEACACHE_CONFIG",) + RETURN_NAMES = ("teacache_config",) + FUNCTION = "create_config" + CATEGORY = "LightX2V/Config" + + def create_config(self, enable, threshold, cache_key_steps_only): + """Create TeaCache configuration.""" + config = { + "enable": enable, + "threshold": threshold, + "cache_key_steps_only": cache_key_steps_only, + } + return (config,) + + +class LightX2VQuantization: + """Quantization configuration node.""" + + @classmethod + def INPUT_TYPES(cls): + # Get available quantization backends + available_ops = get_available_quant_ops() + quant_backends = [] + + for op_name, is_available in available_ops: + if is_available: + quant_backends.append(op_name) + + # Always have at least one option + if not quant_backends: + quant_backends = ["none"] + + return { + "required": { + "dit_precision": (["bf16", "int8", "fp8"], {"default": "bf16", "tooltip": "DIT模型量化精度"}), + "t5_precision": (["bf16", "int8", "fp8"], {"default": "bf16", "tooltip": "T5编码器量化精度"}), + "clip_precision": (["fp16", "int8", "fp8"], {"default": "fp16", "tooltip": "CLIP编码器量化精度"}), + "quant_backend": (quant_backends, {"default": quant_backends[0], "tooltip": "量化计算后端"}), + "sensitive_layers_precision": (["fp32", "bf16"], {"default": "fp32", "tooltip": "敏感层(归一化和嵌入层)精度"}), + } + } + + RETURN_TYPES = ("QUANT_CONFIG",) + RETURN_NAMES = ("quantization_config",) + FUNCTION = "create_config" + CATEGORY = "LightX2V/Config" + + def create_config(self, dit_precision, t5_precision, clip_precision, quant_backend, sensitive_layers_precision): + """Create quantization configuration.""" + config = { + "dit_precision": dit_precision, + "t5_precision": t5_precision, + "clip_precision": clip_precision, + "quant_backend": quant_backend, + "sensitive_layers_precision": sensitive_layers_precision, + } + return (config,) + + +class LightX2VMemoryOptimization: + """Memory optimization configuration node.""" + + @classmethod + def INPUT_TYPES(cls): + # Get available attention types + available_attn = get_available_attn_ops() + attn_types = [] + + for op_name, is_available in available_attn: + if is_available: + attn_types.append(op_name) + + # Always include fallback + if "torch_sdpa" not in attn_types: + attn_types.append("torch_sdpa") + + return { + "required": { + "optimization_level": ( + ["none", "low", "medium", "high", "extreme"], + {"default": "none", "tooltip": "内存优化级别,越高越省内存但可能影响速度"}, + ), + "attention_type": (attn_types, {"default": attn_types[0], "tooltip": "注意力机制类型"}), + }, + "optional": { + # GPU optimization + "enable_rotary_chunk": ("BOOLEAN", {"default": False, "tooltip": "启用旋转编码分块"}), + "rotary_chunk_size": ("INT", {"default": 100, "min": 100, "max": 10000, "step": 100}), + "clean_cuda_cache": ("BOOLEAN", {"default": False, "tooltip": "及时清理CUDA缓存"}), + # CPU offloading + "enable_cpu_offload": ("BOOLEAN", {"default": False, "tooltip": "启用CPU卸载"}), + "offload_granularity": (["block", "phase"], {"default": "phase", "tooltip": "卸载粒度"}), + "offload_ratio": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.1}), + # Module management + "lazy_load": ("BOOLEAN", {"default": False, "tooltip": "延迟加载模型"}), + "unload_after_inference": ("BOOLEAN", {"default": False, "tooltip": "推理后卸载模块"}), + }, + } + + RETURN_TYPES = ("MEMORY_CONFIG",) + RETURN_NAMES = ("memory_config",) + FUNCTION = "create_config" + CATEGORY = "LightX2V/Config" + + def create_config( + self, + optimization_level, + attention_type, + enable_rotary_chunk=False, + rotary_chunk_size=100, + clean_cuda_cache=False, + enable_cpu_offload=False, + offload_granularity="phase", + offload_ratio=1.0, + lazy_load=False, + unload_after_inference=False, + ): + """Create memory optimization configuration.""" + config = { + "optimization_level": optimization_level, + "attention_type": attention_type, + "enable_rotary_chunk": enable_rotary_chunk, + "rotary_chunk_size": rotary_chunk_size, + "clean_cuda_cache": clean_cuda_cache, + "enable_cpu_offload": enable_cpu_offload, + "offload_granularity": offload_granularity, + "offload_ratio": offload_ratio, + "lazy_load": lazy_load, + "unload_after_inference": unload_after_inference, + } + return (config,) + + +class LightX2VLightweightVAE: + """Lightweight VAE configuration node.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "use_tiny_vae": ("BOOLEAN", {"default": False, "tooltip": "使用轻量级VAE加速解码"}), + "use_tiling_vae": ("BOOLEAN", {"default": False, "tooltip": "使用VAE分块推理减少显存"}), + } + } + + RETURN_TYPES = ("VAE_CONFIG",) + RETURN_NAMES = ("vae_config",) + FUNCTION = "create_config" + CATEGORY = "LightX2V/Config" + + def create_config(self, use_tiny_vae, use_tiling_vae): + """Create VAE configuration.""" + config = { + "use_tiny_vae": use_tiny_vae, + "use_tiling_vae": use_tiling_vae, + } + return (config,) + + +class LightX2VModularInference: + """Modular inference node that combines all configurations.""" + + def __init__(self): + self.config_manager = ModularConfigManager() + self._current_runner = None + self._current_config_hash = None + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "inference_config": ("INFERENCE_CONFIG", {"tooltip": "基础推理配置"}), + "prompt": ("STRING", {"multiline": True, "default": "", "tooltip": "生成提示词"}), + "negative_prompt": ("STRING", {"multiline": True, "default": "", "tooltip": "负面提示词"}), + }, + "optional": { + "image": ("IMAGE", {"tooltip": "i2v任务的输入图像"}), + "teacache_config": ("TEACACHE_CONFIG", {"tooltip": "TeaCache配置"}), + "quantization_config": ("QUANT_CONFIG", {"tooltip": "量化配置"}), + "memory_config": ("MEMORY_CONFIG", {"tooltip": "内存优化配置"}), + "vae_config": ("VAE_CONFIG", {"tooltip": "VAE配置"}), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "generate" + CATEGORY = "LightX2V/Inference" + + def _get_config_hash(self, configs: Dict[str, Any]) -> str: + """Generate a hash for configuration to detect changes.""" + import hashlib + import json + + # Only hash model-related configs + relevant_configs = { + "model_cls": configs.get("inference", {}).get("model_cls"), + "model_path": configs.get("inference", {}).get("model_path"), + "quantization": configs.get("quantization"), + "memory_lazy_load": configs.get("memory", {}).get("lazy_load"), + } + + config_str = json.dumps(relevant_configs, sort_keys=True) + return hashlib.md5(config_str.encode()).hexdigest() + + def generate( + self, + inference_config, + prompt, + negative_prompt, + image=None, + teacache_config=None, + quantization_config=None, + memory_config=None, + vae_config=None, + **kwargs, + ): + """Generate video using modular configuration.""" + + # Set environment variables + os.environ["TOKENIZERS_PARALLELISM"] = "false" + if "DTYPE" not in os.environ: + os.environ["DTYPE"] = "BF16" + if "ENABLE_GRAPH_MODE" not in os.environ: + os.environ["ENABLE_GRAPH_MODE"] = "false" + if "ENABLE_PROFILING_DEBUG" not in os.environ: + os.environ["ENABLE_PROFILING_DEBUG"] = "true" + + # Collect all configurations + configs = { + "inference": inference_config, + } + + if teacache_config: + configs["teacache"] = teacache_config + if quantization_config: + configs["quantization"] = quantization_config + if memory_config: + configs["memory"] = memory_config + if vae_config: + configs["vae"] = vae_config + + # Build final configuration + config = self.config_manager.build_final_config(configs) + + # Add prompt and negative prompt + config.prompt = prompt + config.negative_prompt = negative_prompt + + # Check if task requires image + if config.task == "i2v" and image is None: + raise ValueError("i2v task requires input image") + + temp_files = [] + + try: + # Handle image input for i2v + if config.task == "i2v" and image is not None: + # Convert ComfyUI image to PIL + image_np = (image[0].cpu().numpy() * 255).astype(np.uint8) + pil_image = Image.fromarray(image_np) + + # Save to temporary file + with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp: + pil_image.save(tmp.name) + config.image_path = tmp.name + temp_files.append(tmp.name) + + # Check if we need to reinitialize runner + config_hash = self._get_config_hash(configs) + needs_reinit = ( + self._current_runner is None or self._current_config_hash != config_hash or configs.get("memory", {}).get("lazy_load", False) + ) + + if needs_reinit: + # Clear old runner + if self._current_runner is not None: + del self._current_runner + torch.cuda.empty_cache() + gc.collect() + + # Initialize new runner + self._current_runner = init_runner(config) + self._current_config_hash = config_hash + else: + # Update config for existing runner + self._current_runner.config = config + + # Set up progress callback + total_steps = config.get("infer_steps", 40) + progress = ProgressBar(total_steps) + + def update_progress(current_step, total): + progress.update_absolute(current_step) + + self._current_runner.set_progress_callback(update_progress) + + # Run inference + images = asyncio.run(self._current_runner.run_pipeline(save_video=False)) + + # Clean up if requested + if configs.get("memory", {}).get("unload_after_inference", False): + del self._current_runner + self._current_runner = None + self._current_config_hash = None + + torch.cuda.empty_cache() + gc.collect() + + # Convert output to ComfyUI format + images = (images + 1) / 2 + images = images.squeeze(0).permute(1, 2, 3, 0).cpu() + images = torch.clamp(images, 0, 1) + + return (images,) + + except Exception as e: + logging.error(f"Error during inference: {e}") + raise + + finally: + # Clean up temporary files + for temp_file in temp_files: + if os.path.exists(temp_file): + try: + os.unlink(temp_file) + except Exception: + pass + + +# Node mappings +NODE_CLASS_MAPPINGS = { + "LightX2VInferenceConfig": LightX2VInferenceConfig, + "LightX2VTeaCache": LightX2VTeaCache, + "LightX2VQuantization": LightX2VQuantization, + "LightX2VMemoryOptimization": LightX2VMemoryOptimization, + "LightX2VLightweightVAE": LightX2VLightweightVAE, + "LightX2VModularInference": LightX2VModularInference, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LightX2VInferenceConfig": "LightX2V 推理配置", + "LightX2VTeaCache": "LightX2V TeaCache缓存", + "LightX2VQuantization": "LightX2V 低精度量化", + "LightX2VMemoryOptimization": "LightX2V 内存优化", + "LightX2VLightweightVAE": "LightX2V 轻量VAE", + "LightX2VModularInference": "LightX2V 模块化推理", +}