diff --git a/source/Makefile b/source/Makefile new file mode 100644 index 0000000..036a68f --- /dev/null +++ b/source/Makefile @@ -0,0 +1,70 @@ +# Makefile for ComfyUI-AudioBatch project + +# Use the python from the current virtual environment +# This assumes you have activated your ComfyUI venv before running `make` +# If not, you might need to specify the path to the python executable +PYTHON := python + +# Define test command with necessary pytest options +# --import-mode=append: Prevents pytest from prepending the test root to sys.path, +# solving many relative/absolute import conflicts. +# -v: Verbose output. +# -s: Show print statements (useful for debugging tests). +# --ignore=path/to/ignore: Example of ignoring a directory if needed. +PYTEST_CMD := pytest --import-mode=append -v -s + +# --- Targets --- + +.PHONY: help test test-one lint format check + +help: + @echo "Makefile for ComfyUI-AudioBatch" + @echo "" + @echo "Usage:" + @echo " make help Show this help message." + @echo " make install Install development dependencies from requirements-dev.txt." + @echo " make test Run all tests in the 'tests/' directory." + @echo " make test-one Run a specific test file. Usage: make test-one file=tests/test_audio_blend.py" + @echo " make lint Run flake8 linter on the source code." + @echo " make format Run black code formatter on the source code." + @echo " make check Run all checks (lint, type check, tests)." + @echo " make type-check Run mypy for static type checking." + + +# Target to install development dependencies +install: + @echo "Installing development dependencies..." + @$(PYTHON) -m pip install -r requirements-dev.txt + +# Target to run all tests +test: + @echo "Running all tests..." + @$(PYTEST_CMD) tests/ + +# Target to run a single test file, specified with `make test-one file=...` +test-one: + @if [ -z "$(file)" ]; then \ + echo "Error: Please specify a file to test. Usage: make test-one file=tests/your_test_file.py"; \ + exit 1; \ + fi + @echo "Running test for file: $(file)" + @$(PYTEST_CMD) $(file) + +# Target to run the linter (e.g., flake8) +lint: + @echo "Running linter (flake8)..." + flake8 . + +# Target to run mypy for type checking +type-check: + @echo "Running static type checker (mypy)..." + @$(PYTHON) -m mypy . + +# Target to format the code (e.g., black) +format: + @echo "Formatting code (black)..." + @$(PYTHON) -m black . + +# Target to run all checks together +check: lint type-check test + @echo "All checks completed." \ No newline at end of file diff --git a/source/tests/bootstrap/__init__.py b/source/tests/bootstrap/__init__.py new file mode 100644 index 0000000..fae4571 --- /dev/null +++ b/source/tests/bootstrap/__init__.py @@ -0,0 +1,55 @@ +# Copyright (c) 2025 Salvador E. Tropea +# Copyright (c) 2025 Instituto Nacional de TecnologĂ­a Industrial +# License: GPLv3 +# Project: ComfyUI-AudioSeparation +# +# Why such a complex thing? +# Python imports are broken by design, and here we hit a huge limitation: +# 1. ComfyUI nodes MUST use relative imports, if you don't do it things like "import utils" becomes ambiguous +# You might think this can be overcome polluting the sys.path, but this isn't true. If ComfyUI, or some +# other node, already imported an module named "utils" you'll get the already imported module, not the one +# you want. And polluting sys.path can make other nodes, or ComfyUI itself, import the wrong module when +# they use a "non top-level import". +# Conclusion: The only safe way to import from a ComfyUI node is by using relative imports. +# 2. We have tools in the "tool" subdir, they are intended to run as standalone scripts. They can't use +# relative imports to access the node sub modules because you'll hit the error "ImportError: attempted +# relative import with no known parent package". So imports in tools MUST be absolute. This can be solved +# adding the node root to sys.path. Here we are not polluting the sys.path because we are the top-level. +# Conclusion: This script is used to solve adding the correct path to sys.path +# 3. As you MUST use relative imports in the node sub-modules, when a sub-module depends on another sub-module +# it will do something like "from ..XXXX", this is OK when all is relative. But tools are using absolute +# imports, so you'll hit the error "ImportError: attempted relative import beyond top-level package" +# Conclusion: All sub-modules must be wrapped by an umbrella sub-module, this is what "source" is. +# +# Structure: +# Node-name/ <-- is a module from ComfyUI point of view, dynamically imported using importlib +# \-- __init__.py <-- needed because we are a module, relative imports +# \-- nodes.py <-- relative imports +# | +# \-- source/ <-- umbrella package, just to solve relative imports +# | \-- __init__.py +# | | +# | \-- utils/ <-- a submodule, uses relative imports +# | | \-- __init__.py +# | | \-- misc.py +# | | +# | \-- db/ <-- another submodule, uses relative imports +# | \-- __init__.py +# | \-- models_db.py <-- Must use relative imports, can access ..utils.misc +# | +# \-- tool/ +# \-- batch_convert.py <-- a tool, must use absolute import of "source" +# | +# \-- bootstrap/ +# \-- __init__.py <-- THIS file, adds to sys.path to make "source" available + +# --- BOOTSTRAP: Make the script aware of the project root --- +import os +import sys +# 1. Get the absolute path of the current script's directory (tool/) +script_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +# 2. Get the project root by going one directory up (MyAwesomeNode/) +project_root = os.path.dirname(script_dir) +# 3. Add the project root to the Python path +if project_root not in sys.path: + sys.path.insert(0, project_root) diff --git a/source/tests/test_audio_batch.py b/source/tests/test_audio_batch.py new file mode 100644 index 0000000..2a2dc8d --- /dev/null +++ b/source/tests/test_audio_batch.py @@ -0,0 +1,142 @@ +""" +Regression tests for the AudioBatch node in ComfyUI-AudioBatch. +""" + +import bootstrap # noqa: F401 +import torch +import pytest +from nodes.nodes_audio import AudioBatch + + +# Helper function (can be moved to a shared conftest.py or test_utils.py later) +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 audio_batch_node(): + return AudioBatch() + + +def test_batch_simple_match(audio_batch_node): + """Tests batching two identical audio inputs.""" + sr = 44100 + audio1 = create_dummy_audio(1, 2, 1000, sr) + audio2 = create_dummy_audio(1, 2, 1000, sr) + + (result_audio,) = audio_batch_node.batch_audio(audio1, audio2) + + expected_waveform = torch.cat((audio1['waveform'], audio2['waveform']), dim=0) + + assert result_audio['sample_rate'] == sr + assert result_audio['waveform'].shape == (2, 2, 1000) + assert torch.allclose(result_audio['waveform'], expected_waveform) + + +def test_batch_different_batch_sizes(audio_batch_node): + """Tests batching inputs with B=2 and B=3, expecting an output of B=5.""" + sr = 44100 + audio1 = create_dummy_audio(2, 2, 1000, sr) + audio2 = create_dummy_audio(3, 2, 1000, sr) + + (result_audio,) = audio_batch_node.batch_audio(audio1, audio2) + + assert result_audio['waveform'].shape[0] == 5 # 2 + 3 + + +def test_batch_different_lengths_padding(audio_batch_node): + """Tests if the shorter audio is padded to match the longer one.""" + sr = 44100 + audio1 = create_dummy_audio(1, 2, 500, sr) + audio2 = create_dummy_audio(1, 2, 1000, sr) + + (result_audio,) = audio_batch_node.batch_audio(audio1, audio2) + + # Both items in the output batch should have the length of the longer audio + assert result_audio['waveform'].shape[2] == 1000 + # The first item should have been padded + assert result_audio['waveform'][0].shape[1] == 1000 + + +def test_batch_mono_and_stereo(audio_batch_node): + """Tests if the mono audio is converted to stereo to match the other input.""" + sr = 44100 + audio1_mono = create_dummy_audio(1, 1, 1000, sr) + audio2_stereo = create_dummy_audio(1, 2, 1000, sr) + + (result_audio,) = audio_batch_node.batch_audio(audio1_mono, audio2_stereo) + + # Both items in the output batch should be stereo + assert result_audio['waveform'].shape[1] == 2 + assert result_audio['waveform'].shape[0] == 2 + + +def test_batch_different_sample_rates(audio_batch_node): + """Tests if audio2 is resampled to match audio1's sample rate.""" + sr1, sr2 = 44100, 22050 + samples1, samples2 = 2000, 1000 + audio1 = create_dummy_audio(1, 2, samples1, sr1) + audio2 = create_dummy_audio(1, 2, samples2, sr2) + + (result_audio,) = audio_batch_node.batch_audio(audio1, audio2) + + # Output SR should be sr1 + assert result_audio['sample_rate'] == sr1 + # Length of audio2 after resampling is approx samples2 * (sr1/sr2) = 1000 * 2 = 2000 + # So both items should have length 2000 + assert result_audio['waveform'].shape[2] == 2000 + + +def test_batch_full_alignment_and_content(audio_batch_node): + """ + Strictly tests the end-to-end alignment (SR, channels, length) and content + of the batching process. + """ + # --- Input 1: A stereo signal at 48000 Hz, 1 second long --- + sr1 = 48000 + len1 = sr1 * 1 + # Create L/R channels that are easy to identify (ramps) + wf1_l = torch.linspace(0.1, 0.2, len1) + wf1_r = torch.linspace(0.3, 0.4, len1) + wf1_orig = torch.stack((wf1_l, wf1_r), dim=0).unsqueeze(0) # (1, 2, 48000) + audio1 = {"waveform": wf1_orig, "sample_rate": sr1} + + # --- Input 2: A mono signal at 16000 Hz, 1.5 seconds long --- + sr2 = 16000 + len2 = int(sr2 * 1.5) + wf2_orig = (torch.ones(1, 1, len2) * 0.5) # Constant value mono signal, (1, 1, 24000) + audio2 = {"waveform": wf2_orig, "sample_rate": sr2} + + # --- Action --- + (result_audio,) = audio_batch_node.batch_audio(audio1, audio2) + result_wf = result_audio['waveform'] + + # --- Assertions --- + # 1. Final parameters should match audio1's SR, be stereo, and have the longest length after resampling. + target_sr = sr1 # 48000 + target_channels = 2 + len2_resampled = int(len2 * (target_sr / sr2)) # 24000 * 3 = 72000 + target_samples = max(len1, len2_resampled) # max(48000, 72000) = 72000 + + assert result_audio['sample_rate'] == target_sr + assert result_wf.shape == (2, target_channels, target_samples) # B=2, C=2, N=72000 + + # 2. Verify content of the first item (was audio1) + item1_output = result_wf[0] # (2, 72000) + # The first part should match the original audio1 + assert torch.allclose(item1_output[:, :len1], wf1_orig.squeeze(0)) + # The padded part should be zeros + assert torch.all(item1_output[:, len1:] == 0) + + # 3. Verify content of the second item (was audio2) + item2_output = result_wf[1] # (2, 72000) + # It was mono, so both output channels should be identical + assert torch.allclose(item2_output[0, :], item2_output[1, :]) + # The content should be a resampled version of the constant 0.5 signal. + # A resampled constant signal should still be (roughly) constant. + # We check the first part, up to where the original content was. + resampled_part = item2_output[0, :len2_resampled] + # The average value should be very close to 0.5. Due to filtering (Gibbs effect), + # it won't be perfect, especially at the edges. + assert torch.mean(resampled_part[100:-100]).item() == pytest.approx(0.5, abs=1e-3) diff --git a/source/tests/test_audio_blend.py b/source/tests/test_audio_blend.py new file mode 100644 index 0000000..2953d78 --- /dev/null +++ b/source/tests/test_audio_blend.py @@ -0,0 +1,173 @@ +""" +Regression tests for the AudioBlend node in ComfyUI-AudioBatch. +""" + +import torch +import pytest +import logging + +import bootstrap # noqa: F401 +# Now we can import the classes and functions to be tested +# We need AudioBlend, but it uses AudioBatchAligner internally, which is in the same file. +# We also import the logger to see node output during tests. +from nodes.nodes_audio import AudioBlend, logger as node_logger + +# Configure logging for tests to see output from the node +node_logger.setLevel(logging.DEBUG) + + +# --- Pytest Fixtures (Reusable Setup) --- + +@pytest.fixture +def audio_blend_node(): + """Provides an instance of the AudioBlend node for tests.""" + return AudioBlend() + + +# --- Helper Functions --- + +def create_dummy_audio(batch_size: int, channels: int, samples: int, sr: int, + device: str = 'cpu') -> dict: + """Creates a standard ComfyUI AUDIO dictionary for testing.""" + # Using torch.ones for predictable values, multiplied by a small float + # to avoid pure 1s which can mask errors in blending. + waveform = torch.ones(batch_size, channels, samples, device=device, dtype=torch.float32) * 0.5 + return {"waveform": waveform, "sample_rate": sr} + + +# --- Test Cases for AudioBlend --- + +def test_blend_simple_match(audio_blend_node): + """ + Tests blending two identical audio inputs. + This is the "happy path" where no resampling, channel conversion, or padding is needed. + """ + sr = 44100 + audio1 = create_dummy_audio(batch_size=1, channels=2, samples=1000, sr=sr) + audio2 = create_dummy_audio(batch_size=1, channels=2, samples=1000, sr=sr) + gain1, gain2 = 0.5, 0.8 + + (result_audio,) = audio_blend_node.blend_audio(audio1, gain1, gain2, audio2) + + # Assertions + expected_waveform = (audio1['waveform'] * gain1) + (audio2['waveform'] * gain2) + assert result_audio['sample_rate'] == sr + assert result_audio['waveform'].shape == expected_waveform.shape + assert torch.allclose(result_audio['waveform'], expected_waveform) + + +def test_blend_no_audio2(audio_blend_node): + """ + Tests the case where the optional `audio2` input is not provided (is None). + """ + sr = 44100 + audio1 = create_dummy_audio(batch_size=2, channels=1, samples=1000, sr=sr) + gain1, gain2 = 0.7, 1.0 # gain2 should be ignored + + (result_audio,) = audio_blend_node.blend_audio(audio1, gain1, gain2, audio2=None) + + # Assertions + expected_waveform = audio1['waveform'] * gain1 + assert result_audio['sample_rate'] == sr + assert result_audio['waveform'].shape == expected_waveform.shape + assert torch.allclose(result_audio['waveform'], expected_waveform) + + +def test_blend_mismatched_sample_rate(audio_blend_node): + """ + Tests if audio2 is correctly resampled to match audio1's sample rate. + """ + sr1, sr2 = 44100, 22050 + audio1 = create_dummy_audio(batch_size=1, channels=2, samples=2000, sr=sr1) + audio2 = create_dummy_audio(batch_size=1, channels=2, samples=1000, sr=sr2) + gain1, gain2 = 1.0, 1.0 + + (result_audio,) = audio_blend_node.blend_audio(audio1, gain1, gain2, audio2) + + # Assertions + # Output SR should match audio1 + assert result_audio['sample_rate'] == sr1 + # Output length should match audio1, as it's longer after resampling audio2 + assert result_audio['waveform'].shape[2] == audio1['waveform'].shape[2] + # Output batch and channels should match + assert result_audio['waveform'].shape[0] == 1 + assert result_audio['waveform'].shape[1] == 2 + + +def test_blend_mono_and_stereo(audio_blend_node): + """ + Tests if a mono input is correctly converted to stereo when blended with a stereo input. + """ + sr = 44100 + audio1_mono = create_dummy_audio(batch_size=1, channels=1, samples=1000, sr=sr) + audio2_stereo = create_dummy_audio(batch_size=1, channels=2, samples=1000, sr=sr) + gain1, gain2 = 1.0, 1.0 + + (result_audio,) = audio_blend_node.blend_audio(audio1_mono, gain1, gain2, audio2_stereo) + + # Assertions + assert result_audio['sample_rate'] == sr + # Output should be stereo (2 channels) + assert result_audio['waveform'].shape[1] == 2 + assert result_audio['waveform'].shape[2] == 1000 + + # Check if the mono part was correctly duplicated and blended. + # The two output channels should differ only by the difference in audio2's stereo channels. + output_diff = result_audio['waveform'][:, 0, :] - result_audio['waveform'][:, 1, :] + input2_diff = audio2_stereo['waveform'][:, 0, :] - audio2_stereo['waveform'][:, 1, :] + # Since gain2 is 1.0, the difference should be the same. + assert torch.allclose(output_diff, input2_diff) + + +def test_blend_mismatched_length(audio_blend_node): + """ + Tests if the shorter audio is correctly padded to match the longer one. + """ + sr = 44100 + len_short, len_long = 500, 1000 + audio1_short = create_dummy_audio(batch_size=1, channels=2, samples=len_short, sr=sr) + audio2_long = create_dummy_audio(batch_size=1, channels=2, samples=len_long, sr=sr) + gain1, gain2 = 1.0, 1.0 + + (result_audio,) = audio_blend_node.blend_audio(audio1_short, gain1, gain2, audio2_long) + + # Assertions + assert result_audio['sample_rate'] == sr + # Output length should match the longer audio + assert result_audio['waveform'].shape[2] == len_long + + # Check the tail end of the blended audio. It should only contain audio2's contribution. + tail_start_index = len_short + output_tail = result_audio['waveform'][..., tail_start_index:] + expected_tail = audio2_long['waveform'][..., tail_start_index:] * gain2 + assert torch.allclose(output_tail, expected_tail) + + +def test_blend_mismatched_batch_size(audio_blend_node): + """ + Tests blending inputs with different batch sizes. The output batch size should be the max. + """ + sr = 44100 + b1, b2 = 2, 3 + samples = 1000 + channels = 2 + audio1 = create_dummy_audio(batch_size=b1, channels=channels, samples=samples, sr=sr) + audio2 = create_dummy_audio(batch_size=b2, channels=channels, samples=samples, sr=sr) + gain1, gain2 = 1.0, 1.0 + + (result_audio,) = audio_blend_node.blend_audio(audio1, gain1, gain2, audio2) + output_waveform = result_audio['waveform'] + + # Assertions + # Output batch size should be the max of the two inputs + assert output_waveform.shape[0] == max(b1, b2) + assert output_waveform.shape[1] == channels + assert output_waveform.shape[2] == samples + + # The first two items should be a blend of audio1 and audio2 + expected_blend_part = (audio1['waveform'] * gain1) + (audio2['waveform'][:b1] * gain2) + assert torch.allclose(output_waveform[:b1], expected_blend_part) + + # The last item (index 2) should only contain the contribution from audio2 + expected_last_item = audio2['waveform'][2] * gain2 + assert torch.allclose(output_waveform[2], expected_last_item) diff --git a/source/tests/test_audio_processing_nodes.py b/source/tests/test_audio_processing_nodes.py new file mode 100644 index 0000000..7fa8285 --- /dev/null +++ b/source/tests/test_audio_processing_nodes.py @@ -0,0 +1,230 @@ +""" +Regression tests for the AudioChannelConverter, AudioResampler, AudioProcessAdvanced node in ComfyUI-AudioBatch. +""" + +import bootstrap # noqa: F401 +import logging +import torch +import pytest +from nodes.nodes_audio import AudioChannelConverter, AudioResampler, AudioProcessAdvanced + + +# 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} + + +# --- Tests for AudioChannelConverter --- + +@pytest.fixture +def converter_node(): + return AudioChannelConverter() + + +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") + 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") + assert result_audio['waveform'].shape[1] == 1 + + +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") + 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, :]) + assert torch.allclose(result_audio['waveform'][:, 0, :], audio_5_1['waveform'][:, 0, :]) + + +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") + 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): + """ + Strictly tests if stereo_to_mono correctly averages the L and R channels. + """ + sr = 44100 + samples = sr * 1 # 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) + 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") + + # 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) + + +# --- Tests for AudioResampler --- + +@pytest.fixture +def resampler_node(): + return AudioResampler() + + +def test_resampler_upsample(resampler_node): + sr_orig, sr_target = 22050, 44100 + samples_orig = 1000 + audio = create_dummy_audio(2, 2, samples_orig, sr_orig) + (result_audio,) = resampler_node.resample_audio(audio, sr_target) + + expected_samples = int(samples_orig * (sr_target / sr_orig)) + assert result_audio['sample_rate'] == sr_target + assert result_audio['waveform'].shape[2] == expected_samples + + +def test_resampler_downsample(resampler_node): + sr_orig, sr_target = 48000, 16000 + samples_orig = 3000 + audio = create_dummy_audio(1, 1, samples_orig, sr_orig) + (result_audio,) = resampler_node.resample_audio(audio, sr_target) + + expected_samples = int(samples_orig * (sr_target / sr_orig)) + assert result_audio['sample_rate'] == sr_target + assert abs(result_audio['waveform'].shape[2] - expected_samples) < 2 # Resampling can have off-by-one + + +def test_resampler_no_op(resampler_node): + """Tests if resampling is skipped when target SR is 0 or same as original.""" + sr = 44100 + audio = create_dummy_audio(1, 2, 1000, sr) + + # Test with target_sr = 0 + (result_audio_zero,) = resampler_node.resample_audio(audio, 0) + assert torch.allclose(result_audio_zero['waveform'], audio['waveform']) + assert result_audio_zero['sample_rate'] == sr + + # Test with target_sr = original_sr + (result_audio_same,) = resampler_node.resample_audio(audio, sr) + assert torch.allclose(result_audio_same['waveform'], audio['waveform']) + assert result_audio_same['sample_rate'] == sr + + +def get_peak_frequency(waveform: torch.Tensor, sample_rate: int) -> float: + """Helper to find the dominant frequency in a waveform using FFT.""" + if waveform.ndim > 1: + # Use the first channel and first batch item for analysis + waveform = waveform.squeeze() + if waveform.ndim > 1: + waveform = waveform[0] + + # Perform Real Fast Fourier Transform + fft_result = torch.fft.rfft(waveform) + # Get frequency bins for the FFT result + freq_bins = torch.fft.rfftfreq(n=waveform.size(-1), d=1./sample_rate) + # Find the index of the maximum magnitude in the FFT + peak_index = torch.argmax(torch.abs(fft_result)) + # Get the frequency corresponding to that peak index + peak_freq = freq_bins[peak_index].item() + return peak_freq + + +def test_resampler_preserves_frequency_content(resampler_node): + """ + Strictly tests if resampling correctly preserves the pitch (frequency) of a signal. + """ + # Original signal: A 440 Hz sine wave at 44100 Hz SR + sr_orig = 44100 + target_freq = 440.0 # A4 note + duration = 1.0 + samples_orig = int(sr_orig * duration) + t = torch.linspace(0., duration, samples_orig) + original_waveform = torch.sin(2 * torch.pi * target_freq * t).unsqueeze(0).unsqueeze(0) # (1,1,N) + original_audio = {"waveform": original_waveform, "sample_rate": sr_orig} + + # Check the frequency of the original signal to ensure our helper works + original_peak_freq = get_peak_frequency(original_waveform, sr_orig) + assert abs(original_peak_freq - target_freq) < 1.0 # Allow for slight FFT bin inaccuracy + + # --- Test Downsampling --- + sr_down = 16000 + (downsampled_audio,) = resampler_node.resample_audio(original_audio, sr_down) + + # The frequency content should still be centered at 440 Hz + downsampled_peak_freq = get_peak_frequency(downsampled_audio['waveform'], sr_down) + logging.info(f"Downsampled from {sr_orig}Hz to {sr_down}Hz. Original peak: {original_peak_freq:.2f}Hz, " + f"New peak: {downsampled_peak_freq:.2f}Hz") + assert abs(downsampled_peak_freq - target_freq) < 2.0 # Allow slightly larger tolerance for resamplers + + # --- Test Upsampling --- + sr_up = 48000 + (upsampled_audio,) = resampler_node.resample_audio(original_audio, sr_up) + + # The frequency content should still be centered at 440 Hz + upsampled_peak_freq = get_peak_frequency(upsampled_audio['waveform'], sr_up) + logging.info(f"Upsampled from {sr_orig}Hz to {sr_up}Hz. Original peak: {original_peak_freq:.2f}Hz, " + f"New peak: {upsampled_peak_freq:.2f}Hz") + assert abs(upsampled_peak_freq - target_freq) < 2.0 + + +# --- Test for AudioProcessAdvanced --- + +@pytest.fixture +def advanced_processor_node(): + return AudioProcessAdvanced() + + +def test_advanced_processor_integration_smoke_test(advanced_processor_node): + """ + A simple "smoke test" to ensure AudioProcessAdvanced correctly calls its + sub-components and produces an output of the expected shape and SR. + We don't need to re-test all permutations, as those are covered + by the unit tests for the individual converter and resampler nodes. + """ + # Setup: Start with stereo audio at a high sample rate + sr_orig = 44100 + audio_stereo_orig = create_dummy_audio(1, 2, sr_orig * 2, sr_orig) # 2 seconds + + # Action: Convert to mono and downsample to 16k + target_sr = 16000 + (result_audio,) = advanced_processor_node.process_audio( + audio=audio_stereo_orig, + channel_conversion="force_mono", + target_sample_rate=target_sr + ) + + # Assertions: + # 1. Did it successfully produce an output? + assert result_audio is not None + assert "waveform" in result_audio + assert "sample_rate" in result_audio + + # 2. Is the final sample rate correct? + assert result_audio['sample_rate'] == target_sr + + # 3. Is the final channel count correct? + assert result_audio['waveform'].shape[1] == 1 # Should be mono + + # 4. Is the final number of samples roughly correct? + original_samples = sr_orig * 2 + expected_samples = int(original_samples * (target_sr / sr_orig)) + # Check that the actual number of samples is close to the expected number + assert abs(result_audio['waveform'].shape[2] - expected_samples) < 2 diff --git a/source/tests/test_select_audio_from_batch.py b/source/tests/test_select_audio_from_batch.py new file mode 100644 index 0000000..6dfacfa --- /dev/null +++ b/source/tests/test_select_audio_from_batch.py @@ -0,0 +1,75 @@ +""" +Regression tests for the SelectAudioFromBatch node in ComfyUI-AudioBatch. +""" + +import bootstrap # noqa: F401 +import torch +import pytest +from nodes.nodes_audio import SelectAudioFromBatch + + +# 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 select_node(): + return SelectAudioFromBatch() + + +def test_select_valid_index(select_node): + """Tests selecting a valid index from the batch.""" + sr = 44100 + audio_batch = create_dummy_audio(5, 2, 1000, sr) + + (result_audio,) = select_node.select_audio(audio_batch, index=2, behavior_out_of_range="error", + silence_duration_seconds=1.0) + + # Output should be a batch of 1 + assert result_audio['waveform'].shape[0] == 1 + # It should be the 3rd item from the original batch + assert torch.allclose(result_audio['waveform'], audio_batch['waveform'][2:3, :, :]) + assert result_audio['sample_rate'] == sr + + +def test_select_out_of_range_error(select_node): + """Tests if an error is raised when index is out of range and behavior is 'error'.""" + sr = 44100 + audio_batch = create_dummy_audio(3, 2, 1000, sr) + + with pytest.raises(ValueError, match="Index 5 is out of range for batch of size 3"): + select_node.select_audio(audio_batch, index=5, behavior_out_of_range="error", silence_duration_seconds=1.0) + + +def test_select_out_of_range_silence_original_length(select_node): + """Tests if silent audio of original length is returned for an out-of-range index.""" + sr = 44100 + samples = 1234 + channels = 2 + audio_batch = create_dummy_audio(3, channels, samples, sr) + + (result_audio,) = select_node.select_audio(audio_batch, index=3, behavior_out_of_range="silence_original_length", + silence_duration_seconds=1.0) + + assert result_audio['waveform'].shape == (1, channels, samples) + assert result_audio['sample_rate'] == sr + # Check if it's all zeros (silence) + assert torch.all(result_audio['waveform'] == 0) + + +def test_select_out_of_range_silence_fixed_length(select_node): + """Tests if silent audio of a fixed length is returned for an out-of-range index.""" + sr = 48000 + audio_batch = create_dummy_audio(3, 1, 1000, sr) + silence_duration = 2.5 + + (result_audio,) = select_node.select_audio(audio_batch, index=10, behavior_out_of_range="silence_fixed_length", + silence_duration_seconds=silence_duration) + + expected_samples = int(sr * silence_duration) + + assert result_audio['waveform'].shape == (1, 1, expected_samples) + assert result_audio['sample_rate'] == sr + assert torch.all(result_audio['waveform'] == 0)