Files
WildAi f71eef465a refactor(log): rename console prefix to [VibeVoice TTS]
Partial rename completed across 24 source files (73 sites). Tests pinned
the old literal; test_logging_idiom.py now derives from PREFIX.
2026-10-03 16:03:14 +03:00

1585 lines
58 KiB
Python

"""External model loading for VibeVoice.
Loads a VibeVoice model from a standalone weight file (safetensors / .bin /
.gguf) placed in ComfyUI's ``diffusion_models`` folder, binding the required
config / tokenizer / preprocessor JSONs via sidecar files or packaged defaults.
This module bypasses ComfyUI's ``model_detection.detect_unet_config()`` entirely
(VibeVoice's state-dict keys are not recognized by it) and instead reuses the
existing :class:`~modules.loader.VibeVoiceLoader` internals for config,
tokenizer, processor, and model instantiation.
Sidecar file convention (place next to the weight file):
- ``<weight>.config.json`` → architecture config (preferred)
- ``config.json`` → architecture config (same directory)
- ``<weight>.preprocessor.json`` → audio preprocessor config (preferred)
- ``preprocessor_config.json`` → audio preprocessor config (same dir)
- ``tokenizer.json`` → Qwen2.5 text tokenizer (same dir)
The ``config_name`` dropdown selects the loading BRANCH (TTS vs ASR). Inside
the ASR branch the checkpoint's own ``model_type`` selects the model FAMILY,
because two incompatible ASR families are loadable there and the config JSON
is the only thing that can tell them apart:
- ``"vibevoice_asr"`` → the transformers >= 5.3.0 classes
(``AutoConfig`` / ``VibeVoiceAsrForConditionalGeneration``), which build
the ``language_model.model.*`` tree the published ``VibeVoice-ASR-HF``
checkpoints store. See :mod:`modules.asr_native`.
- ``"vibevoice"`` → the vendored ``src/vibevoice`` classes
(``VibeVoiceASRConfig`` / ``VibeVoiceASRForConditionalGeneration``),
which build the ``model.language_model.*`` tree of the original
``microsoft/VibeVoice-ASR`` checkpoints.
"""
import os
import json
import struct
import logging
import contextlib
import gc
import torch
import comfy.utils
import comfy.model_management as model_management
from transformers import BitsAndBytesConfig
from ..src.vibevoice.modular.configuration_vibevoice import VibeVoiceConfig, VibeVoiceASRConfig
from ..src.vibevoice.modular.configuration_vibevoice_streaming import VibeVoiceStreamingConfig
from ..src.vibevoice.modular.modeling_vibevoice_asr import VibeVoiceASRForConditionalGeneration
from ..src.vibevoice.modular.modular_vibevoice_text_tokenizer import VibeVoiceASRTextTokenizerFast
from ..src.vibevoice.processor.vibevoice_asr_processor import VibeVoiceASRProcessor
from ..src.vibevoice.processor.vibevoice_tokenizer_processor import VibeVoiceTokenizerProcessor
from .loader import VibeVoiceLoader, QUANT_STORAGE_DTYPES
from .base_loader import iter_safetensors_tensors, place_tensor_on_device
from .asr_native import (
NATIVE_ASR_MODEL_TYPE,
build_native_asr_processor,
is_native_asr_config_path,
load_native_asr_config,
instantiate_native_asr_model,
)
from .attention_utils import (
SAGE_ATTENTION_AVAILABLE,
resolve_attention_mode,
resolve_realtime_attention_mode,
resolve_asr_attention_mode,
get_attn_implementation_for_load,
check_sage_attention_compatible,
)
from .dtype_utils import resolve_dtype, cast_model_to_dtype_if_needed, set_config_dtype
from .convrot_quant import UnsupportedQuantFormat
from .quant_common import validate_weight_plan
from .diagnostics import diagnostics_enabled
from pathlib import Path
from .patcher import (
VibeVoiceASRPatcher,
VibeVoicePatcher,
dynamic_vram_available,
resolve_core_patcher_class,
select_patcher_class,
)
from .memory_census import measured_load, report_census
if SAGE_ATTENTION_AVAILABLE:
from ..src.vibevoice.modular.sage_attention_patch import set_sage_attention
_PACKAGED_CONFIG_FILES = {
"VibeVoice-1.5B": "default_VibeVoice-1.5B_config.json",
"VibeVoice-7B": "default_VibeVoice-Large_config.json",
"VibeVoice-ASR": "default_VibeVoice-ASR_config.json",
}
AUTO_CONFIG_NAME = "Auto-detect"
EXTERNAL_CONFIG_OPTIONS = [
AUTO_CONFIG_NAME,
"VibeVoice-1.5B",
"VibeVoice-7B",
"VibeVoice-Realtime-0.5B",
"VibeVoice-ASR",
]
_LEGACY_CONFIG_ALIASES = {
"vibevoice-large": "VibeVoice-7B",
}
def normalize_config_name(config_name: str) -> str:
"""Map legacy/alias config_name values onto their canonical option."""
if not config_name:
return config_name
if config_name in EXTERNAL_CONFIG_OPTIONS:
return config_name
return _LEGACY_CONFIG_ALIASES.get(config_name.lower(), config_name)
def resolve_auto_config_name(weight_path: str, gguf_reader=None, weights_fp=None) -> str:
"""Resolve ``AUTO_CONFIG_NAME`` to a concrete family from the weights."""
if weights_fp is None:
from .config_detect import fingerprint_weights
weights_fp = fingerprint_weights(weight_path, gguf_reader=gguf_reader)
detected = weights_fp.config_name if weights_fp is not None else None
if not detected:
explicit = [o for o in EXTERNAL_CONFIG_OPTIONS if o != AUTO_CONFIG_NAME]
if weights_fp is None:
observed = "Observed dimensions: unavailable (no embedding fingerprint)."
else:
observed = (
"Observed dimensions: "
f"hidden={weights_fp.hidden_size}, vocab={weights_fp.vocab_size}."
)
raise ValueError(
f"config_name '{AUTO_CONFIG_NAME}': could not determine the "
f"VibeVoice architecture from "
f"'{os.path.basename(weight_path)}'. {observed} Auto-detect reads "
f"the checkpoint's embedding fingerprint; unmatched dimensions "
f"are not guessed. Select config_name explicitly (one of: "
f"{', '.join(explicit)}) or place a sidecar config.json next to "
f"the weight file."
)
logging.debug(
f"[VibeVoice TTS] Auto-detected architecture '{detected}' from "
f"'{os.path.basename(weight_path)}'"
)
return detected
def reconcile_config(selected_name: str, resolved_config, weights_fp):
"""Reconcile the selected config against the weights' fingerprint."""
from .config_detect import config_fingerprint
if weights_fp is None:
return selected_name, False
cfg_fp = config_fingerprint(resolved_config)
if cfg_fp is None or cfg_fp == (weights_fp.hidden_size, weights_fp.vocab_size):
return selected_name, False
detected = weights_fp.config_name
if not detected or detected == selected_name:
return selected_name, False
logging.warning(
f"[VibeVoice TTS] Config mismatch: weights contain '{weights_fp.source_key}' with "
f"shape (vocab={weights_fp.vocab_size}, hidden={weights_fp.hidden_size}) "
f"-> {detected}, but config '{selected_name}' "
f"(hidden={cfg_fp[0]}, vocab={cfg_fp[1]}) was selected. "
f"Using '{detected}'."
)
return detected, True
# ====================================================================
# Reconciliation memo
# ====================================================================
_RECONCILIATION_MEMO: dict = {}
def _reconciliation_memo_key(weight_path: str, selected_name: str):
abspath = os.path.abspath(weight_path) if weight_path else ""
mtime_ns = 0
size = 0
try:
stat = os.stat(abspath)
mtime_ns = stat.st_mtime_ns
size = stat.st_size
except OSError:
pass
return (abspath, mtime_ns, size, selected_name)
def reconciled_config_name(weight_path: str, selected_name: str) -> str:
return _RECONCILIATION_MEMO.get(
_reconciliation_memo_key(weight_path, selected_name), selected_name
)
def remember_reconciled_config(
weight_path: str, selected_name: str, effective_name: str
) -> None:
if not effective_name:
return
_RECONCILIATION_MEMO[
_reconciliation_memo_key(weight_path, selected_name)
] = effective_name
def clear_reconciliation_memo() -> None:
_RECONCILIATION_MEMO.clear()
# ====================================================================
# Path resolution helpers
# ====================================================================
def _packaged_configs_dir() -> str:
return os.path.join(
os.path.dirname(__file__), "..", "src", "vibevoice", "configs"
)
def _get_packaged_config_path(config_name: str) -> str:
filename = _PACKAGED_CONFIG_FILES.get(config_name)
if filename is None:
return ""
return os.path.normpath(os.path.join(_packaged_configs_dir(), filename))
def resolve_sidecar_config(weight_path: str, config_name: str) -> str:
sidecar_path = weight_path + ".config.json"
if os.path.exists(sidecar_path):
logging.debug(f"[VibeVoice TTS] Using sidecar config: {sidecar_path}")
return sidecar_path
dir_sidecar = os.path.join(os.path.dirname(weight_path), "config.json")
if os.path.exists(dir_sidecar):
logging.debug(f"[VibeVoice TTS] Using directory sidecar config: {dir_sidecar}")
return dir_sidecar
packaged = _get_packaged_config_path(config_name)
if packaged and os.path.exists(packaged):
logging.debug(f"[VibeVoice TTS] Using packaged default config for '{config_name}': {packaged}")
return packaged
raise FileNotFoundError(
f"No architecture config found for external model "
f"'{os.path.basename(weight_path)}'.\n"
f"Place a sidecar config next to the weight file:\n"
f" - '{sidecar_path}' (preferred), or\n"
f" - 'config.json' in '{os.path.dirname(weight_path)}'\n"
f"Alternatively select a config_name with a packaged default "
f"({', '.join(_PACKAGED_CONFIG_FILES.keys())})."
)
def resolve_sidecar_preprocessor(weight_path: str) -> str:
sidecar_path = weight_path + ".preprocessor.json"
if os.path.exists(sidecar_path):
logging.debug(f"[VibeVoice TTS] Using sidecar preprocessor config: {sidecar_path}")
return sidecar_path
dir_sidecar = os.path.join(os.path.dirname(weight_path), "preprocessor_config.json")
if os.path.exists(dir_sidecar):
logging.debug(f"[VibeVoice TTS] Using directory sidecar preprocessor config: {dir_sidecar}")
return dir_sidecar
return ""
def resolve_sidecar_tokenizer_dir(weight_path: str) -> str:
return os.path.dirname(weight_path)
# ====================================================================
# ASR-specific helpers
# ====================================================================
ASR_CONFIG_NAMES = {"VibeVoice-ASR"}
def is_asr_config_name(config_name: str) -> bool:
return config_name in ASR_CONFIG_NAMES
def _load_asr_config(config_path: str) -> "VibeVoiceASRConfig":
if not os.path.exists(config_path):
raise FileNotFoundError(f"ASR config not found: {config_path}")
if is_native_asr_config_path(config_path):
return load_native_asr_config(config_path)
return VibeVoiceASRConfig.from_pretrained(config_path)
def _load_asr_tokenizer(tokenizer_dir: str) -> "VibeVoiceASRTextTokenizerFast":
tokenizer_file_path = os.path.join(tokenizer_dir, "tokenizer.json")
if not os.path.exists(tokenizer_file_path):
packaged_tokenizer_path = os.path.join(
_packaged_configs_dir(), "tokenizer.json"
)
if os.path.exists(packaged_tokenizer_path):
logging.debug(
f"[VibeVoice TTS] Using packaged tokenizer.json fallback for ASR: {packaged_tokenizer_path}"
)
tokenizer_file_path = packaged_tokenizer_path
else:
raise FileNotFoundError(
f"No 'tokenizer.json' found for external ASR model. "
f"Place a 'tokenizer.json' next to the weight file in "
f"'{tokenizer_dir}'."
)
return VibeVoiceASRTextTokenizerFast(tokenizer_file=tokenizer_file_path)
def _load_asr_processor(
tokenizer,
preprocessor_config_path: str,
) -> "VibeVoiceASRProcessor":
config = {}
if preprocessor_config_path and os.path.exists(preprocessor_config_path):
with open(preprocessor_config_path, "r", encoding="utf-8") as f:
config = json.load(f)
speech_tok_compress_ratio = config.get("speech_tok_compress_ratio", 3200)
target_sample_rate = config.get("target_sample_rate", 24000)
normalize_audio = config.get("normalize_audio", True)
audio_processor = VibeVoiceTokenizerProcessor(
sampling_rate=target_sample_rate,
normalize_audio=normalize_audio,
target_dB_FS=config.get("target_dB_FS", -25),
eps=config.get("eps", 1e-6),
)
return VibeVoiceASRProcessor(
tokenizer=tokenizer,
audio_processor=audio_processor,
speech_tok_compress_ratio=speech_tok_compress_ratio,
target_sample_rate=target_sample_rate,
normalize_audio=normalize_audio,
)
def _instantiate_asr_model(
config,
attn_implementation: str,
final_load_dtype: torch.dtype,
use_meta: bool = True,
):
if getattr(config, "model_type", None) == NATIVE_ASR_MODEL_TYPE:
return instantiate_native_asr_model(
config,
attn_implementation=attn_implementation,
final_load_dtype=final_load_dtype,
use_meta=use_meta,
)
if hasattr(config, "decoder_config"):
config.decoder_config._attn_implementation = attn_implementation
set_config_dtype(config, final_load_dtype)
if hasattr(config, "decoder_config"):
set_config_dtype(config.decoder_config, final_load_dtype)
ctx = torch.device("meta") if use_meta else contextlib.nullcontext()
with ctx:
return VibeVoiceASRForConditionalGeneration(config)
# ====================================================================
# Weight file state-dict loading (safetensors / bin / gguf dispatch)
# ====================================================================
def _load_gguf_state_dict(weight_path: str, device=None) -> dict:
if device is None:
device = torch.device("cpu")
try:
import gguf
except ImportError as e:
raise RuntimeError(
"Loading .gguf weights requires the 'gguf' Python package. "
"Install it with: pip install gguf"
) from e
logging.debug(f"[VibeVoice TTS] Loading GGUF state dict from: {weight_path}")
from .gguf_quant import dequantize_reader_tensor, open_gguf_reader
reader = open_gguf_reader(weight_path)
state_dict = {}
for tensor in reader.tensors:
dequantized = dequantize_reader_tensor(tensor)
state_dict[tensor.name] = dequantized.to(device)
logging.debug(f"[VibeVoice TTS] Loaded {len(state_dict)} tensors from GGUF file")
return state_dict
def _load_weight_state_dict(weight_path: str, device) -> dict:
if weight_path.lower().endswith(".gguf"):
return _load_gguf_state_dict(weight_path, device=device)
return comfy.utils.load_torch_file(weight_path, device=device)
# ====================================================================
# Quant-resident loading (GGUF raw-block + ConvRot INT8)
# ====================================================================
def _open_gguf_reader(weight_path: str):
try:
import gguf # noqa: F401
except ImportError as e:
raise RuntimeError(
"Loading .gguf weights requires the 'gguf' Python package. "
"Install it with: pip install gguf"
) from e
from .gguf_quant import open_gguf_reader
return open_gguf_reader(weight_path)
def _gguf_kquant_present(reader) -> bool:
from .gguf_quant import _T
kquants = {_T.Q4_K, _T.Q5_K, _T.Q6_K}
return any(t.tensor_type in kquants for t in reader.tensors)
_QUANT_STORAGE_DTYPES = QUANT_STORAGE_DTYPES
def _plan_quantized_safetensors_load(quant_map: dict):
from .convrot_quant import make_convrot_linear
from .fp8_quant import make_fp8_linear
layer_plan = {}
n_rowwise = 0
n_fp8_resident = 0
for prefix, info in quant_map.items():
if info.convrot:
layer_plan[prefix] = make_convrot_linear(info)
elif info.resident_fp8:
layer_plan[prefix] = make_fp8_linear(info)
n_fp8_resident += 1
else:
n_rowwise += 1
return layer_plan, n_rowwise, n_fp8_resident
def _dequantize_rowwise_weight(prefix: str, info, w, s):
from .convrot_quant import resolve_orig_dtype
from .quant_common import QuantTargetMismatch
storage_dtype = info.rowwise_dtype or torch.int8
if w.dtype != storage_dtype and w.dtype != torch.uint8:
raise QuantTargetMismatch(
f"Rowwise layer '{prefix}': expected {storage_dtype} storage, "
f"checkpoint has {w.dtype}"
)
if w.dim() != 2 or w.shape[0] != info.out_features \
or (info.in_features and w.shape[1] != info.in_features):
raise QuantTargetMismatch(
f"Rowwise layer '{prefix}': weight shape {tuple(w.shape)} "
f"disagrees with metadata [out={info.out_features}, "
f"in={info.in_features}]"
)
orig_dtype = resolve_orig_dtype(info.orig_dtype)
if info.group_size:
gs = int(info.group_size)
if w.shape[0] % gs or w.shape[1] % gs:
raise QuantTargetMismatch(
f"Blockwise layer '{prefix}': weight {tuple(w.shape)} "
f"not divisible by group_size {gs} on both dims"
)
og, ig = w.shape[0] // gs, w.shape[1] // gs
if tuple(s.shape) != (og, ig):
raise QuantTargetMismatch(
f"Blockwise layer '{prefix}': scale shape "
f"{tuple(s.shape)} disagrees with weight "
f"{tuple(w.shape)} at group_size {gs} (expected "
f"({og}, {ig}))"
)
return (
(
w.to(torch.float32).view(og, gs, ig, gs)
* s.to(torch.float32).view(og, 1, ig, 1)
).reshape(w.shape).to(orig_dtype)
)
if not (s.dim() == 0 or (s.dim() == 2 and s.shape[1] == 1
and s.shape[0] in (w.shape[0], 1))):
raise QuantTargetMismatch(
f"Rowwise layer '{prefix}': scale shape {tuple(s.shape)} "
f"is neither a scalar nor [out,1] (weight "
f"{tuple(w.shape)})"
)
return (w.to(torch.float32) * s.to(torch.float32)).to(orig_dtype)
def _prepare_quantized_safetensors_load(state_dict: dict, quant_map: dict):
from .convrot_quant import QUANT_META_SUFFIX
from .quant_common import QuantTargetMismatch
layer_plan, n_rowwise, n_fp8_resident = _plan_quantized_safetensors_load(
quant_map
)
for prefix, info in quant_map.items():
if info.convrot:
continue
w_key = f"{prefix}.weight"
s_key = f"{prefix}.weight_scale"
w = state_dict.get(w_key)
s = state_dict.get(s_key)
if w is None or s is None:
raise QuantTargetMismatch(
f"Rowwise-quantized layer '{prefix}' is missing its "
f"'{w_key}' / '{s_key}' tensors in the checkpoint"
)
if info.resident_fp8:
if w.dtype != info.rowwise_dtype:
raise QuantTargetMismatch(
f"Resident fp8 layer '{prefix}': expected "
f"{info.rowwise_dtype} storage, checkpoint has {w.dtype}"
)
if not (s.dim() == 0 or (s.dim() == 1 and s.numel() == 1)):
raise QuantTargetMismatch(
f"Resident fp8 layer '{prefix}': scale shape "
f"{tuple(s.shape)} is not a per-tensor scalar"
)
continue
state_dict[w_key] = _dequantize_rowwise_weight(prefix, info, w, s)
state_dict.pop(s_key, None)
for prefix in quant_map:
state_dict.pop(f"{prefix}.{QUANT_META_SUFFIX}", None)
return layer_plan, n_rowwise, n_fp8_resident
# ====================================================================
# Streaming safetensors load
# ====================================================================
_SAFETENSORS_DTYPE_TO_TORCH = {
"F64": torch.float64, "F32": torch.float32, "F16": torch.float16,
"BF16": torch.bfloat16, "I64": torch.int64, "I32": torch.int32,
"I16": torch.int16, "I8": torch.int8, "U8": torch.uint8, "BOOL": torch.bool,
}
def read_safetensors_tensors_by_name(weight_path: str, names) -> dict:
"""Read the named tensors out of the header's byte ranges, one at a time.
``safe_open`` maps the entire file, and on Windows the first read through
such a mapping commits roughly one file size as private, untouched memory
that stays pinned for the lifetime of the returned tensor — charging a
9.5 GB checkpoint to host RAM to look up a few 4-byte scales. The header
already records each tensor's ``data_offsets``, so reading exactly those
byte ranges costs only the bytes asked for and touches no mapping.
"""
wanted = list(dict.fromkeys(names))
if not wanted:
return {}
remaining = set(wanted)
found = {}
with open(weight_path, "rb") as f:
(header_len,) = struct.unpack("<Q", f.read(8))
header = json.loads(f.read(header_len))
data_start = 8 + header_len
for name, info in header.items():
if name not in remaining:
continue
dtype = _SAFETENSORS_DTYPE_TO_TORCH.get(info.get("dtype"))
if dtype is None:
raise ValueError(
f"Checkpoint tensor '{name}' has unsupported safetensors "
f"dtype {info.get('dtype')!r}"
)
begin, end = info["data_offsets"]
buf = bytearray(end - begin)
f.seek(data_start + begin)
view = memoryview(buf)
while view:
read = f.readinto(view)
if not read:
raise ValueError(
f"Checkpoint tensor '{name}' is truncated: wanted "
f"{end - begin} bytes at offset {data_start + begin}, "
f"got {end - begin - len(view)}"
)
view = view[read:]
found[name] = torch.frombuffer(buf, dtype=dtype).reshape(
tuple(info["shape"])
)
remaining.discard(name)
return found
def _stream_apply_safetensors(
model, weight_path: str, quant_map: dict, target_device=None,
):
"""Assign a quantized safetensors checkpoint per-tensor.
Reads only the required scale tensors in pass 1, straight from their byte
ranges, and places every tensor onto target_device as it is assigned.
"""
from .convrot_quant import QUANT_META_SUFFIX
from .quant_common import QuantTargetMismatch
from .loader import VibeVoiceLoader, _shape_mismatch_error
resident_prefixes = {
p for p, info in quant_map.items()
if info.convrot or info.resident_fp8
}
dequant_infos = {
p: info for p, info in quant_map.items()
if not info.convrot and not info.resident_fp8
}
# Pass 1: read ONLY the required scale keys, straight from their byte
# ranges. safe_open here would commit the whole file as private RAM.
scales = {}
if dequant_infos:
scale_keys = {f"{prefix}.weight_scale" for prefix in dequant_infos}
read = read_safetensors_tensors_by_name(weight_path, scale_keys)
for prefix in dequant_infos:
s_key = f"{prefix}.weight_scale"
if s_key in read:
scales[prefix] = read[s_key]
else:
raise QuantTargetMismatch(
f"Rowwise-quantized layer '{prefix}' is missing its "
f"'{s_key}' tensor in the checkpoint"
)
params = dict(model.named_parameters())
buffers = dict(model.named_buffers())
unexpected = []
meta_suffix = f".{QUANT_META_SUFFIX}"
def _assign(key, tensor):
target = params.get(key)
if target is not None:
if tuple(tensor.shape) != tuple(target.shape):
raise _shape_mismatch_error(
model, [(key, tuple(tensor.shape), tuple(target.shape))]
)
tensor = place_tensor_on_device(tensor, target_device)
comfy.utils.set_attr_param(model, key, tensor)
return True
target_buf = buffers.get(key)
if target_buf is not None:
if tuple(tensor.shape) != tuple(target_buf.shape):
raise _shape_mismatch_error(
model, [(key, tuple(tensor.shape), tuple(target_buf.shape))]
)
tensor = place_tensor_on_device(tensor, target_device)
comfy.utils.set_attr_buffer(model, key, tensor)
return True
return False
with measured_load("stream-apply-safetensors") as _rss:
_rss.mark("stream-begin")
assigned = set()
for key, tensor in iter_safetensors_tensors(weight_path):
if not assigned:
_rss.mark("first-view")
if key.endswith(meta_suffix):
continue
prefix, _, leaf = key.rpartition(".")
if prefix in dequant_infos and leaf == "weight_scale":
continue
if prefix in dequant_infos and leaf == "weight":
tensor = _dequantize_rowwise_weight(
prefix, dequant_infos[prefix], tensor, scales[prefix]
)
else:
if tensor.dtype in _QUANT_STORAGE_DTYPES:
if prefix not in resident_prefixes \
or leaf not in ("weight", "weight_scale"):
raise ValueError(
f"Checkpoint contains quantized-weight tensors "
f"({key}: {tensor.dtype}) that no comfy_quant metadata "
f"declares. This node only loads quant formats it can "
f"execute; re-export the checkpoint or use a dense "
f"(bf16/fp16) file."
)
target = params.get(key)
if target is not None and tensor.dtype != target.dtype:
raise QuantTargetMismatch(
f"Resident layer '{prefix}': expected {target.dtype} "
f"storage for '{leaf}', checkpoint has {tensor.dtype}"
)
if _assign(key, tensor):
assigned.add(key)
else:
unexpected.append(key)
_rss.mark("stream-end")
report_census(model, phase=f"stream-end:{Path(weight_path).stem}")
expected = set(model.state_dict().keys())
missing_keys = [k for k in expected if k not in assigned]
return VibeVoiceLoader._post_assign_fixups(
model, missing_keys, unexpected
)
def _config_ties_word_embeddings(config) -> bool:
decoder_config = getattr(config, "decoder_config", None)
for cfg in (decoder_config, config):
if cfg is None:
continue
flag = getattr(cfg, "tie_word_embeddings", False)
if isinstance(flag, bool) and flag:
return True
return False
def _assert_lm_head_not_tied(config, quant_map: dict) -> None:
from .quant_common import QuantTargetMismatch
if "lm_head" in quant_map and _config_ties_word_embeddings(config):
raise QuantTargetMismatch(
"Checkpoint quantizes 'lm_head' but the config sets "
"tie_word_embeddings=True: a tied lm_head shares the input "
"embedding's storage and cannot carry quantized weights. The "
"sidecar config.json likely does not match this checkpoint — "
"fix the config or re-export without quantizing lm_head."
)
def _assert_gguf_lm_head_not_tied(config, reader) -> None:
from .gguf_quant import FLOAT_GGML_TYPES
if not _config_ties_word_embeddings(config):
return
for t in reader.tensors:
if t.name in ("lm_head.weight", "output.weight") and t.tensor_type not in FLOAT_GGML_TYPES:
from .quant_common import QuantTargetMismatch
raise QuantTargetMismatch(
"GGUF file carries a quantized "
f"'{t.name}' (lm_head) but the "
"config sets tie_word_embeddings=True: a tied lm_head shares "
"the input embedding's storage and cannot carry quantized "
"weights. The sidecar config.json likely does not match "
"this checkpoint — fix the config or re-export without "
"quantizing lm_head."
)
def _demote_nonlinear_fp8_residents(model, quant_map: dict) -> dict:
from dataclasses import replace as dc_replace
from .quant_common import resolve_module
resolved = {}
for prefix, info in quant_map.items():
if info.resident_fp8:
try:
target = resolve_module(model, prefix)
except KeyError:
target = None
if target is not None and not isinstance(target, torch.nn.Linear):
logging.debug(
f"[VibeVoice TTS] fp8-resident layer '{prefix}' targets "
f"{type(target).__name__}, not nn.Linear — falling back "
f"to dequant-at-load"
)
info = dc_replace(info, resident_fp8=False)
resolved[prefix] = info
return resolved
def _assert_dense_loadable(state_dict: dict) -> None:
bad = [
(k, str(v.dtype))
for k, v in state_dict.items()
if v.dtype in _QUANT_STORAGE_DTYPES
]
if bad:
sample = ", ".join(f"{k} [{d}]" for k, d in bad[:5])
raise ValueError(
"Checkpoint contains quantized-weight tensors but carries no "
f"executable quantization metadata ({len(bad)} tensors, e.g. "
f"{sample}). Loading them as floats would corrupt the model. "
"Re-export it with *.comfy_quant metadata (comfy-model-tools) "
"or use the dense/BF16 checkpoint."
)
def _stream_apply_dense_safetensors(model, weight_path: str, target_device=None):
"""Assign a dense safetensors checkpoint per-tensor straight onto ``target_device``.
This is the dense counterpart of :func:`_stream_apply_safetensors` and the
single dense entry point for both the TTS and ASR routes: the same
zero-copy iterator, the same per-tensor placement, the same post-assign
fixups. Nothing ever materialises the whole checkpoint as a CPU model, so a
BF16/FP16 file reaches VRAM the same way an FP8 one does.
"""
from .loader import VibeVoiceLoader
with measured_load("dense-stream-apply") as _rss:
_rss.mark("stream-begin")
VibeVoiceLoader._stream_apply_dense(
model,
iter_safetensors_tensors(weight_path),
preserve_file_views=True,
target_device=target_device,
)
_rss.mark("stream-end")
report_census(model, phase=f"stream-end:{Path(weight_path).stem}")
def _install_gguf_weights(model, reader, target_device=None) -> dict:
import numpy as np
from .gguf_quant import (
FLOAT_GGML_TYPES,
GGUFTensor,
SUPPORTED_GGML_TYPES,
UnsupportedGGMLType,
dequantize_reader_tensor,
gguf_linear_factory,
map_keys,
)
from .quant_common import (
QuantTargetMismatch,
replace_linears_for_quant,
resolve_module,
)
tensors = list(reader.tensors)
mapping = map_keys([t.name for t in tensors])
modules_by_name = dict(model.named_modules())
resident_plan = {}
plan = []
unsupported = []
dequant_load = []
dequant_bytes = 0
for t in tensors:
target_key = mapping[t.name]
logical_shape = tuple(int(s) for s in reversed(t.shape))
tt = t.tensor_type
if tt in FLOAT_GGML_TYPES:
plan.append((t, "dense", target_key, None, 0))
continue
if tt not in SUPPORTED_GGML_TYPES:
unsupported.append((t.name, tt))
continue
if not target_key.endswith(".weight"):
raise QuantTargetMismatch(
f"Quantized GGUF tensor '{t.name}' maps to '{target_key}', "
f"which is not a Linear weight; this checkpoint cannot be "
f"executed with quant-resident linears."
)
module_path = target_key[: -len(".weight")]
target_module = modules_by_name.get(module_path)
if target_module is None:
raise QuantTargetMismatch(
f"GGUF-quantized tensor '{t.name}' maps to '{module_path}' "
f"(missing), expected an nn.Linear. The sidecar config may "
f"not match this checkpoint."
)
if isinstance(target_module, torch.nn.Linear):
expected_shape = tuple(target_module.weight.shape)
if logical_shape != expected_shape:
raise QuantTargetMismatch(
f"GGUF tensor '{t.name}' shape {logical_shape} disagrees "
f"with model {type(target_module).__name__} shape "
f"{expected_shape}"
)
plan.append((t, "resident", module_path, None, 0))
resident_plan[module_path] = gguf_linear_factory(tt)
continue
target_param = getattr(target_module, "weight", None)
if isinstance(target_param, torch.Tensor):
expected_shape = tuple(target_param.shape)
if logical_shape != expected_shape:
raise QuantTargetMismatch(
f"GGUF tensor '{t.name}' shape {logical_shape} disagrees "
f"with model {type(target_module).__name__} shape "
f"{expected_shape}"
)
dst = (
target_param.dtype
if target_param.dtype.is_floating_point
else torch.float32
)
nbytes = 1
for s in expected_shape:
nbytes *= int(s)
plan.append((t, "dense", target_key, dst, nbytes * dst.itemsize))
dequant_load.append(target_key)
dequant_bytes += nbytes * dst.itemsize
continue
kind = type(target_module).__name__
raise QuantTargetMismatch(
f"GGUF-quantized tensor '{t.name}' maps to '{module_path}' "
f"({kind}), expected an nn.Linear. The sidecar config may "
f"not match this checkpoint."
)
if unsupported:
name, tt = unsupported[0]
raise UnsupportedGGMLType(tt, f"{name}" + (
f" (+{len(unsupported) - 1} more)" if len(unsupported) > 1 else ""
))
replaced = replace_linears_for_quant(model, resident_plan)
resident_set = set(replaced)
total_raw = 0
for t, kind, target, _dst, _nbytes in plan:
if kind != "resident" or target not in resident_set:
continue
module = resolve_module(model, target)
raw = GGUFTensor.from_reader_tensor(t).raw
# The raw blocks ARE the model's storage for a quant-resident Linear,
# so they belong where the rest of the model belongs. Left on the host
# they cost ~8 GB of private RAM and every forward takes the paged
# `cast_bias_weight` path; in VRAM they cost the same bytes, load
# flat, and the forward dequantizes in place. Dequantizing to float at
# load instead would cost twice the VRAM and lose the point of a GGUF
# checkpoint.
module.set_raw_weight(place_tensor_on_device(raw, target_device))
total_raw += module.weight.numel()
def _dense_pairs():
for t, kind, target, dst, _nbytes in plan:
if kind != "dense":
continue
if dst is None:
arr = np.ascontiguousarray(t.data)
tensor = torch.from_numpy(
arr if arr.flags.writeable else arr.copy()
).clone()
if t.data.dtype == np.uint8:
tensor = tensor.view(torch.bfloat16)
shape = tuple(int(s) for s in reversed(t.shape))
yield target, tensor.reshape(shape)
else:
yield target, dequantize_reader_tensor(t, dst)
known_missing = {f"{p}.weight" for p in replaced}
VibeVoiceLoader._stream_apply_dense(
model, _dense_pairs(), known_missing=known_missing,
target_device=target_device,
)
if dequant_load:
logging.debug(
f"[VibeVoice TTS] Dequantized {len(dequant_load)} non-Linear GGUF weight(s) at "
f"load (embeddings/conv heads): "
f"{', '.join(dequant_load[:5])}"
+ (f" (+{len(dequant_load) - 5} more)" if len(dequant_load) > 5 else "")
)
return {
"n_resident_layers": len(replaced),
"raw_bytes": total_raw,
"n_dequant_load": len(dequant_load),
"dequant_load_bytes": int(dequant_bytes),
"weight_family": "gguf_block",
}
# ====================================================================
# Low-bit / naive quantization detection (defensive warning)
# ====================================================================
_SAFETENSORS_INT_DTYPES = {
"I8", "U8", "I16", "U16", "I32", "U32", "I64", "U64", "I4", "U4",
}
_QUANT_SCALE_NAME_HINTS = ("scale", "zero_point", "qzero", "absmax")
_LOWBIT_FRACTION_THRESHOLD = 0.10
def _inspect_safetensors_quantization(weight_path: str):
import struct
try:
with open(weight_path, "rb") as f:
(header_len,) = struct.unpack("<Q", f.read(8))
header = json.loads(f.read(header_len))
except Exception:
return None
header.pop("__metadata__", None)
total = len(header)
int_count = 0
has_scale_meta = False
for name, info in header.items():
if info.get("dtype", "") in _SAFETENSORS_INT_DTYPES:
int_count += 1
low = name.lower()
if any(hint in low for hint in _QUANT_SCALE_NAME_HINTS):
has_scale_meta = True
return {"total": total, "int_count": int_count, "has_scale_meta": has_scale_meta}
def _inspect_gguf_quantization(weight_path: str):
try:
import gguf # noqa: F401
from gguf.constants import GGMLQuantizationType
except ImportError:
return None
try:
from .gguf_quant import open_gguf_reader
reader = open_gguf_reader(weight_path)
except Exception:
return None
low_bit_types = {
GGMLQuantizationType.IQ1_S,
GGMLQuantizationType.IQ1_M,
GGMLQuantizationType.IQ2_XXS,
GGMLQuantizationType.IQ2_XS,
GGMLQuantizationType.IQ2_S,
GGMLQuantizationType.IQ3_XXS,
GGMLQuantizationType.IQ3_S,
}
total = len(reader.tensors)
low_bit_count = sum(1 for t in reader.tensors if t.tensor_type in low_bit_types)
return {"total": total, "low_bit_count": low_bit_count}
def warn_if_lowbit_quantization(weight_path: str) -> None:
lower = weight_path.lower()
if lower.endswith(".safetensors"):
info = _inspect_safetensors_quantization(weight_path)
if info is None or info["total"] == 0:
return
frac = info["int_count"] / info["total"]
if frac >= _LOWBIT_FRACTION_THRESHOLD and not info["has_scale_meta"]:
logging.warning(
f"[VibeVoice TTS] [low-bit check] '{os.path.basename(weight_path)}' stores "
f"{info['int_count']}/{info['total']} tensors as raw integers with "
f"NO dequantization scale/zero-point metadata. Use a proper "
f"quantized or full-precision (BF16/FP16) checkpoint instead."
)
return
if lower.endswith(".gguf"):
info = _inspect_gguf_quantization(weight_path)
if info is None or info["total"] == 0:
return
frac = info["low_bit_count"] / info["total"]
if frac >= _LOWBIT_FRACTION_THRESHOLD:
logging.warning(
f"[VibeVoice TTS] [low-bit check] '{os.path.basename(weight_path)}' contains "
f"{info['low_bit_count']}/{info['total']} tensors at sub-4-bit "
f"I-quant precision. Consider a higher-quality quant "
f"(Q4_K_M / Q5_K_M / Q8_0 / BF16)."
)
return
# ====================================================================
# In-memory state dict loading
# ====================================================================
def _load_state_dict_into_model_from_memory(
model, state_dict: dict, preserve_file_views: bool = True, target_device=None,
):
"""Load an in-memory state dict into an already-instantiated model.
Preserves zero-copy memory-mapped file views directly in parameters.
When target_device is specified (CUDA), transfers parameters directly
to GPU VRAM without keeping an intermediate copy in system RAM.
"""
if not preserve_file_views:
for key in list(state_dict.keys()):
tensor = state_dict[key]
if isinstance(tensor, torch.Tensor):
state_dict[key] = tensor.clone()
VibeVoiceLoader._apply_state_dict(model, state_dict)
if target_device is not None and target_device.type == "cuda":
model.to(target_device)
return model
# ====================================================================
# Core external loading function
# ====================================================================
def _log_load_diagnostics(
*,
config_name: str,
requested_attention_mode: str,
resolved_attention_mode: str,
weight_family: str,
load_device,
) -> None:
if not diagnostics_enabled():
return
logging.info(
"[VibeVoice TTS] Load diagnostics: model='%s' family=%s requested_attention=%s "
"resolved_attention=%s device=%s",
config_name, weight_family, requested_attention_mode,
resolved_attention_mode, load_device,
)
def _reset_gguf_forward_counters() -> None:
try:
from .gguf_quant import reset_gguf_forward_counters
reset_gguf_forward_counters()
except Exception:
pass
def load_external_vibevoice_model(
weight_path: str,
config_name: str,
attention_mode: str = "eager",
use_llm_4bit: bool = False,
dtype_str: str = "auto",
device=None,
) -> dict:
"""Load a VibeVoice model from an external weight file directly into VRAM."""
config_name = normalize_config_name(config_name)
_reset_gguf_forward_counters()
if not os.path.isfile(weight_path):
raise FileNotFoundError(f"External weight file not found: {weight_path}")
try:
_stat = os.stat(weight_path)
source_mtime_ns = _stat.st_mtime_ns
source_size = _stat.st_size
except OSError:
source_mtime_ns = 0
source_size = 0
warn_if_lowbit_quantization(weight_path)
if is_asr_config_name(config_name):
return load_external_vibevoice_asr_model(
weight_path=weight_path,
config_name=config_name,
attention_mode=attention_mode,
dtype_str=dtype_str,
device=device,
)
requested_attention_mode = attention_mode
attention_mode = resolve_attention_mode(attention_mode, use_llm_4bit)
lower_path = weight_path.lower()
is_gguf_file = lower_path.endswith(".gguf")
gguf_reader = _open_gguf_reader(weight_path) if is_gguf_file else None
convrot_quant_map = {}
if gguf_reader is None and lower_path.endswith(".safetensors"):
try:
from .convrot_quant import scan_checkpoint_quantization
convrot_quant_map = scan_checkpoint_quantization(weight_path)
except UnsupportedQuantFormat:
raise
except Exception as e:
logging.debug(f"[VibeVoice TTS] ConvRot scan skipped for {weight_path}: {e}")
convrot_quant_map = {}
validate_weight_plan(
is_gguf_file=is_gguf_file,
convrot_quant_map=convrot_quant_map,
use_llm_4bit=use_llm_4bit,
attention_mode=attention_mode,
gguf_kquant_present=_gguf_kquant_present(gguf_reader) if gguf_reader else False,
)
from .config_detect import fingerprint_weights
weights_fp = fingerprint_weights(weight_path, gguf_reader=gguf_reader)
if config_name == AUTO_CONFIG_NAME:
config_name = resolve_auto_config_name(weight_path, weights_fp=weights_fp)
config_path = resolve_sidecar_config(weight_path, config_name)
config = VibeVoiceLoader._load_config(config_path, config_name)
effective_name, config_changed = reconcile_config(config_name, config, weights_fp)
if config_changed:
config_name = effective_name
config_path = resolve_sidecar_config(weight_path, config_name)
config = VibeVoiceLoader._load_config(config_path, config_name)
is_streaming = isinstance(config, VibeVoiceStreamingConfig)
if is_streaming:
logging.debug(f"[VibeVoice TTS] External model '{config_name}' detected as streaming model")
attention_mode = resolve_realtime_attention_mode(attention_mode)
tokenizer_dir = resolve_sidecar_tokenizer_dir(weight_path)
vibevoice_tokenizer = VibeVoiceLoader._load_tokenizer(tokenizer_dir, config_name)
preprocessor_path = resolve_sidecar_preprocessor(weight_path)
processor = VibeVoiceLoader._load_processor(
vibevoice_tokenizer, preprocessor_path, is_streaming=is_streaming
)
load_device = (
model_management.get_torch_device()
if not isinstance(device, torch.device)
else device
)
model_dtype = resolve_dtype(dtype_str, load_device)
quant_config = None
final_load_dtype = model_dtype
if use_llm_4bit:
bnb_compute_dtype = model_dtype
if attention_mode == "sage":
bnb_compute_dtype, final_load_dtype = torch.float32, torch.float32
quant_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_use_double_quant=True,
bnb_4bit_compute_dtype=bnb_compute_dtype,
)
attn_implementation_for_load = get_attn_implementation_for_load(attention_mode)
try:
logging.debug(
f"[VibeVoice TTS] Instantiating external VibeVoice model '{config_name}' with "
f"dtype={final_load_dtype}, attention='{attn_implementation_for_load}'"
)
model = VibeVoiceLoader._instantiate_model(
config=config,
is_streaming=is_streaming,
attn_implementation=attn_implementation_for_load,
final_load_dtype=final_load_dtype,
)
weight_family = "dense"
quant_stats = {}
# convrot_quant_map is only ever populated from a .safetensors scan, so
# this is bool(convrot_quant_map). The batch fallback below is kept
# only because its guard helpers are still exercised directly by
# tests; every real load goes through the streaming assign.
stream_quant_load = (
bool(convrot_quant_map)
and weight_path.lower().endswith(".safetensors")
)
if gguf_reader is not None:
_assert_gguf_lm_head_not_tied(config, gguf_reader)
quant_stats = _install_gguf_weights(
model, gguf_reader, target_device=load_device
)
weight_family = "gguf_block"
del gguf_reader
gguf_reader = None
gc.collect()
elif convrot_quant_map:
from .quant_common import replace_linears_for_quant
_assert_lm_head_not_tied(config, convrot_quant_map)
convrot_quant_map = _demote_nonlinear_fp8_residents(
model, convrot_quant_map
)
if stream_quant_load:
layer_plan, n_rowwise, n_fp8_resident = (
_plan_quantized_safetensors_load(convrot_quant_map)
)
replaced = replace_linears_for_quant(model, layer_plan)
_stream_apply_safetensors(
model,
weight_path,
convrot_quant_map,
target_device=load_device,
)
else:
with measured_load("dense-read-state-dict"):
state_dict = _load_weight_state_dict(weight_path, torch.device("cpu"))
layer_plan, n_rowwise, n_fp8_resident = (
_prepare_quantized_safetensors_load(state_dict, convrot_quant_map)
)
replaced = replace_linears_for_quant(model, layer_plan)
model = _load_state_dict_into_model_from_memory(
model,
state_dict,
preserve_file_views=True,
target_device=load_device,
)
del state_dict
state_dict = None
weight_family = (
"fp8_resident"
if n_fp8_resident and len(replaced) == n_fp8_resident
else "convrot_int8"
)
quant_stats = {
"n_resident_layers": len(replaced),
"n_rowwise_layers": n_rowwise,
"n_fp8_resident_layers": n_fp8_resident,
}
else:
_stream_apply_dense_safetensors(
model, weight_path, target_device=load_device
)
cast_model_to_dtype_if_needed(model, final_load_dtype)
if quant_config is not None:
try:
from transformers.integrations.bitsandbytes import replace_with_bnb_linear
except ImportError:
from transformers.utils.bitsandbytes import replace_with_bnb_linear
replace_with_bnb_linear(
model,
quantization_config=quant_config,
modules_to_not_convert=None,
)
if attention_mode == "sage":
if check_sage_attention_compatible():
set_sage_attention(model)
else:
raise RuntimeError("Incompatible hardware/setup for SageAttention.")
model.eval()
setattr(model, "_llm_4bit", bool(quant_config))
report_census(model, phase=f"pre-h2d:{config_name}")
_log_load_diagnostics(
config_name=config_name,
requested_attention_mode=requested_attention_mode,
resolved_attention_mode=attention_mode,
weight_family=weight_family,
load_device=load_device,
)
return {
"config": config,
"processor": processor,
"model": model,
"model_name": config_name,
"source_path": weight_path,
"source_mtime_ns": source_mtime_ns,
"source_size": source_size,
"attention_mode": attention_mode,
"use_llm_4bit": bool(use_llm_4bit),
"dtype_str": dtype_str,
"is_streaming": is_streaming,
"is_asr": False,
"weight_family": weight_family,
"dynamic_vram_route": False,
"quant_stats": quant_stats,
}
except Exception as e:
logging.error(
f"[VibeVoice TTS] Failed to load external VibeVoice model '{config_name}' "
f"from {weight_path}: {e}"
)
raise RuntimeError(
f"Failed to load external VibeVoice model '{config_name}': {e}"
)
# ====================================================================
# ASR external loading branch
# ====================================================================
def load_external_vibevoice_asr_model(
weight_path: str,
config_name: str,
attention_mode: str = "sdpa",
dtype_str: str = "auto",
device=None,
) -> dict:
config_name = normalize_config_name(config_name)
_reset_gguf_forward_counters()
if not os.path.isfile(weight_path):
raise FileNotFoundError(f"External weight file not found: {weight_path}")
try:
_stat = os.stat(weight_path)
source_mtime_ns = _stat.st_mtime_ns
source_size = _stat.st_size
except OSError:
source_mtime_ns = 0
source_size = 0
warn_if_lowbit_quantization(weight_path)
requested_attention_mode = attention_mode
attention_mode = resolve_asr_attention_mode(
resolve_attention_mode(attention_mode, quantize_4bit=False)
)
lower_path = weight_path.lower()
is_gguf_file = lower_path.endswith(".gguf")
gguf_reader = _open_gguf_reader(weight_path) if is_gguf_file else None
convrot_quant_map = {}
if gguf_reader is None and lower_path.endswith(".safetensors"):
try:
from .convrot_quant import scan_checkpoint_quantization
convrot_quant_map = scan_checkpoint_quantization(weight_path)
except UnsupportedQuantFormat:
raise
except Exception as e:
logging.debug(f"[VibeVoice TTS] ConvRot scan skipped for {weight_path}: {e}")
convrot_quant_map = {}
validate_weight_plan(
is_gguf_file=is_gguf_file,
convrot_quant_map=convrot_quant_map,
use_llm_4bit=False,
attention_mode=attention_mode,
gguf_kquant_present=_gguf_kquant_present(gguf_reader) if gguf_reader else False,
)
config_path = resolve_sidecar_config(weight_path, config_name)
config = _load_asr_config(config_path)
tokenizer_dir = resolve_sidecar_tokenizer_dir(weight_path)
preprocessor_path = resolve_sidecar_preprocessor(weight_path)
if is_native_asr_config_path(config_path):
processor = build_native_asr_processor(tokenizer_dir, preprocessor_path)
else:
processor = _load_asr_processor(
_load_asr_tokenizer(tokenizer_dir), preprocessor_path
)
load_device = (
model_management.get_torch_device()
if not isinstance(device, torch.device)
else device
)
model_dtype = resolve_dtype(dtype_str, load_device)
attn_implementation_for_load = get_attn_implementation_for_load(attention_mode)
try:
logging.debug(
f"[VibeVoice TTS] Instantiating external VibeVoice ASR model '{config_name}' with "
f"dtype={model_dtype}, attention='{attn_implementation_for_load}'"
)
model = _instantiate_asr_model(
config=config,
attn_implementation=attn_implementation_for_load,
final_load_dtype=model_dtype,
)
weight_family = "dense"
quant_stats = {}
# convrot_quant_map is only ever populated from a .safetensors scan, so
# this is bool(convrot_quant_map). The batch fallback below is kept
# only because its guard helpers are still exercised directly by
# tests; every real load goes through the streaming assign.
stream_quant_load = (
bool(convrot_quant_map)
and weight_path.lower().endswith(".safetensors")
)
if gguf_reader is not None:
_assert_gguf_lm_head_not_tied(config, gguf_reader)
quant_stats = _install_gguf_weights(
model, gguf_reader, target_device=load_device
)
weight_family = "gguf_block"
del gguf_reader
gguf_reader = None
gc.collect()
elif convrot_quant_map:
from .quant_common import replace_linears_for_quant
_assert_lm_head_not_tied(config, convrot_quant_map)
convrot_quant_map = _demote_nonlinear_fp8_residents(
model, convrot_quant_map
)
if stream_quant_load:
layer_plan, n_rowwise, n_fp8_resident = (
_plan_quantized_safetensors_load(convrot_quant_map)
)
replaced = replace_linears_for_quant(model, layer_plan)
_stream_apply_safetensors(
model,
weight_path,
convrot_quant_map,
target_device=load_device,
)
else:
with measured_load("dense-read-state-dict"):
state_dict = _load_weight_state_dict(weight_path, torch.device("cpu"))
layer_plan, n_rowwise, n_fp8_resident = (
_prepare_quantized_safetensors_load(state_dict, convrot_quant_map)
)
replaced = replace_linears_for_quant(model, layer_plan)
model = _load_state_dict_into_model_from_memory(
model,
state_dict,
preserve_file_views=True,
target_device=load_device,
)
del state_dict
state_dict = None
weight_family = (
"fp8_resident"
if n_fp8_resident and len(replaced) == n_fp8_resident
else "convrot_int8"
)
quant_stats = {
"n_resident_layers": len(replaced),
"n_rowwise_layers": n_rowwise,
"n_fp8_resident_layers": n_fp8_resident,
}
else:
_stream_apply_dense_safetensors(
model, weight_path, target_device=load_device
)
cast_model_to_dtype_if_needed(model, model_dtype)
if attention_mode == "sage":
if check_sage_attention_compatible():
set_sage_attention(model)
else:
raise RuntimeError("Incompatible hardware/setup for SageAttention.")
model.eval()
_log_load_diagnostics(
config_name=config_name,
requested_attention_mode=requested_attention_mode,
resolved_attention_mode=attention_mode,
weight_family=weight_family,
load_device=load_device,
)
return {
"config": config,
"processor": processor,
"model": model,
"model_name": config_name,
"source_path": weight_path,
"source_mtime_ns": source_mtime_ns,
"source_size": source_size,
"attention_mode": attention_mode,
"use_llm_4bit": False,
"dtype_str": dtype_str,
"is_streaming": False,
"is_asr": True,
"weight_family": weight_family,
"dynamic_vram_route": False,
"quant_stats": quant_stats,
}
except Exception as e:
logging.error(
f"[VibeVoice TTS] Failed to load external VibeVoice ASR model '{config_name}' "
f"from {weight_path}: {e}"
)
raise RuntimeError(
f"Failed to load external VibeVoice ASR model '{config_name}': {e}"
)