Add DG_VideoAudioMixer
This commit is contained in:
@@ -1,8 +1,10 @@
|
||||
import os
|
||||
import tempfile
|
||||
import torchaudio
|
||||
import uuid
|
||||
import sys
|
||||
import shutil
|
||||
from collections.abc import Mapping
|
||||
|
||||
# Function to find ComfyUI directories
|
||||
def get_comfyui_temp_dir():
|
||||
@@ -733,14 +735,338 @@ class VideoLengthAdjuster:
|
||||
{"waveform": padded_audio.unsqueeze(0), "sample_rate": sample_rate}
|
||||
)
|
||||
|
||||
class DG_VideoAudioMixer:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images1": ("IMAGE", ),
|
||||
"video_info1": ("VHS_VIDEOINFO", ),
|
||||
"images2": ("IMAGE", ),
|
||||
"video_info2": ("VHS_VIDEOINFO", ),
|
||||
"bgm": ("AUDIO", ), # Add BGM as required input first
|
||||
"bgm_volume": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.05}),
|
||||
},
|
||||
"optional": {
|
||||
"audio1": ("AUDIO", ),
|
||||
"audio2": ("AUDIO", ),
|
||||
"fade_in_sec": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.1}),
|
||||
}
|
||||
}
|
||||
|
||||
CATEGORY = "LatentSyncNode"
|
||||
FUNCTION = "VideoAudioMixer"
|
||||
TITLE = "DG_VideoAudioMixer"
|
||||
RETURN_NAMES = ("images_output", "audio_output", "video_info_output")
|
||||
RETURN_TYPES = ("IMAGE", "AUDIO", "VHS_VIDEOINFO")
|
||||
|
||||
def VideoAudioMixer(self, images1, video_info1, images2, video_info2, bgm=None, bgm_volume=0.3, audio1=None, audio2=None, fade_in_sec=1.0):
|
||||
print(f"DEBUG: bgm={bgm is not None}, bgm_volume={bgm_volume}, fade_in_sec={fade_in_sec}")
|
||||
|
||||
# Verify frames
|
||||
if not isinstance(images1, torch.Tensor) or not isinstance(images2, torch.Tensor):
|
||||
raise ValueError("images1 and images2 must be frame tensors")
|
||||
|
||||
# Handle frame dimensions (assume [frames, h, w, c])
|
||||
if images1.shape[1:] != images2.shape[1:]:
|
||||
raise ValueError(f"Incompatible resolutions: images1 {images1.shape}, images2 {images2.shape}")
|
||||
|
||||
# Concatenate frames
|
||||
concatenated_frames = torch.cat([images1, images2], dim=0) # [frames_total, h, w, c]
|
||||
|
||||
# Extract FPS from video_info
|
||||
fps1 = video_info1["loaded_fps"]
|
||||
fps2 = video_info2["loaded_fps"]
|
||||
if fps1 != fps2:
|
||||
print(f"Warning: Different FPS (video1: {fps1}, video2: {fps2}), using {fps1}")
|
||||
output_fps = fps1
|
||||
|
||||
# Handle audio
|
||||
audio_output = None
|
||||
sample_rate = None
|
||||
|
||||
# Function to extract waveform and sample_rate from audio
|
||||
def get_audio_data(audio_input, label=""):
|
||||
print(f"{label}: Type of audio_input = {type(audio_input)}, value = {audio_input}")
|
||||
if audio_input is None:
|
||||
print(f"{label}: No audio provided")
|
||||
return None, None
|
||||
|
||||
if isinstance(audio_input, Mapping):
|
||||
try:
|
||||
waveform = audio_input["waveform"].squeeze(0)
|
||||
sample_rate = audio_input["sample_rate"]
|
||||
print(f"{label}: Audio extracted from Mapping, waveform shape={waveform.shape}, sample_rate={sample_rate}")
|
||||
return waveform, sample_rate
|
||||
except KeyError as e:
|
||||
print(f"{label}: Error - Missing key in Mapping: {e}")
|
||||
return None, None
|
||||
except Exception as e:
|
||||
print(f"{label}: Error extracting from Mapping: {e}")
|
||||
return None, None
|
||||
|
||||
elif callable(audio_input):
|
||||
try:
|
||||
audio_data = audio_input()
|
||||
if isinstance(audio_data, dict) and "waveform" in audio_data:
|
||||
waveform = audio_data["waveform"].squeeze(0)
|
||||
print(f"{label}: Audio extracted from function, waveform shape={waveform.shape}, sample_rate={audio_data['sample_rate']}")
|
||||
return waveform, audio_data["sample_rate"]
|
||||
else:
|
||||
print(f"{label}: Invalid function result: {audio_data}")
|
||||
return None, None
|
||||
except Exception as e:
|
||||
print(f"{label}: Error evaluating function: {e}")
|
||||
return None, None
|
||||
|
||||
elif isinstance(audio_input, dict) and "waveform" in audio_input:
|
||||
waveform = audio_input["waveform"].squeeze(0)
|
||||
print(f"{label}: Audio extracted from dictionary, waveform shape={waveform.shape}, sample_rate={audio_input['sample_rate']}")
|
||||
return waveform, audio_input["sample_rate"]
|
||||
|
||||
else:
|
||||
print(f"{label}: Audio format not recognized: {type(audio_input)}")
|
||||
return None, None
|
||||
|
||||
# Extract audio data from all audio inputs
|
||||
audio_waveform1, sample_rate1 = get_audio_data(audio1, "Audio1")
|
||||
audio_waveform2, sample_rate2 = get_audio_data(audio2, "Audio2")
|
||||
bgm_waveform, bgm_sample_rate = get_audio_data(bgm, "BGM")
|
||||
|
||||
# Determine target sample rate (prefer audio1, then audio2, then BGM, fallback to 44100)
|
||||
sample_rate = sample_rate1 or sample_rate2 or bgm_sample_rate or 44100
|
||||
print(f"Using sample rate: {sample_rate}Hz")
|
||||
|
||||
# Import torchaudio for resampling
|
||||
import torchaudio
|
||||
|
||||
# Resample all audio to the target sample rate
|
||||
if audio_waveform1 is not None and sample_rate1 != sample_rate:
|
||||
resampler = torchaudio.transforms.Resample(
|
||||
orig_freq=sample_rate1,
|
||||
new_freq=sample_rate
|
||||
)
|
||||
audio_waveform1 = resampler(audio_waveform1)
|
||||
print(f"Resampled audio1 from {sample_rate1}Hz to {sample_rate}Hz")
|
||||
|
||||
if audio_waveform2 is not None and sample_rate2 != sample_rate:
|
||||
resampler = torchaudio.transforms.Resample(
|
||||
orig_freq=sample_rate2,
|
||||
new_freq=sample_rate
|
||||
)
|
||||
audio_waveform2 = resampler(audio_waveform2)
|
||||
print(f"Resampled audio2 from {sample_rate2}Hz to {sample_rate}Hz")
|
||||
|
||||
if bgm_waveform is not None and bgm_sample_rate != sample_rate:
|
||||
resampler = torchaudio.transforms.Resample(
|
||||
orig_freq=bgm_sample_rate,
|
||||
new_freq=sample_rate
|
||||
)
|
||||
bgm_waveform = resampler(bgm_waveform)
|
||||
print(f"Resampled BGM from {bgm_sample_rate}Hz to {sample_rate}Hz")
|
||||
|
||||
# Calculate durations
|
||||
duration1 = images1.shape[0] / fps1
|
||||
duration2 = images2.shape[0] / fps2
|
||||
total_duration = duration1 + duration2
|
||||
total_samples = int(total_duration * sample_rate)
|
||||
|
||||
print(f"Video durations: Video1={duration1:.2f}s, Video2={duration2:.2f}s, Total={total_duration:.2f}s")
|
||||
|
||||
# Generate silent audio for missing primary audio inputs
|
||||
if audio_waveform1 is None:
|
||||
audio_waveform1 = torch.zeros((1, int(duration1 * sample_rate)))
|
||||
print(f"Audio1: Silence generated, shape={audio_waveform1.shape}")
|
||||
|
||||
if audio_waveform2 is None:
|
||||
audio_waveform2 = torch.zeros((1, int(duration2 * sample_rate)))
|
||||
print(f"Audio2: Silence generated, shape={audio_waveform2.shape}")
|
||||
|
||||
# Match channel counts between primary audio streams
|
||||
if audio_waveform1.shape[0] != audio_waveform2.shape[0]:
|
||||
print(f"Channel count mismatch: audio1 has {audio_waveform1.shape[0]} channels, audio2 has {audio_waveform2.shape[0]} channels")
|
||||
|
||||
# If audio1 is mono and audio2 is stereo
|
||||
if audio_waveform1.shape[0] == 1 and audio_waveform2.shape[0] == 2:
|
||||
# Convert audio1 to stereo by duplicating the channel
|
||||
audio_waveform1 = audio_waveform1.repeat(2, 1)
|
||||
print(f"Converted audio1 to stereo: new shape={audio_waveform1.shape}")
|
||||
|
||||
# If audio1 is stereo and audio2 is mono
|
||||
elif audio_waveform1.shape[0] == 2 and audio_waveform2.shape[0] == 1:
|
||||
# Convert audio2 to stereo by duplicating the channel
|
||||
audio_waveform2 = audio_waveform2.repeat(2, 1)
|
||||
print(f"Converted audio2 to stereo: new shape={audio_waveform2.shape}")
|
||||
|
||||
# Concatenate the primary audio streams
|
||||
primary_audio = torch.cat([audio_waveform1, audio_waveform2], dim=1)
|
||||
print(f"Concatenated primary audio: shape={primary_audio.shape}")
|
||||
|
||||
# Check if primary audio has actual content (not just silence)
|
||||
has_actual_audio = primary_audio.abs().max() > 0.01
|
||||
print(f"Primary audio has actual content: {has_actual_audio}")
|
||||
|
||||
# Process background music if provided
|
||||
if bgm_waveform is not None:
|
||||
print(f"Processing BGM: shape={bgm_waveform.shape}")
|
||||
|
||||
# Match channel count with primary audio
|
||||
primary_channels = primary_audio.shape[0]
|
||||
if bgm_waveform.shape[0] != primary_channels:
|
||||
if bgm_waveform.shape[0] == 1 and primary_channels == 2:
|
||||
# Convert mono BGM to stereo
|
||||
bgm_waveform = bgm_waveform.repeat(2, 1)
|
||||
print(f"Converted mono BGM to stereo: new shape={bgm_waveform.shape}")
|
||||
elif bgm_waveform.shape[0] == 2 and primary_channels == 1:
|
||||
# Convert stereo BGM to mono
|
||||
bgm_waveform = bgm_waveform.mean(dim=0, keepdim=True)
|
||||
print(f"Converted stereo BGM to mono: new shape={bgm_waveform.shape}")
|
||||
|
||||
# Loop or trim BGM to match total audio length
|
||||
if bgm_waveform.shape[1] < total_samples:
|
||||
# BGM is shorter than needed, loop it
|
||||
repeats_needed = (total_samples + bgm_waveform.shape[1] - 1) // bgm_waveform.shape[1]
|
||||
bgm_repeated = bgm_waveform.repeat(1, repeats_needed)
|
||||
bgm_waveform = bgm_repeated[:, :total_samples]
|
||||
print(f"Looped BGM {repeats_needed} times to match duration, new shape={bgm_waveform.shape}")
|
||||
elif bgm_waveform.shape[1] > total_samples:
|
||||
# BGM is longer than needed, trim it
|
||||
bgm_waveform = bgm_waveform[:, :total_samples]
|
||||
print(f"Trimmed BGM to match duration, new shape={bgm_waveform.shape}")
|
||||
|
||||
# Apply fade-in effect
|
||||
if fade_in_sec > 0:
|
||||
fade_samples = int(fade_in_sec * sample_rate)
|
||||
if fade_samples > 0 and fade_samples < bgm_waveform.shape[1]:
|
||||
fade_curve = torch.linspace(0, 1, fade_samples)
|
||||
for c in range(bgm_waveform.shape[0]):
|
||||
bgm_waveform[c, :fade_samples] *= fade_curve
|
||||
print(f"Applied {fade_in_sec}s fade-in to BGM")
|
||||
|
||||
# Mix BGM with primary audio
|
||||
if has_actual_audio:
|
||||
# For speech audio, we need a smoother approach than instant volume changes
|
||||
# Use a sliding window average to detect audio presence, then smooth the volume control
|
||||
|
||||
# Step 1: Calculate audio energy over time with a sliding window
|
||||
window_size = int(0.3 * sample_rate) # 300ms window, good for speech
|
||||
primary_energy = torch.zeros(primary_audio.shape[1])
|
||||
|
||||
# Calculate energy profile
|
||||
for i in range(primary_audio.shape[1]):
|
||||
start = max(0, i - window_size//2)
|
||||
end = min(primary_audio.shape[1], i + window_size//2)
|
||||
window_data = primary_audio[:, start:end]
|
||||
primary_energy[i] = window_data.abs().mean()
|
||||
|
||||
# Step 2: Apply smoothing to the energy profile
|
||||
smoothing_window = int(0.5 * sample_rate) # 500ms smoothing window
|
||||
smoothed_energy = torch.zeros_like(primary_energy)
|
||||
for i in range(len(primary_energy)):
|
||||
start = max(0, i - smoothing_window//2)
|
||||
end = min(len(primary_energy), i + smoothing_window//2)
|
||||
smoothed_energy[i] = primary_energy[start:end].mean()
|
||||
|
||||
# Step 3: Convert energy to volume level
|
||||
# Set threshold - below this energy level, BGM will be at full volume
|
||||
threshold = 0.005
|
||||
# Set range - how quickly it transitions from min to max volume
|
||||
range_factor = 0.01
|
||||
|
||||
# Create the volume mask
|
||||
bgm_volume_mask = torch.ones(bgm_waveform.shape[1])
|
||||
|
||||
for i in range(min(len(smoothed_energy), bgm_volume_mask.shape[0])):
|
||||
# Map energy to volume: higher energy = lower BGM volume
|
||||
energy = smoothed_energy[i]
|
||||
if energy > threshold:
|
||||
# Linear mapping from energy to volume
|
||||
volume_factor = max(bgm_volume, 1.0 - (energy - threshold) / range_factor)
|
||||
bgm_volume_mask[i] = volume_factor
|
||||
else:
|
||||
# Below threshold, full volume
|
||||
bgm_volume_mask[i] = 1.0
|
||||
|
||||
# Apply fade-in at the beginning of BGM
|
||||
fade_samples = int(fade_in_sec * sample_rate)
|
||||
if fade_samples > 0 and fade_samples < bgm_volume_mask.shape[0]:
|
||||
fade_curve = torch.linspace(0, 1, fade_samples)
|
||||
for i in range(fade_samples):
|
||||
bgm_volume_mask[i] *= fade_curve[i]
|
||||
|
||||
# Apply the volume mask to all BGM channels
|
||||
volume_adjusted_bgm = bgm_waveform.clone()
|
||||
for c in range(bgm_waveform.shape[0]):
|
||||
volume_adjusted_bgm[c, :min(bgm_waveform.shape[1], len(bgm_volume_mask))] *= bgm_volume_mask[:min(bgm_waveform.shape[1], len(bgm_volume_mask))]
|
||||
|
||||
print(f"Applied smooth BGM volume control with speech detection")
|
||||
|
||||
# Ensure primary_audio and BGM are the same length
|
||||
if primary_audio.shape[1] < volume_adjusted_bgm.shape[1]:
|
||||
# Pad primary audio with zeros
|
||||
padding = torch.zeros(primary_audio.shape[0], volume_adjusted_bgm.shape[1] - primary_audio.shape[1])
|
||||
primary_audio = torch.cat([primary_audio, padding], dim=1)
|
||||
elif primary_audio.shape[1] > volume_adjusted_bgm.shape[1]:
|
||||
# Pad BGM with zeros
|
||||
padding = torch.zeros(volume_adjusted_bgm.shape[0], primary_audio.shape[1] - volume_adjusted_bgm.shape[1])
|
||||
volume_adjusted_bgm = torch.cat([volume_adjusted_bgm, padding], dim=1)
|
||||
|
||||
audio_output = primary_audio + volume_adjusted_bgm
|
||||
print(f"Mixed BGM with smooth volume control for speech audio")
|
||||
else:
|
||||
# If there's no actual audio content, use BGM at full volume
|
||||
audio_output = bgm_waveform
|
||||
print(f"Using BGM at full volume (no primary audio content)")
|
||||
else:
|
||||
# No BGM provided, just use primary audio
|
||||
audio_output = primary_audio
|
||||
print("No BGM provided, using only primary audio")
|
||||
|
||||
# Ensure audio_output has correct batch dimension
|
||||
audio_output = audio_output.unsqueeze(0)
|
||||
|
||||
# Update video_info for output
|
||||
video_info_output = {
|
||||
"loaded_fps": output_fps,
|
||||
"loaded_frame_count": concatenated_frames.shape[0],
|
||||
"loaded_duration": concatenated_frames.shape[0] / output_fps,
|
||||
"loaded_width": images1.shape[2], # Width
|
||||
"loaded_height": images1.shape[1], # Height
|
||||
"source_fps": video_info1["source_fps"],
|
||||
"source_frame_count": video_info1["source_frame_count"] + video_info2["source_frame_count"],
|
||||
"source_duration": video_info1["source_duration"] + video_info2["source_duration"],
|
||||
"source_width": video_info1["source_width"],
|
||||
"source_height": video_info1["source_height"],
|
||||
}
|
||||
|
||||
# Prepare audio output
|
||||
audio_output_dict = None
|
||||
if audio_output is not None:
|
||||
audio_output_dict = {
|
||||
"waveform": audio_output,
|
||||
"sample_rate": sample_rate
|
||||
}
|
||||
print(f"Final audio prepared: waveform shape={audio_output_dict['waveform'].shape}, sample_rate={sample_rate}")
|
||||
else:
|
||||
print("No final audio generated")
|
||||
|
||||
print(f"Output: frames={concatenated_frames.shape}, audio={audio_output.shape if audio_output is not None else 'none'}")
|
||||
|
||||
return (concatenated_frames, audio_output_dict, video_info_output)
|
||||
|
||||
# Node Mappings for ComfyUI
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"LatentSyncNode": LatentSyncNode,
|
||||
"VideoLengthAdjuster": VideoLengthAdjuster,
|
||||
"DG_VideoAudioMixer": DG_VideoAudioMixer,
|
||||
}
|
||||
|
||||
# Display Names for ComfyUI
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LatentSyncNode": "LatentSync1.5 Node",
|
||||
"VideoLengthAdjuster": "Video Length Adjuster",
|
||||
"DG_VideoAudioMixer": "DG Video Audio Mixer",
|
||||
}
|
||||
Reference in New Issue
Block a user