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:
+3
-1
@@ -1,2 +1,4 @@
|
||||
line-length = 150
|
||||
indent-width = 4
|
||||
indent-width = 4
|
||||
|
||||
extend-select = ["I"]
|
||||
|
||||
+1
-9
@@ -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"]
|
||||
|
||||
@@ -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)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
@@ -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 模块化推理",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user