From cf6b9d1596621e2a78bd027373d6da253fb2a5b3 Mon Sep 17 00:00:00 2001 From: "Salvador E. Tropea" Date: Thu, 17 Jul 2025 11:04:21 -0300 Subject: [PATCH] [Nodes][Added] Spectral downmix for mono --- README.md | 20 +++- source/nodes/nodes_audio.py | 76 ++++++++++-- source/nodes/utils/downmix.py | 126 ++++++++++++++++++++ source/tests/test_audio_processing_nodes.py | 105 +++++++++++++--- source/tests/test_downmix_util.py | 75 ++++++++++++ 5 files changed, 373 insertions(+), 29 deletions(-) create mode 100644 source/nodes/utils/downmix.py create mode 100644 source/tests/test_downmix_util.py diff --git a/README.md b/README.md index 12ae940..2a26603 100644 --- a/README.md +++ b/README.md @@ -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. diff --git a/source/nodes/nodes_audio.py b/source/nodes/nodes_audio.py index 94a5f48..998383f 100644 --- a/source/nodes/nodes_audio.py +++ b/source/nodes/nodes_audio.py @@ -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. diff --git a/source/nodes/utils/downmix.py b/source/nodes/utils/downmix.py new file mode 100644 index 0000000..fb269e3 --- /dev/null +++ b/source/nodes/utils/downmix.py @@ -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) diff --git a/source/tests/test_audio_processing_nodes.py b/source/tests/test_audio_processing_nodes.py index 07f797a..40d2851 100644 --- a/source/tests/test_audio_processing_nodes.py +++ b/source/tests/test_audio_processing_nodes.py @@ -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: diff --git a/source/tests/test_downmix_util.py b/source/tests/test_downmix_util.py new file mode 100644 index 0000000..39ede59 --- /dev/null +++ b/source/tests/test_downmix_util.py @@ -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)