From 53d7edc1b4541e9ce681dd74be27ab7e13f6daeb Mon Sep 17 00:00:00 2001 From: "Salvador E. Tropea" Date: Tue, 15 Jul 2025 13:01:59 -0300 Subject: [PATCH] [Added] Audio Blend --- README.md | 14 ++ source/nodes/__init__.py | 0 nodes_audio.py => source/nodes/nodes_audio.py | 229 +++++------------- source/nodes/utils/aligner.py | 142 +++++++++++ 4 files changed, 220 insertions(+), 165 deletions(-) create mode 100644 source/nodes/__init__.py rename nodes_audio.py => source/nodes/nodes_audio.py (65%) create mode 100644 source/nodes/utils/aligner.py diff --git a/README.md b/README.md index 83f47d0..faa656d 100644 --- a/README.md +++ b/README.md @@ -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: diff --git a/source/nodes/__init__.py b/source/nodes/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/nodes_audio.py b/source/nodes/nodes_audio.py similarity index 65% rename from nodes_audio.py rename to source/nodes/nodes_audio.py index 7ae03af..1ec67ec 100644 --- a/nodes_audio.py +++ b/source/nodes/nodes_audio.py @@ -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,) diff --git a/source/nodes/utils/aligner.py b/source/nodes/utils/aligner.py new file mode 100644 index 0000000..8040150 --- /dev/null +++ b/source/nodes/utils/aligner.py @@ -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