Files
Retrieval-based-Voice-Conve…/infer/hubert.py
2026-07-19 21:19:30 +08:00

97 lines
3.0 KiB
Python

import logging
from functools import lru_cache
from pathlib import Path
import torch
from torch import nn
from transformers import AutoFeatureExtractor, HubertModel
logger = logging.getLogger(__name__)
PROJECT_ROOT = Path(__file__).resolve().parent.parent
class HubertModelWithFinalProj(HubertModel):
def __init__(self, config):
super().__init__(config)
self.final_proj = nn.Linear(config.hidden_size, config.classifier_proj_size)
HUBERT_MODEL_PATH = (PROJECT_ROOT / "assets" / "hubert_base").resolve()
def _device_type(device):
if isinstance(device, torch.device):
return device.type
return str(device).split(":", 1)[0]
def load_hubert_model(device, is_half=False):
"""Load the local Transformers HuBERT/ContentVec model for RVC."""
if not (HUBERT_MODEL_PATH / "config.json").is_file():
raise FileNotFoundError(
f"Transformers HuBERT model not found: {HUBERT_MODEL_PATH}"
)
dtype = torch.float16 if is_half else torch.float32
load_options = {
"local_files_only": True,
"torch_dtype": dtype,
}
# DirectML does not implement every SDPA kernel used by Transformers.
if _device_type(device) == "privateuseone":
load_options["attn_implementation"] = "eager"
logger.info(
"Loading Transformers HuBERT from %s (%s on %s)",
HUBERT_MODEL_PATH,
dtype,
device,
)
model = HubertModelWithFinalProj.from_pretrained(
str(HUBERT_MODEL_PATH), **load_options
)
model = model.to(device)
return model.eval()
@lru_cache(maxsize=1)
def hubert_audio_requires_normalization():
feature_extractor = AutoFeatureExtractor.from_pretrained(
str(HUBERT_MODEL_PATH), local_files_only=True
)
return bool(feature_extractor.do_normalize)
def extract_hubert_features(model, source, version, padding_mask=None):
"""Return the RVC v1 (256-D) or v2 (768-D) HuBERT representation.
Transformers hidden_states[N] is numerically equivalent to the source checkpoint's
output_layer=N for this converted checkpoint. RVC v1 uses layer 9 followed
by final_proj; RVC v2 uses the final (12th) encoder layer directly.
"""
if version not in {"v1", "v2"}:
raise ValueError(f"Unsupported RVC feature version: {version!r}")
attention_mask = None
if padding_mask is not None and bool(torch.any(padding_mask).item()):
attention_mask = (~padding_mask.bool()).long()
if version == "v1":
outputs = model(
input_values=source,
attention_mask=attention_mask,
output_hidden_states=True,
return_dict=True,
)
features = outputs.hidden_states[9]
return model.final_proj(features)
outputs = model(
input_values=source,
attention_mask=attention_mask,
output_hidden_states=False,
return_dict=True,
)
return outputs.last_hidden_state