[Added] Audio Blend
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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,)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user