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