Add a dedicated loader node that loads VibeVoice checkpoints from user-provided .safetensors/.pt files in models/diffusion_models, since ComfyUI's stock Load Diffusion Model cannot detect VibeVoice architecture. The loader emits a VIBEVOICE_MODEL custom type bundle (model/processor/config/state_dict) consumed by an optional external_model input on the TTS, Realtime, and ASR nodes.
- modules/custom_types.py: VibeVoiceModel = io.Custom("VIBEVOICE_MODEL")
- modules/external_loader.py: sidecar config/preprocessor/tokenizer resolution + load_external_vibevoice_model() (TTS/streaming) and load_external_vibevoice_asr_model() (ASR); CPU-first load, in-memory state-dict injection, dtype cast, optional 4-bit quant (TTS only), SageAttention
- nodes/external_loader_node.py: VibeVoiceLoadExternalModel node
- modules/generation.py: ExternalVibeVoiceModelHandler + load_vibevoice_from_external()
- modules/asr_generation.py: ExternalVibeVoiceASRModelHandler + load_asr_from_external()
- nodes/tts_node.py, realtime_node.py, asr_node.py: optional external_model input with kind guards (streaming/ASR/TTS mismatch rejection)
- README: Loading External Models section + v2.2.0 changelog
- example_workflows/VibeVoice_external_model_example.json
- docs/plans/2026-08-16-external-model-input.md (status: COMPLETED)
Tests: +122 new/extended (test_custom_types, test_external_loader, test_external_loader_node, plus extensions to node_schema, generation, asr_generation, realtime_node, asr_node, patcher_behavioral, integration, docs_consistency, workflow, imports, extension). Full suite: 687 passed, 5 pre-existing failures, 4 skipped.
162 lines
6.5 KiB
Python
162 lines
6.5 KiB
Python
"""Tests for nodes/external_loader_node.py - Load VibeVoice Model node."""
|
|
|
|
import pytest
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
from comfy_api.latest import io
|
|
|
|
from ComfyUI_VibeVoice.nodes.external_loader_node import VibeVoiceExternalLoaderNode
|
|
|
|
|
|
class TestExternalLoaderNodeSchema:
|
|
"""Test the VibeVoiceExternalLoaderNode schema."""
|
|
|
|
def test_node_schema_id(self):
|
|
"""define_schema().node_id == 'VibeVoiceLoadExternalModel'."""
|
|
schema = VibeVoiceExternalLoaderNode.define_schema()
|
|
assert schema.node_id == "VibeVoiceLoadExternalModel"
|
|
|
|
def test_node_display_name(self):
|
|
"""Display name is 'Load VibeVoice Model'."""
|
|
schema = VibeVoiceExternalLoaderNode.define_schema()
|
|
assert schema.display_name == "Load VibeVoice Model"
|
|
|
|
def test_node_has_model_file_input(self):
|
|
"""'model_file' is in the input ids."""
|
|
schema = VibeVoiceExternalLoaderNode.define_schema()
|
|
input_ids = [inp.id for inp in schema.inputs]
|
|
assert "model_file" in input_ids
|
|
|
|
def test_node_has_config_name_input(self):
|
|
"""'config_name' is in the input ids."""
|
|
schema = VibeVoiceExternalLoaderNode.define_schema()
|
|
input_ids = [inp.id for inp in schema.inputs]
|
|
assert "config_name" in input_ids
|
|
|
|
def test_node_has_attention_mode_input(self):
|
|
"""'attention_mode' is in the input ids."""
|
|
schema = VibeVoiceExternalLoaderNode.define_schema()
|
|
input_ids = [inp.id for inp in schema.inputs]
|
|
assert "attention_mode" in input_ids
|
|
|
|
def test_node_has_quantize_input(self):
|
|
"""'quantize_llm_4bit' is in the input ids."""
|
|
schema = VibeVoiceExternalLoaderNode.define_schema()
|
|
input_ids = [inp.id for inp in schema.inputs]
|
|
assert "quantize_llm_4bit" in input_ids
|
|
|
|
def test_node_has_dtype_input(self):
|
|
"""'dtype' is in the input ids."""
|
|
schema = VibeVoiceExternalLoaderNode.define_schema()
|
|
input_ids = [inp.id for inp in schema.inputs]
|
|
assert "dtype" in input_ids
|
|
|
|
def test_node_output_is_vibevoice_model(self):
|
|
"""Output type string is 'VIBEVOICE_MODEL'."""
|
|
schema = VibeVoiceExternalLoaderNode.define_schema()
|
|
assert len(schema.outputs) == 1
|
|
assert schema.outputs[0].io_type == "VIBEVOICE_MODEL"
|
|
|
|
def test_node_category(self):
|
|
"""Node category is 'audio/tts'."""
|
|
assert VibeVoiceExternalLoaderNode.CATEGORY == "audio/tts"
|
|
|
|
|
|
class TestExternalLoaderNodeExecute:
|
|
"""Test the VibeVoiceExternalLoaderNode.execute() method."""
|
|
|
|
def test_node_execute_calls_load_external(self):
|
|
"""execute() calls load_external_vibevoice_model with correct kwargs."""
|
|
fake_bundle = {"model": MagicMock(), "processor": MagicMock()}
|
|
|
|
with patch(
|
|
"ComfyUI_VibeVoice.nodes.external_loader_node.load_external_vibevoice_model",
|
|
return_value=fake_bundle,
|
|
) as mock_load, patch(
|
|
"ComfyUI_VibeVoice.nodes.external_loader_node.folder_paths.get_full_path_or_raise",
|
|
return_value="/fake/path/model.safetensors",
|
|
):
|
|
result = VibeVoiceExternalLoaderNode.execute(
|
|
model_file="model.safetensors",
|
|
config_name="VibeVoice-1.5B",
|
|
attention_mode="sdpa",
|
|
quantize_llm_4bit=False,
|
|
dtype="auto",
|
|
)
|
|
|
|
mock_load.assert_called_once()
|
|
call_kwargs = mock_load.call_args[1]
|
|
assert call_kwargs["weight_path"] == "/fake/path/model.safetensors"
|
|
assert call_kwargs["config_name"] == "VibeVoice-1.5B"
|
|
assert call_kwargs["attention_mode"] == "sdpa"
|
|
assert call_kwargs["use_llm_4bit"] is False
|
|
assert call_kwargs["dtype_str"] == "auto"
|
|
|
|
def test_node_execute_resolves_path_via_folder_paths(self):
|
|
"""execute() resolves the path via folder_paths.get_full_path_or_raise."""
|
|
fake_bundle = {"model": MagicMock()}
|
|
|
|
with patch(
|
|
"ComfyUI_VibeVoice.nodes.external_loader_node.load_external_vibevoice_model",
|
|
return_value=fake_bundle,
|
|
), patch(
|
|
"ComfyUI_VibeVoice.nodes.external_loader_node.folder_paths.get_full_path_or_raise",
|
|
return_value="/fake/path/model.safetensors",
|
|
) as mock_resolve:
|
|
VibeVoiceExternalLoaderNode.execute(
|
|
model_file="model.safetensors",
|
|
config_name="VibeVoice-1.5B",
|
|
attention_mode="sdpa",
|
|
quantize_llm_4bit=False,
|
|
dtype="auto",
|
|
)
|
|
|
|
mock_resolve.assert_called_once_with("diffusion_models", "model.safetensors")
|
|
|
|
def test_node_execute_returns_node_output(self):
|
|
"""execute() returns an io.NodeOutput wrapping the bundle dict."""
|
|
fake_bundle = {"model": MagicMock(), "processor": MagicMock()}
|
|
|
|
with patch(
|
|
"ComfyUI_VibeVoice.nodes.external_loader_node.load_external_vibevoice_model",
|
|
return_value=fake_bundle,
|
|
), patch(
|
|
"ComfyUI_VibeVoice.nodes.external_loader_node.folder_paths.get_full_path_or_raise",
|
|
return_value="/fake/path/model.safetensors",
|
|
):
|
|
result = VibeVoiceExternalLoaderNode.execute(
|
|
model_file="model.safetensors",
|
|
config_name="VibeVoice-1.5B",
|
|
attention_mode="sdpa",
|
|
quantize_llm_4bit=False,
|
|
dtype="auto",
|
|
)
|
|
|
|
assert isinstance(result, io.NodeOutput)
|
|
# The bundle should be the first output value
|
|
assert result[0] is fake_bundle
|
|
|
|
def test_node_execute_passes_quantize_flag(self):
|
|
"""execute() passes quantize_llm_4bit=True through to the loader."""
|
|
fake_bundle = {"model": MagicMock()}
|
|
|
|
with patch(
|
|
"ComfyUI_VibeVoice.nodes.external_loader_node.load_external_vibevoice_model",
|
|
return_value=fake_bundle,
|
|
) as mock_load, patch(
|
|
"ComfyUI_VibeVoice.nodes.external_loader_node.folder_paths.get_full_path_or_raise",
|
|
return_value="/fake/path/model.safetensors",
|
|
):
|
|
VibeVoiceExternalLoaderNode.execute(
|
|
model_file="model.safetensors",
|
|
config_name="VibeVoice-Large",
|
|
attention_mode="eager",
|
|
quantize_llm_4bit=True,
|
|
dtype="bf16",
|
|
)
|
|
|
|
call_kwargs = mock_load.call_args[1]
|
|
assert call_kwargs["use_llm_4bit"] is True
|
|
assert call_kwargs["dtype_str"] == "bf16"
|
|
assert call_kwargs["config_name"] == "VibeVoice-Large"
|