ellie-Bert-VITS2 / processing_bert_vits2.py
hans00's picture
Fix transformers v5 compatibility and rework BERT feature expansion
a85a126 verified
Raw History Blame Contribute Delete
15.8 kB
# coding=utf-8
# Copyright 2023 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Processor class for Bert VITS2."""
import os
import re
from typing import Any, Dict, List, Optional, Sequence, Union
from transformers import AutoTokenizer, PreTrainedTokenizerBase
from transformers.processing_utils import ProcessorMixin
from transformers.tokenization_utils_base import BatchEncoding
from transformers.utils import logging
try:
# `trust_remote_code` loads these files as a package, and only relative imports make transformers fetch the
# sibling modules alongside this one.
from .configuration_bert_vits2 import BertVits2Config
from .tokenization_bert_vits2 import BertVits2Tokenizer
except ImportError:
# Imported as a standalone module (e.g. after putting this directory on sys.path).
from configuration_bert_vits2 import BertVits2Config
from tokenization_bert_vits2 import BertVits2Tokenizer
logger = logging.get_logger(__name__)
# Make the tokenizer resolvable by `AutoTokenizer` when this package is imported directly (rather than through
# `trust_remote_code`, which wires the mapping up from `auto_map` instead).
try:
AutoTokenizer.register(BertVits2Config, slow_tokenizer_class=BertVits2Tokenizer)
except (AttributeError, TypeError, ValueError) as error: # pragma: no cover - depends on transformers version
logger.debug(f"Could not register BertVits2Tokenizer with AutoTokenizer: {error}")
CHINESE_DIGITS = ("零", "一", "二", "三", "四", "五", "六", "七", "八", "九")
CHINESE_SMALL_UNITS = ("", "十", "百", "千")
CHINESE_LARGE_UNITS = ("", "萬", "億", "兆", "京")
MAX_SUPPORTED_DIGITS = len(CHINESE_LARGE_UNITS) * 4
# Digits are read wherever they appear, including flush against Chinese text ("共123元"). A leading sign only
# counts when it starts a token, so "A-1" keeps its hyphen as punctuation while "-1" reads as 負一.
NUMBER_PATTERN = re.compile(r"(?:(?<![\w.])(?P<sign>[+-]))?(?<![\d.])(?P<integer>\d+)(?:\.(?P<decimal>\d+))?")
THOUSANDS_PATTERN = re.compile(r"(?<!\d)(\d{1,3}(?:,\d{3})+)(?!\d)")
PUNCTUATION_MAP = {
",": ",",
"、": ",",
";": ",",
":": ",",
"。": ".",
"?": "?",
"!": "!",
"(": ",",
")": ",",
"「": ",",
"」": ",",
"《": ",",
"》": ",",
"-": "-",
"~": "…",
}
FULLWIDTH_DIGITS = {chr(0xFF10 + i): str(i) for i in range(10)}
def chinese_number_to_words(text: str, style: str = "cardinal") -> str:
"""Read an ASCII number as Mandarin.
Args:
text: A number such as `"254"`, `"-3.5"` or `"+7"`.
style: `"cardinal"` reads 254 as 二百五十四; `"digit"` reads it as 二五四, which is how serial numbers,
years and phone numbers are usually announced.
"""
match = NUMBER_PATTERN.fullmatch(text)
if match is None:
raise ValueError(f"{text!r} is not a number this function can read.")
sign, integer, decimal = match.group("sign"), match.group("integer"), match.group("decimal")
out = {"-": "負", "+": "正"}.get(sign, "")
if style == "digit":
out += _read_digits(integer)
elif style == "cardinal":
out += _read_cardinal(integer)
else:
raise ValueError(f"Unknown number style {style!r}; expected 'cardinal' or 'digit'.")
if decimal:
out += "點" + _read_digits(decimal)
return out
def _read_digits(digits: str) -> str:
return "".join(CHINESE_DIGITS[int(digit)] for digit in digits)
def _read_cardinal(digits: str) -> str:
stripped = digits.lstrip("0")
if not stripped:
return CHINESE_DIGITS[0]
if len(stripped) > MAX_SUPPORTED_DIGITS:
# Beyond 京 there is no agreed reading; fall back to digit-by-digit rather than emit nonsense.
return _read_digits(stripped)
length = len(stripped)
out: List[str] = []
pending_zero = False
for i, char in enumerate(stripped):
power = length - i - 1
small, large = power % 4, power // 4
digit = int(char)
if digit == 0:
pending_zero = True
else:
if pending_zero and out:
out.append(CHINESE_DIGITS[0])
pending_zero = False
out.append(CHINESE_DIGITS[digit] + CHINESE_SMALL_UNITS[small])
# A 萬/億/兆 unit is only spoken when its four-digit group holds something.
if small == 0 and large > 0:
group = stripped[max(length - (large + 1) * 4, 0) : length - large * 4]
if any(c != "0" for c in group):
out.append(CHINESE_LARGE_UNITS[large])
pending_zero = False
spoken = "".join(out)
# 10-19 are read 十/十三, not 一十/一十三. Higher numbers keep the 一 (一百一十).
if length == 2 and stripped[0] == "1":
spoken = spoken[1:]
return spoken
class BertVits2Processor(ProcessorMixin):
r"""
Constructs a Bert VITS2 processor, which wraps the phoneme tokenizer and the per-language BERT tokenizers into a
single object.
The processor owns text normalization. Everything downstream assumes the phoneme sequence and the BERT token
sequence describe the *same* normalized string, so normalization has to happen once, here, before either
tokenizer sees the text.
Args:
tokenizer ([`BertVits2Tokenizer`]):
The phoneme tokenizer.
bert_tokenizers (`Dict[str, PreTrainedTokenizerBase]`):
One BERT tokenizer per supported language, keyed by language code.
number_style (`str`, *optional*, defaults to `"cardinal"`):
How digits are read aloud. `"cardinal"` reads `254` as 二百五十四, `"digit"` as 二五四.
"""
tokenizer_class = "AutoTokenizer"
attributes = ["tokenizer"]
def __init__(self, tokenizer: PreTrainedTokenizerBase, **kwargs):
# `bert_tokenizers` is deliberately not a ProcessorMixin attribute: it is a *mapping* of tokenizers, which
# the sub-processor machinery cannot type-check or serialise. It is saved through `to_dict` instead.
bert_tokenizers = kwargs.pop("bert_tokenizers", None) or {}
self.number_style = kwargs.pop("number_style", "cardinal")
kwargs.pop("languages", None) # accepted and ignored: superseded by the tokenizer's `languages`
super().__init__(tokenizer, **kwargs)
self.bert_tokenizers: Dict[str, PreTrainedTokenizerBase] = dict(bert_tokenizers)
# ------------------------------------------------------------------
# text normalization
# ------------------------------------------------------------------
def normalize(self, text: str, language: Optional[str] = None) -> str:
"""Normalize raw text into the form both tokenizers see."""
text = "".join(FULLWIDTH_DIGITS.get(char, char) for char in text)
for source, target in PUNCTUATION_MAP.items():
text = text.replace(source, target)
text = text.replace("...", "…")
text = re.sub(r"\s+", " ", text).strip()
if language is not None and language.startswith("zh"):
text = THOUSANDS_PATTERN.sub(lambda m: m.group(1).replace(",", ""), text)
text = NUMBER_PATTERN.sub(lambda m: chinese_number_to_words(m.group(0), self.number_style), text)
return text
def to_g2p_text(self, text: str, language: Optional[str] = None) -> str:
"""Replace whitespace with the phoneme the model uses for a short pause."""
space_token = getattr(self.tokenizer, "space_token", None)
if space_token is None:
return text
return re.sub(r"\s", space_token, text)
# Kept for backwards compatibility with the previously published processor.
def preprocess_stage1(self, text: str, language: Optional[str] = None) -> str:
return self.normalize(text, language)
def preprocess_stage2(self, text: str, language: Optional[str] = None) -> str:
return self.to_g2p_text(text, language)
# ------------------------------------------------------------------
# __call__
# ------------------------------------------------------------------
def __call__(
self,
text: Union[str, Sequence[str]] = None,
language: str = None,
return_tensors: Optional[str] = "pt",
add_special_tokens: bool = True,
**kwargs,
) -> BatchEncoding:
"""Prepare one or more sentences for [`BertVits2Model`].
Args:
text (`str` or `List[str]`): The sentence(s) to synthesise.
language (`str`): Language code, which must be one of the processor's BERT languages.
return_tensors (`str`, *optional*, defaults to `"pt"`): `"pt"`, `"np"`, or `None` for plain lists.
add_special_tokens (`bool`, *optional*, defaults to `True`):
Whether the BERT tokenizer adds its `[CLS]`/`[SEP]` tokens. The phoneme side compensates with
zero-width entries in `word_to_phoneme`.
Returns:
[`BatchEncoding`] with `input_ids`, `attention_mask`, `tone_ids`, `language_ids`, `word_to_phoneme`,
`bert_input_ids` and `bert_attention_mask`.
"""
if text is None:
raise ValueError("`text` is required.")
if language is None:
raise ValueError("`language` is required for BertVits2Processor.")
if language not in self.bert_tokenizers:
raise ValueError(
f"Language '{language}' is not supported by this processor. "
f"Available languages: {sorted(self.bert_tokenizers)}."
)
texts = [text] if isinstance(text, str) else list(text)
if not texts:
raise ValueError("`text` must contain at least one sentence.")
normalized = [self.normalize(sentence, language) for sentence in texts]
phoneme_ids: List[List[int]] = []
tone_ids: List[List[int]] = []
language_ids: List[List[int]] = []
word_to_phoneme: List[List[int]] = []
for sentence in normalized:
phonemes, tones, languages, word2ph = self.tokenizer.convert_g2p(
self.to_g2p_text(sentence, language), language, add_special_tokens
)
phoneme_ids.append(self.tokenizer.convert_phonemes_to_ids(phonemes))
tone_ids.append(tones)
language_ids.append(languages)
word_to_phoneme.append(word2ph)
bert_tokenizer = self.bert_tokenizers[language]
bert_encoded = bert_tokenizer(
normalized,
padding="longest",
padding_side="right",
add_special_tokens=add_special_tokens,
return_attention_mask=True,
return_token_type_ids=False,
**kwargs,
)
for sentence, word2ph, bert_ids, mask in zip(
normalized, word_to_phoneme, bert_encoded["input_ids"], bert_encoded["attention_mask"]
):
self._check_alignment(sentence, word2ph, bert_ids, mask, bert_tokenizer)
pad_id = self.tokenizer.pad_token_id or 0
phone_length = max(len(ids) for ids in phoneme_ids)
bert_length = max(len(ids) for ids in bert_encoded["input_ids"])
data = {
"input_ids": [_pad_to(ids, phone_length, pad_id) for ids in phoneme_ids],
"attention_mask": [_pad_to([1] * len(ids), phone_length, 0) for ids in phoneme_ids],
"tone_ids": [_pad_to(ids, phone_length, 0) for ids in tone_ids],
"language_ids": [_pad_to(ids, phone_length, 0) for ids in language_ids],
# Padding entries repeat their BERT feature zero times, so they contribute nothing downstream.
"word_to_phoneme": [_pad_to(ids, bert_length, 0) for ids in word_to_phoneme],
"bert_input_ids": bert_encoded["input_ids"],
"bert_attention_mask": bert_encoded["attention_mask"],
}
return BatchEncoding(data, tensor_type=return_tensors)
def _check_alignment(
self,
sentence: str,
word2ph: Sequence[int],
bert_ids: Sequence[int],
attention_mask: Sequence[int],
bert_tokenizer: PreTrainedTokenizerBase,
) -> None:
"""Fail loudly when grapheme-to-phoneme and BERT disagree about how many units the sentence has.
`word_to_phoneme` must have one entry per BERT token, because the model uses it to repeat each BERT token's
feature across that token's phonemes. A mismatch used to surface as an opaque shape error deep inside the
model (or, in ONNX Runtime, as `invalid expand shape`).
"""
num_tokens = sum(attention_mask) # padding is not part of the alignment
if len(word2ph) == num_tokens:
return
tokens = bert_tokenizer.convert_ids_to_tokens(list(bert_ids))[:num_tokens]
raise ValueError(
f"Text cannot be aligned with the BERT tokenizer: grapheme-to-phoneme produced {len(word2ph)} units "
f"but the BERT tokenizer produced {num_tokens} tokens for {sentence!r}.\n"
f"BERT tokens: {tokens}\n"
"This happens when the text contains runs that BERT splits differently from the g2p front-end, "
"typically Latin words, emoji, or unusual symbols. Remove or transliterate them before synthesising."
)
# ------------------------------------------------------------------
# serialization
# ------------------------------------------------------------------
def to_dict(self) -> Dict[str, Any]:
output = super().to_dict()
output.pop("bert_tokenizers", None)
output["bert_tokenizers"] = {language: f"bert_{language}" for language in sorted(self.bert_tokenizers)}
output["number_style"] = self.number_style
return output
def save_pretrained(self, save_directory: Union[str, os.PathLike], **kwargs):
os.makedirs(save_directory, exist_ok=True)
for language, tokenizer in self.bert_tokenizers.items():
tokenizer.save_pretrained(os.path.join(save_directory, f"bert_{language}"))
return super().save_pretrained(save_directory, **kwargs)
@classmethod
def from_pretrained(cls, pretrained_model_name_or_path: Union[str, os.PathLike], **kwargs):
processor_dict, _ = cls.get_processor_dict(pretrained_model_name_or_path, **kwargs)
subfolders = processor_dict.get("bert_tokenizers", {})
tokenizer_kwargs = {k: v for k, v in kwargs.items() if k not in ("config", "subfolder")}
tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path, **tokenizer_kwargs)
bert_tokenizers = {
language: AutoTokenizer.from_pretrained(
pretrained_model_name_or_path, subfolder=subfolder, **tokenizer_kwargs
)
for language, subfolder in subfolders.items()
}
return cls(
tokenizer,
bert_tokenizers=bert_tokenizers,
number_style=processor_dict.get("number_style", "cardinal"),
)
def _pad_to(values: Sequence[int], length: int, pad_value: int) -> List[int]:
values = list(values)
return values + [pad_value] * (length - len(values))
__all__ = ["BertVits2Processor", "chinese_number_to_words"]