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

206 lines
6.3 KiB
Python

"""Dtype utilities for VibeVoice nodes.
Provides dtype resolution and casting helpers using ComfyUI's
model_management functions for consistent precision selection.
"""
import torch
import logging
from typing import List
from .device_utils import should_use_bf16, should_use_fp16, _get_model_management
# Dtype string constants
DTYPE_AUTO = "auto"
DTYPE_BF16 = "bf16"
DTYPE_FP16 = "fp16"
DTYPE_FP32 = "fp32"
# Mapping from string to torch.dtype
_DTYPE_MAP = {
DTYPE_BF16: torch.bfloat16,
DTYPE_FP16: torch.float16,
DTYPE_FP32: torch.float32,
}
# Reverse mapping
_DTYPE_STR_MAP = {
torch.bfloat16: DTYPE_BF16,
torch.float16: DTYPE_FP16,
torch.float32: DTYPE_FP32,
}
def get_dtype_options() -> List[str]:
"""Get list of available dtype options for node dropdown.
Returns:
List of dtype option strings.
"""
return [DTYPE_AUTO, DTYPE_BF16, DTYPE_FP16, DTYPE_FP32]
def _torch_dtype_is_deprecated_property(config) -> bool:
"""True when ``type(config).torch_dtype`` is the transformers >= 5.0
deprecated property alias (which warns on every get/set)."""
return isinstance(getattr(type(config), "torch_dtype", None), property)
def set_config_dtype(config, dtype) -> None:
"""Record the load dtype on a transformers config without triggering the
transformers v5 deprecation warning.
"""
if _torch_dtype_is_deprecated_property(config):
config.dtype = dtype
else:
config.torch_dtype = dtype
def get_config_dtype(config):
"""Return the torch dtype recorded on a transformers config, or None."""
if _torch_dtype_is_deprecated_property(config):
dtype = getattr(config, "dtype", None)
else:
dtype = getattr(config, "torch_dtype", None)
if dtype is None:
dtype = getattr(config, "dtype", None)
if isinstance(dtype, str):
dtype = getattr(torch, dtype, None)
return dtype
def resolve_dtype(dtype_str: str, device: torch.device = None) -> torch.dtype:
"""Resolve a dtype string to a torch.dtype.
Args:
dtype_str: Dtype string ("auto", "bf16", "fp16", "fp32").
device: torch.device for "auto" resolution. If None, uses default device.
Returns:
torch.dtype instance.
Raises:
ValueError: If dtype_str is not recognized.
"""
if dtype_str == DTYPE_AUTO or dtype_str is None:
mm = _get_model_management()
if device is None:
device = mm.get_torch_device()
if should_use_bf16(device):
return torch.bfloat16
elif should_use_fp16(device):
return torch.float16
else:
return torch.float32
if dtype_str in _DTYPE_MAP:
return _DTYPE_MAP[dtype_str]
raise ValueError(f"Unknown dtype: {dtype_str}. Valid options: {list(_DTYPE_MAP.keys())}")
def get_dtype_str(dtype: torch.dtype) -> str:
"""Convert a torch.dtype to its string representation.
Args:
dtype: torch.dtype instance.
Returns:
Dtype string (e.g., "bf16", "fp16", "fp32").
Raises:
ValueError: If dtype is not recognized.
"""
if dtype in _DTYPE_STR_MAP:
return _DTYPE_STR_MAP[dtype]
raise ValueError(f"Unknown torch.dtype: {dtype}. Supported: {list(_DTYPE_STR_MAP.values())}")
def _quant_protected_names(model) -> set:
"""Parameter names excluded from bulk dtype casts.
- Parameters of modules marked ``_quant_resident`` (GGUFLinear / ConvRotInt8Linear / FP8Linear).
- Any parameter whose leaf name is ``weight_scale``.
"""
protected = set()
for name, mod in model.named_modules():
if getattr(mod, "_quant_resident", False):
prefix = f"{name}." if name else ""
protected.add(f"{prefix}weight")
protected.add(f"{prefix}weight_scale")
return protected
def representative_dtype(model):
"""First FLOATING parameter dtype (None when the tree has none)."""
try:
resident_modules = {
name for name, mod in model.named_modules()
if getattr(mod, "_quant_resident", False)
}
for name, p in model.named_parameters():
if not p.dtype.is_floating_point:
continue
owner = name.rpartition(".")[0]
if owner in resident_modules:
continue
return p.dtype
except Exception:
pass
return getattr(model, "dtype", None)
def _cast_mismatched_params(model, dtype: torch.dtype) -> None:
"""Cast only floating, unprotected, mismatched parameters, then re-tie."""
protected = _quant_protected_names(model)
for name, param in model.named_parameters():
if name in protected:
continue
if not param.dtype.is_floating_point:
continue
if param.dtype != dtype:
param.data = param.data.to(dtype)
config = getattr(model, "config", None)
if config is not None and hasattr(model, "tie_weights"):
decoder_config = getattr(config, "decoder_config", None)
tied = bool(getattr(decoder_config, "tie_word_embeddings", False)) or \
bool(getattr(config, "tie_word_embeddings", False))
if tied:
model.tie_weights()
def cast_model_to_dtype(model, dtype: torch.dtype) -> None:
"""Cast a model's FLOATING parameters to the specified dtype.
Quant-resident storage (raw uint8 GGUF blocks, int8 ConvRot weights,
fp32 ``weight_scale``) is NEVER touched.
"""
if dtype is None:
return
_cast_mismatched_params(model, dtype)
logging.debug(f"[VibeVoice TTS] Model cast to dtype (filtered): {dtype}")
def cast_model_to_dtype_if_needed(model, dtype: torch.dtype) -> None:
"""Cast a model to ``dtype`` only when some castable parameter mismatches."""
if dtype is None:
return
protected = _quant_protected_names(model)
mismatched = []
for name, param in model.named_parameters():
if name in protected or not param.dtype.is_floating_point:
continue
if param.dtype != dtype:
mismatched.append((name, param.dtype))
if not mismatched:
return
_cast_mismatched_params(model, dtype)
sources = ", ".join(sorted({str(src) for _, src in mismatched}))
logging.debug(
f"[VibeVoice TTS] Model cast {sources} -> {dtype} ({len(mismatched)} mismatched params)"
)