diff --git a/ChatTTS/core.py b/ChatTTS/core.py index 4639c88..e3891cd 100644 --- a/ChatTTS/core.py +++ b/ChatTTS/core.py @@ -220,7 +220,7 @@ class Chat: for i, t in enumerate(text): invalid_characters = count_invalid_characters(t) if len(invalid_characters): - self.logger.log(logging.WARNING, f'Invalid characters found! : {invalid_characters}') + self.logger.warn(f'Invalid characters found! : {invalid_characters}') text[i] = apply_character_map(t) if do_homophone_replacement and self._init_homophones_replacer(): text[i], replaced_words = self.homophones_replacer.replace(text[i]) @@ -357,10 +357,10 @@ class Chat: else: try: self.homophones_replacer = HomophonesReplacer(os.path.join(os.path.dirname(__file__), 'res', 'homophones_map.json')) - self.logger.log(logging.INFO, 'homophones_replacer loaded.') + self.logger.log(logging.INFO, 'successfully loaded HomophonesReplacer.') return True except (IOError, json.JSONDecodeError) as e: - self.logger.log(logging.WARNING, f'Error loading homophones map: {e}') + self.logger.log(logging.WARNING, f'error loading homophones map: {e}') except Exception as e: - self.logger.log(logging.WARNING, f'Error loading homophones_replacer: {e}') + self.logger.log(logging.WARNING, f'error loading HomophonesReplacer: {e}') return False diff --git a/ChatTTS/utils/infer.py b/ChatTTS/utils/infer.py index ddebd60..63fcf9b 100644 --- a/ChatTTS/utils/infer.py +++ b/ChatTTS/utils/infer.py @@ -1,8 +1,12 @@ - +import json import re +from typing import Dict, Tuple, List +import sys + +from numba import jit +import numpy as np import torch import torch.nn.functional as F -import json class CustomRepetitionPenaltyLogitsProcessorRepeat(): @@ -47,6 +51,26 @@ class CustomRepetitionPenaltyLogitsProcessor(): return scores +@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 @@ -65,37 +89,40 @@ class HomophonesReplacer: [python-pinyin](https://github.com/mozillazg/python-pinyin) """ - def __init__(self, map_file_path): - self.homophones_map = self.load_homophones_map(map_file_path) + 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): + def _load_homophones_map(self, map_file_path: str) -> np.ndarray: with open(map_file_path, 'r', encoding='utf-8') as f: - homophones_map = json.load(f) - return homophones_map + 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): - result = [] - replaced_words = [] - for char in text: - if char in self.homophones_map: - repl_char = self.homophones_map[char] - result.append(repl_char) - replaced_words.append((char, repl_char)) - else: - result.append(char) - return ''.join(result), replaced_words + def replace(self, text: str): + arr, lst = _fast_replace( + self.homophones_map, + text.encode(self.coding), + ) + return arr.tobytes().decode(self.coding), lst -def count_invalid_characters(s): - - s = re.sub(r'\[uv_break\]|\[laugh\]|\[lbreak\]', '', s) - pattern = re.compile(r'[^\u4e00-\u9fffA-Za-z,。、,\. ]') - non_alphabetic_chinese_chars = pattern.findall(s) +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) -def detect_language(sentence): +chinese_char_pattern = re.compile(r'[\u4e00-\u9fff]') +english_word_pattern = re.compile(r'\b[A-Za-z]+\b') - 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) @@ -104,9 +131,9 @@ def detect_language(sentence): return "zh" else: return "en" - - -character_map = { + + +character_simplifier = str.maketrans({ ':': ',', ';': ',', '!': '。', @@ -135,9 +162,9 @@ character_map = { '>': ',', '<': ',', '-': ',', -} +}) -halfwidth_2_fullwidth_map = { +halfwidth_2_fullwidth = str.maketrans({ '!': '!', '"': '“', "'": '‘', @@ -170,12 +197,10 @@ halfwidth_2_fullwidth_map = { '|': '|', '}': '}', '~': '~' - } + }) -def apply_half2full_map(text): - translation_table = str.maketrans(halfwidth_2_fullwidth_map) - return text.translate(translation_table) +def apply_half2full_map(text: str) -> str: + return text.translate(halfwidth_2_fullwidth) -def apply_character_map(text): - translation_table = str.maketrans(character_map) - return text.translate(translation_table) +def apply_character_map(text: str) -> str: + return text.translate(character_simplifier) diff --git a/tools/audio/np.py b/tools/audio/np.py index 3dfb46d..9812f78 100644 --- a/tools/audio/np.py +++ b/tools/audio/np.py @@ -7,8 +7,7 @@ def unsafe_float_to_int16(audio: np.ndarray) -> np.ndarray: This function will destroy audio, use only once. """ am = np.abs(audio).max() * 32768 - if am > 32768: - am = 32768 * 32768 / am + am = 32767 * 32768 / am np.multiply(audio, am, audio) audio16 = audio.astype(np.int16) return audio16