[Added][Audio Join][Audio Split] New nodes with test and example

This commit is contained in:
Salvador E. Tropea
2025-07-16 08:55:26 -03:00
parent c5e110e58b
commit 20f2e30df2
5 changed files with 270 additions and 0 deletions
+115
View File
@@ -659,3 +659,118 @@ class AudioMusicalNote:
# A UI warning could also be sent if this were a generator.
logger.error(f"Error parsing note: {e}. Defaulting to 440.0 Hz.")
return (440.0,)
class AudioJoin2Channels:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"audio_left": ("AUDIO", {"tooltip": "The audio signal for the left channel. Will be converted to mono."}),
"audio_right": ("AUDIO", {"tooltip": "The audio signal for the right channel. Will be converted to mono."}),
},
}
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio_out",)
FUNCTION = "join_channels"
CATEGORY = BASE_CATEGORY + "/" + MANIPULATION_CATEGORY
DESCRIPTION = "Joins two audio signals (L/R) into a single stereo audio signal."
UNIQUE_NAME = "SET_AudioJoin2Channels"
DISPLAY_NAME = "Audio Join 2 Channels"
def join_channels(self, audio_left: dict, audio_right: dict):
# 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")
# 2. Align the two mono signals in terms of sample rate and length.
# This reuses the robust logic from AudioBatchAligner.
# The output of the aligner will have channels=1 since both inputs are mono.
aligner = AudioBatchAligner(mono_left_audio, mono_right_audio)
aligned_left_wf, aligned_right_wf, target_sr = aligner.get_aligned_waveforms()
# aligned_left_wf is (B_left, 1, N), aligned_right_wf is (B_right, 1, N)
# 3. Handle mismatched batch sizes.
b_left, b_right = aligned_left_wf.shape[0], aligned_right_wf.shape[0]
target_batch_size = max(b_left, b_right)
# Create final waveforms with the target batch size, repeating the last item if needed.
final_left_wf = aligned_left_wf
if b_left < target_batch_size:
last_left_item = aligned_left_wf[-1:, :, :] # (1, 1, N)
repeats_needed = target_batch_size - b_left
final_left_wf = torch.cat([aligned_left_wf, last_left_item.repeat(repeats_needed, 1, 1)], dim=0)
final_right_wf = aligned_right_wf
if b_right < target_batch_size:
last_right_item = aligned_right_wf[-1:, :, :] # (1, 1, N)
repeats_needed = target_batch_size - b_right
final_right_wf = torch.cat([aligned_right_wf, last_right_item.repeat(repeats_needed, 1, 1)], dim=0)
# At this point, final_left_wf and final_right_wf are both (target_batch_size, 1, N)
# 4. Concatenate the mono channels into a stereo signal.
# `torch.cat` along the channel dimension (dim=1).
stereo_waveform = torch.cat((final_left_wf, final_right_wf), dim=1) # (B, 2, N)
logger.info(f"Joined L/R channels into stereo audio. Final shape: {stereo_waveform.shape}, SR: {target_sr}")
output_audio = {
"waveform": stereo_waveform,
"sample_rate": target_sr
}
return (output_audio,)
class AudioSplit2Channels:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"audio": ("AUDIO", {"tooltip": "A stereo audio signal to split into separate channels."}),
},
}
RETURN_TYPES = ("AUDIO", "AUDIO")
RETURN_NAMES = ("audio_left", "audio_right")
FUNCTION = "split_channels"
CATEGORY = BASE_CATEGORY + "/" + MANIPULATION_CATEGORY
DESCRIPTION = "Splits a stereo audio signal into two separate mono audio signals (L/R)."
UNIQUE_NAME = "SET_AudioSplit2Channels"
DISPLAY_NAME = "Audio Split 2 Channels"
def split_channels(self, audio: dict):
waveform = audio['waveform'] # (B, C, N)
sample_rate = audio['sample_rate']
num_channels = waveform.shape[1]
# 1. Validate that the input is stereo.
if num_channels != 2:
msg = f"Input audio must be stereo (2 channels) to be split. Got {num_channels} channels."
logger.error(msg)
# This is a hard requirement, so raising an error is appropriate.
raise ValueError(msg)
# 2. Slice the tensor along the channel dimension.
# Slicing with [:, 0:1, :] keeps the channel dimension as 1, so the output is (B, 1, N)
left_channel_wf = waveform[:, 0:1, :]
right_channel_wf = waveform[:, 1:2, :]
logger.info(f"Split stereo audio into L/R channels. Output shape for each: {left_channel_wf.shape}")
# 3. Package each channel into its own ComfyUI AUDIO dict.
audio_left = {
"waveform": left_channel_wf,
"sample_rate": sample_rate
}
audio_right = {
"waveform": right_channel_wf,
"sample_rate": sample_rate
}
return (audio_left, audio_right)
+126
View File
@@ -0,0 +1,126 @@
"""
Regression tests for the AudioJoin2Channels and AudioSplit2Channels nodes in ComfyUI-AudioBatch.
"""
import bootstrap # noqa: F401
import torch
import pytest
from nodes.nodes_audio import AudioJoin2Channels, AudioSplit2Channels
# Helper function
def create_dummy_audio(batch_size, channels, samples, sr, device='cpu'):
waveform = torch.randn(batch_size, channels, samples, device=device, dtype=torch.float32)
return {"waveform": waveform, "sample_rate": sr}
@pytest.fixture
def join_node():
return AudioJoin2Channels()
@pytest.fixture
def split_node():
return AudioSplit2Channels()
# --- Tests for AudioJoin2Channels ---
def test_join_simple(join_node):
"""Tests joining two perfectly matched mono signals."""
sr = 44100
samples = 1000
left_audio = create_dummy_audio(1, 1, samples, sr)
right_audio = create_dummy_audio(1, 1, samples, sr)
(stereo_audio,) = join_node.join_channels(left_audio, right_audio)
assert stereo_audio['waveform'].shape == (1, 2, samples)
# Check if L/R channels match the original mono inputs
assert torch.allclose(stereo_audio['waveform'][:, 0:1, :], left_audio['waveform'])
assert torch.allclose(stereo_audio['waveform'][:, 1:2, :], right_audio['waveform'])
def test_join_converts_inputs_to_mono(join_node):
"""Tests that stereo inputs are correctly converted to mono before joining."""
sr = 44100
samples = 1000
# Left input is stereo, Right is mono
left_stereo_in = create_dummy_audio(1, 2, samples, sr)
right_mono_in = create_dummy_audio(1, 1, samples, sr)
(result_audio,) = join_node.join_channels(left_stereo_in, right_mono_in)
# Left channel of the output should be the average of the stereo input
expected_left_mono = torch.mean(left_stereo_in['waveform'], dim=1, keepdim=True)
assert result_audio['waveform'].shape == (1, 2, samples)
assert torch.allclose(result_audio['waveform'][:, 0:1, :], expected_left_mono)
assert torch.allclose(result_audio['waveform'][:, 1:2, :], right_mono_in['waveform'])
def test_join_aligns_and_batches(join_node):
"""Tests that mismatched SR, length, and batch sizes are handled."""
# Left: B=2, 1ch, 44.1k SR, 1s long
audio_left = create_dummy_audio(2, 1, 44100, 44100)
# Right: B=3, 2ch, 22.05k SR, 0.5s long
audio_right = create_dummy_audio(3, 2, 11025, 22050)
(result_audio,) = join_node.join_channels(audio_left, audio_right)
# Expected output shape:
# Batch size = max(2, 3) = 3
# Channels = 2 (stereo)
# SR = 44100 (from left)
# Samples = 44100 (from left, since right is 11025*2=22050 after resampling)
assert result_audio['waveform'].shape == (3, 2, 44100)
assert result_audio['sample_rate'] == 44100
# Check the repeated last item for the left channel
assert torch.allclose(result_audio['waveform'][1, 0, :], result_audio['waveform'][2, 0, :])
# --- Tests for AudioSplit2Channels ---
def test_split_simple(split_node):
"""Tests splitting a standard stereo signal."""
sr = 44100
samples = 1000
stereo_audio = create_dummy_audio(1, 2, samples, sr)
(left_out, right_out) = split_node.split_channels(stereo_audio)
# Check left channel output
assert left_out['sample_rate'] == sr
assert left_out['waveform'].shape == (1, 1, samples)
assert torch.allclose(left_out['waveform'], stereo_audio['waveform'][:, 0:1, :])
# Check right channel output
assert right_out['sample_rate'] == sr
assert right_out['waveform'].shape == (1, 1, samples)
assert torch.allclose(right_out['waveform'], stereo_audio['waveform'][:, 1:2, :])
def test_split_with_batch(split_node):
"""Tests that splitting preserves the batch dimension."""
sr = 44100
samples = 1000
stereo_batch = create_dummy_audio(5, 2, samples, sr)
(left_out, right_out) = split_node.split_channels(stereo_batch)
assert left_out['waveform'].shape == (5, 1, samples)
assert right_out['waveform'].shape == (5, 1, samples)
def test_split_raises_error_on_non_stereo(split_node):
"""Tests that an error is raised if the input is not stereo."""
# Test with mono
mono_audio = create_dummy_audio(1, 1, 1000, 44100)
with pytest.raises(ValueError, match="Input audio must be stereo"):
split_node.split_channels(mono_audio)
# Test with 3 channels
three_ch_audio = create_dummy_audio(1, 3, 1000, 44100)
with pytest.raises(ValueError, match="Input audio must be stereo"):
split_node.split_channels(three_ch_audio)