Files
wildminder-ComfyUI-VibeVoice/tests/test_generate_progress_callback.py
T
WildAi b4fce48cff 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.
2026-08-16 02:24:52 +03:00

320 lines
12 KiB
Python

"""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])