[Tests][Added] Reegression tests
This commit is contained in:
@@ -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."
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user