[Nodes][Added] Spectral downmix for mono

This commit is contained in:
Salvador E. Tropea
2025-07-17 11:04:21 -03:00
parent 111b58a615
commit cf6b9d1596
5 changed files with 373 additions and 29 deletions
+19 -1
View File
@@ -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
View File
@@ -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.
+126
View File
@@ -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)
+88 -17
View File
@@ -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:
+75
View File
@@ -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)