feat: add normalizer interface & move instances out (#420)

This commit is contained in:
源文雨
2024-06-24 17:28:14 +09:00
committed by GitHub
parent b62e0dc3c5
commit c8cb6bd327
8 changed files with 270 additions and 247 deletions
+12 -57
View File
@@ -1,7 +1,6 @@
import os
import logging
import tempfile
from functools import partial
from typing import Literal, Optional, List, Callable
import numpy as np
@@ -13,12 +12,13 @@ from huggingface_hub import snapshot_download
from .model.dvae import DVAE
from .model.gpt import GPT
from .utils.gpu import select_device
from .utils.infer import count_invalid_characters, detect_language, apply_character_map, apply_half2full_map, HomophonesReplacer
from .utils.io import get_latest_modified_file, del_all
from .infer.api import refine_text, infer_code
from .utils.dl import check_all_assets, download_all_assets
from .utils.log import logger as utils_logger
from .norm import Normalizer
class Chat:
def __init__(self, logger=logging.getLogger(__name__)):
@@ -26,9 +26,10 @@ class Chat:
utils_logger.set_logger(logger)
self.pretrain_models = {}
self.normalizer = {}
self.homophones_replacer = self.homophones_replacer = HomophonesReplacer(os.path.join(os.path.dirname(__file__), 'res', 'homophones_map.json'))
self.normalizer = Normalizer(
os.path.join(os.path.dirname(__file__), 'res', 'homophones_map.json'),
logger,
)
def has_loaded(self, use_decoder = False):
not_finish = False
@@ -188,6 +189,8 @@ class Chat:
def unload(self):
logger = self.logger
del_all(self.pretrain_models)
self.normalizer.destroy()
del self.normalizer
del_list = ["vocos", "_vocos_decode", 'gpt', 'decoder', 'dvae']
for module in del_list:
if hasattr(self, module):
@@ -212,23 +215,10 @@ class Chat:
if not isinstance(text, list):
text = [text]
if do_text_normalization:
for i, t in enumerate(text):
_lang = detect_language(t) if lang is None else lang
if self._init_normalizer(_lang):
text[i] = self.normalizer[_lang](t)
if _lang == 'zh':
text[i] = apply_half2full_map(text[i])
for i, t in enumerate(text):
invalid_characters = count_invalid_characters(t)
if len(invalid_characters):
self.logger.warn(f'Invalid characters found! : {invalid_characters}')
text[i] = apply_character_map(t)
if do_homophone_replacement:
text[i], replaced_words = self.homophones_replacer.replace(text[i])
if replaced_words:
repl_res = ', '.join([f'{_[0]}->{_[1]}' for _ in replaced_words])
self.logger.log(logging.INFO, f'Homophones replace: {repl_res}')
text = [self.normalizer(
t, do_text_normalization, do_homophone_replacement, lang,
) for t in text]
if not skip_refine_text:
refined = refine_text(
@@ -314,38 +304,3 @@ class Chat:
result.destroy()
del_all(x)
return wavs
def _init_normalizer(self, lang) -> bool:
if lang in self.normalizer:
return True
if lang == 'zh':
try:
from tn.chinese.normalizer import Normalizer
self.normalizer[lang] = Normalizer().normalize
return True
except:
self.logger.log(
logging.WARNING,
'Package WeTextProcessing not found!',
)
self.logger.log(
logging.WARNING,
'Run: conda install -c conda-forge pynini=2.1.5 && pip install WeTextProcessing',
)
else:
try:
from nemo_text_processing.text_normalization.normalize import Normalizer
self.normalizer[lang] = partial(Normalizer(input_case='cased', lang=lang).normalize, verbose=False, punct_post_process=True)
return True
except:
self.logger.log(
logging.WARNING,
'Package nemo_text_processing not found!',
)
self.logger.log(
logging.WARNING,
'Run: conda install -c conda-forge pynini=2.1.5 && pip install nemo_text_processing',
)
return False
+199
View File
@@ -0,0 +1,199 @@
import json
import logging
import re
from typing import Dict, Tuple, List, Literal, Callable, Optional
import sys
from numba import jit
import numpy as np
from .utils.io import del_all
@jit
def _find_index(table: np.ndarray, val: np.uint16):
for i in range(table.size):
if table[i] == val:
return i
return -1
@jit
def _fast_replace(table: np.ndarray, text: bytes) -> Tuple[np.ndarray, List[Tuple[str, str]]]:
result = np.frombuffer(text, dtype=np.uint16).copy()
replaced_words = []
for i in range(result.size):
ch = result[i]
p = _find_index(table[0], ch)
if p >= 0:
repl_char = table[1][p]
result[i] = repl_char
replaced_words.append((chr(ch), chr(repl_char)))
return result, replaced_words
class Normalizer:
def __init__(self, map_file_path: str, logger=logging.getLogger(__name__)):
self.logger = logger
self.normalizers: Dict[str, Callable[[str], str]] = {}
self.homophones_map = self._load_homophones_map(map_file_path)
"""
homophones_map
Replace the mispronounced characters with correctly pronounced ones.
Creation process of homophones_map.json:
1. Establish a word corpus using the [Tencent AI Lab Embedding Corpora v0.2.0 large] with 12 million entries. After cleaning, approximately 1.8 million entries remain. Use ChatTTS to infer the text.
2. Record discrepancies between the inferred and input text, identifying about 180,000 misread words.
3. Create a pinyin to common characters mapping using correctly read characters by ChatTTS.
4. For each discrepancy, extract the correct pinyin using [python-pinyin] and find homophones with the correct pronunciation from the mapping.
Thanks to:
[Tencent AI Lab Embedding Corpora for Chinese and English Words and Phrases](https://ai.tencent.com/ailab/nlp/en/embedding.html)
[python-pinyin](https://github.com/mozillazg/python-pinyin)
"""
self.coding = "utf-16-le" if sys.byteorder == "little" else "utf-16-be"
self.accept_pattern = re.compile(r'[^\u4e00-\u9fffA-Za-z,。、,\. ]')
self.sub_pattern = re.compile(r'\[uv_break\]|\[laugh\]|\[lbreak\]')
self.chinese_char_pattern = re.compile(r'[\u4e00-\u9fff]')
self.english_word_pattern = re.compile(r'\b[A-Za-z]+\b')
self.character_simplifier = str.maketrans({
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
':': ',',
';': ',',
'!': '.',
'(': ',',
')': ',',
'[': ',',
']': ',',
'>': ',',
'<': ',',
'-': ',',
})
self.halfwidth_2_fullwidth = str.maketrans({
'!': '',
'"': '',
"'": '',
'#': '',
'$': '',
'%': '',
'&': '',
'(': '',
')': '',
',': '',
'-': '',
'*': '',
'+': '',
'.': '',
'/': '',
':': '',
';': '',
'<': '',
'=': '',
'>': '',
'?': '',
'@': '',
# '[': '',
'\\': '',
# ']': '',
'^': '',
# '_': '_',
'`': '',
'{': '',
'|': '',
'}': '',
'~': ''
})
def __call__(
self,
text: str,
do_text_normalization=True,
do_homophone_replacement=True,
lang: Optional[Literal["zh", "en"]] = None,
) -> str:
if do_text_normalization:
_lang = self._detect_language(text) if lang is None else lang
if _lang in self.normalizers:
text = self.normalizers[_lang](text)
if _lang == 'zh':
text = self._apply_half2full_map(text)
invalid_characters = self._count_invalid_characters(text)
if len(invalid_characters):
self.logger.warn(f'found invalid characters: {invalid_characters}')
text = self._apply_character_map(text)
if do_homophone_replacement:
arr, replaced_words = _fast_replace(
self.homophones_map,
text.encode(self.coding),
)
if replaced_words:
text = arr.tobytes().decode(self.coding)
repl_res = ', '.join([f'{_[0]}->{_[1]}' for _ in replaced_words])
self.logger.info(f'replace homophones: {repl_res}')
return text
def register(self, name: str, normalizer: Callable[[str], str]) -> bool:
if name in self.normalizers:
self.logger.warn(f"name {name} has been registered")
return False
if not isinstance(normalizer, Callable[[str], str]):
self.logger.warn("normalizer must have caller type (str) -> str")
return False
self.normalizers[name] = normalizer
return True
def unregister(self, name: str):
if name in self.normalizers:
del self.normalizers[name]
def destroy(self):
del_all(self.normalizers)
del self.homophones_map
def _load_homophones_map(self, map_file_path: str) -> np.ndarray:
with open(map_file_path, 'r', encoding='utf-8') as f:
homophones_map: Dict[str, str] = json.load(f)
map = np.empty((2, len(homophones_map)), dtype=np.uint32)
for i, k in enumerate(homophones_map.keys()):
map[:, i] = (ord(k), ord(homophones_map[k]))
del homophones_map
return map
def _count_invalid_characters(self, s: str):
s = self.sub_pattern.sub('', s)
non_alphabetic_chinese_chars = self.accept_pattern.findall(s)
return set(non_alphabetic_chinese_chars)
def _apply_half2full_map(self, text: str) -> str:
return text.translate(self.halfwidth_2_fullwidth)
def _apply_character_map(self, text: str) -> str:
return text.translate(self.character_simplifier)
def _detect_language(self, sentence: str) -> Literal["zh", "en"]:
chinese_chars = self.chinese_char_pattern.findall(sentence)
english_words = self.english_word_pattern.findall(sentence)
if len(chinese_chars) > len(english_words):
return "zh"
else:
return "en"
-162
View File
@@ -1,162 +0,0 @@
import json
import re
from typing import Dict, Tuple, List
import sys
from numba import jit
import numpy as np
@jit
def _find_index(table: np.ndarray, val: np.uint16):
for i in range(table.size):
if table[i] == val:
return i
return -1
@jit
def _fast_replace(table: np.ndarray, text: bytes) -> Tuple[np.ndarray, List[Tuple[str, str]]]:
result = np.frombuffer(text, dtype=np.uint16).copy()
replaced_words = []
for i in range(result.size):
ch = result[i]
p = _find_index(table[0], ch)
if p >= 0:
repl_char = table[1][p]
result[i] = repl_char
replaced_words.append((chr(ch), chr(repl_char)))
return result, replaced_words
class HomophonesReplacer:
"""
Homophones Replacer
Replace the mispronounced characters with correctly pronounced ones.
Creation process of homophones_map.json:
1. Establish a word corpus using the [Tencent AI Lab Embedding Corpora v0.2.0 large] with 12 million entries. After cleaning, approximately 1.8 million entries remain. Use ChatTTS to infer the text.
2. Record discrepancies between the inferred and input text, identifying about 180,000 misread words.
3. Create a pinyin to common characters mapping using correctly read characters by ChatTTS.
4. For each discrepancy, extract the correct pinyin using [python-pinyin] and find homophones with the correct pronunciation from the mapping.
Thanks to:
[Tencent AI Lab Embedding Corpora for Chinese and English Words and Phrases](https://ai.tencent.com/ailab/nlp/en/embedding.html)
[python-pinyin](https://github.com/mozillazg/python-pinyin)
"""
def __init__(self, map_file_path: str):
self.homophones_map = self._load_homophones_map(map_file_path)
self.coding = "utf-16-le" if sys.byteorder == "little" else "utf-16-be"
def _load_homophones_map(self, map_file_path: str) -> np.ndarray:
with open(map_file_path, 'r', encoding='utf-8') as f:
homophones_map: Dict[str, str] = json.load(f)
map = np.empty((2, len(homophones_map)), dtype=np.uint32)
for i, k in enumerate(homophones_map.keys()):
map[:, i] = (ord(k), ord(homophones_map[k]))
del homophones_map
return map
def replace(self, text: str):
arr, lst = _fast_replace(
self.homophones_map,
text.encode(self.coding),
)
return arr.tobytes().decode(self.coding), lst
accept_pattern = re.compile(r'[^\u4e00-\u9fffA-Za-z,。、,\. ]')
sub_pattern = re.compile(r'\[uv_break\]|\[laugh\]|\[lbreak\]')
def count_invalid_characters(s: str):
global accept_pattern, sub_pattern
s = sub_pattern.sub('', s)
non_alphabetic_chinese_chars = accept_pattern.findall(s)
return set(non_alphabetic_chinese_chars)
chinese_char_pattern = re.compile(r'[\u4e00-\u9fff]')
english_word_pattern = re.compile(r'\b[A-Za-z]+\b')
def detect_language(sentence):
global chinese_char_pattern, english_word_pattern
chinese_chars = chinese_char_pattern.findall(sentence)
english_words = english_word_pattern.findall(sentence)
if len(chinese_chars) > len(english_words):
return "zh"
else:
return "en"
character_simplifier = str.maketrans({
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
'': '',
':': ',',
';': ',',
'!': '.',
'(': ',',
')': ',',
'[': ',',
']': ',',
'>': ',',
'<': ',',
'-': ',',
})
halfwidth_2_fullwidth = str.maketrans({
'!': '',
'"': '',
"'": '',
'#': '',
'$': '',
'%': '',
'&': '',
'(': '',
')': '',
',': '',
'-': '',
'*': '',
'+': '',
'.': '',
'/': '',
':': '',
';': '',
'<': '',
'=': '',
'>': '',
'?': '',
'@': '',
# '[': '',
'\\': '',
# ']': '',
'^': '',
# '_': '_',
'`': '',
'{': '',
'|': '',
'}': '',
'~': ''
})
def apply_half2full_map(text: str) -> str:
return text.translate(halfwidth_2_fullwidth)
def apply_character_map(text: str) -> str:
return text.translate(character_simplifier)
+36 -13
View File
@@ -10,6 +10,7 @@ from tools.logger import get_logger
logger = get_logger(" WebUI ")
from tools.seeder import TorchSeedContext
from tools.normalizer import normalizer_en_nemo_text, normalizer_zh_tn
import ChatTTS
chat = ChatTTS.Chat(get_logger("ChatTTS"))
@@ -37,25 +38,47 @@ def generate_seed():
def on_voice_change(vocie_selection):
return voices.get(vocie_selection)['seed']
def load_chat(cust_path: Optional[str], coef: Optional[str]) -> bool:
if cust_path == None:
ret = chat.load_models(coef=coef, compile=sys.platform != 'win32')
else:
logger.info('local model path: %s', cust_path)
ret = chat.load_models('custom', custom_path=cust_path, coef=coef, compile=sys.platform != 'win32')
global custom_path
custom_path = cust_path
if ret:
try:
chat.normalizer.register("en", normalizer_en_nemo_text())
except:
logger.warn('Package nemo_text_processing not found!')
logger.warn(
'Run: conda install -c conda-forge pynini=2.1.5 && pip install nemo_text_processing',
)
try:
chat.normalizer.register("zh", normalizer_zh_tn())
except:
logger.warn('Package WeTextProcessing not found!')
logger.warn(
'Run: conda install -c conda-forge pynini=2.1.5 && pip install WeTextProcessing',
)
return ret
def reload_chat(coef: Optional[str]) -> str:
global custom_path
chat.unload()
gr.Info("Model unloaded.")
if len(coef) != 230:
gr.Warning("Ingore invalid DVAE coefficient.")
coef = None
try:
if len(coef) != 230:
gr.Warning("Ingore invalid DVAE coefficient.")
coef = None
if custom_path == None:
ret = chat.load_models(coef=coef, compile=sys.platform != 'win32')
else:
logger.info('local model path: %s', custom_path)
ret = chat.load_models('custom', custom_path=custom_path, coef=coef, compile=sys.platform != 'win32')
if not ret:
raise gr.Error("Unable to load model.")
gr.Info("Reload succeess.")
return chat.coef
global custom_path
ret = load_chat(custom_path, coef)
except Exception as e:
raise gr.Error(str(e))
if not ret:
raise gr.Error("Unable to load model.")
gr.Info("Reload succeess.")
return chat.coef
def refine_text(text, text_seed_input, refine_text_flag):
if not refine_text_flag:
+7 -15
View File
@@ -93,29 +93,21 @@ def main():
)
parser = argparse.ArgumentParser(description='ChatTTS demo Launch')
parser.add_argument('--server_name', type=str, default='0.0.0.0', help='Server name')
parser.add_argument('--server_port', type=int, default=8080, help='Server port')
parser.add_argument('--root_path', type=str, default=None, help='Root Path')
parser.add_argument('--custom_path', type=str, default=None, help='the custom model path')
parser.add_argument('--server_name', type=str, default='0.0.0.0', help='server name')
parser.add_argument('--server_port', type=int, default=8080, help='server port')
parser.add_argument('--root_path', type=str, default=None, help='root path')
parser.add_argument('--custom_path', type=str, default=None, help='custom model path')
parser.add_argument('--coef', type=str, default=None, help='custom dvae coefficient')
args = parser.parse_args()
logger.info("loading ChatTTS model...")
global chat, custom_path
if args.custom_path == None:
ret = chat.load_models(compile=sys.platform != 'win32')
else:
logger.info('local model path: %s', args.custom_path)
ret = chat.load_models('custom', custom_path=args.custom_path, compile=sys.platform != 'win32')
if ret:
if load_chat(args.custom_path, args.coef):
logger.info("Models loaded successfully.")
else:
logger.error("Models load failed.")
sys.exit(1)
custom_path = args.custom_path
dvae_coef_text.value = chat.coef
demo.launch(server_name=args.server_name, server_port=args.server_port, root_path=args.root_path, inbrowser=True)
+2
View File
@@ -0,0 +1,2 @@
from .en import normalizer_en_nemo_text
from .zh import normalizer_zh_tn
+9
View File
@@ -0,0 +1,9 @@
from typing import Callable
from functools import partial
def normalizer_en_nemo_text() -> Callable[[str], str]:
from nemo_text_processing.text_normalization.normalize import Normalizer
return partial(
Normalizer(input_case='cased', lang="en").normalize,
verbose=False, punct_post_process=True,
)
+5
View File
@@ -0,0 +1,5 @@
from typing import Callable
def normalizer_zh_tn() -> Callable[[str], str]:
from tn.chinese.normalizer import Normalizer
return Normalizer().normalize