# 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"(?:(?[+-]))?(?\d+)(?:\.(?P\d+))?") THOUSANDS_PATTERN = re.compile(r"(? 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"]