feat: add audio processing utilities in LightX2VConfigCombinerV2 for MP3 extension, duration retrieval, and background noise generation

This commit is contained in:
LazyBusyYang
2025-11-24 08:36:43 +00:00
parent d89597c2c1
commit 91d0c06847
+380 -8
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
@@ -952,6 +955,321 @@ class LightX2VConfigCombinerV2:
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 {
@@ -1058,22 +1376,24 @@ class LightX2VConfigCombinerV2:
# Handle talk objects
if hasattr(config, "talk_objects") and config.talk_objects:
talk_objects = config.talk_objects
processed_talk_objects = []
src_talk_objects = []
for talk_obj in talk_objects:
processed_obj = {}
src_obj = {}
if "audio" in talk_obj:
processed_obj["audio"] = talk_obj["audio"]
src_obj["audio"] = talk_obj["audio"]
if "mask" in talk_obj:
processed_obj["mask"] = talk_obj["mask"]
src_obj["mask"] = talk_obj["mask"]
if "audio" in processed_obj:
processed_talk_objects.append(processed_obj)
if "audio" in src_obj:
src_talk_objects.append(src_obj)
# Resolve paths and download URLs
for obj in processed_talk_objects:
# 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"]
@@ -1094,6 +1414,10 @@ class LightX2VConfigCombinerV2:
# 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"]
@@ -1116,6 +1440,54 @@ class LightX2VConfigCombinerV2:
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"]