1665 lines
61 KiB
Python
1665 lines
61 KiB
Python
import gc
|
|
import io
|
|
import json
|
|
import logging
|
|
import os
|
|
import subprocess as sp
|
|
import wave
|
|
|
|
import numpy as np
|
|
import torch
|
|
from comfy.utils import ProgressBar
|
|
from PIL import Image
|
|
|
|
from .bridge import get_available_attn_ops, get_available_quant_ops
|
|
from .config_builder import (
|
|
ConfigBuilder,
|
|
InferenceConfigBuilder,
|
|
LoRAChainBuilder,
|
|
TalkObjectConfigBuilder,
|
|
)
|
|
from .data_models import (
|
|
InferenceConfig,
|
|
MemoryOptimizationConfig,
|
|
QuantizationConfig,
|
|
TalkObjectsConfig,
|
|
TeaCacheConfig,
|
|
)
|
|
from .file_handlers import (
|
|
AudioFileHandler,
|
|
ComfyUIFileResolver,
|
|
HTTPFileDownloader,
|
|
ImageFileHandler,
|
|
TempFileManager,
|
|
)
|
|
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
|
|
from .model_utils import scan_loras, 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(),)
|
|
|
|
|
|
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,)
|
|
|
|
|
|
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,)
|
|
|
|
|
|
class LightX2VConfigCombiner:
|
|
def __init__(self):
|
|
self.config_builder = ConfigBuilder()
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"inference_config": (
|
|
"INFERENCE_CONFIG",
|
|
{"tooltip": "Basic inference configuration"},
|
|
),
|
|
},
|
|
"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"}),
|
|
},
|
|
}
|
|
|
|
RETURN_TYPES = ("COMBINED_CONFIG",)
|
|
RETURN_NAMES = ("combined_config",)
|
|
FUNCTION = "combine_configs"
|
|
CATEGORY = "LightX2V/Config"
|
|
|
|
def combine_configs(
|
|
self,
|
|
inference_config,
|
|
teacache_config=None,
|
|
quantization_config=None,
|
|
memory_config=None,
|
|
lora_chain=None,
|
|
talk_objects_config=None,
|
|
):
|
|
"""Combine multiple configurations into final config."""
|
|
# Convert dict configs back to objects if needed
|
|
|
|
# Create objects from dicts
|
|
inf_config = InferenceConfig(**inference_config) if isinstance(inference_config, dict) else None
|
|
tea_config = TeaCacheConfig(**teacache_config) if teacache_config and isinstance(teacache_config, dict) else None
|
|
quant_config = QuantizationConfig(**quantization_config) if quantization_config and isinstance(quantization_config, dict) else None
|
|
mem_config = MemoryOptimizationConfig(**memory_config) if memory_config and isinstance(memory_config, dict) else None
|
|
|
|
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,
|
|
)
|
|
|
|
return (config,)
|
|
|
|
|
|
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."""
|
|
|
|
# Convert dict configs back to objects if needed
|
|
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
|
|
|
|
# Build combined 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,
|
|
)
|
|
|
|
# Add prompts to config
|
|
config.prompt = prompt
|
|
config.negative_prompt = negative_prompt
|
|
|
|
# Validate task requirements
|
|
if config.task in ["i2v", "s2v", "rs2v"] and image is None:
|
|
raise ValueError("i2v or s2v or rs2v task requires input image")
|
|
|
|
# Handle image input
|
|
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}")
|
|
|
|
# Handle audio input for seko models
|
|
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}")
|
|
|
|
# Handle talk objects
|
|
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)
|
|
|
|
# Resolve paths and download URLs
|
|
for obj in processed_talk_objects:
|
|
if "audio" in obj and obj["audio"]:
|
|
audio_path = obj["audio"]
|
|
|
|
# Check if it's a URL and download if needed
|
|
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
|
|
# Handle relative paths
|
|
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']}")
|
|
|
|
# Check if file exists
|
|
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"]
|
|
|
|
# Check if it's a URL and download if needed
|
|
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}")
|
|
# Don't skip the object if mask download fails (mask is optional)
|
|
# Handle relative paths
|
|
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']}")
|
|
|
|
# Check if file exists
|
|
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:
|
|
"""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()
|
|
|
|
@staticmethod
|
|
def extend_mp3(input_path: str, output_path: str, duration: float) -> bool:
|
|
"""Extend or truncate MP3 audio file.
|
|
|
|
Extend or truncate the input audio based on its duration and target
|
|
duration:
|
|
- 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
|
|
|
|
|
|
Args:
|
|
input_path (str):
|
|
Path to the input MP3 file.
|
|
output_path (str):
|
|
Path to the output MP3 file.
|
|
duration (float):
|
|
Target duration in seconds.
|
|
|
|
Returns:
|
|
bool:
|
|
Returns True if the operation succeeds.
|
|
|
|
Raises:
|
|
ValueError:
|
|
Raised when input audio duration exceeds duration + 0.1.
|
|
"""
|
|
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 as e:
|
|
raise ValueError(f"Failed to parse audio information: {input_path}")
|
|
except Exception as e:
|
|
raise
|
|
|
|
@staticmethod
|
|
def get_audio_duration(input_path: str) -> float:
|
|
"""Get the duration of an audio file.
|
|
|
|
Uses ffprobe to extract audio stream information and returns the
|
|
duration in seconds.
|
|
|
|
|
|
Args:
|
|
input_path (str):
|
|
Path to the audio file.
|
|
|
|
Returns:
|
|
float:
|
|
Audio duration in seconds.
|
|
|
|
Raises:
|
|
ValueError:
|
|
Raised when audio stream information cannot be retrieved or
|
|
parsed.
|
|
"""
|
|
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))
|
|
return input_duration
|
|
|
|
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.
|
|
|
|
Generate white noise audio data with optional normalization using
|
|
RMS or standard deviation. The noise is generated using a normal
|
|
distribution and can be normalized to a target RMS value or standard
|
|
deviation.
|
|
|
|
|
|
Args:
|
|
duration (float):
|
|
Audio duration in seconds.
|
|
framerate (int):
|
|
Sample rate in Hz.
|
|
n_channels (int, optional):
|
|
Number of audio channels. Defaults to 1 (mono).
|
|
rms (float, optional):
|
|
Target RMS value for normalization. If provided, the noise
|
|
will be normalized to this RMS value. Defaults to None.
|
|
std_dev (float, optional):
|
|
Target standard deviation for normalization. If provided, the
|
|
noise will be normalized to this standard deviation.
|
|
Defaults to None.
|
|
seed (int, optional):
|
|
Random seed for reproducible generation. Defaults to None.
|
|
|
|
Returns:
|
|
np.ndarray:
|
|
Generated audio data with shape (n_samples, n_channels) for
|
|
multi-channel or (n_samples,) for mono channel, where
|
|
n_samples = duration * framerate.
|
|
"""
|
|
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: str | io.BytesIO, framerate: int, sample_width: int = 2) -> None:
|
|
"""Save audio data as WAV file or BytesIO object.
|
|
|
|
Convert normalized float audio data to integer format and save as
|
|
WAV file. Supports mono and multi-channel audio with configurable
|
|
sample width.
|
|
|
|
|
|
Args:
|
|
audio_data (np.ndarray):
|
|
Audio data with shape (n_samples,) for mono or
|
|
(n_samples, n_channels) for multi-channel. Values should
|
|
be in the range [-1.0, 1.0].
|
|
output_path (str | io.BytesIO):
|
|
Output file path as string or BytesIO object.
|
|
framerate (int):
|
|
Sample rate in Hz.
|
|
sample_width (int, optional):
|
|
Sample width in bytes. Supported values are 1 (8-bit),
|
|
2 (16-bit), and 4 (32-bit). Defaults to 2.
|
|
"""
|
|
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: list[str]) -> io.BytesIO:
|
|
"""Generate background mask from positive mask images.
|
|
|
|
Generate a background mask by finding pixels that are zero (or
|
|
below threshold) in all input positive mask images. The resulting
|
|
mask marks background regions (all masks are zero) as white (255)
|
|
and foreground regions (any mask has non-zero values) as black (0).
|
|
|
|
|
|
Args:
|
|
positive_mask_paths (list[str]):
|
|
List of paths to positive mask image files. All images
|
|
must have the same width and height.
|
|
|
|
Returns:
|
|
io.BytesIO:
|
|
BytesIO object containing the background mask image in JPEG
|
|
format. The mask is a grayscale image where white (255)
|
|
represents background regions and black (0) represents
|
|
foreground regions.
|
|
|
|
Raises:
|
|
ValueError:
|
|
Raised when mask images have different dimensions.
|
|
"""
|
|
width = None
|
|
height = None
|
|
opened_imgs = list()
|
|
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."""
|
|
|
|
# Convert dict configs back to objects if needed
|
|
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
|
|
|
|
# Build combined 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,
|
|
)
|
|
|
|
# Add prompts to config
|
|
config.prompt = prompt
|
|
config.negative_prompt = negative_prompt
|
|
|
|
# Validate task requirements
|
|
if config.task in ["i2v", "s2v", "rs2v"] and image is None:
|
|
raise ValueError("i2v or s2v or rs2v task requires input image")
|
|
|
|
# Handle image input
|
|
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}")
|
|
|
|
# Handle audio input for seko models
|
|
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}")
|
|
|
|
# Handle talk objects
|
|
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 and download URLs,
|
|
# record the max duration of the src talk objects
|
|
max_src_duration = None
|
|
for obj in src_talk_objects:
|
|
if "audio" in obj and obj["audio"]:
|
|
audio_path = obj["audio"]
|
|
|
|
# Check if it's a URL and download if needed
|
|
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
|
|
# Handle relative paths
|
|
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']}")
|
|
|
|
# Check if file exists
|
|
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"]
|
|
|
|
# Check if it's a URL and download if needed
|
|
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}")
|
|
# Don't skip the object if mask download fails (mask is optional)
|
|
# Handle relative paths
|
|
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']}")
|
|
|
|
# Check if file exists
|
|
if not os.path.exists(obj["mask"]):
|
|
logging.warning(f"Mask file not found: {obj['mask']}")
|
|
|
|
if len(src_talk_objects) > 1:
|
|
# extend audio duration to the max duration of the src talk objects
|
|
processed_talk_objects: list[dict[str, str]] = list()
|
|
mask_img_paths = list()
|
|
extend_count = 0
|
|
for obj in src_talk_objects:
|
|
dst_obj = dict()
|
|
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")
|
|
# generate background mask and audio
|
|
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())
|
|
bg_obj = dict(
|
|
audio=bg_audio_path,
|
|
mask=bg_mask_path,
|
|
)
|
|
processed_talk_objects.append(bg_obj)
|
|
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,)
|
|
|
|
|
|
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)
|
|
# return dict(
|
|
# seed=config.get("seed", 42),
|
|
# image_path=config.get("image_path", ""),
|
|
# audio_path=config.get("audio_path", ""),
|
|
# prompt=config.get("prompt", ""),
|
|
# negative_prompt=config.get("negative_prompt", ""),
|
|
# save_result_path=config.get("save_result_path", ""),
|
|
# clip_configs=clip_configs,
|
|
# target_shape=config.get("target_shape", []),
|
|
# )
|
|
|
|
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:
|
|
# current_runner.end_run()
|
|
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":
|
|
# rs2v 使用 shot_runner 管线
|
|
# current_runner.set_config(config)
|
|
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
|
|
|
|
finally:
|
|
# Cleanup is handled by TempFileManager destructor
|
|
pass
|
|
|
|
|
|
NODE_CLASS_MAPPINGS = {
|
|
"LightX2VInferenceConfig": LightX2VInferenceConfig,
|
|
"LightX2VTeaCache": LightX2VTeaCache,
|
|
"LightX2VQuantization": LightX2VQuantization,
|
|
"LightX2VMemoryOptimization": LightX2VMemoryOptimization,
|
|
"LightX2VLoRALoader": LightX2VLoRALoader,
|
|
"LightX2VConfigCombiner": LightX2VConfigCombiner,
|
|
"LightX2VConfigCombinerV2": LightX2VConfigCombinerV2,
|
|
"LightX2VConfigCombinerV3": LightX2VConfigCombinerV3,
|
|
"LightX2VModularInferenceV2": LightX2VModularInferenceV2,
|
|
"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",
|
|
"LightX2VConfigCombiner": "LightX2V Config Combiner",
|
|
"LightX2VConfigCombinerV2": "LightX2V Config Combiner V2",
|
|
"LightX2VConfigCombinerV3": "LightX2V Config Combiner V3",
|
|
"LightX2VModularInferenceV2": "LightX2V Modular Inference V2",
|
|
"LightX2VTalkObjectInput": "LightX2V Talk Object Input (Single)",
|
|
"LightX2VTalkObjectsCombiner": "LightX2V Talk Objects Combiner",
|
|
"LightX2VTalkObjectsFromFiles": "LightX2V Talk Objects From Files",
|
|
"LightX2VTalkObjectsFromJSON": "LightX2V Talk Objects From JSON (API)",
|
|
}
|