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.
This commit is contained in:
WildAi
2026-08-16 02:24:52 +03:00
parent b2feff00d4
commit b4fce48cff
12 changed files with 1330 additions and 5 deletions
+32
View File
@@ -160,6 +160,38 @@ This node features a sophisticated system for managing performance, memory, and
## Changelog
<details open>
<summary><strong>v2.1.1 - Standard ComfyUI Progress Bar During Inference</strong></summary>
### ✨ 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).
</details>
<details>
<summary><strong>v2.1.0 - torchaudio-Primary Audio Backend (librosa now optional)</strong></summary>
### ✨ Highlights
+48
View File
@@ -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:
+35 -4
View File
@@ -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:
+1 -1
View File
@@ -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"]
@@ -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).
# ------------------------------------------------------------------
@@ -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 = {
+105
View File
@@ -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,)
+319
View File
@@ -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])
+144
View File
@@ -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,)
+178
View File
@@ -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}"
+150
View File
@@ -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,)
+286
View File
@@ -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()"