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

134 lines
5.1 KiB
Python

"""Shared quantized-linear replacement and weight-plan validation.
Used by BOTH quant-resident families:
- GGUF raw-block residency (:mod:`modules.gguf_quant`) — ``nn.Linear`` targets
are swapped for :class:`~modules.gguf_quant.GGUFLinear` before any weights
are assigned.
- ConvRot INT8 checkpoints (:mod:`modules.convrot_quant`) — targets are swapped
for :class:`~modules.convrot_quant.ConvRotInt8Linear`.
Replacement must run BEFORE weight assignment (the swapped modules own
differently-shaped/dtyped parameters). All checks are strict: a target that is
missing, not an ``nn.Linear``, or shape-inconsistent raises immediately rather
than silently misloading weights.
"""
from __future__ import annotations
import logging
import torch
class QuantTargetMismatch(RuntimeError):
"""Raised when a planned quantized-linear replacement cannot be applied."""
def resolve_module(model: torch.nn.Module, path: str) -> torch.nn.Module:
"""Resolve a dotted module path against ``model`` (KeyError if absent)."""
mod = model
for part in path.split("."):
if not hasattr(mod, part):
raise KeyError(path)
mod = getattr(mod, part)
return mod
def replace_linears_for_quant(model: torch.nn.Module, layer_plan: dict) -> list:
"""Swap every module named in ``layer_plan`` for a factory-built module.
Args:
model: The model tree to mutate.
layer_plan: mapping of dotted module path ->
callable ``(in_features, out_features, has_bias) -> nn.Module``.
Returns:
List of replaced module paths (order follows the plan).
Raises:
QuantTargetMismatch: On missing / non-Linear / shape-mismatched
targets. Nothing is replaced unless ALL targets validate first.
"""
modules = dict(model.named_modules())
built = []
for prefix, factory in layer_plan.items():
if prefix not in modules:
raise QuantTargetMismatch(
f"Quantized layer '{prefix}' not found in model tree"
)
target = modules[prefix]
if isinstance(target, torch.nn.Linear):
in_f = target.in_features
out_f = target.out_features
bias = target.bias is not None
else:
kind = type(target).__name__
raise QuantTargetMismatch(
f"Quantized layer '{prefix}' is {kind}, expected nn.Linear. "
f"Rotated INT8 quantization can only execute on Linear "
f"modules. If this checkpoint quantizes/rotates non-Linear "
f"modules (e.g. embeddings, norms), re-export it excluding "
f"those layers (e.g. the converter's skip-embeddings / "
f"--heur option)."
)
try:
# Construct replacements under meta context so zero RAM is committed.
# Real storage is assigned when weights are loaded.
with torch.device("meta"):
new_mod = factory(in_f, out_f, bias)
except Exception as e:
raise QuantTargetMismatch(
f"Failed to build quantized replacement for '{prefix}': {e}"
) from e
built.append((prefix, new_mod))
parent_cache = {}
replaced = []
for prefix, new_mod in built:
parent_name, _, child_name = prefix.rpartition(".")
if parent_name not in parent_cache:
parent_cache[parent_name] = (
model if not parent_name else resolve_module(model, parent_name)
)
setattr(parent_cache[parent_name], child_name, new_mod)
replaced.append(prefix)
logging.debug(f"[VibeVoice TTS] Replaced {len(replaced)} nn.Linear(s) with quant-resident modules")
return replaced
def validate_weight_plan(
*,
is_gguf_file: bool,
convrot_quant_map: dict | None,
use_llm_4bit: bool,
attention_mode: str,
gguf_kquant_present: bool = False,
) -> None:
"""Single choke point for quantization-family exclusivity rules.
Rules:
- GGUF file + ConvRot metadata present -> hard error (mutually exclusive).
- ConvRot checkpoint + bnb 4-bit requested -> hard error.
- SageAttention over GGUF K-quants -> allowed, warning only (sage patches
attention math only; K-quant linears are unaffected).
"""
if is_gguf_file and convrot_quant_map:
raise ValueError(
"Weight plan conflict: the file contains both GGUF tensors and "
"ConvRot (*.comfy_quant) metadata. These formats are mutually "
"exclusive; the checkpoint is likely corrupted or misconverted."
)
if convrot_quant_map and use_llm_4bit:
raise ValueError(
"Weight plan conflict: the checkpoint already carries ConvRot INT8 "
"quantized linears; 'quantize_llm_4bit' cannot be combined with it."
)
if attention_mode == "sage" and is_gguf_file and gguf_kquant_present:
logging.warning(
"[VibeVoice TTS] SageAttention requested alongside GGUF K-quants: sage patches "
"attention computation only, quantized linear layers still run "
"through per-matmul dequantization."
)