Loading problem closed for every model type. Live, 2026-09-30 (RTX 4070 Ti SUPER 16GB), 1.5B bf16 and 7B fp8 both load straight into VRAM with +0.69 GB / +1.45 GB of machine RAM and no sustained SSD traffic. What was wrong -------------- The per-tensor placement used `Tensor.to(cuda)` on a memory-mapped view. That is a HOST-side read: the copy engine faults every page in through the CPU, at 0.55 GB/s and +2.18 GB of machine RAM per 2 GB (report 2026-09-30 section F6). Every route paid it, on every model type. The 7B fp8 load had looked SSD-free only because pass 1's `safe_open` was committing the whole file as private memory first, so those host-side copies were served from RAM. The 19.7 -> 28 -> 20 GB spike was the receipt for that. What changed ------------ * base_loader.place_tensor_on_device(): one placement helper for every route. With an aimdo mapping (main.py sets aimdo_enabled at startup) it uses core's own `read_tensor_file_slice_into` to DMA the file byte range straight into a preallocated CUDA tensor -- the same primitive core already uses to page weights into VRAM, measured at 2.6 GB/s with the page cache left clean. Falls back to `.to()` when core declines or when the aimdo native library raises; a one-time `[vvload]` log line confirms the DMA is live. * One streaming assign for all routes. `_stream_apply_dense` gained `target_device`, and a new `_stream_apply_dense_safetensors` is the dense twin of `_stream_apply_safetensors`. The TTS and ASR dense branches no longer build a whole CPU state dict and call `model.to(cuda)`; GGUF and the internal official-model loader joined the same path. * `read_safetensors_tensors_by_name()`: pass-1 dequant scales are read from the byte ranges the safetensors header already records, instead of `safe_open`. That drops a genuine ~1x-file private commit. * select_patcher_class() always returns the standard ModelPatcher, so core can full-load instead of trapping weights in VBAR for an autoregressive model. Tests ----- 5 new tests for DMA routing and both fallbacks; load-path tests updated to the new read seam. Targeted run: 373 passed / 34 failed, all 34 pre-existing and asserting the ModelPatcherDynamic/vbar architecture this change removes (26 in test_dynamic_patcher_selection.py). Full suite not run, per user rule.
952 lines
42 KiB
Python
952 lines
42 KiB
Python
"""Tests for modules/loader.py - Model loading and caching."""
|
|
|
|
import os
|
|
import json
|
|
import torch
|
|
import pytest
|
|
from unittest.mock import patch, MagicMock
|
|
|
|
from ComfyUI_VibeVoice.modules.loader import (
|
|
VibeVoiceModelHandler,
|
|
VibeVoiceLoader,
|
|
LOADED_MODELS_CACHE,
|
|
cleanup_old_models,
|
|
)
|
|
from ComfyUI_VibeVoice.modules.base_loader import BaseVibeVoiceLoader
|
|
from ComfyUI_VibeVoice.modules.model_info import AVAILABLE_VIBEVOICE_MODELS, MODEL_CONFIGS
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def mock_tts_folder():
|
|
"""Register a 'tts' folder with folder_paths for tests."""
|
|
import folder_paths
|
|
tts_path = os.path.join(folder_paths.models_dir, "tts")
|
|
if "tts" not in folder_paths.folder_names_and_paths:
|
|
supported_exts = folder_paths.supported_pt_extensions.union({".safetensors", ".json"})
|
|
folder_paths.folder_names_and_paths["tts"] = ([tts_path], supported_exts)
|
|
yield
|
|
# Cleanup not needed — folder_paths is global
|
|
|
|
|
|
class TestVibeVoiceModelHandler:
|
|
"""Test VibeVoiceModelHandler class."""
|
|
|
|
def test_handler_init(self):
|
|
handler = VibeVoiceModelHandler("VibeVoice-1.5B", attention_mode="sdpa", use_llm_4bit=False)
|
|
assert handler.model_pack_name == "VibeVoice-1.5B"
|
|
assert handler.attention_mode == "sdpa"
|
|
assert handler.use_llm_4bit is False
|
|
assert handler.model is None
|
|
assert handler.processor is None
|
|
|
|
def test_handler_has_device_attribute(self):
|
|
"""Handler must have a device attribute for ComfyUI's ModelPatcher."""
|
|
handler = VibeVoiceModelHandler("VibeVoice-1.5B")
|
|
assert hasattr(handler, "device")
|
|
# Initially None — ModelPatcher.__init__ will set it to offload_device
|
|
|
|
def test_handler_cache_key(self):
|
|
handler = VibeVoiceModelHandler("VibeVoice-1.5B", attention_mode="sdpa", use_llm_4bit=False)
|
|
assert handler.cache_key == "VibeVoice-1.5B_attn_sdpa_q4_0"
|
|
|
|
def test_handler_cache_key_4bit(self):
|
|
handler = VibeVoiceModelHandler("VibeVoice-7B", attention_mode="eager", use_llm_4bit=True)
|
|
assert handler.cache_key == "VibeVoice-7B_attn_eager_q4_1"
|
|
|
|
def test_handler_size_calculation(self):
|
|
handler = VibeVoiceModelHandler("VibeVoice-1.5B")
|
|
assert handler.size == int(3.0 * (1024**3))
|
|
|
|
def test_handler_size_large(self):
|
|
handler = VibeVoiceModelHandler("VibeVoice-7B")
|
|
assert handler.size == int(17.4 * (1024**3))
|
|
|
|
def test_handler_is_torch_module(self):
|
|
handler = VibeVoiceModelHandler("VibeVoice-1.5B")
|
|
assert isinstance(handler, torch.nn.Module)
|
|
|
|
|
|
class TestHandlerSizeRefinement:
|
|
"""Plan 2026-08-18, Phase 6 (D7/RC-7): handler.size is refined from the
|
|
real parameters after load, replacing the config-based size_gb estimate."""
|
|
|
|
def test_handler_size_refined_after_load(self):
|
|
"""After load_model, handler.size equals the real parameter byte total."""
|
|
handler = VibeVoiceModelHandler("VibeVoice-1.5B")
|
|
# Config-based estimate before load.
|
|
assert handler.size == int(3.0 * (1024**3))
|
|
|
|
# A tiny model with a known parameter byte total.
|
|
tiny = torch.nn.Linear(16, 16, bias=False) # 16*16 = 256 floats
|
|
expected_bytes = 256 * tiny.weight.element_size()
|
|
|
|
with patch.object(VibeVoiceLoader, "load_model", return_value=(tiny, MagicMock())):
|
|
handler.load_model(torch.device("cpu"), attention_mode="sdpa")
|
|
|
|
assert handler.size == expected_bytes
|
|
|
|
def test_handler_size_kept_when_params_empty(self):
|
|
"""If the model has no parameters, the config estimate is kept."""
|
|
handler = VibeVoiceModelHandler("VibeVoice-1.5B")
|
|
original_size = handler.size
|
|
|
|
empty = torch.nn.Module() # no parameters
|
|
|
|
with patch.object(VibeVoiceLoader, "load_model", return_value=(empty, MagicMock())):
|
|
handler.load_model(torch.device("cpu"), attention_mode="sdpa")
|
|
|
|
assert handler.size == original_size
|
|
|
|
def test_handler_size_kept_on_exception(self):
|
|
"""If parameter iteration raises, the config estimate is kept."""
|
|
handler = VibeVoiceModelHandler("VibeVoice-1.5B")
|
|
original_size = handler.size
|
|
|
|
broken = MagicMock()
|
|
broken.parameters.side_effect = RuntimeError("boom")
|
|
|
|
with patch.object(VibeVoiceLoader, "load_model", return_value=(broken, MagicMock())):
|
|
handler.load_model(torch.device("cpu"), attention_mode="sdpa")
|
|
|
|
assert handler.size == original_size
|
|
|
|
|
|
class TestVibeVoiceLoaderResolvePaths:
|
|
"""Test VibeVoiceLoader._resolve_model_paths."""
|
|
|
|
def test_resolve_local_dir(self, tmp_path):
|
|
model_dir = tmp_path / "MyLocalModel"
|
|
model_dir.mkdir()
|
|
with patch("ComfyUI_VibeVoice.modules.loader.AVAILABLE_VIBEVOICE_MODELS", {
|
|
"MyLocalModel": {"type": "local_dir", "path": str(model_dir)}
|
|
}):
|
|
model_path, config_path, preproc_path, tokenizer_dir = \
|
|
VibeVoiceLoader._resolve_model_paths("MyLocalModel")
|
|
assert model_path == str(model_dir)
|
|
assert config_path == str(model_dir / "config.json")
|
|
|
|
def test_resolve_standalone(self, tmp_path):
|
|
model_file = tmp_path / "model.safetensors"
|
|
with patch("ComfyUI_VibeVoice.modules.loader.AVAILABLE_VIBEVOICE_MODELS", {
|
|
"model": {"type": "standalone", "path": str(model_file)}
|
|
}):
|
|
model_path, config_path, preproc_path, tokenizer_dir = \
|
|
VibeVoiceLoader._resolve_model_paths("model")
|
|
assert model_path is None
|
|
assert config_path.endswith(".config.json")
|
|
assert tokenizer_dir == str(tmp_path)
|
|
|
|
def test_resolve_official(self, tmp_path):
|
|
"""Test official model path resolution (with download mocked)."""
|
|
with patch("ComfyUI_VibeVoice.modules.loader.AVAILABLE_VIBEVOICE_MODELS", {
|
|
"TestModel": {"type": "official", "repo_id": "test/repo"}
|
|
}), patch("ComfyUI_VibeVoice.modules.loader.folder_paths") as mock_fp, \
|
|
patch("huggingface_hub.snapshot_download") as mock_dl, \
|
|
patch("os.path.exists", return_value=True):
|
|
mock_fp.get_folder_paths.return_value = [str(tmp_path)]
|
|
model_path, config_path, preproc_path, tokenizer_dir = \
|
|
VibeVoiceLoader._resolve_model_paths("TestModel")
|
|
assert "TestModel" in model_path
|
|
assert config_path.endswith("config.json")
|
|
# snapshot_download should NOT be called since we mocked os.path.exists to True
|
|
mock_dl.assert_not_called()
|
|
|
|
|
|
class TestVibeVoiceLoaderInstantiateModel:
|
|
"""Test VibeVoiceLoader._instantiate_model."""
|
|
|
|
def test_instantiate_model_non_streaming(self):
|
|
"""Test that non-streaming model is instantiated correctly."""
|
|
config = MagicMock()
|
|
config.decoder_config = MagicMock()
|
|
config.torch_dtype = None
|
|
with patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceForConditionalGeneration") as mock_cls:
|
|
mock_model = MagicMock()
|
|
mock_cls.return_value = mock_model
|
|
|
|
model = VibeVoiceLoader._instantiate_model(
|
|
config=config,
|
|
is_streaming=False,
|
|
attn_implementation="sdpa",
|
|
final_load_dtype=torch.float16,
|
|
)
|
|
|
|
mock_cls.assert_called_once_with(config)
|
|
assert model == mock_model
|
|
# Verify attn_implementation was set on decoder_config
|
|
assert config.decoder_config._attn_implementation == "sdpa"
|
|
|
|
def test_instantiate_model_streaming(self):
|
|
"""Test that streaming model uses the correct class."""
|
|
config = MagicMock()
|
|
config.decoder_config = MagicMock()
|
|
config.torch_dtype = None
|
|
with patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceStreamingForConditionalGenerationInference") as mock_cls:
|
|
mock_model = MagicMock()
|
|
mock_cls.return_value = mock_model
|
|
|
|
model = VibeVoiceLoader._instantiate_model(
|
|
config=config,
|
|
is_streaming=True,
|
|
attn_implementation="sdpa",
|
|
final_load_dtype=torch.float16,
|
|
)
|
|
|
|
mock_cls.assert_called_once_with(config)
|
|
assert model == mock_model
|
|
|
|
def test_instantiate_model_sets_dtype_on_config(self):
|
|
"""Test that dtype is set on config."""
|
|
config = MagicMock()
|
|
config.decoder_config = MagicMock()
|
|
config.torch_dtype = None
|
|
|
|
with patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceForConditionalGeneration"):
|
|
VibeVoiceLoader._instantiate_model(
|
|
config=config,
|
|
is_streaming=False,
|
|
attn_implementation="eager",
|
|
final_load_dtype=torch.bfloat16,
|
|
)
|
|
|
|
assert config.torch_dtype == torch.bfloat16
|
|
assert config.decoder_config.torch_dtype == torch.bfloat16
|
|
|
|
def test_instantiate_model_real_config_no_deprecation(self):
|
|
"""With a REAL transformers PretrainedConfig (v5: torch_dtype is a
|
|
deprecated property), _instantiate_model must record the dtype on the
|
|
canonical ``dtype`` attribute and emit no torch_dtype deprecation.
|
|
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("ComfyUI_VibeVoice.modules.loader.VibeVoiceForConditionalGeneration"):
|
|
VibeVoiceLoader._instantiate_model(
|
|
config=config,
|
|
is_streaming=False,
|
|
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"
|
|
|
|
|
|
class TestResolveCheckpointPath:
|
|
"""Test VibeVoiceLoader._resolve_checkpoint_path."""
|
|
|
|
def test_standalone_returns_path_directly(self, tmp_path):
|
|
"""Standalone model returns the path directly with is_sharded=False."""
|
|
ckpt_file = tmp_path / "model.safetensors"
|
|
ckpt_file.write_text("dummy")
|
|
model_info = {"path": str(ckpt_file)}
|
|
result_path, is_sharded = VibeVoiceLoader._resolve_checkpoint_path(
|
|
model_path=None, model_type="standalone", model_info=model_info
|
|
)
|
|
assert result_path == str(ckpt_file)
|
|
assert is_sharded is False
|
|
|
|
def test_standalone_raises_if_file_not_found(self):
|
|
"""Standalone model raises FileNotFoundError if file doesn't exist."""
|
|
model_info = {"path": "/nonexistent/path.safetensors"}
|
|
with pytest.raises(FileNotFoundError, match="Standalone checkpoint not found"):
|
|
VibeVoiceLoader._resolve_checkpoint_path(
|
|
model_path=None, model_type="standalone", model_info=model_info
|
|
)
|
|
|
|
def test_official_single_safetensors(self, tmp_path):
|
|
"""Official model with single model.safetensors file."""
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
(model_dir / "model.safetensors").write_text("dummy")
|
|
result_path, is_sharded = VibeVoiceLoader._resolve_checkpoint_path(
|
|
model_path=str(model_dir), model_type="official", model_info={}
|
|
)
|
|
assert result_path == str(model_dir / "model.safetensors")
|
|
assert is_sharded is False
|
|
|
|
def test_official_sharded_safetensors(self, tmp_path):
|
|
"""Official model with sharded safetensors (index.json present)."""
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
(model_dir / "model.safetensors.index.json").write_text('{"weight_map": {}}')
|
|
result_path, is_sharded = VibeVoiceLoader._resolve_checkpoint_path(
|
|
model_path=str(model_dir), model_type="official", model_info={}
|
|
)
|
|
assert result_path == str(model_dir / "model.safetensors.index.json")
|
|
assert is_sharded is True
|
|
|
|
def test_official_single_pytorch_bin(self, tmp_path):
|
|
"""Official model with single pytorch_model.bin file."""
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
(model_dir / "pytorch_model.bin").write_text("dummy")
|
|
result_path, is_sharded = VibeVoiceLoader._resolve_checkpoint_path(
|
|
model_path=str(model_dir), model_type="official", model_info={}
|
|
)
|
|
assert result_path == str(model_dir / "pytorch_model.bin")
|
|
assert is_sharded is False
|
|
|
|
def test_official_sharded_pytorch_bin(self, tmp_path):
|
|
"""Official model with sharded pytorch_model (index.json present)."""
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
(model_dir / "pytorch_model.bin.index.json").write_text('{"weight_map": {}}')
|
|
result_path, is_sharded = VibeVoiceLoader._resolve_checkpoint_path(
|
|
model_path=str(model_dir), model_type="official", model_info={}
|
|
)
|
|
assert result_path == str(model_dir / "pytorch_model.bin.index.json")
|
|
assert is_sharded is True
|
|
|
|
def test_official_prioritizes_safetensors_over_bin(self, tmp_path):
|
|
"""When both safetensors and bin exist, safetensors is preferred."""
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
(model_dir / "model.safetensors").write_text("dummy")
|
|
(model_dir / "pytorch_model.bin").write_text("dummy")
|
|
result_path, is_sharded = VibeVoiceLoader._resolve_checkpoint_path(
|
|
model_path=str(model_dir), model_type="official", model_info={}
|
|
)
|
|
assert result_path == str(model_dir / "model.safetensors")
|
|
assert is_sharded is False
|
|
|
|
def test_official_no_checkpoint_raises(self, tmp_path):
|
|
"""Raises FileNotFoundError when no checkpoint file found."""
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
with pytest.raises(FileNotFoundError, match="No checkpoint file found"):
|
|
VibeVoiceLoader._resolve_checkpoint_path(
|
|
model_path=str(model_dir), model_type="official", model_info={}
|
|
)
|
|
|
|
def test_official_dir_not_found_raises(self):
|
|
"""Raises FileNotFoundError when model directory doesn't exist."""
|
|
with pytest.raises(FileNotFoundError, match="Model directory not found"):
|
|
VibeVoiceLoader._resolve_checkpoint_path(
|
|
model_path="/nonexistent/dir", model_type="official", model_info={}
|
|
)
|
|
|
|
|
|
class TestLoadShardedStateDict:
|
|
"""Test VibeVoiceLoader._load_sharded_state_dict."""
|
|
|
|
def test_load_sharded_merges_shards(self, tmp_path):
|
|
"""Test that sharded state dict is loaded and merged correctly."""
|
|
# Create shard files (dummy content)
|
|
(tmp_path / "model-00001-of-00002.safetensors").write_text("dummy1")
|
|
(tmp_path / "model-00002-of-00002.safetensors").write_text("dummy2")
|
|
|
|
# Create index file
|
|
index_data = {
|
|
"weight_map": {
|
|
"layer1.weight": "model-00001-of-00002.safetensors",
|
|
"layer1.bias": "model-00001-of-00002.safetensors",
|
|
"layer2.weight": "model-00002-of-00002.safetensors",
|
|
"layer2.bias": "model-00002-of-00002.safetensors",
|
|
}
|
|
}
|
|
index_path = tmp_path / "model.safetensors.index.json"
|
|
index_path.write_text(json.dumps(index_data))
|
|
|
|
# Mock load_torch_file to return different dicts per shard
|
|
shard1_data = {"layer1.weight": "tensor1", "layer1.bias": "tensor2"}
|
|
shard2_data = {"layer2.weight": "tensor3", "layer2.bias": "tensor4"}
|
|
|
|
def mock_load(path, device=None):
|
|
if "00001" in path:
|
|
return shard1_data
|
|
elif "00002" in path:
|
|
return shard2_data
|
|
return {}
|
|
|
|
with patch("ComfyUI_VibeVoice.modules.loader.comfy.utils.load_torch_file", side_effect=mock_load):
|
|
result = VibeVoiceLoader._load_sharded_state_dict(
|
|
index_path=str(index_path),
|
|
model_dir=str(tmp_path),
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
assert "layer1.weight" in result
|
|
assert "layer1.bias" in result
|
|
assert "layer2.weight" in result
|
|
assert "layer2.bias" in result
|
|
assert result["layer1.weight"] == "tensor1"
|
|
assert result["layer2.weight"] == "tensor3"
|
|
|
|
def test_load_sharded_empty_weight_map_raises(self, tmp_path):
|
|
"""Test that empty weight_map raises ValueError."""
|
|
index_path = tmp_path / "model.safetensors.index.json"
|
|
index_path.write_text('{"weight_map": {}}')
|
|
|
|
with pytest.raises(ValueError, match="empty weight_map"):
|
|
VibeVoiceLoader._load_sharded_state_dict(
|
|
index_path=str(index_path),
|
|
model_dir=str(tmp_path),
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
def test_load_sharded_missing_shard_raises(self, tmp_path):
|
|
"""Test that missing shard file raises FileNotFoundError."""
|
|
index_data = {
|
|
"weight_map": {
|
|
"layer1.weight": "model-00001-of-00002.safetensors",
|
|
"layer2.weight": "model-00002-of-00002.safetensors",
|
|
}
|
|
}
|
|
index_path = tmp_path / "model.safetensors.index.json"
|
|
index_path.write_text(json.dumps(index_data))
|
|
|
|
# Only create shard 1, not shard 2
|
|
(tmp_path / "model-00001-of-00002.safetensors").write_text("dummy")
|
|
|
|
with patch("ComfyUI_VibeVoice.modules.loader.comfy.utils.load_torch_file") as mock_load:
|
|
mock_load.return_value = {}
|
|
with pytest.raises(FileNotFoundError, match="Shard file not found"):
|
|
VibeVoiceLoader._load_sharded_state_dict(
|
|
index_path=str(index_path),
|
|
model_dir=str(tmp_path),
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
|
|
class _TwoParam(torch.nn.Module):
|
|
"""Tiny real module for streaming-assign tests (no mocks).
|
|
|
|
``config`` carries the (untied) gate read by the post-assign fixups.
|
|
"""
|
|
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.layer1 = torch.nn.Linear(2, 2, bias=False)
|
|
self.layer2 = torch.nn.Linear(2, 2, bias=False)
|
|
self.config = MagicMock()
|
|
self.config.decoder_config.tie_word_embeddings = False
|
|
self.config.tie_word_embeddings = False
|
|
|
|
|
|
class TestVibeVoiceLoaderLoadStateDict:
|
|
"""Test VibeVoiceLoader._load_state_dict_into_model (streaming assign).
|
|
|
|
Plan 2026-08-28: dense checkpoints (sharded or single-file) are applied
|
|
per-tensor via ``_stream_apply_dense`` — no merged state dict, and every
|
|
tensor is cloned into private memory, severing the checkpoint's file
|
|
mapping (the mmap ghost behind the sharded 7B RAM bloat + Pin errors).
|
|
These tests use REAL safetensors files so the mmap semantics are
|
|
exercised end-to-end.
|
|
"""
|
|
|
|
def test_load_state_dict_single_safetensors(self, tmp_path):
|
|
"""Single-file safetensors: streamed per-tensor into the model."""
|
|
from safetensors.torch import save_file
|
|
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
w1 = torch.arange(4, dtype=torch.float32).reshape(2, 2)
|
|
w2 = torch.full((2, 2), 7.0)
|
|
save_file(
|
|
{"layer1.weight": w1, "layer2.weight": w2},
|
|
str(model_dir / "model.safetensors"),
|
|
)
|
|
|
|
model = _TwoParam()
|
|
result = VibeVoiceLoader._load_state_dict_into_model(
|
|
model=model,
|
|
model_path=str(model_dir),
|
|
model_type="official",
|
|
model_info={},
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
assert result is model
|
|
assert torch.equal(model.layer1.weight.data, w1)
|
|
assert torch.equal(model.layer2.weight.data, w2)
|
|
assert model.layer1.weight.device.type == "cpu"
|
|
|
|
def test_load_state_dict_standalone(self, tmp_path):
|
|
"""Standalone safetensors file: streamed per-tensor into the model."""
|
|
from safetensors.torch import save_file
|
|
|
|
ckpt_file = tmp_path / "checkpoint.safetensors"
|
|
w1 = torch.ones(2, 2)
|
|
w2 = torch.full((2, 2), 2.0)
|
|
save_file({"layer1.weight": w1, "layer2.weight": w2}, str(ckpt_file))
|
|
|
|
model = _TwoParam()
|
|
result = VibeVoiceLoader._load_state_dict_into_model(
|
|
model=model,
|
|
model_path=None,
|
|
model_type="standalone",
|
|
model_info={"path": str(ckpt_file)},
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
assert result is model
|
|
assert torch.equal(model.layer1.weight.data, w1)
|
|
assert torch.equal(model.layer2.weight.data, w2)
|
|
|
|
def test_load_state_dict_sharded(self, tmp_path):
|
|
"""Sharded safetensors: every shard streams into the model."""
|
|
from safetensors.torch import save_file
|
|
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
|
|
w1 = torch.ones(2, 2)
|
|
w2 = torch.full((2, 2), 2.0)
|
|
save_file({"layer1.weight": w1}, str(model_dir / "model-00001-of-00002.safetensors"))
|
|
save_file({"layer2.weight": w2}, str(model_dir / "model-00002-of-00002.safetensors"))
|
|
index_data = {
|
|
"weight_map": {
|
|
"layer1.weight": "model-00001-of-00002.safetensors",
|
|
"layer2.weight": "model-00002-of-00002.safetensors",
|
|
}
|
|
}
|
|
(model_dir / "model.safetensors.index.json").write_text(json.dumps(index_data))
|
|
|
|
model = _TwoParam()
|
|
result = VibeVoiceLoader._load_state_dict_into_model(
|
|
model=model,
|
|
model_path=str(model_dir),
|
|
model_type="official",
|
|
model_info={},
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
assert result is model
|
|
assert torch.equal(model.layer1.weight.data, w1)
|
|
assert torch.equal(model.layer2.weight.data, w2)
|
|
|
|
def test_load_state_dict_logs_missing_keys(self, tmp_path, caplog):
|
|
"""Keys the checkpoint omits are reported as missing."""
|
|
import logging
|
|
from safetensors.torch import save_file
|
|
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
save_file({"layer1.weight": torch.ones(2, 2)}, str(model_dir / "model.safetensors"))
|
|
|
|
model = _TwoParam()
|
|
with caplog.at_level(logging.WARNING):
|
|
VibeVoiceLoader._load_state_dict_into_model(
|
|
model=model,
|
|
model_path=str(model_dir),
|
|
model_type="official",
|
|
model_info={},
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
assert any("Missing keys" in r.message for r in caplog.records)
|
|
|
|
def test_assigned_weights_own_private_storage(self, tmp_path):
|
|
"""mmap severing: assigned params are private clones, not file views.
|
|
|
|
safetensors tensors are zero-copy views into the mapped file; a view
|
|
retained by an offloaded parameter would pin the whole file mapping
|
|
in the process working set (ghost RAM + unstable cudaHostRegister
|
|
pins). A private clone's storage covers exactly its own bytes, while
|
|
a view's storage covers the file's entire data section.
|
|
"""
|
|
from safetensors.torch import save_file
|
|
|
|
model_dir = tmp_path / "TestModel"
|
|
model_dir.mkdir()
|
|
w1 = torch.ones(2, 2)
|
|
w2 = torch.full((2, 2), 2.0)
|
|
save_file(
|
|
{"layer1.weight": w1, "layer2.weight": w2},
|
|
str(model_dir / "model.safetensors"),
|
|
)
|
|
|
|
model = _TwoParam()
|
|
VibeVoiceLoader._load_state_dict_into_model(
|
|
model=model,
|
|
model_path=str(model_dir),
|
|
model_type="official",
|
|
model_info={},
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
for param in (model.layer1.weight, model.layer2.weight):
|
|
own_bytes = param.numel() * param.element_size()
|
|
assert param.untyped_storage().nbytes() == own_bytes, (
|
|
"assigned parameter still sits in the checkpoint's file "
|
|
"mapping — the streaming assign must clone into private memory"
|
|
)
|
|
|
|
|
|
class TestStreamApplyDense:
|
|
"""Direct unit tests for VibeVoiceLoader._stream_apply_dense."""
|
|
|
|
def test_assigns_params_and_buffers_preserving_file_views(self):
|
|
"""Default keeps the source tensor, so a file view stays a file view.
|
|
|
|
Every production caller now passes a ``target_device``, and a
|
|
``.to(device)`` already produces private storage — the clone only
|
|
matters for the no-device path, which is why the default is
|
|
``preserve_file_views=True``.
|
|
"""
|
|
model = _TwoParam()
|
|
model.register_buffer("pos", torch.zeros(3))
|
|
w = torch.ones(2, 2)
|
|
pairs = [
|
|
("layer1.weight", w),
|
|
("layer2.weight", torch.full((2, 2), 2.0)),
|
|
("pos", torch.ones(3)),
|
|
]
|
|
missing, unexpected = VibeVoiceLoader._stream_apply_dense(model, iter(pairs))
|
|
|
|
assert unexpected == []
|
|
assert missing == []
|
|
assert torch.equal(model.layer1.weight.data, w)
|
|
assert model.layer1.weight.data_ptr() == w.data_ptr()
|
|
assert torch.equal(model.pos, torch.ones(3))
|
|
|
|
def test_preserve_file_views_false_clones_into_private_storage(self):
|
|
model = _TwoParam()
|
|
w = torch.ones(2, 2)
|
|
VibeVoiceLoader._stream_apply_dense(
|
|
model, iter([("layer1.weight", w)]), preserve_file_views=False
|
|
)
|
|
|
|
assert torch.equal(model.layer1.weight.data, w)
|
|
assert model.layer1.weight.data_ptr() != w.data_ptr()
|
|
|
|
def test_target_device_cuda_places_params_on_the_accelerator(self):
|
|
if not torch.cuda.is_available():
|
|
pytest.skip("no GPU in this environment")
|
|
dev = torch.device("cuda", 0)
|
|
model = _TwoParam()
|
|
w = torch.ones(2, 2)
|
|
VibeVoiceLoader._stream_apply_dense(
|
|
model, iter([("layer1.weight", w), ("layer2.weight", w)]),
|
|
target_device=dev,
|
|
)
|
|
|
|
assert model.layer1.weight.device.type == "cuda"
|
|
assert model.layer2.weight.device.type == "cuda"
|
|
assert torch.equal(model.layer1.weight.data.cpu(), w)
|
|
# A .to(device) is already private storage; no extra clone needed.
|
|
assert model.layer1.weight.data_ptr() != w.data_ptr()
|
|
|
|
def test_shape_mismatch_raises_friendly_error(self):
|
|
model = _TwoParam()
|
|
pairs = [("layer1.weight", torch.ones(3, 3))]
|
|
with pytest.raises(ValueError, match="shapes do not match"):
|
|
VibeVoiceLoader._stream_apply_dense(model, iter(pairs))
|
|
|
|
def test_unexpected_and_missing_reported(self):
|
|
model = _TwoParam()
|
|
pairs = [("layer1.weight", torch.ones(2, 2)), ("bogus.key", torch.zeros(1))]
|
|
missing, unexpected = VibeVoiceLoader._stream_apply_dense(model, iter(pairs))
|
|
|
|
assert unexpected == ["bogus.key"]
|
|
assert "layer2.weight" in missing
|
|
|
|
|
|
class TestCleanupOldModels:
|
|
"""Test cleanup_old_models function."""
|
|
|
|
def test_cleanup_keeps_specified_key(self):
|
|
from ComfyUI_VibeVoice.modules.loader import LOADED_MODELS_CACHE
|
|
from ComfyUI_VibeVoice.modules.utils import VIBEVOICE_PATCHER_CACHE
|
|
|
|
LOADED_MODELS_CACHE.clear()
|
|
VIBEVOICE_PATCHER_CACHE.clear()
|
|
|
|
LOADED_MODELS_CACHE["key1"] = "model1"
|
|
LOADED_MODELS_CACHE["key2"] = "model2"
|
|
|
|
with patch("ComfyUI_VibeVoice.modules.loader.model_management"):
|
|
cleanup_old_models(keep_cache_key="key1")
|
|
|
|
assert "key1" in LOADED_MODELS_CACHE
|
|
assert "key2" not in LOADED_MODELS_CACHE
|
|
|
|
def test_cleanup_clears_all(self):
|
|
from ComfyUI_VibeVoice.modules.loader import LOADED_MODELS_CACHE
|
|
from ComfyUI_VibeVoice.modules.utils import VIBEVOICE_PATCHER_CACHE
|
|
|
|
LOADED_MODELS_CACHE.clear()
|
|
VIBEVOICE_PATCHER_CACHE.clear()
|
|
|
|
LOADED_MODELS_CACHE["key1"] = "model1"
|
|
LOADED_MODELS_CACHE["key2"] = "model2"
|
|
|
|
with patch("ComfyUI_VibeVoice.modules.loader.model_management"):
|
|
cleanup_old_models(keep_cache_key=None)
|
|
|
|
assert len(LOADED_MODELS_CACHE) == 0
|
|
|
|
|
|
class TestTTSLoaderBaseInheritance:
|
|
"""IMP-004: TTS loader shares the BaseVibeVoiceLoader."""
|
|
|
|
def test_tts_loader_uses_base(self):
|
|
assert isinstance(VibeVoiceLoader(), BaseVibeVoiceLoader)
|
|
|
|
|
|
# ====================================================================
|
|
# AUDIT PHASE B — B2: config loading & streaming detection
|
|
# ====================================================================
|
|
class TestLoadConfig:
|
|
"""B2: _load_config must pick the right config class + fallback."""
|
|
|
|
def test_streaming_json_uses_streaming_config_class(self, tmp_path):
|
|
config_path = tmp_path / "config.json"
|
|
config_path.write_text(json.dumps({"model_type": "vibevoice_streaming"}))
|
|
|
|
with patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceStreamingConfig") as mock_stream, \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceConfig") as mock_base:
|
|
VibeVoiceLoader._load_config(str(config_path), "SomeModel")
|
|
mock_stream.from_pretrained.assert_called_once_with(str(config_path))
|
|
mock_base.from_pretrained.assert_not_called()
|
|
|
|
def test_non_streaming_json_uses_base_config_class(self, tmp_path):
|
|
config_path = tmp_path / "config.json"
|
|
config_path.write_text(json.dumps({"model_type": "vibevoice"}))
|
|
|
|
with patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceStreamingConfig") as mock_stream, \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceConfig") as mock_base:
|
|
VibeVoiceLoader._load_config(str(config_path), "SomeModel")
|
|
mock_base.from_pretrained.assert_called_once_with(str(config_path))
|
|
mock_stream.from_pretrained.assert_not_called()
|
|
|
|
def test_missing_config_falls_back_to_large_default(self):
|
|
with patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceConfig") as mock_base:
|
|
VibeVoiceLoader._load_config("/nonexistent/config.json", "VibeVoice-Large")
|
|
mock_base.from_pretrained.assert_called_once()
|
|
fallback_arg = mock_base.from_pretrained.call_args[0][0]
|
|
assert "default_VibeVoice-Large_config.json" in fallback_arg
|
|
|
|
def test_missing_config_falls_back_to_15b_default(self):
|
|
with patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceConfig") as mock_base:
|
|
VibeVoiceLoader._load_config("/nonexistent/config.json", "VibeVoice-1.5B")
|
|
fallback_arg = mock_base.from_pretrained.call_args[0][0]
|
|
assert "default_VibeVoice-1.5B_config.json" in fallback_arg
|
|
|
|
|
|
# ====================================================================
|
|
# AUDIT PHASE B — B3: tokenizer acquisition order
|
|
# ====================================================================
|
|
class TestLoadTokenizer:
|
|
"""B3 acquisition order, copy-free: existing file -> packaged direct
|
|
load -> HF download -> RuntimeError."""
|
|
|
|
def test_existing_tokenizer_no_packaged_no_download(self, tmp_path):
|
|
(tmp_path / "tokenizer.json").write_text("{}")
|
|
with patch("ComfyUI_VibeVoice.modules.loader.hf_hub_download") as mock_dl, patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceTextTokenizerFast") as mock_tok:
|
|
VibeVoiceLoader._load_tokenizer(str(tmp_path), "TestModel")
|
|
mock_dl.assert_not_called()
|
|
assert mock_tok.call_args[1]["tokenizer_file"] == str(tmp_path / "tokenizer.json")
|
|
|
|
def test_packaged_fallback_loaded_directly_no_side_effects(self, tmp_path):
|
|
"""Packaged tokenizer loads straight from the node folder; the
|
|
user's model directory is never written to."""
|
|
with patch("os.path.exists", side_effect=lambda p: (
|
|
True if "configs" in p else os.path.isfile(p)
|
|
)), patch("ComfyUI_VibeVoice.modules.loader.hf_hub_download") as mock_dl, patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceTextTokenizerFast") as mock_tok:
|
|
VibeVoiceLoader._load_tokenizer(str(tmp_path), "TestModel")
|
|
mock_dl.assert_not_called()
|
|
used_path = mock_tok.call_args[1]["tokenizer_file"]
|
|
assert "configs" in used_path
|
|
assert used_path.endswith("tokenizer.json")
|
|
assert not (tmp_path / "tokenizer.json").exists()
|
|
|
|
def test_download_fallback_second_repo_succeeds(self, tmp_path):
|
|
calls = []
|
|
|
|
def fake_download(repo_id=None, filename=None, local_dir=None):
|
|
calls.append(repo_id)
|
|
if repo_id == "Qwen/Qwen2.5-1.5B":
|
|
raise RuntimeError("offline")
|
|
with open(os.path.join(local_dir, "tokenizer.json"), "w") as f:
|
|
f.write("{}")
|
|
|
|
# Hide the packaged tokenizer so the download path runs.
|
|
with patch("os.path.exists",
|
|
side_effect=lambda p: (False if "configs" in p
|
|
else os.path.isfile(p))), patch("ComfyUI_VibeVoice.modules.loader.hf_hub_download",
|
|
side_effect=fake_download), patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceTextTokenizerFast"):
|
|
VibeVoiceLoader._load_tokenizer(str(tmp_path), "TestModel")
|
|
assert calls == ["Qwen/Qwen2.5-1.5B", "Qwen/Qwen2.5-7B"]
|
|
|
|
def test_all_sources_fail_raises_runtime_error(self, tmp_path):
|
|
with patch("os.path.exists", return_value=False), patch("ComfyUI_VibeVoice.modules.loader.hf_hub_download",
|
|
side_effect=RuntimeError("offline")), patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceTextTokenizerFast"):
|
|
with pytest.raises(RuntimeError, match="Could not get 'tokenizer.json'"):
|
|
VibeVoiceLoader._load_tokenizer(str(tmp_path), "TestModel")
|
|
|
|
|
|
# ====================================================================
|
|
# AUDIT PHASE B — B6: load_model orchestration (fully mocked)
|
|
# ====================================================================
|
|
class TestLoadModelOrchestration:
|
|
"""B6: full load_model sequence, cache, and failure isolation."""
|
|
|
|
def _model_registry(self, model_type="official"):
|
|
return {"TestModel": {"type": model_type, "repo_id": "test/repo", "path": "x"}}
|
|
|
|
def test_unknown_model_raises_value_error(self):
|
|
with patch("ComfyUI_VibeVoice.modules.loader.AVAILABLE_VIBEVOICE_MODELS", {}):
|
|
with pytest.raises(ValueError, match="Unknown VibeVoice model"):
|
|
VibeVoiceLoader.load_model("Nope", torch.device("cpu"))
|
|
|
|
def test_cache_hit_short_circuits(self):
|
|
LOADED_MODELS_CACHE.clear()
|
|
sentinel = ("cached_model", "cached_processor")
|
|
LOADED_MODELS_CACHE["TestModel_attn_sdpa_q4_0"] = sentinel
|
|
try:
|
|
with patch("ComfyUI_VibeVoice.modules.loader.AVAILABLE_VIBEVOICE_MODELS",
|
|
self._model_registry()), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._resolve_model_paths") as mock_rp:
|
|
result = VibeVoiceLoader.load_model(
|
|
"TestModel", torch.device("cpu"), attention_mode="sdpa"
|
|
)
|
|
assert result == sentinel
|
|
mock_rp.assert_not_called() # no work done on cache hit
|
|
finally:
|
|
LOADED_MODELS_CACHE.clear()
|
|
|
|
def test_full_sequence_and_cache_store(self):
|
|
LOADED_MODELS_CACHE.clear()
|
|
ledger = []
|
|
|
|
fake_model = MagicMock()
|
|
fake_model.to.return_value = fake_model
|
|
fake_processor = MagicMock()
|
|
fake_config = MagicMock(spec=[]) # not a VibeVoiceStreamingConfig instance
|
|
|
|
# isinstance() needs a real class, not a MagicMock.
|
|
class _FakeStreamingCfg:
|
|
pass
|
|
|
|
try:
|
|
with patch("ComfyUI_VibeVoice.modules.loader.AVAILABLE_VIBEVOICE_MODELS",
|
|
self._model_registry()), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceStreamingConfig", _FakeStreamingCfg), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._resolve_model_paths",
|
|
side_effect=lambda n: (ledger.append("paths"), ("mp", "cp", "pp", "td"))[1]), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_config",
|
|
side_effect=lambda cp, n: (ledger.append("config"), fake_config)[1]), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_tokenizer",
|
|
side_effect=lambda td, n: (ledger.append("tokenizer"), MagicMock())[1]), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_processor",
|
|
side_effect=lambda tok, pp, is_streaming=False: (
|
|
ledger.append("processor"), fake_processor)[1]), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._instantiate_model",
|
|
side_effect=lambda **kw: (ledger.append("instantiate"), fake_model)[1]), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_state_dict_into_model",
|
|
side_effect=lambda **kw: (ledger.append("state_dict"), fake_model)[1]), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.cast_model_to_dtype_if_needed",
|
|
side_effect=lambda m, d: ledger.append("cast")):
|
|
model, processor = VibeVoiceLoader.load_model(
|
|
"TestModel", torch.device("cpu"), attention_mode="sdpa"
|
|
)
|
|
|
|
assert model is fake_model
|
|
assert processor is fake_processor
|
|
assert ledger == ["paths", "config", "tokenizer", "processor",
|
|
"instantiate", "state_dict", "cast"]
|
|
# Plan 2026-08-18 D4/RC-3: dtype applied via conditional cast
|
|
# helper (not an unconditional .to()), eval'd, cached
|
|
fake_model.eval.assert_called_once()
|
|
assert LOADED_MODELS_CACHE["TestModel_attn_sdpa_q4_0"] == (fake_model, fake_processor)
|
|
finally:
|
|
LOADED_MODELS_CACHE.clear()
|
|
|
|
def test_exception_mid_load_wraps_and_does_not_pollute_cache(self):
|
|
LOADED_MODELS_CACHE.clear()
|
|
|
|
class _FakeStreamingCfg:
|
|
pass
|
|
|
|
try:
|
|
with patch("ComfyUI_VibeVoice.modules.loader.AVAILABLE_VIBEVOICE_MODELS",
|
|
self._model_registry()), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceStreamingConfig", _FakeStreamingCfg), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._resolve_model_paths",
|
|
return_value=("mp", "cp", "pp", "td")), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_config",
|
|
return_value=MagicMock(spec=[])), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_tokenizer",
|
|
return_value=MagicMock()), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_processor",
|
|
return_value=MagicMock()), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._instantiate_model",
|
|
side_effect=RuntimeError("boom")):
|
|
with pytest.raises(RuntimeError, match="Failed to load model"):
|
|
VibeVoiceLoader.load_model(
|
|
"TestModel", torch.device("cpu"), attention_mode="sdpa"
|
|
)
|
|
# Cache must NOT contain a partial entry.
|
|
assert "TestModel_attn_sdpa_q4_0" not in LOADED_MODELS_CACHE
|
|
finally:
|
|
LOADED_MODELS_CACHE.clear()
|
|
|
|
def test_4bit_builds_bnb_config_and_replaces_linears(self):
|
|
LOADED_MODELS_CACHE.clear()
|
|
fake_model = MagicMock()
|
|
fake_model.to.return_value = fake_model
|
|
|
|
class _FakeStreamingCfg:
|
|
pass
|
|
|
|
# AUD-014: patch the import target that actually resolves on this
|
|
# transformers version (integrations on 5.x, utils on 4.x).
|
|
try:
|
|
import transformers.integrations.bitsandbytes as _bnb_mod
|
|
_bnb_target = "transformers.integrations.bitsandbytes.replace_with_bnb_linear"
|
|
except ImportError:
|
|
_bnb_target = "transformers.utils.bitsandbytes.replace_with_bnb_linear"
|
|
|
|
try:
|
|
with patch("ComfyUI_VibeVoice.modules.loader.AVAILABLE_VIBEVOICE_MODELS",
|
|
self._model_registry()), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceStreamingConfig", _FakeStreamingCfg), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._resolve_model_paths",
|
|
return_value=("mp", "cp", "pp", "td")), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_config",
|
|
return_value=MagicMock(spec=[])), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_tokenizer",
|
|
return_value=MagicMock()), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_processor",
|
|
return_value=MagicMock()), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._instantiate_model",
|
|
return_value=fake_model), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.VibeVoiceLoader._load_state_dict_into_model",
|
|
return_value=fake_model), \
|
|
patch("ComfyUI_VibeVoice.modules.loader.BitsAndBytesConfig") as mock_bnb, \
|
|
patch(_bnb_target) as mock_replace:
|
|
VibeVoiceLoader.load_model(
|
|
"TestModel", torch.device("cpu"),
|
|
attention_mode="sdpa", use_llm_4bit=True,
|
|
)
|
|
mock_bnb.assert_called_once()
|
|
assert mock_bnb.call_args.kwargs["load_in_4bit"] is True
|
|
mock_replace.assert_called_once()
|
|
assert getattr(fake_model, "_llm_4bit") is True
|
|
finally:
|
|
LOADED_MODELS_CACHE.clear()
|