Files
WildAi 8389ab4ce6 refactor: log through ComfyUI's own root-logger idiom, re-audit every level
The package installed its own handler, Formatter and level on a
module-level logger, then logged through that logger. ComfyUI core already
installs app.logger.ColoredFormatter on the ROOT logger, so a bare
logging.<level>() call whose message starts with "[ComfyUI-VibeVoice] " is
tagged and coloured for free. This removes the duplicate machinery rather
than extending it -- no vv_logging module, Logger subclass, LoggerAdapter,
vendored ANSI table, new handler, setLevel or propagate.

Three hazards that motivated deleting the setup block rather than moving it:

- propagate=False severed pytest's caplog; the tests only passed because
  __init__.py's `if "pytest" in sys.modules` guard skipped the block.
  Bare root calls propagate by default.
- logger.setLevel(INFO) on the package logger pinned every descendant, so
  ComfyUI's --verbose never applied to this node. Deleting the block fixes it.
- The level must stay a literal at the call site; no env-var switch was
  added. modules/diagnostics.py gates diagnostic CONTENT and is untouched.

Every call site (235 across 34 files) is now a bare root call with the
prefix. The vendored src/vibevoice tree used transformers.utils.logging,
which has no module-level info/warning/debug; those files now import stdlib
logging, so a new call site there would fail loudly instead of silently.

Level audit against ERROR=failure / WARNING=degradation / INFO=user-facing
progress / DEBUG=internals. Deleted as noise: the resample notice (44.1 kHz
reference audio is the normal case) and the two "Successfully loaded
external VibeVoice" confirmations, which duplicated the patcher's load line.
Promoted to WARNING: a mixed-naming GGUF, which is resolved by heuristic
majority vote and silently aliases the rest. Demoted to DEBUG: the four
SageAttention kernel-selection lines, model discovery, shard counts,
per-retry download attempts, attention-mode confirmation and the
save_pretrained notice. Kept at INFO: generation complete, transcription
results, model downloads and load starts.

Judgement calls recorded in docs/2026-10-01-vv-logging-cleanup-design.md.
The two I am least sure of: the ungated memory-census profile line was
demoted rather than gated (adding a gate would not be presentation-only),
and patcher.py's "Loading VibeVoice models for..." was kept at INFO against
the ask's example list because it is the only load-start line the package has.

48 new tests in tests/test_logging_idiom.py plus tests/test_audio_utils.py:
byte-exact ColoredFormatter rendering, no-leftover-machinery, AST prefix and
two-direction level policy, and the 44.1 kHz resample regression (the call
still fires with (44100, 24000); nothing is logged). 37 caplog.at_level
pins that named a module logger were stripped -- they only lowered a named
logger and left root at WARNING, so the record was discarded before capture.

Gate: 1924 passed, 30 skipped, 0 failed (1876 before this change). Twelve
deliberate mutations -- re-deleting a message, re-leveling, re-adding
setLevel and propagate=False, stripping a prefix, moving the resample out
of its branch -- were each caught by at least one test.

Not run: ComfyUI was never launched and no checkpoint was loaded.
2026-10-01 23:08:04 +03:00

1136 lines
47 KiB
Python

"""Tests for modules/external_loader.py - External model loading.
Covers:
- Path resolution helpers (sidecar config / preprocessor / tokenizer dir)
- load_external_vibevoice_model() core loading function
"""
import os
import pytest
from unittest.mock import MagicMock, patch
import torch
from ComfyUI_VibeVoice.modules import external_loader
from ComfyUI_VibeVoice.modules.external_loader import (
resolve_sidecar_config,
resolve_sidecar_preprocessor,
resolve_sidecar_tokenizer_dir,
load_external_vibevoice_model,
EXTERNAL_CONFIG_OPTIONS,
)
# Fake streaming config class for isinstance() checks (MagicMock cannot be
# used with isinstance — see tests/test_loader.py for the same pattern).
class _FakeStreamingCfg:
pass
# ====================================================================
# Phase 3 (plan 2026-08-18, D1): external instantiation is meta by default
# ====================================================================
class _TinyCtorModel(torch.nn.Module):
"""Minimal config-ctor stand-in for the vendored model classes."""
def __init__(self, config):
super().__init__()
self.lin = torch.nn.Linear(4, 4)
class TestExternalInstantiationIsMeta:
"""Contract: external instantiation builds under a meta context by default."""
def test_asr_instantiation_defaults_to_meta(self):
with patch.object(
external_loader, "VibeVoiceASRForConditionalGeneration", _TinyCtorModel
):
model = external_loader._instantiate_asr_model(
config=MagicMock(),
attn_implementation="eager",
final_load_dtype=torch.bfloat16,
)
params = list(model.parameters())
assert params
assert all(p.is_meta for p in params), "ASR construction must be meta by default"
def test_asr_instantiation_use_meta_false_eager(self):
with patch.object(
external_loader, "VibeVoiceASRForConditionalGeneration", _TinyCtorModel
):
model = external_loader._instantiate_asr_model(
config=MagicMock(),
attn_implementation="eager",
final_load_dtype=torch.bfloat16,
use_meta=False,
)
params = list(model.parameters())
assert params
assert all(not p.is_meta for p in params), "use_meta=False must be eager"
def test_asr_instantiation_real_config_no_deprecation(self):
"""With a REAL transformers PretrainedConfig (v5: torch_dtype is a
deprecated property), _instantiate_asr_model must record the dtype on
the canonical ``dtype`` attribute and emit no deprecation warning.
A capture handler is attached directly to the emitting transformers
logger (its warning_once is lru_cache-wrapped, so the cache is
cleared to keep the absence assertion non-vacuous)."""
import logging as pylogging
from transformers import PretrainedConfig
config = PretrainedConfig()
config.decoder_config = PretrainedConfig()
logger = pylogging.getLogger("transformers.configuration_utils")
records = []
class _Capture(pylogging.Handler):
def emit(self, record):
records.append(record.getMessage())
handler = _Capture(level=pylogging.WARNING)
logger.addHandler(handler)
pylogging.Logger.warning_once.cache_clear()
try:
with patch.object(
external_loader, "VibeVoiceASRForConditionalGeneration", _TinyCtorModel
):
external_loader._instantiate_asr_model(
config=config,
attn_implementation="eager",
final_load_dtype=torch.bfloat16,
)
finally:
logger.removeHandler(handler)
pylogging.Logger.warning_once.cache_clear()
assert config.dtype == torch.bfloat16
assert config.decoder_config.dtype == torch.bfloat16
assert not any(
"`torch_dtype` is deprecated" in m for m in records
), "torch_dtype deprecation warning was emitted"
# ====================================================================
# Path resolution helpers
# ====================================================================
class TestPathResolution:
"""Test sidecar path resolution helpers."""
def test_sidecar_config_found(self, tmp_path):
"""<weight>.config.json sidecar is preferred when present."""
weight = tmp_path / "foo.safetensors"
weight.write_bytes(b"")
sidecar = tmp_path / "foo.safetensors.config.json"
sidecar.write_text("{}")
result = resolve_sidecar_config(str(weight), "VibeVoice-1.5B")
assert result == str(sidecar)
def test_sidecar_config_dir_fallback(self, tmp_path):
"""config.json in the same directory is used when no <weight>.config.json."""
weight = tmp_path / "foo.safetensors"
weight.write_bytes(b"")
dir_config = tmp_path / "config.json"
dir_config.write_text("{}")
result = resolve_sidecar_config(str(weight), "VibeVoice-1.5B")
assert result == str(dir_config)
def test_sidecar_config_missing_falls_back_to_packaged(self, tmp_path):
"""No sidecar + config_name='VibeVoice-1.5B' → packaged 1.5B default."""
weight = tmp_path / "foo.safetensors"
weight.write_bytes(b"")
result = resolve_sidecar_config(str(weight), "VibeVoice-1.5B")
assert result.endswith("default_VibeVoice-1.5B_config.json")
assert os.path.exists(result)
def test_sidecar_config_missing_falls_back_to_7b(self, tmp_path):
"""No sidecar + config_name='VibeVoice-7B' → packaged 7B default.
Plan 2026-08-27: the option 'VibeVoice-Large' was removed (it was an
alias of 7B); the packaged FILE keeps its historical name.
"""
weight = tmp_path / "foo.safetensors"
weight.write_bytes(b"")
result = resolve_sidecar_config(str(weight), "VibeVoice-7B")
assert result.endswith("default_VibeVoice-Large_config.json")
assert os.path.exists(result)
def test_sidecar_config_no_packaged_default_raises(self, tmp_path):
"""No sidecar + config_name without packaged default → FileNotFoundError."""
weight = tmp_path / "foo.safetensors"
weight.write_bytes(b"")
with pytest.raises(FileNotFoundError):
resolve_sidecar_config(str(weight), "VibeVoice-Realtime-0.5B")
def test_sidecar_preprocessor_found(self, tmp_path):
"""<weight>.preprocessor.json sidecar is returned when present."""
weight = tmp_path / "foo.safetensors"
weight.write_bytes(b"")
sidecar = tmp_path / "foo.safetensors.preprocessor.json"
sidecar.write_text("{}")
result = resolve_sidecar_preprocessor(str(weight))
assert result == str(sidecar)
def test_sidecar_preprocessor_dir_fallback(self, tmp_path):
"""preprocessor_config.json in the same directory is used as fallback."""
weight = tmp_path / "foo.safetensors"
weight.write_bytes(b"")
dir_config = tmp_path / "preprocessor_config.json"
dir_config.write_text("{}")
result = resolve_sidecar_preprocessor(str(weight))
assert result == str(dir_config)
def test_sidecar_preprocessor_missing(self, tmp_path):
"""No preprocessor sidecar → empty string."""
weight = tmp_path / "foo.safetensors"
weight.write_bytes(b"")
result = resolve_sidecar_preprocessor(str(weight))
assert result == ""
def test_sidecar_tokenizer_dir(self, tmp_path):
"""Tokenizer dir is the directory containing the weight file."""
weight = tmp_path / "foo.safetensors"
weight.write_bytes(b"")
result = resolve_sidecar_tokenizer_dir(str(weight))
assert result == str(tmp_path)
def test_external_config_options_contains_expected(self):
"""EXTERNAL_CONFIG_OPTIONS contains the documented config names."""
assert "VibeVoice-1.5B" in EXTERNAL_CONFIG_OPTIONS
assert "VibeVoice-7B" in EXTERNAL_CONFIG_OPTIONS
assert "VibeVoice-Realtime-0.5B" in EXTERNAL_CONFIG_OPTIONS
assert "VibeVoice-ASR" in EXTERNAL_CONFIG_OPTIONS
# Plan 2026-08-27: ambiguous alias removed from the visible list.
assert "VibeVoice-Large" not in EXTERNAL_CONFIG_OPTIONS
# ====================================================================
# load_external_vibevoice_model
# ====================================================================
class TestLoadExternalModel:
"""Test the core load_external_vibevoice_model() function."""
@pytest.fixture
def weight_file(self, tmp_path):
"""Create a dummy weight file."""
weight = tmp_path / "model.safetensors"
weight.write_bytes(b"dummy")
return str(weight)
def _run(self, weight_file, streaming=False, **kwargs):
"""Run load_external_vibevoice_model with all deps mocked.
Returns (result, mocks_dict).
"""
fake_state_dict = {"model.language_model.weight": torch.zeros(2, 2)}
if streaming:
fake_config = _FakeStreamingCfg()
else:
fake_config = MagicMock(spec=[])
fake_tokenizer = MagicMock()
fake_processor = MagicMock()
fake_model = MagicMock()
fake_model.load_state_dict.return_value = ([], [])
# model = model.to(dtype=...) reassigns model to the .to() return
# value; return the same mock so subsequent calls (eval) land on it.
fake_model.to.return_value = fake_model
mocks = {}
with patch.object(
external_loader, "iter_safetensors_tensors",
return_value=iter(list(fake_state_dict.items())),
) as m_iter, patch.object(
external_loader.comfy.utils, "load_torch_file", return_value=fake_state_dict
) as m_load_torch, patch.object(
external_loader.VibeVoiceLoader, "_load_config", return_value=fake_config
) as m_load_config, patch.object(
external_loader.VibeVoiceLoader, "_load_tokenizer", return_value=fake_tokenizer
) as m_load_tokenizer, patch.object(
external_loader.VibeVoiceLoader, "_load_processor", return_value=fake_processor
) as m_load_processor, patch.object(
external_loader.VibeVoiceLoader, "_instantiate_model", return_value=fake_model
) as m_instantiate, patch.object(
external_loader, "resolve_sidecar_config", return_value="/fake/config.json"
), patch.object(
external_loader, "resolve_sidecar_preprocessor", return_value=""
), patch.object(
external_loader, "resolve_sidecar_tokenizer_dir", return_value="/fake/dir"
), patch.object(
external_loader, "resolve_dtype", return_value=torch.float32
), patch.object(
external_loader, "resolve_attention_mode", side_effect=lambda m, q: m
), patch.object(
external_loader, "get_attn_implementation_for_load", return_value="eager"
), patch.object(
external_loader, "VibeVoiceStreamingConfig", _FakeStreamingCfg
):
mocks.update({
"state_dict": fake_state_dict,
"config": fake_config,
"tokenizer": fake_tokenizer,
"processor": fake_processor,
"model": fake_model,
"load_torch": m_load_torch,
"iter_tensors": m_iter,
"load_config": m_load_config,
"load_tokenizer": m_load_tokenizer,
"load_processor": m_load_processor,
"instantiate": m_instantiate,
})
result = load_external_vibevoice_model(weight_file, "VibeVoice-1.5B", **kwargs)
return result, mocks
def test_external_pt_checkpoint_remains_accepted(self, tmp_path):
weight = tmp_path / "realtime-checkpoint.pt"
weight.write_bytes(b"dummy")
result, _ = self._run(str(weight))
assert result["source_path"] == str(weight)
def test_load_external_model_returns_dict_with_required_keys(self, weight_file):
"""Returned bundle has all required keys."""
result, _ = self._run(weight_file)
assert isinstance(result, dict)
for key in ("config", "processor", "model", "model_name", "source_path", "is_streaming"):
assert key in result, f"Missing key: {key}"
def test_bundle_does_not_retain_state_dict(self, weight_file):
"""Plan 2026-08-18 D7/RC-4: the bundle must NOT retain the state dict."""
result, _ = self._run(weight_file)
assert "state_dict" not in result, (
"bundle must not retain the state dict (dead RAM, RC-4)")
def test_dense_route_consumes_the_file_tensor_iterator(self, weight_file):
"""The dense route streams from the shared file iterator.
It used to build a whole CPU state dict and then ``model.to(cuda)``,
which page-faults the file mapping on the way to VRAM. Every route now
hands a per-tensor generator to ``_stream_apply_dense``.
"""
_, mocks = self._run(weight_file)
mocks["iter_tensors"].assert_called_once_with(weight_file)
# No batch read of the whole file into host memory.
mocks["load_torch"].assert_not_called()
def test_load_external_model_assigns_per_tensor_into_model(self, weight_file):
"""Weights land through the streaming assign, not load_state_dict."""
from ComfyUI_VibeVoice.modules.loader import VibeVoiceLoader
seen = {}
real = VibeVoiceLoader._stream_apply_dense
def _spy(model, tensor_pairs, **kw):
seen["pairs"] = list(tensor_pairs)
seen["kwargs"] = kw
return real(model, tensor_pairs, **kw)
with patch.object(
external_loader.VibeVoiceLoader, "_stream_apply_dense",
staticmethod(_spy),
):
self._run(weight_file)
assert [k for k, _ in seen["pairs"]] == ["model.language_model.weight"]
# A device is always requested, so nothing is left as a file view.
assert seen["kwargs"].get("target_device") is not None
def test_external_load_uses_the_streaming_assign(self, weight_file):
"""The external path never falls back to the batch state-dict assign."""
with patch.object(
external_loader.VibeVoiceLoader, "_apply_state_dict",
) as m_apply, patch.object(
external_loader.VibeVoiceLoader, "_stream_apply_dense",
return_value=([], []),
) as m_stream:
self._run(weight_file)
m_apply.assert_not_called()
m_stream.assert_called_once()
def test_load_external_model_applies_dtype(self, weight_file):
"""Plan 2026-08-18 D4/RC-3: conditional cast helper is invoked with the final dtype."""
with patch.object(
external_loader, "cast_model_to_dtype_if_needed"
) as m_cast:
self._run(weight_file, dtype_str="fp32")
m_cast.assert_called_once()
args, _ = m_cast.call_args
assert args[1] == torch.float32
def test_load_external_model_eval_mode(self, weight_file):
"""model.eval() is called."""
_, mocks = self._run(weight_file)
mocks["model"].eval.assert_called_once()
def test_load_external_model_streaming_detection(self, weight_file):
"""Streaming config → is_streaming=True in the bundle."""
result, _ = self._run(weight_file, streaming=True)
assert result["is_streaming"] is True
def test_load_external_model_non_streaming_detection(self, weight_file):
"""Non-streaming config → is_streaming=False in the bundle."""
result, _ = self._run(weight_file, streaming=False)
assert result["is_streaming"] is False
def test_load_external_model_model_name_in_bundle(self, weight_file):
"""model_name in the bundle matches config_name."""
result, _ = self._run(weight_file)
assert result["model_name"] == "VibeVoice-1.5B"
def test_load_external_model_source_path_in_bundle(self, weight_file):
"""source_path in the bundle matches the weight path."""
result, _ = self._run(weight_file)
assert result["source_path"] == weight_file
def test_load_external_model_missing_weight_raises(self, tmp_path):
"""Non-existent weight file → FileNotFoundError."""
missing = str(tmp_path / "nonexistent.safetensors")
with pytest.raises(FileNotFoundError):
load_external_vibevoice_model(missing, "VibeVoice-1.5B")
def test_load_external_model_quantize_4bit(self, weight_file):
"""use_llm_4bit=True → replace_with_bnb_linear is called."""
with patch(
"transformers.integrations.bitsandbytes.replace_with_bnb_linear"
) as mock_bnb:
self._run(weight_file, use_llm_4bit=True)
mock_bnb.assert_called_once()
def test_load_external_model_sage_attention(self, weight_file):
"""attention_mode='sage' + compatible → set_sage_attention called."""
with patch.object(
external_loader, "check_sage_attention_compatible", return_value=True
), patch.object(
external_loader, "set_sage_attention", create=True
) as mock_sage:
self._run(weight_file, attention_mode="sage")
mock_sage.assert_called_once()
def test_load_external_model_instantiate_called_with_config(self, weight_file):
"""_instantiate_model is called with the resolved config."""
_, mocks = self._run(weight_file)
mocks["instantiate"].assert_called_once()
call_kwargs = mocks["instantiate"].call_args[1]
assert call_kwargs["config"] is mocks["config"]
def test_load_external_model_processor_called_with_tokenizer(self, weight_file):
"""_load_processor is called with the loaded tokenizer."""
_, mocks = self._run(weight_file)
mocks["load_processor"].assert_called_once()
call_args = mocks["load_processor"].call_args[0]
assert call_args[0] is mocks["tokenizer"]
def test_load_external_model_tts_bundle_is_asr_false(self, weight_file):
"""TTS/streaming bundles are explicitly marked is_asr=False."""
result, _ = self._run(weight_file)
assert result["is_asr"] is False
# ====================================================================
# ASR config-name detection
# ====================================================================
class TestIsASRConfigName:
"""Test is_asr_config_name() dispatch helper."""
def test_asr_config_name_detected(self):
from ComfyUI_VibeVoice.modules.external_loader import is_asr_config_name
assert is_asr_config_name("VibeVoice-ASR") is True
def test_tts_config_names_not_asr(self):
from ComfyUI_VibeVoice.modules.external_loader import is_asr_config_name
for name in ("VibeVoice-1.5B", "VibeVoice-Large", "VibeVoice-Realtime-0.5B"):
assert is_asr_config_name(name) is False
# ====================================================================
# ASR external loading branch
# ====================================================================
class TestLoadExternalASRModel:
"""Test the ASR branch: load_external_vibevoice_asr_model()."""
@pytest.fixture
def weight_file(self, tmp_path):
"""Create a dummy weight file."""
weight = tmp_path / "asr_model.safetensors"
weight.write_bytes(b"dummy")
return str(weight)
def _run(self, weight_file, **kwargs):
"""Run load_external_vibevoice_asr_model with all deps mocked.
Returns (result, mocks_dict).
"""
fake_state_dict = {"model.language_model.weight": torch.zeros(2, 2)}
fake_config = MagicMock(spec=[])
fake_tokenizer = MagicMock()
fake_processor = MagicMock()
fake_model = MagicMock()
fake_model.load_state_dict.return_value = ([], [])
fake_model.to.return_value = fake_model
mocks = {}
with patch.object(
external_loader, "iter_safetensors_tensors",
return_value=iter(list(fake_state_dict.items())),
) as m_iter, patch.object(
external_loader.comfy.utils, "load_torch_file", return_value=fake_state_dict
) as m_load_torch, patch.object(
external_loader, "_load_asr_config", return_value=fake_config
) as m_load_config, patch.object(
external_loader, "_load_asr_tokenizer", return_value=fake_tokenizer
) as m_load_tokenizer, patch.object(
external_loader, "_load_asr_processor", return_value=fake_processor
) as m_load_processor, patch.object(
external_loader, "_instantiate_asr_model", return_value=fake_model
) as m_instantiate, patch.object(
external_loader, "resolve_sidecar_config", return_value="/fake/config.json"
), patch.object(
external_loader, "resolve_sidecar_preprocessor", return_value=""
), patch.object(
external_loader, "resolve_sidecar_tokenizer_dir", return_value="/fake/dir"
), patch.object(
external_loader, "resolve_dtype", return_value=torch.float32
), patch.object(
external_loader, "resolve_attention_mode",
side_effect=lambda m, quantize_4bit=False: m
), patch.object(
external_loader, "get_attn_implementation_for_load", return_value="sdpa"
):
mocks.update({
"state_dict": fake_state_dict,
"config": fake_config,
"tokenizer": fake_tokenizer,
"processor": fake_processor,
"model": fake_model,
"load_torch": m_load_torch,
"iter_tensors": m_iter,
"load_config": m_load_config,
"load_tokenizer": m_load_tokenizer,
"load_processor": m_load_processor,
"instantiate": m_instantiate,
})
result = external_loader.load_external_vibevoice_asr_model(
weight_file, "VibeVoice-ASR", **kwargs
)
return result, mocks
def test_asr_bundle_has_required_keys(self, weight_file):
"""ASR bundle has all required keys plus is_asr=True.
Plan 2026-08-18 D7/RC-4: the bundle no longer retains ``state_dict``.
"""
result, _ = self._run(weight_file)
assert isinstance(result, dict)
for key in ("config", "processor", "model", "model_name", "source_path", "is_streaming", "is_asr"):
assert key in result, f"Missing key: {key}"
assert "state_dict" not in result, "bundle must not retain the state dict (RC-4)"
assert result["is_asr"] is True
assert result["is_streaming"] is False
def test_asr_dispatch_from_main_loader(self, weight_file):
"""config_name='VibeVoice-ASR' routes to the ASR branch."""
with patch.object(
external_loader, "load_external_vibevoice_asr_model", return_value={"is_asr": True}
) as mock_asr:
result = load_external_vibevoice_model(weight_file, "VibeVoice-ASR")
mock_asr.assert_called_once()
assert result["is_asr"] is True
def test_asr_consumes_the_file_tensor_iterator(self, weight_file):
"""The ASR branch streams exactly like the TTS branch."""
_, mocks = self._run(weight_file)
mocks["iter_tensors"].assert_called_once_with(weight_file)
mocks["load_torch"].assert_not_called()
def test_asr_assigns_per_tensor_into_model(self, weight_file):
"""ASR weights land through the streaming assign, not load_state_dict."""
_, mocks = self._run(weight_file)
fake_model = mocks["model"]
fake_model.load_state_dict.assert_not_called()
assert mocks["iter_tensors"].call_count == 1
def test_asr_applies_dtype(self, weight_file):
"""Plan 2026-08-18 D4/RC-3: ASR branch uses the conditional cast helper."""
with patch.object(
external_loader, "cast_model_to_dtype_if_needed"
) as m_cast:
self._run(weight_file, dtype_str="fp32")
m_cast.assert_called_once()
args, _ = m_cast.call_args
assert args[1] == torch.float32
def test_asr_eval_mode(self, weight_file):
"""ASR model.eval() is called."""
_, mocks = self._run(weight_file)
mocks["model"].eval.assert_called_once()
def test_asr_instantiate_called_with_config(self, weight_file):
"""_instantiate_asr_model is called with the resolved ASR config."""
_, mocks = self._run(weight_file)
mocks["instantiate"].assert_called_once()
call_kwargs = mocks["instantiate"].call_args[1]
assert call_kwargs["config"] is mocks["config"]
def test_asr_processor_called_with_tokenizer(self, weight_file):
"""_load_asr_processor is called with the loaded ASR tokenizer."""
_, mocks = self._run(weight_file)
mocks["load_processor"].assert_called_once()
call_args = mocks["load_processor"].call_args[0]
assert call_args[0] is mocks["tokenizer"]
def test_asr_missing_weight_raises(self, tmp_path):
"""Non-existent weight file → FileNotFoundError."""
missing = str(tmp_path / "nonexistent.safetensors")
with pytest.raises(FileNotFoundError):
external_loader.load_external_vibevoice_asr_model(missing, "VibeVoice-ASR")
def test_asr_no_4bit_quantization(self, weight_file):
"""ASR branch never applies 4-bit quantization (no bnb call)."""
with patch(
"transformers.integrations.bitsandbytes.replace_with_bnb_linear"
) as mock_bnb:
self._run(weight_file)
mock_bnb.assert_not_called()
def test_asr_model_name_in_bundle(self, weight_file):
"""model_name in the ASR bundle matches config_name."""
result, _ = self._run(weight_file)
assert result["model_name"] == "VibeVoice-ASR"
def test_asr_source_path_in_bundle(self, weight_file):
"""source_path in the ASR bundle matches the weight path."""
result, _ = self._run(weight_file)
assert result["source_path"] == weight_file
# ====================================================================
# Fix B: low-bit / naive quantization detection
# ====================================================================
import json
import struct
import logging
from ComfyUI_VibeVoice.modules.external_loader import (
_inspect_safetensors_quantization,
_inspect_gguf_quantization,
warn_if_lowbit_quantization,
)
def _write_safetensors_header(path, tensor_specs):
"""Write a minimal safetensors file containing ONLY the header (no tensor data).
``_inspect_safetensors_quantization`` reads only the header, so no data bytes
are needed. ``tensor_specs`` is a list of (name, dtype, shape) tuples.
"""
header = {}
offset = 0
for name, dtype, shape in tensor_specs:
header[name] = {"dtype": dtype, "shape": shape, "data_offsets": [offset, offset]}
header_bytes = json.dumps(header).encode("utf-8")
with open(path, "wb") as f:
f.write(struct.pack("<Q", len(header_bytes)))
f.write(header_bytes)
return str(path)
class TestInspectSafetensorsQuantization:
"""Test _inspect_safetensors_quantization header parsing."""
def test_naive_int8_cast_detected(self, tmp_path):
"""Many I8 tensors with no scale metadata -> naive cast detected."""
specs = [(f"model.layers.{i}.weight", "I8", [16, 16]) for i in range(10)]
specs += [(f"model.layers.{i}.norm", "BF16", [16]) for i in range(2)]
path = _write_safetensors_header(tmp_path / "naive.safetensors", specs)
info = _inspect_safetensors_quantization(path)
assert info is not None
assert info["total"] == 12
assert info["int_count"] == 10
assert info["has_scale_meta"] is False
def test_proper_quant_with_scales_not_flagged(self, tmp_path):
"""I8 tensors WITH scale tensors -> has_scale_meta True (proper quant)."""
specs = [(f"model.layers.{i}.weight", "I8", [16, 16]) for i in range(10)]
specs += [(f"model.layers.{i}.weight_scale", "BF16", [16]) for i in range(10)]
path = _write_safetensors_header(tmp_path / "proper.safetensors", specs)
info = _inspect_safetensors_quantization(path)
assert info["int_count"] == 10
assert info["has_scale_meta"] is True
def test_full_precision_no_int(self, tmp_path):
"""All BF16/F32 -> int_count 0."""
specs = [(f"model.layers.{i}.weight", "BF16", [16, 16]) for i in range(10)]
path = _write_safetensors_header(tmp_path / "fp.safetensors", specs)
info = _inspect_safetensors_quantization(path)
assert info["int_count"] == 0
assert info["has_scale_meta"] is False
def test_unparseable_file_returns_none(self, tmp_path):
"""A non-safetensors / corrupt file returns None (no crash)."""
path = tmp_path / "corrupt.safetensors"
path.write_bytes(b"not a safetensors file")
assert _inspect_safetensors_quantization(str(path)) is None
class TestWarnIfLowbitQuantization:
"""Test warn_if_lowbit_quantization end-to-end warning behavior."""
def test_naive_int8_warns(self, tmp_path, caplog):
"""Naive int8 cast (no scales) must emit a warning."""
specs = [(f"model.layers.{i}.weight", "I8", [16, 16]) for i in range(10)]
path = _write_safetensors_header(tmp_path / "naive.safetensors", specs)
with caplog.at_level(logging.WARNING):
warn_if_lowbit_quantization(path)
assert any("NO dequantization scale" in r.message for r in caplog.records)
def test_proper_quant_no_warning(self, tmp_path, caplog):
"""Proper quant with scales must NOT warn."""
specs = [(f"model.layers.{i}.weight", "I8", [16, 16]) for i in range(10)]
specs += [(f"model.layers.{i}.weight_scale", "BF16", [16]) for i in range(10)]
path = _write_safetensors_header(tmp_path / "proper.safetensors", specs)
with caplog.at_level(logging.WARNING):
warn_if_lowbit_quantization(path)
assert not any("naive int cast" in r.message for r in caplog.records)
def test_full_precision_no_warning(self, tmp_path, caplog):
"""Full-precision file must NOT warn."""
specs = [(f"model.layers.{i}.weight", "BF16", [16, 16]) for i in range(10)]
path = _write_safetensors_header(tmp_path / "fp.safetensors", specs)
with caplog.at_level(logging.WARNING):
warn_if_lowbit_quantization(path)
assert not any("naive int cast" in r.message for r in caplog.records)
def test_gguf_lowbit_warns(self, tmp_path, caplog):
"""GGUF with many sub-4-bit I-quant tensors must warn."""
# Build a fake GGUFReader with a tensor table of IQ3_XXS tensors.
from gguf.constants import GGMLQuantizationType
fake_tensors = [MagicMock() for _ in range(10)]
for t in fake_tensors:
t.tensor_type = GGMLQuantizationType.IQ3_XXS
fake_reader = MagicMock()
fake_reader.tensors = fake_tensors
gguf_path = str(tmp_path / "lowbit.gguf")
with patch("gguf.GGUFReader", return_value=fake_reader), \
caplog.at_level(logging.WARNING):
warn_if_lowbit_quantization(gguf_path)
assert any("sub-4-bit" in r.message for r in caplog.records)
def test_gguf_highbit_no_warning(self, tmp_path, caplog):
"""GGUF with Q8_0 tensors must NOT warn."""
from gguf.constants import GGMLQuantizationType
fake_tensors = [MagicMock() for _ in range(10)]
for t in fake_tensors:
t.tensor_type = GGMLQuantizationType.Q8_0
fake_reader = MagicMock()
fake_reader.tensors = fake_tensors
gguf_path = str(tmp_path / "highbit.gguf")
with patch("gguf.GGUFReader", return_value=fake_reader), \
caplog.at_level(logging.WARNING):
warn_if_lowbit_quantization(gguf_path)
assert not any("sub-4-bit" in r.message for r in caplog.records)
def test_non_weight_extension_no_crash(self, tmp_path, caplog):
"""An unrecognized extension must be a no-op (no crash, no warning)."""
path = tmp_path / "model.bin"
path.write_bytes(b"")
with caplog.at_level(logging.WARNING):
warn_if_lowbit_quantization(str(path))
assert not any("naive int cast" in r.message for r in caplog.records)
# ====================================================================
# Plan 2026-08-27 (Phase 3): config/weights reconciliation
# ====================================================================
from ComfyUI_VibeVoice.modules.config_detect import WeightsFingerprint
from ComfyUI_VibeVoice.modules.external_loader import reconcile_config
def _cfg_stub(hidden: int, vocab: int):
"""A config object whose decoder_config carries hidden/vocab."""
decoder = MagicMock(spec=["hidden_size", "vocab_size"])
decoder.hidden_size = hidden
decoder.vocab_size = vocab
cfg = MagicMock(spec=["decoder_config"])
cfg.decoder_config = decoder
return cfg
_FP_7B = WeightsFingerprint(
hidden_size=3584, vocab_size=152064,
source_key="model.language_model.embed_tokens.weight",
)
_FP_15B = WeightsFingerprint(
hidden_size=1536, vocab_size=151936,
source_key="model.language_model.embed_tokens.weight",
)
class TestReconcileConfig:
"""Pure decision helper: fingerprint wins over any config source."""
def test_no_fingerprint_keeps_selection(self):
name, changed = reconcile_config("VibeVoice-1.5B", _cfg_stub(1536, 151936), None)
assert (name, changed) == ("VibeVoice-1.5B", False)
def test_matching_fingerprint_keeps_selection(self):
name, changed = reconcile_config(
"VibeVoice-1.5B", _cfg_stub(1536, 151936), _FP_15B
)
assert (name, changed) == ("VibeVoice-1.5B", False)
def test_mismatch_swaps_to_detected_family(self, caplog):
with caplog.at_level(logging.WARNING):
name, changed = reconcile_config(
"VibeVoice-1.5B", _cfg_stub(1536, 151936), _FP_7B
)
assert (name, changed) == ("VibeVoice-7B", True)
assert any("Config mismatch" in r.message for r in caplog.records)
assert any("VibeVoice-7B" in r.message for r in caplog.records)
def test_mismatch_reverse_direction(self, caplog):
with caplog.at_level(logging.WARNING):
name, changed = reconcile_config(
"VibeVoice-7B", _cfg_stub(3584, 152064), _FP_15B
)
assert (name, changed) == ("VibeVoice-1.5B", True)
def test_config_without_fingerprint_keeps_selection(self):
"""Unknown config layouts (no decoder_config) are never swapped."""
cfg = MagicMock(spec=[])
name, changed = reconcile_config("VibeVoice-Realtime-0.5B", cfg, _FP_7B)
assert (name, changed) == ("VibeVoice-Realtime-0.5B", False)
def test_no_warning_when_matching(self, caplog):
with caplog.at_level(logging.WARNING):
reconcile_config("VibeVoice-7B", _cfg_stub(3584, 152064), _FP_7B)
assert not any("Config mismatch" in r.message for r in caplog.records)
class TestConfigFingerprint:
"""config_fingerprint(): decoder_config extraction.
The vendored config classes are mocked in the test env (conftest), so
the packaged JSONs are exercised via the dict form — the same nested
'decoder_config' layout the loader produces.
"""
def test_packaged_7b_config(self):
import json
from ComfyUI_VibeVoice.modules.config_detect import config_fingerprint
path = external_loader._get_packaged_config_path("VibeVoice-7B")
with open(path, "r", encoding="utf-8") as f:
raw = json.load(f)
assert config_fingerprint(raw) == (3584, 152064)
def test_packaged_15b_config(self):
import json
from ComfyUI_VibeVoice.modules.config_detect import config_fingerprint
path = external_loader._get_packaged_config_path("VibeVoice-1.5B")
with open(path, "r", encoding="utf-8") as f:
raw = json.load(f)
assert config_fingerprint(raw) == (1536, 151936)
def test_missing_decoder_config_returns_none(self):
from ComfyUI_VibeVoice.modules.config_detect import config_fingerprint
assert config_fingerprint(MagicMock(spec=[])) is None
def test_dict_form_supported(self):
from ComfyUI_VibeVoice.modules.config_detect import config_fingerprint
raw = {"decoder_config": {"hidden_size": 3584, "vocab_size": 152064}}
assert config_fingerprint(raw) == (3584, 152064)
class TestLoaderReconciliationWiring:
"""load_external_vibevoice_model swaps a contradicting config (D5)."""
@pytest.fixture
def weight_file(self, tmp_path):
weight = tmp_path / "model.safetensors"
weight.write_bytes(b"dummy")
return str(weight)
def _run(self, weight_file, config_name, weights_fp, caplog=None):
"""Run the loader with heavy deps mocked; fingerprint controlled.
resolve_sidecar_config records every config_name it receives and
returns a per-name fake path; _load_config returns a config stub
whose fingerprint matches the requested family.
"""
resolve_calls = []
def fake_resolve(weight_path, name):
resolve_calls.append(name)
return f"/fake/{name}.json"
family_fp = {
"VibeVoice-7B": (3584, 152064),
"VibeVoice-1.5B": (1536, 151936),
}
def fake_load_config(config_path, name):
hidden, vocab = family_fp.get(name, (0, 0))
return _cfg_stub(hidden, vocab)
fake_model = MagicMock()
fake_model.load_state_dict.return_value = ([], [])
fake_model.to.return_value = fake_model
with patch.object(
external_loader, "iter_safetensors_tensors",
return_value=iter([("w", torch.zeros(1))]),
), patch.object(
external_loader.comfy.utils, "load_torch_file",
return_value={"w": torch.zeros(1)},
), patch(
"ComfyUI_VibeVoice.modules.config_detect.fingerprint_weights",
return_value=weights_fp,
), patch.object(
external_loader, "resolve_sidecar_config", side_effect=fake_resolve
), patch.object(
external_loader.VibeVoiceLoader, "_load_config",
side_effect=fake_load_config,
), patch.object(
external_loader.VibeVoiceLoader, "_load_tokenizer",
return_value=MagicMock(),
), patch.object(
external_loader.VibeVoiceLoader, "_load_processor",
return_value=MagicMock(),
), patch.object(
external_loader.VibeVoiceLoader, "_instantiate_model",
return_value=fake_model,
), patch.object(
external_loader, "resolve_sidecar_preprocessor", return_value=""
), patch.object(
external_loader, "resolve_sidecar_tokenizer_dir", return_value="/fake"
), patch.object(
external_loader, "resolve_dtype", return_value=torch.float32
), patch.object(
external_loader, "resolve_attention_mode", side_effect=lambda m, q: m
), patch.object(
external_loader, "get_attn_implementation_for_load", return_value="eager"
), patch.object(
external_loader, "VibeVoiceStreamingConfig", _FakeStreamingCfg
):
result = load_external_vibevoice_model(weight_file, config_name)
return result, resolve_calls
def test_mismatched_selection_is_swapped(self, weight_file, caplog):
with caplog.at_level(logging.WARNING):
result, resolve_calls = self._run(
weight_file, "VibeVoice-1.5B", _FP_7B
)
# First resolution used the selection, second used the detection.
assert resolve_calls == ["VibeVoice-1.5B", "VibeVoice-7B"]
assert result["model_name"] == "VibeVoice-7B"
assert any("Config mismatch" in r.message for r in caplog.records)
def test_matching_selection_not_swapped(self, weight_file, caplog):
with caplog.at_level(logging.WARNING):
result, resolve_calls = self._run(
weight_file, "VibeVoice-1.5B", _FP_15B
)
assert resolve_calls == ["VibeVoice-1.5B"]
assert result["model_name"] == "VibeVoice-1.5B"
assert not any("Config mismatch" in r.message for r in caplog.records)
def test_unfingerprintable_weights_honor_selection(self, weight_file):
result, resolve_calls = self._run(weight_file, "VibeVoice-7B", None)
assert resolve_calls == ["VibeVoice-7B"]
assert result["model_name"] == "VibeVoice-7B"
def test_sidecar_config_mismatch_still_swapped(self, weight_file, caplog):
"""A sidecar-resolved config is not exempt: fingerprint wins (D5)."""
with caplog.at_level(logging.WARNING):
result, resolve_calls = self._run(
weight_file, "VibeVoice-7B", _FP_15B
)
assert resolve_calls == ["VibeVoice-7B", "VibeVoice-1.5B"]
assert result["model_name"] == "VibeVoice-1.5B"
def test_asr_branch_never_runs_detection(self, weight_file):
"""ASR dispatch happens before any fingerprinting."""
with patch(
"ComfyUI_VibeVoice.modules.config_detect.fingerprint_weights"
) as m_fp, patch.object(
external_loader, "load_external_vibevoice_asr_model",
return_value={"model_name": "asr"},
) as m_asr:
result = load_external_vibevoice_model(
weight_file, "VibeVoice-ASR"
)
m_asr.assert_called_once()
m_fp.assert_not_called()
assert result["model_name"] == "asr"
# ====================================================================
# Phase 4 (plan 2026-08-27, D6): Auto-detect semantics
# ====================================================================
from ComfyUI_VibeVoice.modules.external_loader import ( # noqa: E402
AUTO_CONFIG_NAME,
ASR_CONFIG_NAMES,
is_asr_config_name,
resolve_auto_config_name,
)
class TestResolveAutoConfigName:
"""resolve_auto_config_name(): detection -> family or actionable error."""
@pytest.fixture
def weight_file(self, tmp_path):
weight = tmp_path / "model.safetensors"
weight.write_bytes(b"dummy")
return str(weight)
def test_conclusive_fingerprint_returns_family(self, weight_file, caplog):
with caplog.at_level(logging.DEBUG):
name = resolve_auto_config_name(weight_file, weights_fp=_FP_7B)
assert name == "VibeVoice-7B"
assert any("Auto-detected" in r.message for r in caplog.records)
def test_precomputed_fingerprint_skips_recomputation(self, weight_file):
with patch(
"ComfyUI_VibeVoice.modules.config_detect.fingerprint_weights"
) as m_fp:
name = resolve_auto_config_name(weight_file, weights_fp=_FP_15B)
assert name == "VibeVoice-1.5B"
m_fp.assert_not_called()
def test_computes_fingerprint_when_not_provided(self, tmp_path):
weight = tmp_path / "model.gguf"
weight.write_bytes(b"dummy")
sentinel_reader = object()
with patch(
"ComfyUI_VibeVoice.modules.config_detect.fingerprint_weights",
return_value=_FP_7B,
) as m_fp:
name = resolve_auto_config_name(
str(weight), gguf_reader=sentinel_reader
)
assert name == "VibeVoice-7B"
# A caller-provided GGUF reader must be reused (no re-open, D3).
assert m_fp.call_args[1]["gguf_reader"] is sentinel_reader
def test_inconclusive_raises_actionable_error(self, weight_file):
with patch(
"ComfyUI_VibeVoice.modules.config_detect.fingerprint_weights",
return_value=None,
):
with pytest.raises(ValueError, match="could not determine"):
resolve_auto_config_name(weight_file)
def test_error_message_offers_explicit_choices_only(self, weight_file):
with patch(
"ComfyUI_VibeVoice.modules.config_detect.fingerprint_weights",
return_value=None,
):
with pytest.raises(ValueError) as excinfo:
resolve_auto_config_name(weight_file)
message = str(excinfo.value)
offered = message.split("one of:")[1]
for option in EXTERNAL_CONFIG_OPTIONS:
if option != AUTO_CONFIG_NAME:
assert option in offered
# The sentinel is never offered as an explicit choice.
assert AUTO_CONFIG_NAME not in offered
def test_unknown_family_fingerprint_raises(self, weight_file):
foreign = WeightsFingerprint(
hidden_size=999, vocab_size=888, source_key="k"
)
with pytest.raises(ValueError, match="could not determine"):
resolve_auto_config_name(weight_file, weights_fp=foreign)
def test_unmatched_fingerprint_reports_dimensions_and_explicit_selection(self, weight_file):
foreign = WeightsFingerprint(
hidden_size=999, vocab_size=888, source_key="k"
)
with pytest.raises(ValueError) as excinfo:
resolve_auto_config_name(weight_file, weights_fp=foreign)
message = str(excinfo.value)
assert "hidden=999" in message
assert "vocab=888" in message
assert "explicitly" in message
assert "not guessed" in message
class TestAutoDetectLoaderSemantics:
"""load_external_vibevoice_model resolves AUTO_CONFIG_NAME (Step 4.2)."""
@pytest.fixture
def weight_file(self, tmp_path):
weight = tmp_path / "model.safetensors"
weight.write_bytes(b"dummy")
return str(weight)
def test_auto_adopts_detected_family(self, weight_file, caplog):
harness = TestLoaderReconciliationWiring()
with caplog.at_level(logging.DEBUG):
result, resolve_calls = harness._run(
weight_file, AUTO_CONFIG_NAME, _FP_7B
)
assert resolve_calls == ["VibeVoice-7B"]
assert result["model_name"] == "VibeVoice-7B"
assert any("Auto-detected" in r.message for r in caplog.records)
def test_auto_inconclusive_fails_before_state_dict_load(self, weight_file):
"""Fail-fast: no heavy load work before the actionable error (D6)."""
with patch.object(
external_loader.comfy.utils, "load_torch_file",
return_value={"w": torch.zeros(1)},
) as m_load, patch(
"ComfyUI_VibeVoice.modules.config_detect.fingerprint_weights",
return_value=None,
):
with pytest.raises(ValueError, match="could not determine"):
load_external_vibevoice_model(weight_file, AUTO_CONFIG_NAME)
m_load.assert_not_called()
def test_auto_never_dispatches_asr(self):
"""Auto is TTS-branch only; ASR remains an explicit selection."""
assert AUTO_CONFIG_NAME not in ASR_CONFIG_NAMES
assert is_asr_config_name(AUTO_CONFIG_NAME) is False