refactor: streamline LightX2V module by removing unused files and consolidating configuration management into bridge.py for improved maintainability and clarity

This commit is contained in:
gaclove
2025-07-15 20:30:34 +08:00
parent e5a3eca0ce
commit 04170f4ae7
10 changed files with 876 additions and 2574 deletions
+3 -1
View File
@@ -1,2 +1,4 @@
line-length = 150
indent-width = 4
indent-width = 4
extend-select = ["I"]
+1 -9
View File
@@ -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"]
+450
View File
@@ -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)
-44
View File
@@ -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",
]
-232
View File
@@ -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
-237
View File
@@ -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
-179
View File
@@ -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
-855
View File
@@ -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",
}
File diff suppressed because it is too large Load Diff
+422 -9
View File
@@ -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 模块化推理",
}