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:
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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 = {
|
||||
|
||||
@@ -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,)
|
||||
|
||||
@@ -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])
|
||||
@@ -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,)
|
||||
|
||||
@@ -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}"
|
||||
|
||||
@@ -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,)
|
||||
@@ -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()"
|
||||
Reference in New Issue
Block a user