feat(seedvr2): add super-resolution nodes and modularize wrapper

This commit is contained in:
gaclove
2026-08-31 17:21:26 +08:00
parent 4958316d52
commit fc1495f8a0
23 changed files with 2509 additions and 2218 deletions
+40
View File
@@ -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
View File
@@ -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
-543
View File
@@ -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)
+41
View File
@@ -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",
]
+80
View File
@@ -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)
+80
View File
@@ -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",
}
+64
View File
@@ -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}"
)
+26
View File
@@ -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",
]
+85
View File
@@ -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
+48
View File
@@ -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
+111
View File
@@ -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)
+25
View File
@@ -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
+35
View File
@@ -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
View File
@@ -42,6 +42,7 @@ def support_model_cls_list() -> List[str]:
"wan2.2_audio",
"wan2.2_moe_distill",
"qwen_image",
"seedvr2",
]
-1664
View File
File diff suppressed because it is too large Load Diff
+64
View File
@@ -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"]
+639
View File
@@ -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
View File
@@ -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(),)
+148
View File
@@ -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
+48
View File
@@ -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
View File
@@ -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
View File
@@ -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,)