Merge remote-tracking branch 'origin/fix_sekotalk_input'

This commit is contained in:
gaclove
2025-11-25 06:49:49 +00:00
+565
View File
@@ -1,7 +1,10 @@
import gc
import io
import json
import logging
import os
import subprocess as sp
import wave
import numpy as np
import torch
@@ -1132,6 +1135,566 @@ class LightX2VConfigCombinerV2:
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)
data = json.loads(output.decode("utf-8"))
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)
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)
data = json.loads(output.decode("utf-8"))
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 task"}),
"audio": (
"AUDIO",
{"tooltip": "Input audio for audio-driven generation for s2v 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"] and image is None:
raise ValueError("i2v or s2v task requires input image")
# Handle image input
if config.task in ["i2v", "s2v"] 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."""
@@ -1244,6 +1807,7 @@ NODE_CLASS_MAPPINGS = {
"LightX2VConfigCombiner": LightX2VConfigCombiner,
"LightX2VModularInference": LightX2VModularInference,
"LightX2VConfigCombinerV2": LightX2VConfigCombinerV2,
"LightX2VConfigCombinerV3": LightX2VConfigCombinerV3,
"LightX2VModularInferenceV2": LightX2VModularInferenceV2,
"LightX2VTalkObjectInput": TalkObjectInput,
"LightX2VTalkObjectsCombiner": TalkObjectsCombiner,
@@ -1260,6 +1824,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"LightX2VConfigCombiner": "LightX2V Config Combiner",
"LightX2VModularInference": "LightX2V Modular Inference",
"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",