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

436 lines
16 KiB
Python

"""Shared generation utilities for VibeVoice nodes.
Contains common functions for model loading, audio processing,
and generation to avoid code duplication across nodes.
"""
import torch
import gc
import logging
import numpy as np
from typing import Optional, Tuple, Any
import comfy.model_management as model_management
from .progress_utils import ProgressBarWithConsole
from .loader import VibeVoiceModelHandler, VibeVoiceLoader, cleanup_old_models, LOADED_MODELS_CACHE
from .patcher import VibeVoicePatcher, load_to_device, select_patcher_class
from .model_registry import (
FAMILY_TTS,
evict_if_changed,
identity_for_external,
register_model_bundle,
)
from .model_info import is_model_type
from .utils import VIBEVOICE_PATCHER_CACHE
from .audio_utils import parse_script_1_based, preprocess_comfy_audio, set_seed, check_for_interrupt
from .device_utils import get_torch_device, get_offload_device, DEVICE_CPU
from .dtype_utils import resolve_dtype, DTYPE_AUTO
from .attention_utils import resolve_attention_mode, resolve_realtime_attention_mode
from .gguf_quant import log_gguf_forward_counters
from .memory_census import measured_load, report_census
from .diagnostics import diagnostics_enabled
def resolve_generation_family(
model_name: str,
external_model: dict | None = None,
) -> str:
"""Return exactly ``tts`` or ``streaming_tts`` for generation routing."""
if external_model is not None:
if external_model.get("is_asr"):
raise ValueError(
"ASR models cannot generate speech. Use the 'VibeVoice ASR' node."
)
if external_model.get("is_streaming"):
return "streaming_tts"
return "tts"
if not model_name:
raise ValueError(
"No VibeVoice model was selected. Select a TTS model or use the "
"'VibeVoice ASR' node for an ASR model."
)
if is_model_type(model_name, "streaming_tts"):
return "streaming_tts"
if is_model_type(model_name, "tts"):
return "tts"
raise ValueError(
f"Model '{model_name}' has unsupported type 'asr'. Use the "
"'VibeVoice ASR' node for ASR models."
)
class ExternalVibeVoiceModelHandler(torch.nn.Module):
"""Handler for an externally-loaded (pre-instantiated) VibeVoice model."""
def __init__(
self,
model,
processor,
model_pack_name: str,
attention_mode: str = "sdpa",
source_path: str = "",
):
super().__init__()
self.model = model
self.processor = processor
self.model_pack_name = model_pack_name
self.attention_mode = attention_mode
self.source_path = source_path
self.cache_key = f"external_{model_pack_name}_attn_{attention_mode}"
self.device = None
self.size = self._estimate_size(model)
@staticmethod
def _estimate_size(model) -> int:
"""Estimate the model's VRAM footprint in bytes from its parameters."""
try:
total = 0
for p in model.parameters():
total += p.numel() * p.element_size()
if total > 0:
return total
except Exception:
pass
return int(4.0 * (1024**3))
def load_model(self, device, attention_mode: str = "sdpa"):
"""No-op: the model is already loaded."""
logging.debug(
f"[VibeVoice TTS] ExternalVibeVoiceModelHandler.load_model called but model is "
f"already loaded for '{self.model_pack_name}'"
)
def load_vibevoice_from_external(
model_bundle: dict,
device: str = "auto",
dtype: str = DTYPE_AUTO,
attention_mode: str = "sdpa",
) -> Tuple[VibeVoicePatcher, Any, Any]:
"""Wrap an externally-loaded VibeVoice model bundle in a patcher and load to VRAM."""
if model_bundle.get("model") is None and model_bundle.get("source_path"):
from .external_loader import load_external_vibevoice_model
model_bundle = load_external_vibevoice_model(
model_bundle["source_path"],
model_bundle.get("model_name") or "",
attention_mode=model_bundle.get("attention_mode") or "eager",
use_llm_4bit=bool(model_bundle.get("use_llm_4bit", False)),
dtype_str=model_bundle.get("dtype_str") or "auto",
)
for required_key in ("model", "processor", "model_name"):
if model_bundle.get(required_key) is None:
raise ValueError(
f"External VibeVoice model bundle is missing required key "
f"'{required_key}'. Got keys: {list(model_bundle.keys())}"
)
model = model_bundle["model"]
processor = model_bundle["processor"]
model_name = model_bundle["model_name"]
source_path = model_bundle.get("source_path", "")
bundle_attention = model_bundle.get("attention_mode")
if isinstance(bundle_attention, str) and bundle_attention:
actual_attention_mode = bundle_attention
else:
actual_attention_mode = resolve_attention_mode(attention_mode, False)
if device == DEVICE_CPU:
load_device = torch.device(DEVICE_CPU)
offload_device = torch.device(DEVICE_CPU)
else:
load_device = get_torch_device(device)
offload_device = get_offload_device()
target_dtype = resolve_dtype(dtype, load_device)
bundle_use_llm_4bit = bool(model_bundle.get("use_llm_4bit", False))
bundle_dtype_str = model_bundle.get("dtype_str") or dtype
cache_key = identity_for_external(
source_path,
model_name,
actual_attention_mode,
use_llm_4bit=bundle_use_llm_4bit,
dtype_str=bundle_dtype_str,
)
register_model_bundle(cache_key, model_bundle)
evict_if_changed(FAMILY_TTS, cache_key, (VIBEVOICE_PATCHER_CACHE,))
if cache_key not in VIBEVOICE_PATCHER_CACHE:
model_handler = ExternalVibeVoiceModelHandler(
model=model,
processor=processor,
model_pack_name=model_name,
attention_mode=actual_attention_mode,
source_path=source_path,
)
model_handler.cache_key = cache_key
patcher_cls = select_patcher_class(
model_bundle.get("weight_family"), load_device, legacy_cls=VibeVoicePatcher
)
patcher = patcher_cls(
model_handler,
attention_mode=actual_attention_mode,
load_device=load_device,
offload_device=offload_device,
size=model_handler.size,
dtype=target_dtype,
)
VIBEVOICE_PATCHER_CACHE[cache_key] = patcher
logging.debug(
f"[VibeVoice TTS] Created new external patcher for {model_name} with "
f"attn={actual_attention_mode}"
)
patcher = VIBEVOICE_PATCHER_CACHE[cache_key]
with measured_load("load-to-device"):
load_to_device(patcher)
report_census(patcher.model.model, patcher, phase=f"post-h2d:{model_name}")
loaded_model = patcher.model.model
loaded_processor = patcher.model.processor
if loaded_model is None or loaded_processor is None:
raise RuntimeError(
f"External VibeVoice model and processor could not be loaded for "
f"'{model_name}'. Check logs for errors."
)
return patcher, loaded_model, loaded_processor
def load_vibevoice_model(
model_name: str,
device: str = "auto",
dtype: str = DTYPE_AUTO,
attention_mode: str = "sdpa",
quantize_4bit: bool = False,
force_reload: bool = False,
) -> Tuple[VibeVoicePatcher, Any, Any]:
"""Load or retrieve cached VibeVoice model."""
actual_attention_mode = resolve_attention_mode(attention_mode, quantize_4bit)
if is_model_type(model_name, "streaming_tts"):
actual_attention_mode = resolve_realtime_attention_mode(actual_attention_mode)
if device == DEVICE_CPU:
load_device = torch.device(DEVICE_CPU)
offload_device = torch.device(DEVICE_CPU)
else:
load_device = get_torch_device(device)
offload_device = get_offload_device()
target_dtype = resolve_dtype(dtype, load_device)
cache_key = f"{model_name}_attn_{actual_attention_mode}_q4_{int(quantize_4bit)}"
evict_if_changed(FAMILY_TTS, cache_key, (VIBEVOICE_PATCHER_CACHE,))
if cache_key not in VIBEVOICE_PATCHER_CACHE or force_reload:
if force_reload:
cleanup_old_models(keep_cache_key=cache_key)
model_handler = VibeVoiceModelHandler(
model_name,
attention_mode=actual_attention_mode,
use_llm_4bit=quantize_4bit,
dtype_str=dtype,
)
patcher_cls = select_patcher_class(None, load_device, legacy_cls=VibeVoicePatcher)
patcher = patcher_cls(
model_handler,
attention_mode=actual_attention_mode,
load_device=load_device,
offload_device=offload_device,
size=model_handler.size,
dtype=target_dtype,
)
VIBEVOICE_PATCHER_CACHE[cache_key] = patcher
logging.debug(f"[VibeVoice TTS] Created new patcher for {model_name} with attn={actual_attention_mode}, q4={quantize_4bit}")
patcher = VIBEVOICE_PATCHER_CACHE[cache_key]
with measured_load("load-to-device"):
load_to_device(patcher)
report_census(patcher.model.model, patcher, phase=f"post-h2d:{model_name}")
model = patcher.model.model
processor = patcher.model.processor
if model is None or processor is None:
raise RuntimeError(
f"VibeVoice model and processor could not be loaded for '{model_name}'. Check logs for errors."
)
return patcher, model, processor
def generate_audio(
model: Any,
processor: Any,
text: str,
voice_samples: list,
speaker_ids: list,
cfg_scale: float = 1.3,
inference_steps: int = 10,
seed: int = 42,
do_sample: bool = True,
temperature: float = 0.95,
top_p: float = 0.95,
top_k: int = 0,
max_new_tokens: Optional[int] = None,
) -> Tuple[torch.Tensor, int]:
"""Generate audio using the VibeVoice model."""
_streaming_class_names = {
"VibeVoiceStreamingProcessor",
"VibeVoiceStreamingForConditionalGenerationInference",
}
if (
type(processor).__name__ in _streaming_class_names
or type(model).__name__ in _streaming_class_names
):
raise ValueError(
"A streaming (realtime) VibeVoice model was passed to generate_audio(). "
"Use the canonical VibeVoice TTS realtime family path for this model."
)
parsed_lines_0_based, speaker_ids_1_based = parse_script_1_based(text)
if not parsed_lines_0_based:
raise ValueError("Script is empty or invalid. Please provide text to generate.")
voice_samples_np = []
for vs in voice_samples:
processed = preprocess_comfy_audio(vs)
if processed is not None:
processed = np.asarray(processed, dtype=np.float32)
if processed.ndim == 0:
logging.warning("[VibeVoice TTS] Voice sample is a scalar (0-d array), skipping")
continue
if processed.ndim > 1:
processed = np.squeeze(processed)
if processed.ndim != 1:
logging.warning(f"[VibeVoice TTS] Voice sample has unexpected shape {processed.shape}, skipping")
continue
voice_samples_np.append(processed)
if not voice_samples_np:
raise ValueError(
"No valid voice samples provided. Please connect at least one audio input "
"as a voice sample for the speaker(s)."
)
set_seed(seed)
normalized_script = "\n".join(
f"Speaker {speaker_id + 1}:{speaker_text}"
for speaker_id, speaker_text in parsed_lines_0_based
)
inputs = processor(
text=[normalized_script],
voice_samples=[voice_samples_np],
padding=True,
return_tensors="pt",
return_attention_mask=True,
)
for key, value in inputs.items():
if isinstance(value, torch.Tensor):
if torch.any(torch.isnan(value)) or torch.any(torch.isinf(value)):
logging.error(f"[VibeVoice TTS] Input tensor '{key}' contains NaN or Inf values")
raise ValueError(f"Invalid values in input tensor: {key}")
compute_device = model_management.get_torch_device()
inputs = {
k: v.to(compute_device) if isinstance(v, torch.Tensor) else v
for k, v in inputs.items()
}
model.set_ddpm_inference_steps(num_steps=inference_steps)
gen_inputs = {
"input_ids": inputs.get("input_ids"),
"attention_mask": inputs.get("attention_mask"),
"speech_tensors": inputs.get("speech_tensors"),
"speech_masks": inputs.get("speech_masks"),
"acoustic_input_mask": inputs.get("speech_input_mask"),
"cfg_scale": cfg_scale,
"inference_steps": inference_steps,
"return_speech": True,
"tokenizer": processor.tokenizer,
"do_sample": do_sample,
"temperature": temperature,
"top_p": top_p,
"top_k": top_k,
"max_new_tokens": max_new_tokens,
}
gen_inputs = {k: v for k, v in gen_inputs.items() if v is not None}
with torch.no_grad():
pbar = ProgressBarWithConsole(inference_steps)
def _progress(current: int, total: int) -> None:
model_management.throw_exception_if_processing_interrupted()
pbar.update_absolute(current, total=total)
try:
from .comfy_stream import pull_stats_line, reset_pull_stats
reset_pull_stats()
with measured_load("tts-generate") as _rss:
_rss.mark("gen-enter")
outputs = model.generate(**gen_inputs, progress_callback=_progress)
_rss.mark("gen-return")
if diagnostics_enabled():
logging.info(f"[VibeVoice TTS] {pull_stats_line()}")
except model_management.InterruptProcessingException:
logging.info("[VibeVoice TTS] VibeVoice generation interrupted by user")
raise
finally:
pbar.update_absolute(pbar.total)
pbar.close()
log_gguf_forward_counters("tts_generate")
speech_outputs = outputs.speech_outputs
if not speech_outputs or speech_outputs[0] is None:
raise RuntimeError(
"VibeVoice generation produced no audio. The model emitted no speech "
"tokens during the autoregressive loop. This usually means the loaded "
"weights are corrupt or over-quantized (e.g. a naive int8 cast without "
"dequantization scales, or an extremely low-bit GGUF quant). Try a "
"higher-quality checkpoint (BF16 / FP16 / Q8_0 / Q4_K_M)."
)
output_waveform = speech_outputs[0]
if output_waveform.ndim == 1:
output_waveform = output_waveform.unsqueeze(0)
if output_waveform.ndim == 2:
output_waveform = output_waveform.unsqueeze(0)
sample_rate = 24000
return output_waveform.detach().to(device="cpu", dtype=torch.float32), sample_rate
def force_offload_model(patcher: VibeVoicePatcher, model_name: str, warm: bool = False) -> None:
"""Force offload a VibeVoice model from VRAM."""
logging.info(f"[VibeVoice TTS] Force offloading VibeVoice model '{model_name}' from VRAM...")
if patcher.is_loaded:
if warm:
patcher.unpatch_model(unpatch_weights=True, warm=True)
else:
patcher.unpatch_model(unpatch_weights=True, destroy=True)
model_management.unload_all_models()
gc.collect()
model_management.soft_empty_cache()
logging.info("[VibeVoice TTS] Model force offload completed")