Files
WildAi d2390eb476 fix: correct realtime VibeVoice output on transformers 5.3
Garbled, parameter-insensitive speech came from a randomized EOS head:
5.3 re-runs _initialize_weights over acoustic_connector and
tts_eos_classifier because the vendored _init_weights override had no
_is_hf_initialized guard. Guard it; only checkpoint-absent weights are
initialized now.

- MockCacheLayer exposes both the 4.x and 5.x cache APIs, so the
  prefilled voice prompt is visible to 5.3's mask builder.
- _ensure_cache_has_layers covers the container: offload/prefetch,
  batch ops, crop, and a copyable lazy prefetch stream.
- max_new_tokens is a combined text+speech budget, not latents.
- cfg_scale floor of 1.5 for the realtime family.
- voice presets resolve against every registered TTS root.
- sage excluded from the realtime path: it ignores the attention mask.
- bound transformers to >=5.3.0,<5.4, the measured line.
2026-09-26 23:30:37 +03:00

110 lines
3.4 KiB
Python

"""Progress and cancellation coverage for the official realtime adapter."""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
import torch
from ComfyUI_VibeVoice.modules import realtime_generation
from ComfyUI_VibeVoice.modules.realtime_generation import generate_realtime_audio
from ComfyUI_VibeVoice.modules.voice_presets import PRESET_CACHE_KEYS
class _Model:
def __init__(self, side_effect=None):
self.set_ddpm_inference_steps = MagicMock()
self.generate = MagicMock(side_effect=side_effect)
if side_effect is None:
self.generate.return_value = SimpleNamespace(
speech_outputs=[torch.ones(24000)]
)
class _Processor:
def __init__(self):
self.tokenizer = object()
self.process_input_with_cached_prompt = MagicMock(
return_value={
"tts_text_ids": torch.ones(1, 3, dtype=torch.long),
# The adapter budgets the auto length from the script and the
# TTS-LM prompt length, so both keys must be present.
"tts_lm_input_ids": torch.ones(1, 3, dtype=torch.long),
}
)
def __call__(self, *args, **kwargs):
raise AssertionError("streaming processor __call__ must not be used")
class _ProcessorByName(_Processor):
pass
# The adapter intentionally uses class names to identify the official pair.
_ProcessorByName.__name__ = "VibeVoiceStreamingProcessor"
_Model.__name__ = "VibeVoiceStreamingForConditionalGenerationInference"
def _preset():
return {
key: SimpleNamespace(last_hidden_state=torch.ones(1, index + 2, 4))
for index, key in enumerate(PRESET_CACHE_KEYS)
}
def _run(model):
pbar = MagicMock()
pbar.total = 1
with patch.object(realtime_generation, "ProgressBarWithConsole", return_value=pbar), patch.object(
realtime_generation.model_management,
"get_torch_device",
return_value=torch.device("cpu"),
), patch.object(
realtime_generation.model_management,
"throw_exception_if_processing_interrupted",
):
result = generate_realtime_audio(
model,
_ProcessorByName(),
"[1] Hello",
_preset(),
)
return result, pbar
def test_generate_receives_progress_callback():
model = _Model()
_run(model)
assert callable(model.generate.call_args.kwargs["progress_callback"])
def test_callback_updates_total_and_finalizes():
model = _Model()
_result, pbar = _run(model)
callback = model.generate.call_args.kwargs["progress_callback"]
assert pbar.update_absolute.call_args_list[-1].args == (1,)
callback(4, 12)
pbar.update_absolute.assert_any_call(4, total=12)
def test_callback_checks_interruption():
model = _Model()
_run(model)
callback = model.generate.call_args.kwargs["progress_callback"]
with patch.object(
realtime_generation.model_management,
"throw_exception_if_processing_interrupted",
side_effect=realtime_generation.model_management.InterruptProcessingException(),
):
with pytest.raises(realtime_generation.model_management.InterruptProcessingException):
callback(1, 2)
def test_generation_failure_still_finalizes_progress():
model = _Model(side_effect=RuntimeError("boom"))
with pytest.raises(RuntimeError, match="boom"):
_run(model)