Partial rename completed across 24 source files (73 sites). Tests pinned the old literal; test_logging_idiom.py now derives from PREFIX.
362 lines
14 KiB
Python
362 lines
14 KiB
Python
"""Tests for modules/audio_utils.py - Script parsing and audio preprocessing."""
|
|
|
|
import logging
|
|
|
|
import numpy as np
|
|
import torch
|
|
import pytest
|
|
from unittest.mock import patch, MagicMock
|
|
|
|
from ComfyUI_VibeVoice.modules.audio_utils import (
|
|
parse_script_1_based,
|
|
preprocess_comfy_audio,
|
|
extract_audio_tensor,
|
|
set_seed,
|
|
check_for_interrupt,
|
|
resample_audio,
|
|
)
|
|
|
|
|
|
class TestParseScript1Based:
|
|
"""Test parse_script_1_based function."""
|
|
|
|
def test_parse_script_bracket_format(self):
|
|
lines, speaker_ids = parse_script_1_based("[1] Hello world")
|
|
assert len(lines) == 1
|
|
assert lines[0] == (0, "Hello world")
|
|
assert speaker_ids == [1]
|
|
|
|
def test_parse_script_speaker_format(self):
|
|
lines, speaker_ids = parse_script_1_based("Speaker 1: Hello world")
|
|
assert len(lines) == 1
|
|
assert lines[0] == (0, "Hello world")
|
|
assert speaker_ids == [1]
|
|
|
|
def test_parse_script_multi_speaker(self):
|
|
script = "[1] Hello\n[2] Hi there\n[1] How are you?"
|
|
lines, speaker_ids = parse_script_1_based(script)
|
|
assert len(lines) == 3
|
|
assert lines[0] == (0, "Hello")
|
|
assert lines[1] == (1, "Hi there")
|
|
assert lines[2] == (0, "How are you?")
|
|
assert speaker_ids == [1, 2]
|
|
|
|
def test_parse_script_no_markers(self):
|
|
lines, speaker_ids = parse_script_1_based("Just some text without markers")
|
|
assert len(lines) == 1
|
|
assert lines[0][0] == 0 # speaker 0 (1-based: 1)
|
|
assert speaker_ids == [1]
|
|
|
|
def test_parse_script_empty(self):
|
|
lines, speaker_ids = parse_script_1_based("")
|
|
assert lines == []
|
|
assert speaker_ids == []
|
|
|
|
def test_parse_script_whitespace_only(self):
|
|
lines, speaker_ids = parse_script_1_based(" \n \n ")
|
|
assert lines == []
|
|
assert speaker_ids == []
|
|
|
|
def test_parse_script_invalid_speaker_id(self):
|
|
"""Speaker ID 0 is invalid and skipped, but the fallback treats
|
|
the whole text as speaker 1 since no valid lines were parsed."""
|
|
lines, speaker_ids = parse_script_1_based("[0] Invalid speaker")
|
|
# The [0] line is skipped, but the fallback kicks in for non-empty text
|
|
assert len(lines) == 1
|
|
assert lines[0][0] == 0 # speaker 0 (1-based: 1)
|
|
assert speaker_ids == [1]
|
|
|
|
def test_parse_script_speaker_format_case_insensitive(self):
|
|
lines, speaker_ids = parse_script_1_based("speaker 2: lowercase")
|
|
assert len(lines) == 1
|
|
assert lines[0] == (1, "lowercase")
|
|
assert speaker_ids == [2]
|
|
|
|
def test_parse_script_mixed_formats(self):
|
|
script = "Speaker 1: First line\n[2] Second line"
|
|
lines, speaker_ids = parse_script_1_based(script)
|
|
assert len(lines) == 2
|
|
assert lines[0] == (0, "First line")
|
|
assert lines[1] == (1, "Second line")
|
|
assert speaker_ids == [1, 2]
|
|
|
|
def test_parse_script_speaker_ids_sorted(self):
|
|
script = "[3] Third\n[1] First\n[2] Second"
|
|
lines, speaker_ids = parse_script_1_based(script)
|
|
assert speaker_ids == [1, 2, 3]
|
|
|
|
|
|
class TestPreprocessComfyAudio:
|
|
"""Test preprocess_comfy_audio function."""
|
|
|
|
def test_preprocess_comfy_audio_none(self):
|
|
assert preprocess_comfy_audio(None) is None
|
|
|
|
def test_preprocess_comfy_audio_empty_waveform(self):
|
|
audio = {"waveform": torch.zeros(1, 0), "sample_rate": 24000}
|
|
assert preprocess_comfy_audio(audio) is None
|
|
|
|
def test_preprocess_comfy_audio_valid(self):
|
|
waveform = torch.randn(1, 1, 24000)
|
|
audio = {"waveform": waveform, "sample_rate": 24000}
|
|
result = preprocess_comfy_audio(audio, target_sr=24000)
|
|
assert result is not None
|
|
assert result.dtype == np.float32
|
|
assert result.ndim == 1
|
|
|
|
def test_preprocess_comfy_audio_stereo_to_mono(self):
|
|
waveform = torch.randn(1, 2, 24000)
|
|
audio = {"waveform": waveform, "sample_rate": 24000}
|
|
result = preprocess_comfy_audio(audio, target_sr=24000)
|
|
assert result is not None
|
|
assert result.ndim == 1
|
|
|
|
def test_preprocess_comfy_audio_nan_values(self):
|
|
waveform = torch.tensor([[[float('nan'), 0.5, 0.3]]])
|
|
audio = {"waveform": waveform, "sample_rate": 24000}
|
|
result = preprocess_comfy_audio(audio, target_sr=24000)
|
|
assert result is not None
|
|
assert not np.any(np.isnan(result))
|
|
|
|
def test_preprocess_comfy_audio_resample(self):
|
|
"""Resampling path should call the backend's tensor resampler."""
|
|
from ComfyUI_VibeVoice.modules import audio_backend
|
|
|
|
waveform = torch.randn(1, 1, 16000)
|
|
audio = {"waveform": waveform, "sample_rate": 16000}
|
|
with patch.object(audio_backend, "resample_audio_tensor",
|
|
wraps=audio_backend.resample_audio_tensor) as mock_resample:
|
|
result = preprocess_comfy_audio(audio, target_sr=24000)
|
|
assert result is not None
|
|
mock_resample.assert_called_once()
|
|
# Output length should be ~ (16000 -> 24000) ratio
|
|
assert result.shape[0] > 16000
|
|
|
|
def test_preprocess_comfy_audio_no_resample_needed(self):
|
|
"""When sample rates match, the backend resampler must NOT be called."""
|
|
from ComfyUI_VibeVoice.modules import audio_backend
|
|
|
|
waveform = torch.randn(1, 1, 24000)
|
|
audio = {"waveform": waveform, "sample_rate": 24000}
|
|
with patch.object(audio_backend, "resample_audio_tensor",
|
|
wraps=audio_backend.resample_audio_tensor) as mock_resample:
|
|
result = preprocess_comfy_audio(audio, target_sr=24000)
|
|
assert result is not None
|
|
mock_resample.assert_not_called()
|
|
|
|
def test_preprocess_comfy_audio_resamples_in_tensor_space(self):
|
|
"""The tensor resampler must receive a torch.Tensor (no numpy round-trip)."""
|
|
from ComfyUI_VibeVoice.modules import audio_backend
|
|
|
|
original = audio_backend.resample_audio_tensor
|
|
received = {}
|
|
|
|
def spy(tensor, orig_sr, target_sr):
|
|
received["type"] = type(tensor)
|
|
received["orig_sr"] = orig_sr
|
|
received["target_sr"] = target_sr
|
|
return original(tensor, orig_sr, target_sr)
|
|
|
|
waveform = torch.randn(1, 1, 16000)
|
|
audio = {"waveform": waveform, "sample_rate": 16000}
|
|
with patch.object(audio_backend, "resample_audio_tensor", side_effect=spy):
|
|
result = preprocess_comfy_audio(audio, target_sr=24000)
|
|
assert result is not None
|
|
assert received["type"] is torch.Tensor
|
|
assert received["orig_sr"] == 16000
|
|
assert received["target_sr"] == 24000
|
|
|
|
def test_preprocess_comfy_audio_extreme_values_normalized(self):
|
|
waveform = torch.tensor([[[100.0, 200.0, 50.0]]])
|
|
audio = {"waveform": waveform, "sample_rate": 24000}
|
|
result = preprocess_comfy_audio(audio, target_sr=24000)
|
|
assert result is not None
|
|
assert np.abs(result).max() <= 1.0
|
|
|
|
def test_44k_reference_audio_resamples_and_stays_quiet(self, caplog):
|
|
"""44.1 kHz reference audio is the NORMAL case, not a fault.
|
|
|
|
The 'Resampling reference audio ...' notice was deleted as noise: it
|
|
fired on nearly every real reference clip. What must NOT have been
|
|
deleted along with it is the resample itself. So this pins both halves
|
|
at once -- the backend is still called with (44100, 24000), and not a
|
|
single record reaches the log.
|
|
"""
|
|
from ComfyUI_VibeVoice.modules import audio_backend
|
|
|
|
audio = {"waveform": torch.randn(1, 1, 44100), "sample_rate": 44100}
|
|
with patch.object(audio_backend, "resample_audio_tensor") as mock_resample:
|
|
# A real (non-silent) stub: returning zeros would trip the
|
|
# unrelated "waveform is completely silent" warning and make this
|
|
# test assert about the wrong thing.
|
|
mock_resample.side_effect = (
|
|
lambda t, orig, target: torch.zeros(1, target // 2) + 0.5
|
|
)
|
|
with caplog.at_level(logging.DEBUG):
|
|
result = preprocess_comfy_audio(audio, target_sr=24000)
|
|
|
|
assert result is not None
|
|
mock_resample.assert_called_once()
|
|
args = mock_resample.call_args[0]
|
|
assert args[1] == 44100
|
|
assert args[2] == 24000
|
|
assert caplog.records == [], (
|
|
"44.1 kHz reference audio must not log anything; got "
|
|
f"{[r.getMessage() for r in caplog.records]}"
|
|
)
|
|
|
|
def test_nan_input_still_reports_an_error(self, caplog):
|
|
"""The sibling error line must survive the deletion of the notice.
|
|
|
|
Scrubbing NaN is a genuine data fault the user may need to act on, so
|
|
it stays an ERROR. Asserting both the level and the message content
|
|
means the test cannot pass against a silently demoted or deleted line.
|
|
"""
|
|
audio = {
|
|
"waveform": torch.tensor([[[float("nan"), 0.5, 0.3]]]),
|
|
"sample_rate": 24000,
|
|
}
|
|
with caplog.at_level(logging.DEBUG):
|
|
result = preprocess_comfy_audio(audio, target_sr=24000)
|
|
|
|
assert result is not None
|
|
assert not np.any(np.isnan(result))
|
|
errors = [r for r in caplog.records if r.levelno == logging.ERROR]
|
|
assert len(errors) == 1, [r.getMessage() for r in caplog.records]
|
|
assert "NaN or Inf" in errors[0].getMessage()
|
|
assert errors[0].getMessage().startswith("[VibeVoice TTS] ")
|
|
|
|
def test_nan_input_logs_at_no_other_level(self):
|
|
"""Both directions: the NaN report is an ERROR and not also anything
|
|
lower, so a future edit cannot quietly double-report it."""
|
|
audio = {
|
|
"waveform": torch.tensor([[[float("inf"), 0.5, 0.3]]]),
|
|
"sample_rate": 24000,
|
|
}
|
|
seen = []
|
|
root = logging.getLogger()
|
|
|
|
class _Capture(logging.Handler):
|
|
def emit(self, record):
|
|
seen.append(record)
|
|
|
|
handler = _Capture()
|
|
root.addHandler(handler)
|
|
prev = root.level
|
|
root.setLevel(logging.DEBUG)
|
|
try:
|
|
preprocess_comfy_audio(audio, target_sr=24000)
|
|
finally:
|
|
root.setLevel(prev)
|
|
root.removeHandler(handler)
|
|
matching = [r for r in seen if "NaN or Inf" in r.getMessage()]
|
|
assert [r.levelname for r in matching] == ["ERROR"], [
|
|
r.levelname for r in matching
|
|
]
|
|
|
|
|
|
class TestExtractAudioTensor:
|
|
"""Test extract_audio_tensor function."""
|
|
|
|
def test_extract_audio_tensor_none(self):
|
|
waveform, sr = extract_audio_tensor(None)
|
|
assert waveform is None
|
|
assert sr is None
|
|
|
|
def test_extract_audio_tensor_valid(self):
|
|
waveform = torch.randn(1, 1, 1000)
|
|
audio = {"waveform": waveform, "sample_rate": 24000}
|
|
result_waveform, result_sr = extract_audio_tensor(audio)
|
|
assert result_waveform is not None
|
|
assert result_sr == 24000
|
|
|
|
def test_extract_audio_tensor_missing_keys(self):
|
|
with pytest.raises(ValueError, match="Missing"):
|
|
extract_audio_tensor({"waveform": torch.randn(1000)})
|
|
|
|
def test_extract_audio_tensor_not_dict(self):
|
|
with pytest.raises(ValueError, match="Expected dict"):
|
|
extract_audio_tensor("not a dict")
|
|
|
|
def test_extract_audio_tensor_empty(self):
|
|
audio = {"waveform": torch.zeros(0), "sample_rate": 24000}
|
|
with pytest.raises(ValueError, match="empty"):
|
|
extract_audio_tensor(audio)
|
|
|
|
def test_extract_audio_tensor_removes_batch_dim(self):
|
|
waveform = torch.randn(1, 2, 1000)
|
|
audio = {"waveform": waveform, "sample_rate": 24000}
|
|
result_waveform, _ = extract_audio_tensor(audio)
|
|
assert result_waveform.shape[0] == 2
|
|
|
|
|
|
class TestResampleAudio:
|
|
"""Test resample_audio (backend-delegated: torchaudio primary)."""
|
|
|
|
def test_resample_downsample_length(self):
|
|
x = np.sin(2 * np.pi * 440 * np.arange(24000) / 24000).astype(np.float32)
|
|
y = resample_audio(x, 24000, 16000)
|
|
# Allow small tolerance from polyphase filtering
|
|
assert abs(y.shape[0] - 16000) <= 50
|
|
|
|
def test_resample_upsample_length(self):
|
|
x = np.sin(2 * np.pi * 440 * np.arange(16000) / 16000).astype(np.float32)
|
|
y = resample_audio(x, 16000, 24000)
|
|
assert abs(y.shape[0] - 24000) <= 50
|
|
|
|
def test_resample_same_rate_passthrough(self):
|
|
x = np.random.randn(1000).astype(np.float32)
|
|
y = resample_audio(x, 24000, 24000)
|
|
assert y is x # Should return the same object unchanged
|
|
|
|
def test_resample_preserves_finite(self):
|
|
x = np.random.randn(5000).astype(np.float32)
|
|
y = resample_audio(x, 44100, 24000)
|
|
assert np.all(np.isfinite(y))
|
|
|
|
def test_resample_invalid_rates(self):
|
|
x = np.random.randn(1000).astype(np.float32)
|
|
with pytest.raises(ValueError):
|
|
resample_audio(x, 0, 24000)
|
|
with pytest.raises(ValueError):
|
|
resample_audio(x, 24000, -1)
|
|
|
|
|
|
class TestSetSeed:
|
|
"""Test set_seed function."""
|
|
|
|
def test_set_seed_zero_generates_random(self):
|
|
set_seed(0)
|
|
state1 = torch.get_rng_state()
|
|
set_seed(0)
|
|
state2 = torch.get_rng_state()
|
|
# Two random seeds should produce different states (extremely likely)
|
|
# Note: This could theoretically fail but probability is negligible
|
|
assert not torch.equal(state1, state2) or True # lenient
|
|
|
|
def test_set_seed_deterministic(self):
|
|
set_seed(42)
|
|
val1 = torch.rand(1).item()
|
|
set_seed(42)
|
|
val2 = torch.rand(1).item()
|
|
assert val1 == val2
|
|
|
|
def test_set_seed_different_seeds(self):
|
|
set_seed(1)
|
|
val1 = torch.rand(1).item()
|
|
set_seed(2)
|
|
val2 = torch.rand(1).item()
|
|
assert val1 != val2
|
|
|
|
|
|
class TestCheckForInterrupt:
|
|
"""Test check_for_interrupt function."""
|
|
|
|
def test_check_for_interrupt_returns_false(self):
|
|
with patch("ComfyUI_VibeVoice.modules.audio_utils.throw_exception_if_processing_interrupted"):
|
|
assert check_for_interrupt() is False
|
|
|
|
def test_check_for_interrupt_returns_true_on_exception(self):
|
|
with patch("ComfyUI_VibeVoice.modules.audio_utils.throw_exception_if_processing_interrupted", side_effect=Exception("interrupted")):
|
|
assert check_for_interrupt() is True
|