Files
wildminder-ComfyUI-VibeVoice/tests/test_streaming_progress.py
T
WildAi b4fce48cff feat: report standard ComfyUI progress during TTS/ASR inference (v2.1.1)
Vendored generate() (non-streaming + streaming) gains an optional framework-agnostic progress_callback(current, total) fired once per AR loop step. modules/generation.py wraps it with comfy.utils.ProgressBar (interrupt check per step, dynamic total via update_absolute, guaranteed final 100%% in finally). ASR reports per-token progress via an HF BaseStreamer (greedy/sampling only; beam search falls back to 0->100%%). 25 new tests; full regression 560p/5f(pre-existing)/4s, zero new failures.
2026-08-16 02:24:52 +03:00

151 lines
6.4 KiB
Python

"""Tests for ComfyUI progress reporting in ``generate_streaming_audio()``
(Phase 4 of the 2026-08-15 inference-progress-reporting plan).
The streaming wrapper must:
* pass a callable ``progress_callback`` into the vendored streaming
``model.generate``,
* map ``callback(current, total)`` -> ``ProgressBar.update_absolute(current,
total=total)`` (the bar's placeholder total is corrected dynamically),
* check for user interruption inside the callback,
* always send the final 100% event (success and exception paths).
"""
import numpy as np
import torch
import pytest
from unittest.mock import patch, MagicMock
from ComfyUI_VibeVoice.modules.generation import generate_streaming_audio
def _mock_voice_sample(length: int = 24000) -> np.ndarray:
return np.random.randn(length).astype(np.float32)
def _run_streaming(generate_side_effect=None, expect_exception=None):
"""Run generate_streaming_audio with fully mocked model/processor/prefill.
When ``expect_exception`` is given, the call is wrapped in
``pytest.raises``. Returns (mock_model, mock_pbar_cls, mock_pbar,
mock_interrupt).
"""
mock_model = MagicMock()
mock_model.device = torch.device("cpu")
mock_output = MagicMock()
mock_output.speech_outputs = [torch.randn(24000)]
if generate_side_effect is not None:
mock_model.generate.side_effect = generate_side_effect
else:
mock_model.generate.return_value = mock_output
mock_processor = MagicMock()
mock_processor.tokenizer = MagicMock()
mock_processor.tokenizer.encode = MagicMock(return_value=[10, 11, 12])
mock_processor.prepare_speech_inputs = MagicMock(
return_value={
"padded_speeches": torch.randn(1, 100),
"speech_masks": torch.ones(1, 100, dtype=torch.bool),
}
)
with patch("ComfyUI_VibeVoice.modules.generation.ProgressBar") as mock_pbar_cls, \
patch("ComfyUI_VibeVoice.modules.generation.model_management.throw_exception_if_processing_interrupted") as mock_interrupt, \
patch("ComfyUI_VibeVoice.modules.generation.prefill_voice_prompt", return_value={"lm": MagicMock()}), \
patch("ComfyUI_VibeVoice.modules.generation.preprocess_comfy_audio", return_value=_mock_voice_sample()):
mock_pbar = MagicMock()
mock_pbar.total = 1
mock_pbar_cls.return_value = mock_pbar
call = lambda: generate_streaming_audio(
model=mock_model,
processor=mock_processor,
text="[1] Hello world",
voice_samples=[{"waveform": torch.randn(1, 1, 24000), "sample_rate": 24000}],
speaker_ids=[1],
inference_steps=10,
)
if expect_exception is not None:
with pytest.raises(expect_exception):
call()
else:
call()
return mock_model, mock_pbar_cls, mock_pbar, mock_interrupt
class TestStreamingProgressReporting:
def test_generate_receives_callable_progress_callback(self):
"""T4.1: streaming model.generate must receive a callable progress_callback."""
mock_model, _, _, _ = _run_streaming()
gen_kwargs = mock_model.generate.call_args.kwargs
cb = gen_kwargs.get("progress_callback")
assert cb is not None and callable(cb)
def test_callback_maps_to_update_absolute_with_total(self):
"""T4.2: callback(current, total) -> pbar.update_absolute(current, total=total)."""
mock_model, _, mock_pbar, _ = _run_streaming()
cb = mock_model.generate.call_args.kwargs["progress_callback"]
mock_pbar.update_absolute.reset_mock()
cb(42, 512)
mock_pbar.update_absolute.assert_called_once_with(42, total=512)
def test_callback_checks_interrupt(self):
"""T4.3a: the callback must call throw_exception_if_processing_interrupted;
a raised interrupt propagates to the caller."""
import comfy.model_management as mm
mock_model = MagicMock()
mock_model.device = torch.device("cpu")
mock_output = MagicMock()
mock_output.speech_outputs = [torch.randn(24000)]
mock_model.generate.return_value = mock_output
mock_processor = MagicMock()
mock_processor.tokenizer = MagicMock()
mock_processor.tokenizer.encode = MagicMock(return_value=[10, 11, 12])
mock_processor.prepare_speech_inputs = MagicMock(
return_value={
"padded_speeches": torch.randn(1, 100),
"speech_masks": torch.ones(1, 100, dtype=torch.bool),
}
)
with patch("ComfyUI_VibeVoice.modules.generation.ProgressBar") as mock_pbar_cls, \
patch("ComfyUI_VibeVoice.modules.generation.model_management.throw_exception_if_processing_interrupted") as mock_interrupt, \
patch("ComfyUI_VibeVoice.modules.generation.prefill_voice_prompt", return_value={"lm": MagicMock()}), \
patch("ComfyUI_VibeVoice.modules.generation.preprocess_comfy_audio", return_value=_mock_voice_sample()):
mock_pbar = MagicMock()
mock_pbar.total = 1
mock_pbar_cls.return_value = mock_pbar
mock_interrupt.side_effect = mm.InterruptProcessingException()
generate_streaming_audio(
model=mock_model,
processor=mock_processor,
text="[1] Hello world",
voice_samples=[{"waveform": torch.randn(1, 1, 24000), "sample_rate": 24000}],
speaker_ids=[1],
inference_steps=10,
)
cb = mock_model.generate.call_args.kwargs["progress_callback"]
with pytest.raises(mm.InterruptProcessingException):
cb(1, 512)
mock_interrupt.assert_called()
def test_final_update_sent_on_success_and_exception(self):
"""T4.3b: the finally-block final update fires on success and on raise."""
# Success path.
_, mock_pbar_cls, mock_pbar, _ = _run_streaming()
mock_pbar_cls.assert_called_once_with(1) # placeholder total
final_call = mock_pbar.update_absolute.call_args_list[-1]
assert final_call.args == (1,), f"final update must be (total,), got {final_call}"
# Exception path.
mock_model, mock_pbar_cls2, mock_pbar2, _ = _run_streaming(
generate_side_effect=RuntimeError("boom"),
expect_exception=RuntimeError,
)
final_call2 = mock_pbar2.update_absolute.call_args_list[-1]
assert final_call2.args == (1,)