[Tests][Added] Reegression tests

This commit is contained in:
Salvador E. Tropea
2025-07-15 13:02:31 -03:00
parent 53d7edc1b4
commit 6b41a4e86f
6 changed files with 745 additions and 0 deletions
+70
View File
@@ -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."
+55
View File
@@ -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)
+142
View File
@@ -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)
+173
View File
@@ -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)
+230
View File
@@ -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
@@ -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)