[Nodes][Added] Spectral downmix for mono
This commit is contained in:
@@ -84,6 +84,12 @@ workflows, especially when dealing with multiple audio inputs or outputs.
|
||||
- If input is mono, converts to "fake stereo".
|
||||
- If input is stereo, no change.
|
||||
- If input has more than 2 channels, it takes the first channel and duplicates it to create stereo.
|
||||
- `downmix_method` (COMBO): How to convert to mono.
|
||||
- `average`: Simple average ((L+R)/2). Can reduce volume.
|
||||
- `standard_gain` (default): Sums channels with -3dB gain (0.707). Better preserves perceived loudness.
|
||||
- `spectral`: Averages frequency magnitudes to prevent phase cancellation.
|
||||
- `n_fft` (INT, optional): FFT size for spectral downmixing. Higher values give better frequency resolution but worse time resolution.
|
||||
- `hop_length` (INT, optional): Hop length for STFT. Typically n_fft / 4. Controls time resolution.
|
||||
- **Output:**
|
||||
- `audio_out` (AUDIO): The audio with the converted channel layout. The batch size and sample rate are preserved.
|
||||
|
||||
@@ -98,6 +104,12 @@ workflows, especially when dealing with multiple audio inputs or outputs.
|
||||
- `0`: `keep`
|
||||
- `1`: `force_mono`
|
||||
- `2`: `force_stereo`
|
||||
- `downmix_method` (COMBO): How to convert to mono.
|
||||
- `average`: Simple average ((L+R)/2). Can reduce volume.
|
||||
- `standard_gain` (default): Sums channels with -3dB gain (0.707). Better preserves perceived loudness.
|
||||
- `spectral`: Averages frequency magnitudes to prevent phase cancellation.
|
||||
- `n_fft` (INT, optional): FFT size for spectral downmixing. Higher values give better frequency resolution but worse time resolution.
|
||||
- `hop_length` (INT, optional): Hop length for STFT. Typically n_fft / 4. Controls time resolution.
|
||||
- **Output:**
|
||||
- `audio` (AUDIO): The audio with the converted channel layout. The batch size and sample rate are preserved.
|
||||
|
||||
@@ -121,6 +133,12 @@ workflows, especially when dealing with multiple audio inputs or outputs.
|
||||
- `audio` (AUDIO): The input audio.
|
||||
- `channel_conversion` (COMBO): Same options as the "Audio Channel Converter" node.
|
||||
- `target_sample_rate` (INT): Same options as the "Audio Resampler" node.
|
||||
- `downmix_method` (COMBO): How to convert to mono.
|
||||
- `average`: Simple average ((L+R)/2). Can reduce volume.
|
||||
- `standard_gain` (default): Sums channels with -3dB gain (0.707). Better preserves perceived loudness.
|
||||
- `spectral`: Averages frequency magnitudes to prevent phase cancellation.
|
||||
- `n_fft` (INT, optional): FFT size for spectral downmixing. Higher values give better frequency resolution but worse time resolution.
|
||||
- `hop_length` (INT, optional): Hop length for STFT. Typically n_fft / 4. Controls time resolution.
|
||||
- **Output:**
|
||||
- `audio_out` (AUDIO): The audio after both channel conversion and resampling have been applied.
|
||||
|
||||
@@ -223,7 +241,7 @@ workflows, especially when dealing with multiple audio inputs or outputs.
|
||||
- **Output:**
|
||||
- `audio_out` (AUDIO): A stereo audio signal.
|
||||
- **Behavior Details:**
|
||||
- **Channel Conversion:** Both `audio_left` and `audio_right` are first forced into mono to ensure they each represent a single channel stream.
|
||||
- **Channel Conversion:** Both `audio_left` and `audio_right` are first forced into mono to ensure they each represent a single channel stream. Note that `average` method is used, do it manually to select another mechanism.
|
||||
- **Alignment:** The two mono signals are then aligned to have the same sample rate and length, using the same logic as the "Batch Audios" node (resamples to match `audio_left`'s SR, pads to match the longest duration).
|
||||
- **Batch Handling:** If the inputs have different batch sizes, the last item of the shorter batch is repeated to match the length of the longer batch.
|
||||
|
||||
|
||||
+65
-11
@@ -11,6 +11,7 @@ from typing import Optional, Dict, Any
|
||||
from .utils.aligner import AudioBatchAligner
|
||||
from .utils.logger import main_logger
|
||||
from .utils.misc import parse_time_to_seconds, parse_note_to_frequency
|
||||
from .utils.downmix import spectral_downmix
|
||||
|
||||
logger = main_logger
|
||||
BASE_CATEGORY = "audio"
|
||||
@@ -18,6 +19,28 @@ BATCH_CATEGORY = "batch"
|
||||
CONV_CATEGORY = "conversion"
|
||||
MANIPULATION_CATEGORY = "manipulation"
|
||||
GEN_CATEGORY = "generation"
|
||||
DOWNMIX_OPTIONS = (["average", "standard_gain", "spectral"],
|
||||
{"default": "standard_gain",
|
||||
"tooltip": ("Method for stereo/multi-channel to mono conversion:\n"
|
||||
"- average: Simple average ((L+R)/2). Can reduce volume.\n"
|
||||
"- standard_gain: Sums channels with -3dB gain (0.707). "
|
||||
"Better preserves perceived loudness."
|
||||
"- spectral: Averages frequency magnitudes to prevent phase cancellation.")})
|
||||
DOWNMIX_NFFT = ("INT", {
|
||||
"default": 2048,
|
||||
"min": 256,
|
||||
"max": 8192, # Powers of 2 are typical
|
||||
"step": 256,
|
||||
"tooltip": ("FFT size for spectral downmixing. "
|
||||
"Higher values give better frequency resolution but worse time resolution.")
|
||||
})
|
||||
DOWNMIX_HOP = ("INT", {
|
||||
"default": 512,
|
||||
"min": 64,
|
||||
"max": 4096,
|
||||
"step": 64,
|
||||
"tooltip": "Hop length for STFT. Typically n_fft / 4. Controls time resolution."
|
||||
})
|
||||
|
||||
|
||||
class AudioBatch:
|
||||
@@ -156,7 +179,12 @@ class AudioChannelConverter:
|
||||
"tooltip": "keep: maintain same channels,\n"
|
||||
"stereo_to_mono/force_mono: 1 channel,\n"
|
||||
"mono_to_stereo/force_stereo: 2 channels"}),
|
||||
"downmix_method": DOWNMIX_OPTIONS,
|
||||
},
|
||||
"optional": {
|
||||
"n_fft": DOWNMIX_NFFT,
|
||||
"hop_length": DOWNMIX_HOP,
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
@@ -167,7 +195,8 @@ class AudioChannelConverter:
|
||||
UNIQUE_NAME = "SET_AudioChannelConverter"
|
||||
DISPLAY_NAME = "Audio Channel Converter"
|
||||
|
||||
def convert_channels(self, audio: dict, channel_conversion: str):
|
||||
def convert_channels(self, audio: dict, channel_conversion: str, downmix_method: str, n_fft: int = 2048,
|
||||
hop_length: int = 512):
|
||||
waveform = audio['waveform'] # (B, C, T)
|
||||
sample_rate = audio['sample_rate']
|
||||
|
||||
@@ -192,10 +221,22 @@ class AudioChannelConverter:
|
||||
if original_channels > 2 and channel_conversion == "stereo_to_mono":
|
||||
logger.warning(f"Channel mode 'stereo_to_mono': Input has {original_channels} channels. "
|
||||
"Averaging all to mono.")
|
||||
elif channel_conversion == "force_mono":
|
||||
logger.info(f"Channel mode 'force_mono': Input has {original_channels} channels. Averaging all to mono.")
|
||||
# Average across the channel dimension (dim=1)
|
||||
output_waveform = torch.mean(waveform, dim=1, keepdim=True)
|
||||
logger.info(f"Converting {original_channels} channels to mono using '{downmix_method}' method.")
|
||||
|
||||
if downmix_method == "average":
|
||||
# Simple average across the channel dimension
|
||||
output_waveform = torch.mean(waveform, dim=1, keepdim=True)
|
||||
elif downmix_method == "standard_gain":
|
||||
# Sum channels and apply gain compensation.
|
||||
# This is equivalent to (L*0.707 + R*0.707) for stereo.
|
||||
# For multi-channel, it sums all channels and divides by sqrt(num_channels).
|
||||
gain = 1.0 / (original_channels ** 0.5)
|
||||
logger.debug(f"Applying downmix gain of {gain:.4f} (1/sqrt({original_channels}))")
|
||||
# Sum along channel dimension, then apply gain. keepdim=True for (B,1,N) shape.
|
||||
output_waveform = torch.sum(waveform, dim=1, keepdim=True) * gain
|
||||
elif downmix_method == "spectral":
|
||||
output_waveform = spectral_downmix(waveform, n_fft=n_fft, hop_length=hop_length)
|
||||
|
||||
logger.debug(f"Converted to mono. New shape: {output_waveform.shape}")
|
||||
|
||||
elif channel_conversion == "mono_to_stereo" or channel_conversion == "force_stereo":
|
||||
@@ -291,7 +332,12 @@ class AudioProcessAdvanced:
|
||||
"mono_to_stereo/force_stereo: 2 channels"}),
|
||||
"target_sample_rate": ("INT", {"default": 0, "min": 0, "max": 192000, "step": 100.,
|
||||
"tooltip": "Output sample rate, 0 is same as input"}),
|
||||
"downmix_method": DOWNMIX_OPTIONS,
|
||||
},
|
||||
"optional": {
|
||||
"n_fft": DOWNMIX_NFFT,
|
||||
"hop_length": DOWNMIX_HOP,
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("audio_out",)
|
||||
@@ -301,13 +347,15 @@ class AudioProcessAdvanced:
|
||||
UNIQUE_NAME = "SET_AudioChannelConvResampler"
|
||||
DISPLAY_NAME = "Audio Channel Conv and Resampler"
|
||||
|
||||
def process_audio(self, audio: dict, channel_conversion: str, target_sample_rate: int):
|
||||
def process_audio(self, audio: dict, channel_conversion: str, target_sample_rate: int, downmix_method: str,
|
||||
n_fft: int = 2048, hop_length: int = 512):
|
||||
# Instantiate helper classes (or move their logic directly here)
|
||||
channel_converter_node = AudioChannelConverter()
|
||||
resampler_node = AudioResampler()
|
||||
|
||||
# 1. Channel Conversion
|
||||
(audio_after_channels,) = channel_converter_node.convert_channels(audio, channel_conversion)
|
||||
(audio_after_channels,) = channel_converter_node.convert_channels(audio, channel_conversion, downmix_method,
|
||||
n_fft=n_fft, hop_length=hop_length)
|
||||
|
||||
# 2. Resampling
|
||||
(audio_after_resample,) = resampler_node.resample_audio(audio_after_channels, target_sample_rate)
|
||||
@@ -345,7 +393,12 @@ class AudioForceChannels:
|
||||
"required": {
|
||||
"audio": ("AUDIO",),
|
||||
"channels": ("INT", {"default": 0, "min": 0, "max": 2}),
|
||||
"downmix_method": DOWNMIX_OPTIONS,
|
||||
},
|
||||
"optional": {
|
||||
"n_fft": DOWNMIX_NFFT,
|
||||
"hop_length": DOWNMIX_HOP,
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
@@ -357,8 +410,9 @@ class AudioForceChannels:
|
||||
DISPLAY_NAME = "Audio Force Channels"
|
||||
CHANNELS_TO_MODE = ["keep", "force_mono", "force_stereo"]
|
||||
|
||||
def force_channels(self, audio: dict, channels: int):
|
||||
return AudioChannelConverter().convert_channels(audio, self.CHANNELS_TO_MODE[channels])
|
||||
def force_channels(self, audio: dict, channels: int, downmix_method: str, n_fft: int = 2048, hop_length: int = 512):
|
||||
return AudioChannelConverter().convert_channels(audio, self.CHANNELS_TO_MODE[channels], downmix_method, n_fft=n_fft,
|
||||
hop_length=hop_length)
|
||||
|
||||
|
||||
class AudioCut:
|
||||
@@ -685,8 +739,8 @@ class AudioJoin2Channels:
|
||||
# 1. Force both inputs to be mono to ensure they represent single channels.
|
||||
# We reuse the logic from AudioChannelConverter.
|
||||
channel_converter = AudioChannelConverter()
|
||||
(mono_left_audio,) = channel_converter.convert_channels(audio_left, "force_mono")
|
||||
(mono_right_audio,) = channel_converter.convert_channels(audio_right, "force_mono")
|
||||
(mono_left_audio,) = channel_converter.convert_channels(audio_left, "force_mono", "average")
|
||||
(mono_right_audio,) = channel_converter.convert_channels(audio_right, "force_mono", "average")
|
||||
|
||||
# 2. Align the two mono signals in terms of sample rate and length.
|
||||
# This reuses the robust logic from AudioBatchAligner.
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
# 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 torch
|
||||
import logging
|
||||
from .misc import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.DownmixUtil")
|
||||
|
||||
|
||||
# --- Using Demucs robust STFT and iSTFT implementations ---
|
||||
# These are excellent as they handle windowing, device fallbacks, and reshaping.
|
||||
|
||||
def spectro(x: torch.Tensor, n_fft: int, hop_length: int) -> torch.Tensor:
|
||||
"""
|
||||
Computes the Short-Time Fourier Transform (STFT) of a tensor.
|
||||
Handles batching and MPS device fallbacks.
|
||||
"""
|
||||
*other, length = x.shape
|
||||
x = x.reshape(-1, length) # Flatten batch and channel dimensions
|
||||
|
||||
# Handle MPS device which might not support STFT with all options
|
||||
is_mps = x.device.type == 'mps'
|
||||
original_device = x.device
|
||||
if is_mps:
|
||||
x = x.cpu()
|
||||
|
||||
window = torch.hann_window(n_fft, device=x.device)
|
||||
|
||||
z = torch.stft(
|
||||
x,
|
||||
n_fft=n_fft,
|
||||
hop_length=hop_length,
|
||||
win_length=n_fft,
|
||||
window=window,
|
||||
center=True,
|
||||
pad_mode='reflect',
|
||||
normalized=True,
|
||||
return_complex=True
|
||||
)
|
||||
|
||||
if is_mps:
|
||||
z = z.to(original_device) # Move result back to original device
|
||||
|
||||
_, freqs, frames = z.shape
|
||||
return z.view(*other, freqs, frames)
|
||||
|
||||
|
||||
def ispectro(z: torch.Tensor, hop_length: int, length: int) -> torch.Tensor:
|
||||
"""
|
||||
Computes the Inverse Short-Time Fourier Transform (iSTFT) of a complex tensor.
|
||||
"""
|
||||
*other, freqs, frames = z.shape
|
||||
n_fft = 2 * (freqs - 1)
|
||||
z = z.view(-1, freqs, frames)
|
||||
|
||||
is_mps = z.device.type == 'mps'
|
||||
original_device = z.device
|
||||
if is_mps:
|
||||
z = z.cpu()
|
||||
|
||||
window = torch.hann_window(n_fft, device=z.device)
|
||||
|
||||
x = torch.istft(
|
||||
z,
|
||||
n_fft=n_fft,
|
||||
hop_length=hop_length,
|
||||
win_length=n_fft,
|
||||
window=window,
|
||||
center=True,
|
||||
normalized=True,
|
||||
length=length
|
||||
)
|
||||
|
||||
if is_mps:
|
||||
x = x.to(original_device)
|
||||
|
||||
_, out_length = x.shape
|
||||
return x.view(*other, out_length)
|
||||
|
||||
|
||||
def spectral_downmix(
|
||||
waveform: torch.Tensor, # (B, C, N) stereo or multi-channel
|
||||
n_fft: int = 2048,
|
||||
hop_length: int = 512
|
||||
) -> torch.Tensor: # Returns (B, 1, N) mono
|
||||
"""
|
||||
Downmixes a stereo or multi-channel audio to mono in the spectral domain
|
||||
to avoid phase cancellation, using robust STFT/iSTFT functions.
|
||||
"""
|
||||
if waveform.shape[1] <= 1:
|
||||
logger.debug("Input is already mono, skipping spectral downmix.")
|
||||
return waveform
|
||||
|
||||
batch_size, num_channels, num_samples = waveform.shape
|
||||
|
||||
# 1. Perform STFT on the multi-channel waveform
|
||||
# The spectro function handles reshaping (B,C,N) -> (B*C,N) if needed,
|
||||
# but passing it directly as (B,C,N) is also fine as it correctly reshapes.
|
||||
spectrogram_multi_ch = spectro(waveform, n_fft=n_fft, hop_length=hop_length)
|
||||
# Shape is (B, C, F, T_spec)
|
||||
|
||||
# 2. Separate magnitude and phase
|
||||
magnitudes = spectrogram_multi_ch.abs()
|
||||
phases = spectrogram_multi_ch.angle()
|
||||
|
||||
# 3. Average the magnitudes across channels
|
||||
avg_magnitude = torch.mean(magnitudes, dim=1) # (B, F, T_spec)
|
||||
|
||||
# 4. Use the phase from the first channel as the reference
|
||||
ref_phase = phases[:, 0, :, :] # (B, F, T_spec)
|
||||
|
||||
# 5. Reconstruct the new mono complex spectrogram
|
||||
mono_spectrogram_complex = torch.polar(avg_magnitude, ref_phase)
|
||||
|
||||
# 6. Perform inverse STFT
|
||||
# ispectro expects (..., F, T_spec) and returns (..., N)
|
||||
# Our mono_spectrogram_complex is (B, F, T_spec), which is perfect.
|
||||
mono_waveform_time = ispectro(mono_spectrogram_complex, hop_length=hop_length, length=num_samples)
|
||||
|
||||
# Reshape to (B, 1, N) for ComfyUI AUDIO standard
|
||||
return mono_waveform_time.unsqueeze(1)
|
||||
@@ -25,14 +25,14 @@ def converter_node():
|
||||
def test_converter_mono_to_stereo(converter_node):
|
||||
sr = 44100
|
||||
audio_mono = create_dummy_audio(2, 1, 1000, sr)
|
||||
(result_audio,) = converter_node.convert_channels(audio_mono, "mono_to_stereo")
|
||||
(result_audio,) = converter_node.convert_channels(audio_mono, "mono_to_stereo", "average")
|
||||
assert result_audio['waveform'].shape[1] == 2
|
||||
|
||||
|
||||
def test_converter_stereo_to_mono(converter_node):
|
||||
sr = 44100
|
||||
audio_stereo = create_dummy_audio(2, 2, 1000, sr)
|
||||
(result_audio,) = converter_node.convert_channels(audio_stereo, "stereo_to_mono")
|
||||
(result_audio,) = converter_node.convert_channels(audio_stereo, "stereo_to_mono", "average")
|
||||
assert result_audio['waveform'].shape[1] == 1
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ def test_converter_force_stereo_from_multichannel(converter_node):
|
||||
"""Tests if force_stereo takes the first channel of a 5.1 input."""
|
||||
sr = 44100
|
||||
audio_5_1 = create_dummy_audio(1, 6, 1000, sr)
|
||||
(result_audio,) = converter_node.convert_channels(audio_5_1, "force_stereo")
|
||||
(result_audio,) = converter_node.convert_channels(audio_5_1, "force_stereo", "average")
|
||||
assert result_audio['waveform'].shape[1] == 2
|
||||
# Check if the two output channels are identical copies of the first input channel
|
||||
assert torch.allclose(result_audio['waveform'][:, 0, :], result_audio['waveform'][:, 1, :])
|
||||
@@ -50,39 +50,109 @@ def test_converter_force_stereo_from_multichannel(converter_node):
|
||||
def test_converter_keep_channels(converter_node):
|
||||
sr = 44100
|
||||
audio_stereo = create_dummy_audio(1, 2, 1000, sr)
|
||||
(result_audio,) = converter_node.convert_channels(audio_stereo, "keep")
|
||||
(result_audio,) = converter_node.convert_channels(audio_stereo, "keep", "average")
|
||||
assert result_audio['waveform'].shape[1] == 2
|
||||
assert torch.allclose(result_audio['waveform'], audio_stereo['waveform'])
|
||||
|
||||
|
||||
def test_converter_stereo_to_mono_content_is_average(converter_node):
|
||||
@pytest.mark.parametrize("downmix_method", ["average", "standard_gain"])
|
||||
def test_converter_stereo_to_mono_math(converter_node, downmix_method):
|
||||
"""
|
||||
Strictly tests if stereo_to_mono correctly averages the L and R channels.
|
||||
Strictly tests that stereo-to-mono downmixing is mathematically correct
|
||||
for both 'average' and 'standard_gain' methods.
|
||||
"""
|
||||
sr = 44100
|
||||
samples = sr * 1 # 1 second
|
||||
samples = sr # 1 second
|
||||
|
||||
# Create distinct L and R channels
|
||||
# A simple ramp for L and a constant for R makes averaging easy to verify
|
||||
left_channel = torch.linspace(0, 1, samples)
|
||||
left_channel = torch.linspace(0.1, 0.9, samples)
|
||||
right_channel = torch.ones(samples) * 0.5
|
||||
|
||||
# Create the stereo waveform: (B=1, C=2, N)
|
||||
stereo_waveform = torch.stack((left_channel, right_channel), dim=0).unsqueeze(0)
|
||||
stereo_audio = {"waveform": stereo_waveform, "sample_rate": sr}
|
||||
|
||||
# Run the conversion
|
||||
(mono_audio,) = converter_node.convert_channels(stereo_audio, "stereo_to_mono")
|
||||
# Run the conversion with the parameterized downmix method
|
||||
(mono_audio,) = converter_node.convert_channels(
|
||||
audio=stereo_audio,
|
||||
channel_conversion="force_mono",
|
||||
downmix_method=downmix_method
|
||||
)
|
||||
|
||||
# Calculate the expected result based on the method
|
||||
if downmix_method == "average":
|
||||
expected_mono_waveform = ((left_channel + right_channel) / 2.0).unsqueeze(0).unsqueeze(0)
|
||||
elif downmix_method == "standard_gain":
|
||||
gain = 1.0 / (2**0.5) # Gain for stereo
|
||||
expected_mono_waveform = ((left_channel + right_channel) * gain).unsqueeze(0).unsqueeze(0)
|
||||
else:
|
||||
pytest.fail(f"Test case not implemented for downmix method: {downmix_method}")
|
||||
|
||||
# Assertions
|
||||
# The output waveform should be the mathematical average of L and R
|
||||
expected_mono_waveform = ((left_channel + right_channel) / 2.0).unsqueeze(0).unsqueeze(0) # Shape to (1,1,N)
|
||||
|
||||
assert mono_audio['waveform'].shape == (1, 1, samples)
|
||||
# Use a small tolerance (atol) for floating point comparisons
|
||||
assert torch.allclose(mono_audio['waveform'], expected_mono_waveform, atol=1e-6)
|
||||
|
||||
|
||||
def test_converter_antiphase_cancellation_time_domain(converter_node):
|
||||
"""
|
||||
Strictly tests and confirms that time-domain methods (both average and standard_gain)
|
||||
suffer from phase cancellation with anti-phase signals. This is expected behavior.
|
||||
"""
|
||||
sr = 44100
|
||||
samples = sr
|
||||
|
||||
# Create a sine wave for the left channel
|
||||
t = torch.linspace(0., 1., samples)
|
||||
left_channel = torch.sin(2 * torch.pi * 440.0 * t)
|
||||
# Create a perfectly inverted (anti-phase) wave for the right channel
|
||||
right_channel = -left_channel
|
||||
|
||||
stereo_waveform = torch.stack((left_channel, right_channel), dim=0).unsqueeze(0)
|
||||
antiphase_audio = {"waveform": stereo_waveform, "sample_rate": sr}
|
||||
|
||||
# Run conversion using 'average'
|
||||
(mono_audio_avg,) = converter_node.convert_channels(
|
||||
audio=antiphase_audio,
|
||||
channel_conversion="force_mono",
|
||||
downmix_method="average"
|
||||
)
|
||||
|
||||
# The result of (L + (-L)) / 2 should be a tensor of all zeros.
|
||||
assert torch.allclose(mono_audio_avg['waveform'], torch.zeros_like(mono_audio_avg['waveform']), atol=1e-6)
|
||||
|
||||
|
||||
def test_converter_spectral_downmix_avoids_cancellation(converter_node):
|
||||
"""
|
||||
Verifies that the node, when set to 'spectral' downmix, avoids cancelling
|
||||
an anti-phase signal, unlike the time-domain methods.
|
||||
"""
|
||||
sr = 44100
|
||||
samples = sr
|
||||
|
||||
t = torch.linspace(0., 1., samples)
|
||||
left_channel = torch.sin(2 * torch.pi * 440.0 * t)
|
||||
right_channel = -left_channel # Anti-phase
|
||||
|
||||
stereo_waveform = torch.stack((left_channel, right_channel), dim=0).unsqueeze(0)
|
||||
antiphase_audio = {"waveform": stereo_waveform, "sample_rate": sr}
|
||||
|
||||
# Run conversion using the new 'spectral' method
|
||||
(mono_audio_spectral,) = converter_node.convert_channels(
|
||||
audio=antiphase_audio,
|
||||
channel_conversion="force_mono",
|
||||
downmix_method="spectral"
|
||||
)
|
||||
|
||||
# The spectral downmix should preserve the energy of the signal.
|
||||
# Check that the mean of the squared signal (power) is significant.
|
||||
output_power = torch.mean(mono_audio_spectral['waveform'] ** 2)
|
||||
original_power = torch.mean(left_channel ** 2)
|
||||
|
||||
# It won't be identical to original_power due to phase reconstruction,
|
||||
# but it should be much, much greater than zero.
|
||||
assert output_power > original_power * 0.5 # Assert it retains at least 50% of the power
|
||||
assert output_power > 1e-3 # Assert it's not silent
|
||||
|
||||
|
||||
# --- Tests for AudioResampler ---
|
||||
|
||||
@pytest.fixture
|
||||
@@ -208,7 +278,8 @@ def test_advanced_processor_integration_smoke_test(advanced_processor_node):
|
||||
(result_audio,) = advanced_processor_node.process_audio(
|
||||
audio=audio_stereo_orig,
|
||||
channel_conversion="force_mono",
|
||||
target_sample_rate=target_sr
|
||||
target_sample_rate=target_sr,
|
||||
downmix_method="average"
|
||||
)
|
||||
|
||||
# Assertions:
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
"""
|
||||
Functional and regression tests for the AudioTestSignalGenerator node.
|
||||
"""
|
||||
|
||||
import bootstrap # noqa: F401
|
||||
import torch
|
||||
import pytest
|
||||
from nodes.utils.downmix import spectral_downmix
|
||||
|
||||
|
||||
# Helper to calculate Signal-to-Noise Ratio (SNR)
|
||||
def calculate_snr(signal, noisy_signal):
|
||||
noise = noisy_signal - signal
|
||||
signal_power = torch.mean(signal ** 2)
|
||||
noise_power = torch.mean(noise ** 2)
|
||||
if noise_power == 0:
|
||||
return float('inf')
|
||||
return 10 * torch.log10(signal_power / noise_power).item()
|
||||
|
||||
|
||||
def test_spectral_downmix_preserves_in_phase_signal():
|
||||
"""
|
||||
Tests that an in-phase stereo signal (L=R) is preserved correctly,
|
||||
measuring the SNR of the output vs the original mono source.
|
||||
"""
|
||||
sr = 44100
|
||||
samples = sr
|
||||
|
||||
# Create a mono sine wave
|
||||
t = torch.linspace(0., 1., samples)
|
||||
mono_signal = torch.sin(2 * torch.pi * 440.0 * t)
|
||||
|
||||
# Create a perfect stereo signal where L=R=mono_signal
|
||||
stereo_waveform = mono_signal.unsqueeze(0).repeat(1, 2, 1) # (1, 2, N)
|
||||
|
||||
# Run spectral downmix
|
||||
mono_output = spectral_downmix(stereo_waveform).squeeze() # Get 1D tensor
|
||||
|
||||
# The output should be very close to the original mono signal
|
||||
snr = calculate_snr(mono_signal, mono_output)
|
||||
|
||||
print(f"SNR for in-phase signal: {snr:.2f} dB")
|
||||
# A high SNR indicates the signals are very similar. >30dB is very good for this.
|
||||
assert snr > 30.0
|
||||
|
||||
|
||||
def test_spectral_downmix_avoids_antiphase_cancellation():
|
||||
"""
|
||||
Strict Test: Verifies that spectral downmix AVOIDS phase cancellation
|
||||
for anti-phase signals, which is its main advantage.
|
||||
"""
|
||||
sr = 44100
|
||||
samples = sr
|
||||
|
||||
t = torch.linspace(0., 1., samples)
|
||||
left_channel = torch.sin(2 * torch.pi * 440.0 * t)
|
||||
right_channel = -left_channel # Perfectly anti-phase
|
||||
|
||||
stereo_waveform = torch.stack((left_channel, right_channel), dim=0).unsqueeze(0)
|
||||
|
||||
# Run spectral downmix
|
||||
mono_output_spectral = spectral_downmix(stereo_waveform).squeeze()
|
||||
|
||||
# The magnitudes are identical, so the output should have the same magnitude
|
||||
# as one of the original channels. We check the power (mean of squares).
|
||||
output_power = torch.mean(mono_output_spectral ** 2)
|
||||
original_power = torch.mean(left_channel ** 2)
|
||||
|
||||
print(f"Power of original L channel: {original_power:.4f}")
|
||||
print(f"Power of spectral downmix output: {output_power:.4f}")
|
||||
|
||||
# The output power should be very close to the original channel's power.
|
||||
# A simple time-domain average would result in near-zero power.
|
||||
assert output_power > original_power * 0.9 # Should be very close
|
||||
assert output_power != pytest.approx(0.0)
|
||||
Reference in New Issue
Block a user