diff --git a/README.md b/README.md index 919860a..6680d2b 100644 --- a/README.md +++ b/README.md @@ -160,6 +160,38 @@ This node features a sophisticated system for managing performance, memory, and ## Changelog
+v2.1.1 - Standard ComfyUI Progress Bar During Inference + +### ✨ Highlights +* **Live progress bar:** All three nodes (TTS, Realtime TTS, ASR) now drive the standard + ComfyUI frontend progress bar during inference. Previously the bar sat at 0% for the whole + generation and jumped to 100% only at the end. +* **Responsive cancel:** the progress hook checks ComfyUI's interrupt flag on every loop step, + so pressing cancel stops generation promptly instead of waiting for the current blocking call. +* **Guaranteed 100%:** a final progress event is always emitted, even when generation stops + early (EOS) or raises. + +### 🔧 Changes +* Vendored `generate()` (non-streaming + streaming) gained an optional, framework-agnostic + `progress_callback(current, total)` hook fired once per AR loop step (vendored code stays + `comfy`-free; `None` = disabled, fully backward compatible). +* `modules/generation.py`: `generate_audio()` / `generate_streaming_audio()` wrap the hook with + `comfy.utils.ProgressBar` (throttled WebSocket updates; the bar total self-corrects via + `update_absolute(value, total=...)` once the loop reports its real budget). +* `modules/asr_generation.py`: ASR reports per-token progress through an HF `BaseStreamer` + (greedy/sampling only; beam search falls back to a single 0→100% bar). + +### 🧪 Tests +* New `tests/test_generate_progress_callback.py` (5 tests) and + `tests/test_streaming_progress_callback.py` (4 tests): drive the real vendored loops with + scripted mocks and lock the callback contract (monotonic, bounded, call counts, interrupt + propagation, output determinism). +* New `tests/test_streaming_progress.py` (4 tests); extended `tests/test_generation.py` (+5), + `tests/test_asr_generation.py` (+6), `tests/test_integration.py` (+3 node-level tests). + +
+ +
v2.1.0 - torchaudio-Primary Audio Backend (librosa now optional) ### ✨ Highlights diff --git a/modules/asr_generation.py b/modules/asr_generation.py index 7e5e664..67422f2 100644 --- a/modules/asr_generation.py +++ b/modules/asr_generation.py @@ -12,6 +12,8 @@ import logging from typing import Optional, Tuple, List, Dict, Any import comfy.model_management as model_management +from comfy.utils import ProgressBar +from transformers.generation import BaseStreamer from .asr_loader import VibeVoiceASRLoader, VibeVoiceASRModelHandler, LOADED_ASR_MODELS_CACHE, cleanup_asr_models from .patcher import VibeVoiceASRPatcher @@ -24,6 +26,40 @@ from .attention_utils import resolve_attention_mode logger = logging.getLogger(__name__) +class _ASRProgressStreamer(BaseStreamer): + """Token streamer that drives the standard ComfyUI progress bar during ASR. + + HF ``GenerationMixin`` calls ``put(next_tokens)`` once per generated token + (prompt tokens are NOT streamed). Each call advances the bar by the number + of tokens in the batch and checks for user interruption, so cancelling an + ASR transcription becomes responsive. ``end()`` sends the final 100% event. + + Only attached for greedy/sampling decoding (``num_beams <= 1``); beam + search is not compatible with a plain token streamer and falls back to a + single 0->100% bar. + """ + + def __init__(self, pbar: "ProgressBar", total: int): + self.pbar = pbar + self.total = max(1, int(total)) + self.count = 0 + + def put(self, value): + # Responsive cancellation (raises InterruptProcessingException). + model_management.throw_exception_if_processing_interrupted() + if isinstance(value, torch.Tensor): + n = int(value.numel()) + elif isinstance(value, (list, tuple)): + n = len(value) + else: + n = 1 + self.count = min(self.count + n, self.total) + self.pbar.update_absolute(self.count, total=self.total) + + def end(self): + self.pbar.update_absolute(self.total) + + def load_asr_model( model_name: str, device: str = "auto", @@ -226,11 +262,19 @@ def transcribe_audio( # Remove None values generation_config = {k: v for k, v in generation_config.items() if v is not None} + # Standard ComfyUI progress bar. HF generate() reports per-token progress + # through a streamer (greedy/sampling only; beam search falls back to a + # single 0->100% bar because plain token streamers are beam-incompatible). + pbar = ProgressBar(max_new_tokens) + use_streamer = generation_config.get("num_beams", 1) <= 1 + streamer = _ASRProgressStreamer(pbar, total=max_new_tokens) if use_streamer else None + try: with torch.no_grad(): output_ids = model.generate( **inputs, **generation_config, + **({"streamer": streamer} if streamer is not None else {}), ) # Decode output (exclude input tokens) @@ -264,6 +308,10 @@ def transcribe_audio( except Exception as e: logger.error(f"ASR transcription failed: {e}") raise RuntimeError(f"Transcription failed: {e}") + finally: + # Guarantee the final 100% event even when generation stopped early + # (EOS before max_new_tokens) or raised. + pbar.update_absolute(pbar.total) def force_offload_asr_model(model_name: str, patcher=None) -> None: diff --git a/modules/generation.py b/modules/generation.py index 20a32ba..5488690 100644 --- a/modules/generation.py +++ b/modules/generation.py @@ -237,17 +237,28 @@ def generate_audio( # Generate with torch.no_grad(): + # Standard ComfyUI progress bar. The initial total is only an estimate + # (diffusion steps); the vendored AR loop reports its real budget + # (max_steps) through `progress_callback`, and `update_absolute(value, + # total=...)` re-sets the bar's total dynamically on the first callback. pbar = ProgressBar(inference_steps) + def _progress(current: int, total: int) -> None: + # Responsive cancellation: raises InterruptProcessingException when + # the user pressed cancel (checked once per AR step). + model_management.throw_exception_if_processing_interrupted() + pbar.update_absolute(current, total=total) + try: - outputs = model.generate(**gen_inputs) - pbar.update(inference_steps - pbar.current) + outputs = model.generate(**gen_inputs, progress_callback=_progress) except model_management.InterruptProcessingException: logger.info("VibeVoice generation interrupted by user") raise finally: - pbar.update_absolute(inference_steps) + # Guarantee the final 100% event even when the AR loop stopped + # early (EOS before max_steps) or generation raised. + pbar.update_absolute(pbar.total) # Post-process output output_waveform = outputs.speech_outputs[0] @@ -419,7 +430,27 @@ def generate_streaming_audio( gen_kwargs["return_speech"] = True with torch.no_grad(): - outputs = model.generate(**gen_kwargs) + # Standard ComfyUI progress bar. The streaming loop's real total + # (tts_lm max_length) is only known inside the vendored generate(); + # the initial total=1 placeholder is corrected on the first callback + # via update_absolute(value, total=...). + pbar = ProgressBar(1) + + def _progress(current: int, total: int) -> None: + # Responsive cancellation: raises InterruptProcessingException when + # the user pressed cancel (checked once per loop step). + model_management.throw_exception_if_processing_interrupted() + pbar.update_absolute(current, total=total) + + try: + outputs = model.generate(**gen_kwargs, progress_callback=_progress) + except model_management.InterruptProcessingException: + logger.info("VibeVoice streaming generation interrupted by user") + raise + finally: + # Guarantee the final 100% event even when the loop stopped early + # (EOS classifier) or generation raised. + pbar.update_absolute(pbar.total) speech_outputs = outputs.speech_outputs if not speech_outputs or speech_outputs[0] is None: diff --git a/pyproject.toml b/pyproject.toml index 554eacd..d084f0f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "ComfyUI-VibeVoice" description = "VibeVoice TTS. Expressive, long-form, multi-speaker conversational audio" -version = "2.1.0" +version = "2.1.1" license = {file = "LICENSE"} requires-python = ">=3.10" dependencies = ["torch", "torchaudio", "numpy", "huggingface_hub", "einops", "tokenizers", "soundfile", "s3tokenizer", "tqdm", "conformer", "safetensors", "transformers>=4.51.3", "diffusers", "bitsandbytes"] diff --git a/src/vibevoice/modular/modeling_vibevoice.py b/src/vibevoice/modular/modeling_vibevoice.py index 9eb53ed..9498661 100644 --- a/src/vibevoice/modular/modeling_vibevoice.py +++ b/src/vibevoice/modular/modeling_vibevoice.py @@ -705,6 +705,7 @@ class VibeVoiceForConditionalGeneration(VibeVoicePreTrainedModel): top_p: float = 0.95, top_k: int = 0, tokenizer: Optional[Any] = None, + progress_callback: Optional[Callable[[int, int], None]] = None, **kwargs, ) -> "VibeVoiceGenerationOutput": """Non-streaming text-to-speech generation (faithful original-protocol port). @@ -752,6 +753,13 @@ class VibeVoiceForConditionalGeneration(VibeVoicePreTrainedModel): are accepted for interface compatibility but unused -- the original protocol samples with a constrained softmax (multinomial) or argmax. tokenizer: processor tokenizer; supplies the speech control-token ids. + progress_callback: Optional framework-agnostic hook invoked as + ``progress_callback(current, total)`` — once with ``(0, max_steps)`` + before the AR loop and once after each completed AR step with + ``(step, max_steps)``. ``current`` is monotonic non-decreasing and + bounded by ``total``. ``None`` disables reporting. The callback may + raise to interrupt generation (the exception propagates out of + ``generate()``); callers use this for progress UI and cancellation. Returns: VibeVoiceGenerationOutput with ``sequences`` (input ids) and @@ -828,6 +836,10 @@ class VibeVoiceForConditionalGeneration(VibeVoicePreTrainedModel): max_length_times = int(kwargs.get("max_length_times", 2)) max_steps = max(1, min(max_new_tokens, int(max_length_times * seq_len))) + # Progress reporting: announce the loop budget before the first step. + if progress_callback is not None: + progress_callback(0, max_steps) + # Diffusion steps. num_steps = inference_steps or getattr( self, "ddpm_inference_steps", self.config.diffusion_head_config.ddpm_num_inference_steps @@ -1042,6 +1054,10 @@ class VibeVoiceForConditionalGeneration(VibeVoicePreTrainedModel): break step += 1 + # Progress reporting: one callback per completed AR step. + if progress_callback is not None: + progress_callback(step, max_steps) + # ------------------------------------------------------------------ # Assemble the per-sample waveform(s). # ------------------------------------------------------------------ diff --git a/src/vibevoice/modular/modeling_vibevoice_streaming_inference.py b/src/vibevoice/modular/modeling_vibevoice_streaming_inference.py index da9d0b2..e6aa7ff 100644 --- a/src/vibevoice/modular/modeling_vibevoice_streaming_inference.py +++ b/src/vibevoice/modular/modeling_vibevoice_streaming_inference.py @@ -593,6 +593,7 @@ class VibeVoiceStreamingForConditionalGenerationInference(VibeVoiceStreamingPreT return_speech: bool = True, cfg_scale: float = 1.0, stop_check_fn: Optional[Callable[[], bool]] = None, + progress_callback: Optional[Callable[[int, int], None]] = None, **kwargs, ) -> Union[torch.LongTensor, VibeVoiceGenerationOutput]: """ @@ -610,6 +611,15 @@ class VibeVoiceStreamingForConditionalGenerationInference(VibeVoiceStreamingPreT cfg_scale: Classifier-free guidance scale for speech diffusion. return_speech: If False, skips audio decode concatenation. stop_check_fn: External early-stop hook (returns True to halt). + progress_callback: Optional framework-agnostic hook invoked as + ``progress_callback(current, total)`` after each text-window + prefill and after each generated speech token, where + ``total == tts_lm_generation_config.max_length``. ``current`` + is monotonic non-decreasing and bounded by ``total``. ``None`` + disables reporting. The callback may raise to interrupt + generation (the exception propagates out of ``generate()``); + callers use this for progress UI and cancellation. Orthogonal + to the console ``tqdm`` bar (``show_progress_bar``). Returns: VibeVoiceGenerationOutput with: @@ -748,6 +758,9 @@ class VibeVoiceStreamingForConditionalGenerationInference(VibeVoiceStreamingPreT if progress_bar is not None: progress_bar.update(cur_input_tts_text_ids.shape[1]) progress_bar.set_description(f"Prefilled {total_prefilled_text_tokens} text tokens, generated {total_generated_speech_tokens} speech tokens, current step ({step} / {tts_lm_generation_config.max_length})") + # Framework-agnostic progress hook (ComfyUI progress bar). + if progress_callback is not None: + progress_callback(step, tts_lm_generation_config.max_length) model_inputs = self.prepare_inputs_for_generation(input_ids, **model_kwargs) # Forward pass through the model @@ -815,6 +828,9 @@ class VibeVoiceStreamingForConditionalGenerationInference(VibeVoiceStreamingPreT if progress_bar is not None: progress_bar.update(1) progress_bar.set_description(f"Prefilled {total_prefilled_text_tokens} text tokens, generated {total_generated_speech_tokens} speech tokens, current step ({step} / {tts_lm_generation_config.max_length})") + # Framework-agnostic progress hook (ComfyUI progress bar). + if progress_callback is not None: + progress_callback(step, tts_lm_generation_config.max_length) tts_lm_model_inputs = self.prepare_inputs_for_generation(tts_lm_input_ids, **tts_lm_model_kwargs) tts_lm_additional_inputs = { diff --git a/tests/test_asr_generation.py b/tests/test_asr_generation.py index 9ada8fe..be80432 100644 --- a/tests/test_asr_generation.py +++ b/tests/test_asr_generation.py @@ -300,3 +300,108 @@ class TestForceOffloadASRPatcher: LOADED_ASR_MODELS_CACHE.clear() VIBEVOICE_ASR_PATCHER_CACHE.clear() + + +class TestASRProgressReporting: + """Phase 5 (2026-08-15 progress plan): ASR transcription must drive the + standard ComfyUI ProgressBar through an HF token streamer.""" + + @staticmethod + def _make_mocks(): + mock_model = MagicMock() + mock_param = MagicMock() + mock_param.device = torch.device("cpu") + mock_model.parameters.return_value = iter([mock_param]) + mock_output = torch.tensor([[1, 2, 3, 4, 5, 0]]) + mock_model.generate.return_value = mock_output + + mock_processor = MagicMock() + mock_processor.return_value = {"input_ids": torch.tensor([[1, 2, 3]])} + mock_processor.pad_id = 0 + mock_processor.tokenizer.eos_token_id = 0 + mock_processor.decode.return_value = "Hello world" + mock_processor.post_process_transcription.return_value = [] + return mock_model, mock_processor + + def _transcribe(self, mock_model, mock_processor, **kwargs): + with patch("ComfyUI_VibeVoice.modules.asr_generation.extract_audio_tensor") as mock_extract, \ + patch("ComfyUI_VibeVoice.modules.asr_generation.ProgressBar") as mock_pbar_cls, \ + patch("ComfyUI_VibeVoice.modules.asr_generation.model_management.throw_exception_if_processing_interrupted") as mock_interrupt: + mock_extract.return_value = (torch.randn(24000), 24000) + mock_pbar = MagicMock() + mock_pbar.total = kwargs.get("max_new_tokens", 32768) + mock_pbar_cls.return_value = mock_pbar + + result = transcribe_audio( + model=mock_model, + processor=mock_processor, + audio_input={"waveform": torch.randn(1, 1, 24000), "sample_rate": 24000}, + **kwargs, + ) + return result, mock_model, mock_pbar_cls, mock_pbar, mock_interrupt + + def test_streamer_put_advances_progress(self): + """T5.1: streamer.put() of a 2-token tensor advances the bar to 2.""" + from ComfyUI_VibeVoice.modules.asr_generation import _ASRProgressStreamer + + mock_pbar = MagicMock() + streamer = _ASRProgressStreamer(mock_pbar, total=100) + streamer.put(torch.tensor([[7, 8]])) + + assert streamer.count == 2 + mock_pbar.update_absolute.assert_called_once_with(2, total=100) + + def test_streamer_end_sends_final_total(self): + """T5.2: streamer.end() sends the final total update.""" + from ComfyUI_VibeVoice.modules.asr_generation import _ASRProgressStreamer + + mock_pbar = MagicMock() + streamer = _ASRProgressStreamer(mock_pbar, total=50) + streamer.end() + mock_pbar.update_absolute.assert_called_once_with(50) + + def test_generate_receives_streamer_for_sampling(self): + """T5.3a: model.generate receives streamer= when num_beams == 1.""" + mock_model, mock_processor = self._make_mocks() + _, mock_model, _, _, _ = self._transcribe(mock_model, mock_processor, num_beams=1) + + gen_kwargs = mock_model.generate.call_args.kwargs + assert "streamer" in gen_kwargs + assert gen_kwargs["streamer"] is not None + + def test_generate_no_streamer_for_beam_search(self): + """T5.3b: model.generate must NOT receive streamer= when num_beams > 1.""" + mock_model, mock_processor = self._make_mocks() + _, mock_model, _, _, _ = self._transcribe(mock_model, mock_processor, num_beams=4) + + gen_kwargs = mock_model.generate.call_args.kwargs + assert "streamer" not in gen_kwargs + + def test_streamer_put_checks_interrupt(self): + """T5.4: put() calls throw_exception_if_processing_interrupted; a raised + interrupt propagates.""" + import comfy.model_management as mm + from ComfyUI_VibeVoice.modules.asr_generation import _ASRProgressStreamer + + with patch("ComfyUI_VibeVoice.modules.asr_generation.model_management.throw_exception_if_processing_interrupted") as mock_interrupt: + mock_interrupt.side_effect = mm.InterruptProcessingException() + mock_pbar = MagicMock() + streamer = _ASRProgressStreamer(mock_pbar, total=100) + + with pytest.raises(mm.InterruptProcessingException): + streamer.put(torch.tensor([7])) + mock_interrupt.assert_called() + + def test_transcription_output_unchanged_with_progress(self): + """T5.5: progress plumbing must not change the transcription output.""" + mock_model, mock_processor = self._make_mocks() + (raw_text, segments), _, mock_pbar_cls, mock_pbar, _ = self._transcribe( + mock_model, mock_processor, max_new_tokens=100 + ) + + assert raw_text == "Hello world" + assert segments == [] + # Bar created with the max_new_tokens budget and driven to 100% at the end. + mock_pbar_cls.assert_called_once_with(100) + final_call = mock_pbar.update_absolute.call_args_list[-1] + assert final_call.args == (100,) diff --git a/tests/test_generate_progress_callback.py b/tests/test_generate_progress_callback.py new file mode 100644 index 0000000..0acd8e0 --- /dev/null +++ b/tests/test_generate_progress_callback.py @@ -0,0 +1,319 @@ +"""Tests for the ``progress_callback`` hook in the non-streaming +``VibeVoiceForConditionalGeneration.generate()`` (Phase 1 of the +2026-08-15 inference-progress-reporting plan). + +The real vendored module is loaded with the same stub-diffusers / +real-module pattern used by ``tests/test_model_forward.py``. A fully +scripted mock inner model drives the AR loop deterministically: + +* ``lm_head`` emits a fixed control-token script + ``[diffusion, diffusion, speech_end, eos]`` (greedy decoding), +* the inner ``language_model`` returns fixed hidden states, +* diffusion / tokenizer submodules are mocked with fixed-shape outputs. + +Expected loop trace (seq_len=4, max_length_times=2 -> max_steps=8): + + callback(0, 8) # initial budget announcement + step 1: diffusion -> callback(1, 8) + step 2: diffusion -> callback(2, 8) + step 3: speech_end -> callback(3, 8) + step 4: eos -> finished -> break (no step increment) +""" + +import os +import sys +import types +import importlib.util +from unittest.mock import MagicMock + +import torch +import torch.nn as nn + + +# ---------------------------------------------------------------------- +# Real-module loading (mirrors tests/test_model_forward.py) +# ---------------------------------------------------------------------- +def _stub_diffusers(): + """Stub the diffusers package and submodules the vendored code imports.""" + if "diffusers" in sys.modules and not isinstance(sys.modules["diffusers"], MagicMock): + return # already a real diffusers + + d = types.ModuleType("diffusers") + cu = types.ModuleType("diffusers.configuration_utils") + cu.ConfigMixin = object + cu.register_to_config = lambda *a, **k: None + du = types.ModuleType("diffusers.utils") + du.deprecate = lambda *a, **k: None + tu = types.ModuleType("diffusers.utils.torch_utils") + tu.randn_tensor = lambda *a, **k: None + su = types.ModuleType("diffusers.schedulers.scheduling_utils") + su.KarrasDiffusionSchedulers = object + su.SchedulerMixin = object + su.SchedulerOutput = object + d.configuration_utils = cu + d.utils = du + d.schedulers = su + sys.modules.setdefault("diffusers", d) + sys.modules.setdefault("diffusers.configuration_utils", cu) + sys.modules.setdefault("diffusers.utils", du) + sys.modules.setdefault("diffusers.utils.torch_utils", tu) + sys.modules.setdefault("diffusers.schedulers", su) + sys.modules.setdefault("diffusers.schedulers.scheduling_utils", su) + + +def _load_real_modeling_module(): + """Load the real modeling_vibevoice module, bypassing the conftest mock.""" + _stub_diffusers() + for mod in ( + "src.vibevoice", + "src.vibevoice.modular", + "src.vibevoice.modular.modeling_vibevoice", + "ComfyUI_VibeVoice.src.vibevoice.modular.modeling_vibevoice", + ): + sys.modules.pop(mod, None) + + root = os.path.join(os.getcwd(), "src", "vibevoice") + mod_pkg = types.ModuleType("src.vibevoice.modular") + mod_pkg.__path__ = [os.path.join(root, "modular")] + sys.modules.setdefault("src.vibevoice", types.ModuleType("src.vibevoice")) + sys.modules["src.vibevoice"].__path__ = [root] + sys.modules["src.vibevoice.modular"] = mod_pkg + + cfg_mod = types.ModuleType("src.vibevoice.modular.configuration_vibevoice") + + class _FakeConfig: + __name__ = "VibeVoiceConfig" + model_type = "vibevoice" + + cfg_mod.VibeVoiceConfig = _FakeConfig + sys.modules["src.vibevoice.modular.configuration_vibevoice"] = cfg_mod + + spec = importlib.util.spec_from_file_location( + "src.vibevoice.modular.modeling_vibevoice", + "src/vibevoice/modular/modeling_vibevoice.py", + ) + module = importlib.util.module_from_spec(spec) + sys.modules["src.vibevoice.modular.modeling_vibevoice"] = module + spec.loader.exec_module(module) + return module + + +_modeling = _load_real_modeling_module() +VibeVoiceForConditionalGeneration = _modeling.VibeVoiceForConditionalGeneration +VibeVoiceGenerationOutput = _modeling.VibeVoiceGenerationOutput + + +# ---------------------------------------------------------------------- +# Scripted mock model +# ---------------------------------------------------------------------- +VOCAB = 100 +HIDDEN = 64 +SEQ_LEN = 4 + +# Control-token ids used by the scripted lm_head. +START_ID = 1 +END_ID = 2 +DIFFUSION_ID = 3 +EOS_ID = 4 +BOS_ID = 5 + +# The control-token script emitted by the scripted lm_head, one per AR step. +TOKEN_SCRIPT = [DIFFUSION_ID, DIFFUSION_ID, END_ID, EOS_ID] +EXPECTED_COMPLETED_STEPS = 3 # eos step breaks before `step += 1` +EXPECTED_MAX_STEPS = 8 # min(max_new_tokens, max_length_times * seq_len) = min(1020, 2*4) + + +class _ScriptedLMHead(nn.Module): + """``lm_head`` stand-in: real ``.weight`` (for ``.weight.device`` reads) + but scripted logits so greedy decoding follows ``TOKEN_SCRIPT``.""" + + def __init__(self, script): + super().__init__() + self.weight = nn.Parameter(torch.zeros(HIDDEN, VOCAB)) + self._script = list(script) + self._call = 0 + + def forward(self, x): + idx = self._script[min(self._call, len(self._script) - 1)] + self._call += 1 + logits = torch.full((x.shape[0], VOCAB), -10.0) + logits[:, idx] = 10.0 + return logits + + +class _FakeTokenizer: + speech_start_id = START_ID + speech_end_id = END_ID + speech_diffusion_id = DIFFUSION_ID + eos_id = EOS_ID + eos_token_id = EOS_ID + bos_token_id = BOS_ID + + +class _StepResult: + def __init__(self, prev_sample): + self.prev_sample = prev_sample + + +def _make_scripted_model(): + """Build a VibeVoiceForConditionalGeneration whose AR loop is fully driven + by mocks and follows ``TOKEN_SCRIPT`` deterministically.""" + model = VibeVoiceForConditionalGeneration.__new__(VibeVoiceForConditionalGeneration) + torch.nn.Module.__init__(model) + + model.config = MagicMock() + model.config.use_return_dict = True + model.config.diffusion_head_config = MagicMock() + model.config.diffusion_head_config.ddpm_num_inference_steps = 5 + model.config.acoustic_tokenizer_config = MagicMock() + model.config.acoustic_tokenizer_config.vae_dim = 64 + model.config.acoustic_vae_dim = 64 + model.config.decoder_config = MagicMock() + model.config.decoder_config.max_position_embeddings = 1024 + model.vocab_size = VOCAB + + inner = MagicMock() + inner.speech_scaling_factor = torch.tensor(float("nan")) + inner.speech_bias_factor = torch.tensor(float("nan")) + + # Diffusion scheduler: one timestep, fixed prev_sample. + inner.noise_scheduler = MagicMock() + inner.noise_scheduler.timesteps = [torch.tensor(999)] + inner.noise_scheduler.set_timesteps = MagicMock() + inner.noise_scheduler.step = MagicMock( + return_value=_StepResult(torch.zeros(2, 64)) + ) + + # Prediction head: fixed zero noise (CFG combination stays zero). + def _pred_head(noisy, timesteps, condition): + return torch.zeros(condition.shape[0], 64) + + inner.prediction_head = MagicMock(side_effect=_pred_head) + inner.prediction_head.device = torch.device("cpu") + + # Inner language model: fixed hidden states of the right length. + def _lm_forward(inputs_embeds=None, **kwargs): + x = inputs_embeds + if not isinstance(x, torch.Tensor): + x = torch.zeros(1, SEQ_LEN, HIDDEN) + result = MagicMock() + result.last_hidden_state = torch.zeros(x.shape[0], x.shape[1], HIDDEN) + result.past_key_values = None + return result + + inner.language_model = MagicMock(side_effect=_lm_forward) + + # Acoustic tokenizer: deterministic waveform chunk (Nd, 1, T). + inner.acoustic_tokenizer = MagicMock() + inner.acoustic_tokenizer.device = torch.device("cpu") + inner.acoustic_tokenizer.decode = MagicMock( + return_value=torch.full((1, 1, 1600), 0.25) + ) + + # Semantic tokenizer: encode -> object with .mean (Nd, sem_dim, T_sem). + def _sem_encode(audio, **kwargs): + nd = audio.shape[0] + out = MagicMock() + out.mean = torch.zeros(nd, 16, 4) + return out + + inner.semantic_tokenizer = MagicMock() + inner.semantic_tokenizer.encode = MagicMock(side_effect=_sem_encode) + + # Connectors: (Nd, H) outputs. + def _acoustic_connector(f): + if f.dim() == 3: + return torch.zeros(f.shape[0], f.shape[1], HIDDEN) + return torch.zeros(f.shape[0], HIDDEN) + + def _semantic_connector(f): + return torch.zeros(f.shape[0], HIDDEN) + + inner.acoustic_connector = MagicMock(side_effect=_acoustic_connector) + inner.semantic_connector = MagicMock(side_effect=_semantic_connector) + model.model = inner + + model.get_input_embeddings = MagicMock(return_value=nn.Embedding(VOCAB, HIDDEN)) + model.add_module("lm_head", _ScriptedLMHead(TOKEN_SCRIPT)) + return model + + +def _generate(model, progress_callback=None, **overrides): + """Run generate() with the standard scripted inputs.""" + kwargs = dict( + input_ids=torch.randint(0, VOCAB, (1, SEQ_LEN)), + attention_mask=torch.ones(1, SEQ_LEN, dtype=torch.long), + acoustic_input_mask=torch.zeros(1, SEQ_LEN, dtype=torch.bool), + cfg_scale=1.3, + inference_steps=5, + return_speech=True, + do_sample=False, + tokenizer=_FakeTokenizer(), + ) + kwargs.update(overrides) + if progress_callback is not None: + kwargs["progress_callback"] = progress_callback + return model.generate(**kwargs) + + +# ---------------------------------------------------------------------- +# Tests +# ---------------------------------------------------------------------- +class TestProgressCallbackContract: + def test_callback_none_backward_compatible(self): + """T1.1: generate() without progress_callback runs unchanged.""" + model = _make_scripted_model() + out = _generate(model) + assert isinstance(out, VibeVoiceGenerationOutput) + assert out.speech_outputs is not None + assert len(out.speech_outputs) == 1 + assert isinstance(out.speech_outputs[0], torch.Tensor) + + def test_callback_sequence_monotonic_and_bounded(self): + """T1.2: first call is (0, max_steps); currents monotonic, <= total.""" + model = _make_scripted_model() + calls = [] + _generate(model, progress_callback=lambda c, t: calls.append((c, t))) + + assert calls, "progress_callback was never invoked" + assert calls[0] == (0, EXPECTED_MAX_STEPS) + currents = [c for c, _ in calls] + totals = [t for _, t in calls] + assert currents == sorted(currents), "current must be monotonic non-decreasing" + assert all(0 <= c <= EXPECTED_MAX_STEPS for c in currents) + assert set(totals) == {EXPECTED_MAX_STEPS}, "total must stay constant" + + def test_callback_count_matches_completed_steps(self): + """T1.3: exactly one initial call + one call per completed AR step.""" + model = _make_scripted_model() + calls = [] + _generate(model, progress_callback=lambda c, t: calls.append((c, t))) + + # 1 initial + EXPECTED_COMPLETED_STEPS per-step calls. + assert len(calls) == 1 + EXPECTED_COMPLETED_STEPS + assert [c for c, _ in calls] == [0, 1, 2, 3] + + def test_callback_exception_propagates(self): + """T1.4: a raising callback interrupts generation (exception propagates).""" + model = _make_scripted_model() + + def _raiser(current, total): + if current >= 2: + raise RuntimeError("interrupt requested") + + try: + _generate(model, progress_callback=_raiser) + raised = False + except RuntimeError as e: + raised = "interrupt requested" in str(e) + assert raised, "callback exception must propagate out of generate()" + + def test_output_identical_with_and_without_callback(self): + """T1.5: progress plumbing must not change generated audio.""" + model_a = _make_scripted_model() + out_a = _generate(model_a) + + model_b = _make_scripted_model() + out_b = _generate(model_b, progress_callback=lambda c, t: None) + + assert torch.equal(out_a.speech_outputs[0], out_b.speech_outputs[0]) diff --git a/tests/test_generation.py b/tests/test_generation.py index 69fd950..f5ab4fb 100644 --- a/tests/test_generation.py +++ b/tests/test_generation.py @@ -470,3 +470,147 @@ class TestForceOffloadModel: force_offload_model(mock_patcher, "TestModel", warm=True) mock_patcher.unpatch_model.assert_called_once_with(unpatch_weights=True, warm=True) + + +class TestGenerateAudioProgressReporting: + """Phase 2 (2026-08-15 progress plan): generate_audio() must drive the + standard ComfyUI ProgressBar through the vendored progress_callback.""" + + @staticmethod + def _run(progress_callback_capture: list, generate_side_effect=None): + """Run generate_audio with a mocked model/processor; capture the + progress_callback kwarg handed to model.generate.""" + 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.return_value = {"input_ids": torch.randint(0, 100, (1, 10))} + mock_processor.tokenizer = MagicMock() + + 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.preprocess_comfy_audio", return_value=_mock_voice_sample()): + mock_pbar = MagicMock() + mock_pbar.total = 10 + mock_pbar_cls.return_value = mock_pbar + + # Capture the callback the wrapper passes into the vendored loop. + def _capture(**kwargs): + progress_callback_capture.append(kwargs.get("progress_callback")) + return mock_output + + if generate_side_effect is None: + mock_model.generate.side_effect = _capture + + result = generate_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, + ) + return result, mock_model, mock_pbar_cls, mock_pbar, mock_interrupt + + def test_generate_receives_callable_progress_callback(self): + """T2.1: model.generate must receive a callable progress_callback kwarg.""" + captured = [] + self._run(captured) + assert captured, "model.generate was not called" + assert captured[0] is not None and callable(captured[0]) + + def test_callback_maps_to_update_absolute_with_total(self): + """T2.2: callback(current, total) -> pbar.update_absolute(current, total=total).""" + captured = [] + _, _, _, mock_pbar, _ = self._run(captured) + mock_pbar.update_absolute.reset_mock() + + captured[0](3, 10) + mock_pbar.update_absolute.assert_called_once_with(3, total=10) + + def test_callback_checks_interrupt(self): + """T2.3: the callback must call throw_exception_if_processing_interrupted; + a raised interrupt propagates to the caller. The callback is invoked + INSIDE the patch context so the interrupt mock is still active.""" + 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.return_value = {"input_ids": torch.randint(0, 100, (1, 10))} + mock_processor.tokenizer = MagicMock() + + 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.preprocess_comfy_audio", return_value=_mock_voice_sample()): + mock_pbar = MagicMock() + mock_pbar.total = 10 + mock_pbar_cls.return_value = mock_pbar + mock_interrupt.side_effect = mm.InterruptProcessingException() + + generate_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.get("progress_callback") + assert cb is not None and callable(cb) + with pytest.raises(mm.InterruptProcessingException): + cb(1, 10) + mock_interrupt.assert_called() + + def test_final_update_reaches_total_on_success(self): + """T2.4: after successful generation the bar is driven to 100%.""" + captured = [] + _, _, mock_pbar_cls, mock_pbar, _ = self._run(captured) + + # Initial estimate = inference_steps. + mock_pbar_cls.assert_called_once_with(10) + # The very last update_absolute call must be the final (total) event. + final_call = mock_pbar.update_absolute.call_args_list[-1] + assert final_call.args == (10,), f"final update must be (total,), got {final_call}" + + def test_final_update_sent_even_when_generate_raises(self): + """T2.5: the finally-block final update fires when model.generate raises.""" + mock_model = MagicMock() + mock_model.device = torch.device("cpu") + mock_model.generate.side_effect = RuntimeError("boom") + + mock_processor = MagicMock() + mock_processor.return_value = {"input_ids": torch.randint(0, 100, (1, 10))} + mock_processor.tokenizer = MagicMock() + + with patch("ComfyUI_VibeVoice.modules.generation.ProgressBar") as mock_pbar_cls, \ + patch("ComfyUI_VibeVoice.modules.generation.model_management.throw_exception_if_processing_interrupted"), \ + patch("ComfyUI_VibeVoice.modules.generation.preprocess_comfy_audio", return_value=_mock_voice_sample()): + mock_pbar = MagicMock() + mock_pbar.total = 10 + mock_pbar_cls.return_value = mock_pbar + + with pytest.raises(RuntimeError, match="boom"): + generate_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, + ) + + # The final absolute update must still have been sent. + final_call = mock_pbar.update_absolute.call_args_list[-1] + assert final_call.args == (10,) diff --git a/tests/test_integration.py b/tests/test_integration.py index 8920eae..85f9436 100644 --- a/tests/test_integration.py +++ b/tests/test_integration.py @@ -193,3 +193,181 @@ class TestFullPipelineMocked: voice_samples=[_mock_audio_dict()], speaker_ids=[1], ) + + +class TestNodeProgressIntegration: + """Phase 6 (2026-08-15 progress plan): each node's execute path must drive + the standard ComfyUI ProgressBar with increasing values. The mocked model + invokes the progress hook 3x, simulating a 3-step loop.""" + + def test_tts_node_reports_increasing_progress(self): + from ComfyUI_VibeVoice.nodes.tts_node import VibeVoiceTTSNode + + mock_model = MagicMock() + mock_model.device = torch.device("cpu") + mock_output = MagicMock() + mock_output.speech_outputs = [torch.randn(24000)] + + def _generate(**kwargs): + cb = kwargs.get("progress_callback") + if cb is not None: + for i in (1, 2, 3): + cb(i, 10) + return mock_output + + mock_model.generate.side_effect = _generate + mock_processor = MagicMock() + mock_processor.return_value = {"input_ids": torch.randint(0, 100, (1, 10))} + mock_processor.tokenizer = MagicMock() + mock_patcher = MagicMock() + + with patch("ComfyUI_VibeVoice.nodes.tts_node.load_vibevoice_model", + return_value=(mock_patcher, mock_model, mock_processor)), \ + patch("ComfyUI_VibeVoice.nodes.tts_node.ui.PreviewAudio", MagicMock()), \ + patch("ComfyUI_VibeVoice.modules.generation.ProgressBar") as mock_pbar_cls, \ + patch("ComfyUI_VibeVoice.modules.generation.model_management.throw_exception_if_processing_interrupted"), \ + patch("ComfyUI_VibeVoice.modules.generation.preprocess_comfy_audio", + return_value=_mock_voice_sample()): + mock_pbar = MagicMock() + mock_pbar.total = 10 + mock_pbar_cls.return_value = mock_pbar + + VibeVoiceTTSNode.execute( + model_name="TestModel", + text="[1] Hello world", + quantize_llm_4bit=False, + attention_mode="sdpa", + cfg_scale=1.3, + inference_steps=10, + seed=42, + do_sample=True, + temperature=0.95, + top_p=0.95, + top_k=0, + force_offload=False, + device="cpu", + dtype="fp32", + max_new_tokens=0, + speaker_1_voice=_mock_audio_dict(), + ) + + values = [c.args[0] for c in mock_pbar.update_absolute.call_args_list if c.args] + assert 1 in values and 2 in values and 3 in values + assert values == sorted(values), f"progress must be monotonic, got {values}" + + def test_realtime_node_reports_increasing_progress(self): + from ComfyUI_VibeVoice.nodes.realtime_node import VibeVoiceRealtimeNode + + mock_model = MagicMock() + mock_model.device = torch.device("cpu") + mock_output = MagicMock() + mock_output.speech_outputs = [torch.randn(24000)] + + def _generate(**kwargs): + cb = kwargs.get("progress_callback") + if cb is not None: + for i in (1, 2, 3): + cb(i, 512) + return mock_output + + mock_model.generate.side_effect = _generate + 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), + }) + mock_patcher = MagicMock() + + with patch("ComfyUI_VibeVoice.nodes.realtime_node.load_vibevoice_model", + return_value=(mock_patcher, mock_model, mock_processor)), \ + patch("ComfyUI_VibeVoice.nodes.realtime_node.ui.PreviewAudio", MagicMock()), \ + patch("ComfyUI_VibeVoice.modules.generation.ProgressBar") as mock_pbar_cls, \ + patch("ComfyUI_VibeVoice.modules.generation.model_management.throw_exception_if_processing_interrupted"), \ + 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 + + VibeVoiceRealtimeNode.execute( + model_name="TestStreamingModel", + text="[1] Hello world", + quantize_llm_4bit=False, + attention_mode="sdpa", + cfg_scale=1.3, + inference_steps=10, + seed=42, + do_sample=True, + temperature=0.95, + top_p=0.95, + top_k=0, + stream=False, + force_offload=False, + device="cpu", + dtype="fp32", + speaker_1_voice=_mock_audio_dict(), + ) + + values = [c.args[0] for c in mock_pbar.update_absolute.call_args_list if c.args] + # Loop updates (all but the final guaranteed event) must be monotonic. + loop_values = values[:-1] + assert 1 in loop_values and 2 in loop_values and 3 in loop_values + assert loop_values == sorted(loop_values), f"progress must be monotonic, got {loop_values}" + # The final event must always be present (guaranteed 100%). + assert values[-1] == mock_pbar.total + + def test_asr_node_reports_increasing_progress(self): + from ComfyUI_VibeVoice.nodes.asr_node import VibeVoiceASRNode + + mock_model = MagicMock() + mock_param = MagicMock() + mock_param.device = torch.device("cpu") + mock_model.parameters.return_value = iter([mock_param]) + + def _generate(**kwargs): + streamer = kwargs.get("streamer") + if streamer is not None: + for tok in (10, 11, 12): + streamer.put(torch.tensor([tok])) + streamer.end() + return torch.tensor([[1, 2, 3, 10, 11, 12, 0]]) + + mock_model.generate.side_effect = _generate + mock_processor = MagicMock() + mock_processor.return_value = {"input_ids": torch.tensor([[1, 2, 3]])} + mock_processor.pad_id = 0 + mock_processor.tokenizer.eos_token_id = 0 + mock_processor.decode.return_value = "Hello world" + mock_processor.post_process_transcription.return_value = [] + mock_patcher = MagicMock() + + with patch("ComfyUI_VibeVoice.nodes.asr_node.load_asr_model_patched", + return_value=(mock_patcher, mock_model, mock_processor)), \ + patch("ComfyUI_VibeVoice.modules.asr_generation.ProgressBar") as mock_pbar_cls, \ + patch("ComfyUI_VibeVoice.modules.asr_generation.model_management.throw_exception_if_processing_interrupted"): + mock_pbar = MagicMock() + mock_pbar.total = 32768 + mock_pbar_cls.return_value = mock_pbar + + VibeVoiceASRNode.execute( + model_name="VibeVoice-ASR", + audio=_mock_audio_dict(), + context_info="", + max_new_tokens=32768, + temperature=0.0, + top_p=1.0, + do_sample=False, + num_beams=1, + device="cpu", + dtype="fp32", + attention_mode="sdpa", + force_offload=False, + ) + + values = [c.args[0] for c in mock_pbar.update_absolute.call_args_list if c.args] + assert 1 in values and 2 in values and 3 in values + assert values == sorted(values), f"progress must be monotonic, got {values}" diff --git a/tests/test_streaming_progress.py b/tests/test_streaming_progress.py new file mode 100644 index 0000000..5a30142 --- /dev/null +++ b/tests/test_streaming_progress.py @@ -0,0 +1,150 @@ +"""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,) diff --git a/tests/test_streaming_progress_callback.py b/tests/test_streaming_progress_callback.py new file mode 100644 index 0000000..ff521d7 --- /dev/null +++ b/tests/test_streaming_progress_callback.py @@ -0,0 +1,286 @@ +"""Tests for the ``progress_callback`` hook in the streaming +``VibeVoiceStreamingForConditionalGenerationInference.generate()`` (Phase 3 of +the 2026-08-15 inference-progress-reporting plan). + +The real vendored module is loaded with heavy sibling modules stubbed. The +GenerationMixin plumbing (``_build_generate_config_model_kwargs``, +``prepare_inputs_for_generation``, ``forward_lm`` / ``forward_tts_lm``, +``sample_speech_tokens``) is mocked so the REAL windowed AR loop runs +deterministically: + +* one text window of ``TTS_TEXT_WINDOW_SIZE`` (5) tokens -> 1 callback, +* ``TTS_SPEECH_WINDOW_SIZE`` (6) speech tokens -> 6 callbacks + (the inner speech loop does not break on EOS; EOS only stops the outer + while-loop on the next iteration), +* ``tts_eos_classifier`` returns a high logit so the outer loop stops after + the first text window. + +Expected callback trace (initial step = tts_lm_input_ids len = 3, +max_length = 100): + + (8, 100) # text window prefill (3 + 5) + (9, 100) ... (14, 100) # six speech tokens +""" + +import os +import sys +import types +import importlib.util +from types import SimpleNamespace +from unittest.mock import MagicMock + +import torch +import torch.nn as nn + +from transformers.modeling_utils import PreTrainedModel + + +HIDDEN = 64 +MAX_LENGTH = 100 +TEXT_LEN = 5 # one full text window +TTS_LM_INIT_LEN = 3 # initial tts_lm_input_ids length -> initial `step` +EXPECTED_TEXT_CALLBACK_STEP = TTS_LM_INIT_LEN + TEXT_LEN # 8 +EXPECTED_SPEECH_CALLBACKS = 6 # TTS_SPEECH_WINDOW_SIZE + + +# ---------------------------------------------------------------------- +# Stubs + real-module loading +# ---------------------------------------------------------------------- +class _FakeStreamingConfig: + __name__ = "VibeVoiceStreamingConfig" + model_type = "vibevoice_streaming" + + +class _FakeStreamingPreTrainedModel(PreTrainedModel): + """Real PreTrainedModel subclass so the inference class can inherit from it.""" + + config_class = _FakeStreamingConfig + base_model_prefix = "model" + + def _init_weights(self, module): + pass + + +def _install_stub_modules(): + """Replace the conftest MagicMocks for the streaming module's imports with + stubs that provide real classes where inheritance/instantiation requires + them.""" + streaming_pkg_root = os.path.join(os.getcwd(), "src", "vibevoice") + + # Package hierarchy (needed for relative imports during exec). + if "src.vibevoice" not in sys.modules or not hasattr(sys.modules["src.vibevoice"], "__path__"): + pkg = types.ModuleType("src.vibevoice") + pkg.__path__ = [streaming_pkg_root] + sys.modules["src.vibevoice"] = pkg + else: + sys.modules["src.vibevoice"].__path__ = [streaming_pkg_root] + + mod_pkg_name = "src.vibevoice.modular" + if mod_pkg_name not in sys.modules or not hasattr(sys.modules[mod_pkg_name], "__path__"): + mod_pkg = types.ModuleType(mod_pkg_name) + mod_pkg.__path__ = [os.path.join(streaming_pkg_root, "modular")] + sys.modules[mod_pkg_name] = mod_pkg + else: + sys.modules[mod_pkg_name].__path__ = [os.path.join(streaming_pkg_root, "modular")] + + # modeling_vibevoice_streaming: must expose a REAL PreTrainedModel subclass. + ms = types.ModuleType("src.vibevoice.modular.modeling_vibevoice_streaming") + ms.VibeVoiceStreamingPreTrainedModel = _FakeStreamingPreTrainedModel + ms.VibeVoiceStreamingModel = MagicMock() + ms.BinaryClassifier = MagicMock() + sys.modules["src.vibevoice.modular.modeling_vibevoice_streaming"] = ms + + # configuration_vibevoice_streaming: plain class with model_type (register()). + cfg = types.ModuleType("src.vibevoice.modular.configuration_vibevoice_streaming") + cfg.VibeVoiceStreamingConfig = _FakeStreamingConfig + sys.modules["src.vibevoice.modular.configuration_vibevoice_streaming"] = cfg + + # The remaining imports survive as conftest MagicMocks (names only), but + # ensure they exist so `from .x import Y` resolves. + for name in ( + "src.vibevoice.modular.modular_vibevoice_tokenizer", + "src.vibevoice.modular.modular_vibevoice_diffusion_head", + "src.vibevoice.modular.modular_vibevoice_text_tokenizer", + "src.vibevoice.modular.streamer", + "src.vibevoice.schedule.dpm_solver", + ): + sys.modules.setdefault(name, MagicMock()) + + +def _load_real_streaming_module(): + _install_stub_modules() + sys.modules.pop("src.vibevoice.modular.modeling_vibevoice_streaming_inference", None) + + spec = importlib.util.spec_from_file_location( + "src.vibevoice.modular.modeling_vibevoice_streaming_inference", + "src/vibevoice/modular/modeling_vibevoice_streaming_inference.py", + ) + module = importlib.util.module_from_spec(spec) + sys.modules["src.vibevoice.modular.modeling_vibevoice_streaming_inference"] = module + spec.loader.exec_module(module) + return module + + +_streaming = _load_real_streaming_module() +VibeVoiceStreamingForConditionalGenerationInference = ( + _streaming.VibeVoiceStreamingForConditionalGenerationInference +) +VibeVoiceGenerationOutput = _streaming.VibeVoiceGenerationOutput + + +# ---------------------------------------------------------------------- +# Scripted mock model +# ---------------------------------------------------------------------- +class _FakeTokenizer: + bos_token_id = 1 + eos_token_id = 2 + pad_token_id = 0 + speech_start_id = 3 + speech_end_id = 4 + speech_diffusion_id = 5 + + def convert_tokens_to_ids(self, token): + return 7 + + +class _FakeForwardOutput: + def __init__(self, seq_len: int = 4): + self.last_hidden_state = torch.zeros(1, seq_len, HIDDEN) + self.past_key_values = None + + +def _make_streaming_model(): + cls = VibeVoiceStreamingForConditionalGenerationInference + model = cls.__new__(cls) + torch.nn.Module.__init__(model) + # A real parameter so `self.device` resolves (PreTrainedModel.device). + model.register_parameter("_dummy", nn.Parameter(torch.zeros(1))) + + model.config = MagicMock() + model.config.acoustic_vae_dim = HIDDEN + model.config.decoder_config = MagicMock() + model.config.decoder_config.max_position_embeddings = 1024 + model.config.decoder_config.hidden_size = HIDDEN + model.config.diffusion_head_config = MagicMock() + model.config.diffusion_head_config.ddpm_num_inference_steps = 5 + model.ddpm_inference_steps = 5 + + inner = MagicMock() + inner.speech_scaling_factor = torch.tensor(1.0) + inner.speech_bias_factor = torch.tensor(0.0) + inner.acoustic_tokenizer = MagicMock() + inner.acoustic_tokenizer.device = torch.device("cpu") + inner.acoustic_tokenizer.decode = MagicMock(return_value=torch.full((1, 1600), 0.25)) + inner.acoustic_connector = MagicMock(return_value=torch.zeros(1, 1, HIDDEN)) + model.model = inner + + # EOS classifier: high logit -> sigmoid > 0.5 -> finished after window 1. + model.tts_eos_classifier = MagicMock(return_value=torch.tensor([[10.0]])) + + # --- GenerationMixin plumbing mocks ------------------------------- + def _build_cfg(generation_config, inputs, tokenizer, return_processors=False, **kw): + cfg = SimpleNamespace(max_length=MAX_LENGTH, min_length=0) + ids = kw.get("input_ids") + if ids is None: + ids = torch.zeros(1, 1, dtype=torch.long) + mk = { + "input_ids": ids, + "attention_mask": torch.ones(1, ids.shape[1], dtype=torch.long), + "cache_position": torch.arange(ids.shape[1], dtype=torch.long), + "past_key_values": None, + "use_cache": True, + } + if return_processors: + return cfg, mk, ids, [], [] + return cfg, mk, ids + + model._build_generate_config_model_kwargs = MagicMock(side_effect=_build_cfg) + model.prepare_inputs_for_generation = MagicMock(return_value={}) + model.forward_lm = MagicMock(return_value=_FakeForwardOutput()) + model.forward_tts_lm = MagicMock(return_value=_FakeForwardOutput()) + model.sample_speech_tokens = MagicMock(return_value=torch.zeros(1, HIDDEN)) + # Bypass the GenerationMixin cache bookkeeping (module-level + # _update_model_kwargs_for_generation still runs with real tensors). + model._update_model_kwargs_for_generation = ( + lambda outputs, model_kwargs, is_encoder_decoder=False, num_new_tokens=1: model_kwargs + ) + return model + + +def _run_generate(model, progress_callback=None): + kwargs = dict( + input_ids=torch.randint(0, 100, (1, 4)), + tts_text_ids=torch.randint(0, 100, (1, TEXT_LEN)), + tts_lm_input_ids=torch.randint(0, 100, (1, TTS_LM_INIT_LEN)), + tokenizer=_FakeTokenizer(), + all_prefilled_outputs={ + "lm": _FakeForwardOutput(), + "tts_lm": _FakeForwardOutput(), + "neg_lm": _FakeForwardOutput(), + "neg_tts_lm": _FakeForwardOutput(), + }, + max_new_tokens=50, + cfg_scale=1.0, + return_speech=True, + show_progress_bar=False, # silence the console tqdm bar + ) + if progress_callback is not None: + kwargs["progress_callback"] = progress_callback + return model.generate(**kwargs) + + +# ---------------------------------------------------------------------- +# Tests +# ---------------------------------------------------------------------- +class TestStreamingProgressCallbackContract: + def test_callback_none_backward_compatible(self): + """T3.1: generate() without progress_callback runs unchanged.""" + model = _make_streaming_model() + out = _run_generate(model) + assert isinstance(out, VibeVoiceGenerationOutput) + assert out.speech_outputs is not None + assert len(out.speech_outputs) == 1 + assert isinstance(out.speech_outputs[0], torch.Tensor) + + def test_callback_monotonic_with_constant_total(self): + """T3.2: currents monotonic non-decreasing; total == max_length always.""" + model = _make_streaming_model() + calls = [] + _run_generate(model, progress_callback=lambda c, t: calls.append((c, t))) + + assert calls, "progress_callback was never invoked" + currents = [c for c, _ in calls] + totals = [t for _, t in calls] + assert currents == sorted(currents) + assert set(totals) == {MAX_LENGTH} + assert all(0 <= c <= MAX_LENGTH for c in currents) + + def test_callback_fires_for_text_window_and_speech_tokens(self): + """T3.3: one callback after the text-window prefill and one per + generated speech token (TTS_SPEECH_WINDOW_SIZE = 6).""" + model = _make_streaming_model() + calls = [] + _run_generate(model, progress_callback=lambda c, t: calls.append((c, t))) + + assert len(calls) == 1 + EXPECTED_SPEECH_CALLBACKS, f"got {calls}" + # First call: text window prefill advanced step by TEXT_LEN. + assert calls[0] == (EXPECTED_TEXT_CALLBACK_STEP, MAX_LENGTH) + # Speech-token calls increment by exactly 1 each. + speech_currents = [c for c, _ in calls[1:]] + assert speech_currents == [ + EXPECTED_TEXT_CALLBACK_STEP + i + 1 for i in range(EXPECTED_SPEECH_CALLBACKS) + ] + + def test_callback_exception_propagates(self): + """T3.4: a raising callback interrupts generation.""" + model = _make_streaming_model() + + def _raiser(current, total): + raise RuntimeError("interrupt requested") + + try: + _run_generate(model, progress_callback=_raiser) + raised = False + except RuntimeError as e: + raised = "interrupt requested" in str(e) + assert raised, "callback exception must propagate out of generate()"