From b4fce48cfffb60e72d767dc4f7273e1fa6ac3d2e Mon Sep 17 00:00:00 2001
From: WildAi <2853742+wildminder@users.noreply.github.com>
Date: Sun, 16 Aug 2026 02:24:52 +0300
Subject: [PATCH] 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.
---
README.md | 32 ++
modules/asr_generation.py | 48 +++
modules/generation.py | 39 ++-
pyproject.toml | 2 +-
src/vibevoice/modular/modeling_vibevoice.py | 16 +
.../modeling_vibevoice_streaming_inference.py | 16 +
tests/test_asr_generation.py | 105 ++++++
tests/test_generate_progress_callback.py | 319 ++++++++++++++++++
tests/test_generation.py | 144 ++++++++
tests/test_integration.py | 178 ++++++++++
tests/test_streaming_progress.py | 150 ++++++++
tests/test_streaming_progress_callback.py | 286 ++++++++++++++++
12 files changed, 1330 insertions(+), 5 deletions(-)
create mode 100644 tests/test_generate_progress_callback.py
create mode 100644 tests/test_streaming_progress.py
create mode 100644 tests/test_streaming_progress_callback.py
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()"