[Added] Audio Blend

This commit is contained in:
Salvador E. Tropea
2025-07-15 13:01:59 -03:00
parent 956127cbbd
commit 53d7edc1b4
4 changed files with 220 additions and 165 deletions
+14
View File
@@ -15,6 +15,7 @@ workflows, especially when dealing with multiple audio inputs or outputs.
- [6. Audio Channel Conv and Resampler](#6-audio-channel-conv-and-resampler)
- [7. Audio Information](#7-audio-information)
- [8. Audio Cut](#8-audio-cut)
- [9. Audio Blend](#9-audio-blend)
- [🚀 Installation](#-installation)
- [📦 Dependencies](#-dependencies)
- [🖼️ Examples](#️-examples)
@@ -144,6 +145,19 @@ workflows, especially when dealing with multiple audio inputs or outputs.
- **Output:**
- `audio_out` (AUDIO): The selected portion of the audio.
### 9. Audio Blend
- **Display Name:** `Audio Blend`
- **Internal Name:** `SET_AudioBlend`
- **Category:** `audio/manipulation`
- **Description:** Blends two audio inputs by applying gain and adding them. Supports batches. It can be used to amplify just one audio (or batch)
- **Inputs:**
- `audio1` (AUDIO): The first audio input (batch supported).
- `audio2` (AUDIO): The second audio input (batch supported). Is optional.
- `gain1` (FLOAT): Volume gain for the first audio. Can be negative to subtract.
- `gain2` (FLOAT): Volume gain for the second audio. Can be negative to subtract.
- **Output:**
- `audio_out` (AUDIO): The blended audio.
## 🚀 Installation
You can install the nodes from the ComfyUI nodes manager, the name is *Audio Batch*, or just do it manually:
View File
+64 -165
View File
@@ -6,6 +6,7 @@
# From code generated by Gemini 2.5 Pro
import torch
import torchaudio.transforms as T
from .utils.aligner import AudioBatchAligner
from .utils.logger import main_logger
from .utils.misc import parse_time_to_seconds
@@ -16,15 +17,6 @@ CONV_CATEGORY = "conversion"
MANIPULATION_CATEGORY = "manipulation"
def convert_batch_to_stereo_tensor(audio_waveform_mono_batch: torch.Tensor) -> torch.Tensor:
""" Converts a batch of mono audio tensors (B, 1, N) to stereo (B, 2, N). """
if audio_waveform_mono_batch.ndim != 3 or audio_waveform_mono_batch.shape[1] != 1:
# This could also happen if an input was (N) or (1,N) and wasn't unsqueezed to (B,1,N) yet
raise ValueError("Input for stereo conversion must be a batch of mono audio (B, 1, N), "
f"got {audio_waveform_mono_batch.shape}")
return audio_waveform_mono_batch.repeat(1, 2, 1) # (B, 1, N) -> (B, 2, N)
class AudioBatch:
@classmethod
def INPUT_TYPES(cls):
@@ -43,166 +35,17 @@ class AudioBatch:
UNIQUE_NAME = "SET_AudioBatch"
DISPLAY_NAME = "Batch Audios"
def _preprocess_waveform_batch(
self,
waveform: torch.Tensor, # (B, C_in, N_in)
original_sr: int,
target_sr: int,
target_channels: int,
target_samples: int,
reference_dtype: torch.dtype,
reference_device: torch.device
) -> torch.Tensor:
"""
Helper function to resample, change channels, and pad a batch of waveforms.
Returns a tensor of shape (B, target_channels, target_samples).
"""
# Ensure correct device and dtype
processed_wf = waveform.to(device=reference_device, dtype=reference_dtype)
current_batch_size, current_channels, current_samples = processed_wf.shape
# 1. Resample if necessary
if original_sr != target_sr:
logger.debug(f"Resampling batch from {original_sr} to {target_sr}. Input shape: {processed_wf.shape}")
# Resample expects (..., time)
# For (B, C, N), we can reshape to (B*C, N), resample, then reshape back.
# Or, if T.Resample handles batch dims appropriately (some versions might if C=1 or applied per channel).
# Let's reshape for robustness with T.Resample.
resampler = T.Resample(orig_freq=original_sr, new_freq=target_sr, dtype=reference_dtype,
lowpass_filter_width=24).to(reference_device)
if current_channels == 1:
# Reshape (B, 1, N) to (B, N) for resampler, then unsqueeze back
processed_wf_reshaped = processed_wf.squeeze(1) # (B, N)
resampled_wf_reshaped = resampler(processed_wf_reshaped) # (B, N_new)
processed_wf = resampled_wf_reshaped.unsqueeze(1) # (B, 1, N_new)
else: # Multi-channel (e.g., stereo)
# Resample each channel in the batch separately
# This is more complex if T.Resample doesn't broadcast correctly.
# A common way: permute to (C, B, N), reshape to (C*B, N), resample, reshape back.
# Or loop (less efficient for GPU tensors).
# Let's try reshaping to (B*C, N)
original_shape = processed_wf.shape
processed_wf_flat_batch_channel = processed_wf.reshape(-1, current_samples) # (B*C, N)
resampled_wf_flat = resampler(processed_wf_flat_batch_channel) # (B*C, N_new)
# Reshape back to (B, C, N_new)
processed_wf = resampled_wf_flat.reshape(original_shape[0], original_shape[1], -1)
current_samples = processed_wf.shape[2] # Update sample count
logger.debug(f"Batch after resampling: shape={processed_wf.shape}")
# 2. Adjust number of channels
if current_channels == 1 and target_channels == 2:
processed_wf = convert_batch_to_stereo_tensor(processed_wf) # (B, 1, N) -> (B, 2, N)
logger.debug(f"Batch converted to stereo: shape={processed_wf.shape}")
elif current_channels == 2 and target_channels == 1:
# Example: Convert stereo to mono by averaging (can make this an option)
processed_wf = processed_wf.mean(dim=1, keepdim=True) # (B, 2, N) -> (B, 1, N)
logger.debug(f"Batch converted to mono (avg): shape={processed_wf.shape}")
elif current_channels != target_channels:
# Fallback or error for other unsupported channel conversions (e.g. 5.1 to stereo)
# For now, if shapes don't match after mono/stereo adjustment, it might error later.
# A more robust node would handle various conversions or error clearly.
logger.warning(f"Unhandled channel conversion from {current_channels} to {target_channels}."
f" Resulting channels: {processed_wf.shape[1]}")
# If, for example, target is 2, and current is 5, how to downmix?
# For this node's purpose (batching existing audio), major downmixing is out of scope.
# We primarily handle mono <-> stereo alignment.
if processed_wf.shape[1] != target_channels:
raise ValueError(f"Cannot align channels: input has {processed_wf.shape[1]} after initial processing, "
f"target is {target_channels}")
# 3. Pad length if necessary
if current_samples < target_samples:
padding_needed = target_samples - current_samples
# Pad only the last dimension (samples)
processed_wf = torch.nn.functional.pad(processed_wf, (0, padding_needed))
logger.debug(f"Batch padded: shape={processed_wf.shape}")
elif current_samples > target_samples: # Should not happen if target_samples is max length
logger.warning(f"Waveform has {current_samples} samples, but target is {target_samples}. "
"This indicates an issue in target_samples calculation.")
# Truncate as a fallback, though logic should prevent this.
processed_wf = processed_wf[..., :target_samples]
# Final check for shape
if processed_wf.shape[1] != target_channels or processed_wf.shape[2] != target_samples:
raise RuntimeError(f"Internal Preprocessing Error: Waveform shape {processed_wf.shape} "
f"does not match target ({current_batch_size}, {target_channels}, {target_samples}).")
return processed_wf
def batch_audio(self, audio1: dict, audio2: dict):
waveform1_orig = audio1['waveform'] # (B1, C1, N1)
sr1 = audio1['sample_rate']
waveform2_orig = audio2['waveform'] # (B2, C2, N2)
sr2 = audio2['sample_rate']
logger.debug(f"Audio1 input: shape={waveform1_orig.shape}, sr={sr1}, dtype={waveform1_orig.dtype}, "
f"device={waveform1_orig.device}")
logger.debug(f"Audio2 input: shape={waveform2_orig.shape}, sr={sr2}, dtype={waveform2_orig.dtype}, "
f"device={waveform2_orig.device}")
# Use properties of audio1 as the reference for the output batch
# (e.g., device, dtype, and target sample rate)
reference_device = waveform1_orig.device
reference_dtype = waveform1_orig.dtype
target_sr = sr1 # All audio will be resampled to sr1
# Determine target number of channels for the batch (max of inputs, ensure mono/stereo alignment)
c1 = waveform1_orig.shape[1]
c2 = waveform2_orig.shape[1]
# If one is mono and other is stereo, output will be stereo. Otherwise, max (handles both mono or both stereo).
if (c1 == 1 and c2 == 2) or (c1 == 2 and c2 == 1) or (c1 == 2 and c2 == 2):
target_channels = 2
elif c1 == 1 and c2 == 1:
target_channels = 1
else:
# For other multi-channel counts (e.g. 5.1), this simple logic might not be ideal.
# For now, default to max, but this could be an error or require downmixing node.
logger.warning(f"Complex channel counts detected (C1={c1}, C2={c2}). Defaulting to max channels ({max(c1,c2)}) "
"and hoping for downstream compatibility or further processing. "
"Explicit mono/stereo alignment is preferred for this node.")
target_channels = max(c1, c2)
# Determine target number of samples (length) for the batch
# First, calculate lengths *after* potential resampling
n1_orig = waveform1_orig.shape[2]
n2_orig = waveform2_orig.shape[2]
n1_after_resample = n1_orig # No resampling for audio1 as it's the target_sr reference
n2_after_resample = n2_orig
if sr2 != target_sr:
n2_after_resample = int(n2_orig * (target_sr / sr2))
target_samples = max(n1_after_resample, n2_after_resample)
logger.debug(f"Unified Batch Params: SR={target_sr}, Channels={target_channels}, Samples={target_samples}")
# Preprocess both waveform batches
wf1_processed = self._preprocess_waveform_batch(waveform1_orig, sr1, target_sr, target_channels, target_samples,
reference_dtype, reference_device)
wf2_processed = self._preprocess_waveform_batch(waveform2_orig, sr2, target_sr, target_channels, target_samples,
reference_dtype, reference_device)
# wf1_processed is (B1, target_channels, target_samples)
# wf2_processed is (B2, target_channels, target_samples)
# Concatenate along the batch dimension (dim=0)
try:
batched_waveform_final = torch.cat((wf1_processed, wf2_processed), dim=0)
except RuntimeError as e:
logger.error(f"Error during final torch.cat: {e}")
logger.error(f"Processed Waveform1 shape: {wf1_processed.shape}, dtype: {wf1_processed.dtype}, "
f"device: {wf1_processed.device}")
logger.error(f"Processed Waveform2 shape: {wf2_processed.shape}, dtype: {wf2_processed.dtype}, "
f"device: {wf2_processed.device}")
raise # Re-raise to make error visible
# 1. Instantiate the aligner
aligner = AudioBatchAligner(audio1, audio2)
# 2. Get the aligned waveforms
aligned_wf1, aligned_wf2, target_sr = aligner.get_aligned_waveforms()
# 3. Concatenate
batched_waveform_final = torch.cat((aligned_wf1, aligned_wf2), dim=0)
logger.info(f"Final batched audio: shape={batched_waveform_final.shape}, sr={target_sr}")
output_audio = {
"waveform": batched_waveform_final,
"sample_rate": target_sr
}
output_audio = {"waveform": batched_waveform_final, "sample_rate": target_sr}
return (output_audio,)
@@ -558,3 +401,59 @@ class AudioCut:
return ({"waveform": waveform[..., start_frame:end_frame],
"sample_rate": sample_rate},)
class AudioBlend:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"audio1": ("AUDIO", {"tooltip": "The first audio input (batch supported)."}),
"gain1": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01,
"tooltip": "Volume gain for the first audio. Can be negative to subtract."}),
"gain2": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01,
"tooltip": "Volume gain for the second audio. Can be negative to subtract."}),
},
"optional": {
"audio2": ("AUDIO", {"tooltip": "The second audio input (optional, batch supported)."}),
}
}
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio_out",)
FUNCTION = "blend_audio"
CATEGORY = BASE_CATEGORY + "/" + MANIPULATION_CATEGORY
DESCRIPTION = "Blends two audio inputs by applying gain and adding them. Supports batches."
UNIQUE_NAME = "SET_AudioBlend"
DISPLAY_NAME = "Audio Blend"
def blend_audio(self, audio1: dict, gain1: float, gain2: float, audio2: dict = None):
if audio2 is None:
# Handle the simple case where audio2 is not provided
logger.info(f"Blending audio1 only with gain {gain1}.")
blended_waveform = audio1['waveform'] * gain1
return ({"waveform": blended_waveform, "sample_rate": audio1['sample_rate']},)
# If audio2 is provided, align both audio inputs
# 1. Instantiate the aligner with the logger
aligner = AudioBatchAligner(audio1, audio2)
# 2. Get the aligned waveforms
aligned_wf1, aligned_wf2, target_sr = aligner.get_aligned_waveforms()
b1, c, n = aligned_wf1.shape
b2 = aligned_wf2.shape[0]
# 3. Perform the blend operation on aligned tensors
target_batch_size = max(b1, b2)
blended_waveform = torch.zeros(target_batch_size, c, n, dtype=aligned_wf1.dtype, device=aligned_wf1.device)
logger.info(f"Blending {b1} items from audio1 with {b2} items from audio2 into a batch of {target_batch_size}.")
# Add the scaled audio streams to the output tensor
# This correctly handles mismatched batch sizes by only adding where data exists.
blended_waveform[:b1] += aligned_wf1 * gain1
blended_waveform[:b2] += aligned_wf2 * gain2
output_audio = {"waveform": blended_waveform, "sample_rate": target_sr}
return (output_audio,)
+142
View File
@@ -0,0 +1,142 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnologïa Industrial
# License: GPLv3
# Project: ComfyUI-AudioBatch
#
# Audio batch aligner
# Original code from Gemini 2.5 Pro
import logging
import torch
import torchaudio.transforms as T
from .misc import NODES_NAME
logger = logging.getLogger(f"{NODES_NAME}.aligner")
def convert_batch_to_stereo_tensor(audio_waveform_mono_batch: torch.Tensor) -> torch.Tensor:
""" Converts a batch of mono audio tensors (B, 1, N) to stereo (B, 2, N). """
if audio_waveform_mono_batch.ndim != 3 or audio_waveform_mono_batch.shape[1] != 1:
# This could also happen if an input was (N) or (1,N) and wasn't unsqueezed to (B,1,N) yet
raise ValueError("Input for stereo conversion must be a batch of mono audio (B, 1, N), "
f"got {audio_waveform_mono_batch.shape}")
return audio_waveform_mono_batch.repeat(1, 2, 1) # (B, 1, N) -> (B, 2, N)
class AudioBatchAligner:
""" A helper class to align two audio batches for processing. """
def __init__(self, audio1: dict, audio2: dict):
self.waveform1_orig, self.sr1 = audio1['waveform'], audio1['sample_rate']
self.waveform2_orig, self.sr2 = audio2['waveform'], audio2['sample_rate']
logger.debug(f"Audio1 input: shape={self.waveform1_orig.shape}, sr={self.sr1}, dtype={self.waveform1_orig.dtype}, "
f"device={self.waveform1_orig.device}")
logger.debug(f"Audio2 input: shape={self.waveform2_orig.shape}, sr={self.sr2}, dtype={self.waveform2_orig.dtype}, "
f"device={self.waveform2_orig.device}")
# Use properties of audio1 as the reference for the output batch
self.reference_device = self.waveform1_orig.device
self.reference_dtype = self.waveform1_orig.dtype
self.target_sr = self.sr1 # All audio will be resampled to sr1
# This will be populated by _determine_target_params
self.target_channels: int = 0
self.target_samples: int = 0
self._determine_target_params()
def _determine_target_params(self):
""" Determines the target channels and samples for the unified batch. """
c1, c2 = self.waveform1_orig.shape[1], self.waveform2_orig.shape[1]
if (c1 == 1 and c2 == 2) or (c1 == 2 and c2 == 1) or (c1 == 2 and c2 == 2):
self.target_channels = 2
elif c1 == 1 and c2 == 1:
self.target_channels = 1
else:
# For other multi-channel counts (e.g. 5.1), this simple logic might not be ideal.
# For now, default to max, but this could be an error or require downmixing node.
logger.warning(f"Complex channel counts detected (C1={c1}, C2={c2}). Defaulting to max channels "
f"({max(c1,c2)}) and hoping for downstream compatibility. Explicit mono/stereo "
"alignment is preferred for this node.")
self.target_channels = max(c1, c2)
# Determine target number of samples (length) after potential resampling
n1_orig, n2_orig = self.waveform1_orig.shape[2], self.waveform2_orig.shape[2]
n1_after_resample = n1_orig # No resampling for audio1 as it's the target_sr reference
n2_after_resample = n2_orig if self.sr2 == self.target_sr else int(n2_orig * (self.target_sr / self.sr2))
self.target_samples = max(n1_after_resample, n2_after_resample)
logger.debug(f"Unified Batch Params: SR={self.target_sr}, Channels={self.target_channels}, "
f"Samples={self.target_samples}")
def _align_one_batch(self, waveform: torch.Tensor, original_sr: int) -> torch.Tensor:
""" Processes a single waveform batch to match the target parameters. """
# Ensure correct device and dtype
processed_wf = waveform.to(device=self.reference_device, dtype=self.reference_dtype)
current_batch_size, current_channels, current_samples = processed_wf.shape
# 1. Resample if necessary
if original_sr != self.target_sr:
logger.debug(f"Resampling batch from {original_sr} to {self.target_sr}. Input shape: {processed_wf.shape}")
# Resample expects (..., time)
# For (B, C, N), we can reshape to (B*C, N), resample, then reshape back.
# Or, if T.Resample handles batch dims appropriately (some versions might if C=1 or applied per channel).
# Let's reshape for robustness with T.Resample.
resampler = T.Resample(orig_freq=original_sr, new_freq=self.target_sr, dtype=self.reference_dtype,
lowpass_filter_width=24).to(self.reference_device)
if current_channels == 1:
# Reshape (B, 1, N) to (B, N) for resampler, then unsqueeze back
processed_wf = resampler(processed_wf.squeeze(1)).unsqueeze(1) # (B, 1, N_new)
else: # Multi-channel (e.g., stereo)
# Resample each channel in the batch separately
# This is more complex if T.Resample doesn't broadcast correctly.
# A common way: permute to (C, B, N), reshape to (C*B, N), resample, reshape back.
# Or loop (less efficient for GPU tensors).
# Let's try reshaping to (B*C, N)
original_shape = processed_wf.shape
processed_wf_flat = processed_wf.reshape(-1, current_samples) # (B*C, N)
resampled_wf_flat = resampler(processed_wf_flat) # (B*C, N_new)
# Reshape back to (B, C, N_new)
processed_wf = resampled_wf_flat.reshape(original_shape[0], original_shape[1], -1)
current_samples = processed_wf.shape[2] # Update sample count
logger.debug(f"Batch after resampling: shape={processed_wf.shape}")
# 2. Adjust number of channels
if current_channels != self.target_channels:
if current_channels == 1 and self.target_channels == 2:
processed_wf = convert_batch_to_stereo_tensor(processed_wf) # (B, 1, N) -> (B, 2, N)
logger.debug(f"Batch converted to stereo: shape={processed_wf.shape}")
elif current_channels == 2 and self.target_channels == 1:
# Example: Convert stereo to mono by averaging (can make this an option)
processed_wf = processed_wf.mean(dim=1, keepdim=True) # (B, 2, N) -> (B, 1, N)
logger.debug(f"Batch converted to mono (avg): shape={processed_wf.shape}")
else:
raise ValueError(f"Cannot align channels: input has {current_channels}, target is {self.target_channels}")
# 3. Pad length if necessary
if current_samples < self.target_samples:
padding_needed = self.target_samples - current_samples
# Pad only the last dimension (samples)
processed_wf = torch.nn.functional.pad(processed_wf, (0, padding_needed))
logger.debug(f"Batch padded: shape={processed_wf.shape}")
elif current_samples > self.target_samples: # Should not happen if target_samples is max length
logger.warning(f"Waveform has {current_samples} samples, but target is {self.target_samples}. Truncating.")
# Truncate as a fallback, though logic should prevent this.
processed_wf = processed_wf[..., :self.target_samples]
# Final check for shape
if processed_wf.shape[1] != self.target_channels or processed_wf.shape[2] != self.target_samples:
raise RuntimeError(f"Internal Preprocessing Error: Waveform shape {processed_wf.shape} "
f"does not match target ({current_batch_size}, {self.target_channels}, "
f"{self.target_samples}).")
return processed_wf
def get_aligned_waveforms(self) -> tuple[torch.Tensor, torch.Tensor, int]:
""" Aligns both waveforms and returns them. """
# Preprocess both waveform batches
wf1_processed = self._align_one_batch(self.waveform1_orig, self.sr1)
wf2_processed = self._align_one_batch(self.waveform2_orig, self.sr2)
return wf1_processed, wf2_processed, self.target_sr