lenML/ChatTTS-Forge
301
1from typing import Union2 3import gradio as gr4import numpy as np5import torch6import torch.profiler7 8from modules import refiner9from modules.api.impl.handler.SSMLHandler import SSMLHandler10from modules.api.impl.handler.TTSHandler import TTSHandler11from modules.api.impl.model.audio_model import AdjustConfig12from modules.api.impl.model.chattts_model import ChatTTSConfig, InferConfig13from modules.api.impl.model.enhancer_model import EnhancerConfig14from modules.api.utils import calc_spk_style15from modules.data import styles_mgr16from modules.Enhancer.ResembleEnhance import apply_audio_enhance as _apply_audio_enhance17from modules.normalization import text_normalize18from modules.SentenceSplitter import SentenceSplitter19from modules.speaker import Speaker, speaker_mgr20from modules.ssml_parser.SSMLParser import SSMLBreak, SSMLSegment, create_ssml_parser21from modules.utils import audio22from modules.utils.hf import spaces23from modules.webui import webui_config24 25 26def get_speakers():27 return speaker_mgr.list_speakers()28 29 30def get_speaker_names() -> tuple[list[Speaker], list[str]]:31 speakers = get_speakers()32 33 def get_speaker_show_name(spk):34 if spk.gender == "*" or spk.gender == "":35 return spk.name36 return f"{spk.gender} : {spk.name}"37 38 speaker_names = [get_speaker_show_name(speaker) for speaker in speakers]39 speaker_names.sort(key=lambda x: x.startswith("*") and "-1" or x)40 41 return speakers, speaker_names42 43 44def get_styles():45 return styles_mgr.list_items()46 47 48def load_spk_info(file):49 if file is None:50 return "empty"51 try:52 53 spk: Speaker = Speaker.from_file(file)54 infos = spk.to_json()55 return f"""56- name: {infos.name}57- gender: {infos.gender}58- describe: {infos.describe}59 """.strip()60 except:61 return "load failed"62 63 64def segments_length_limit(65 segments: list[Union[SSMLBreak, SSMLSegment]], total_max: int66) -> list[Union[SSMLBreak, SSMLSegment]]:67 ret_segments = []68 total_len = 069 for seg in segments:70 if isinstance(seg, SSMLBreak):71 ret_segments.append(seg)72 continue73 total_len += len(seg["text"])74 if total_len > total_max:75 break76 ret_segments.append(seg)77 return ret_segments78 79 80@torch.inference_mode()81@spaces.GPU(duration=120)82def apply_audio_enhance(audio_data, sr, enable_denoise, enable_enhance):83 return _apply_audio_enhance(audio_data, sr, enable_denoise, enable_enhance)84 85 86@torch.inference_mode()87@spaces.GPU(duration=120)88def synthesize_ssml(89 ssml: str,90 batch_size=4,91 enable_enhance=False,92 enable_denoise=False,93 eos: str = "[uv_break]",94 spliter_thr: int = 100,95 pitch: float = 0,96 speed_rate: float = 1,97 volume_gain_db: float = 0,98 normalize: bool = True,99 headroom: float = 1,100 progress=gr.Progress(track_tqdm=True),101):102 try:103 batch_size = int(batch_size)104 except Exception:105 batch_size = 8106 107 ssml = ssml.strip()108 109 if ssml == "":110 raise gr.Error("SSML is empty, please input some SSML")111 112 parser = create_ssml_parser()113 segments = parser.parse(ssml)114 max_len = webui_config.ssml_max115 segments = segments_length_limit(segments, max_len)116 117 if len(segments) == 0:118 raise gr.Error("No valid segments in SSML")119 120 infer_config = InferConfig(121 batch_size=batch_size,122 spliter_threshold=spliter_thr,123 eos=eos,124 # NOTE: SSML not support `infer_seed` contorl125 # seed=42,126 )127 adjust_config = AdjustConfig(128 pitch=pitch,129 speed_rate=speed_rate,130 volume_gain_db=volume_gain_db,131 normalize=normalize,132 headroom=headroom,133 )134 enhancer_config = EnhancerConfig(135 enabled=enable_denoise or enable_enhance or False,136 lambd=0.9 if enable_denoise else 0.1,137 )138 139 handler = SSMLHandler(140 ssml_content=ssml,141 infer_config=infer_config,142 adjust_config=adjust_config,143 enhancer_config=enhancer_config,144 )145 146 audio_data, sr = handler.enqueue()147 148 # NOTE: 这里必须要加,不然 gradio 没法解析成 mp3 格式149 audio_data = audio.audio_to_int16(audio_data)150 151 return sr, audio_data152 153 154# @torch.inference_mode()155@spaces.GPU(duration=120)156def tts_generate(157 text,158 temperature=0.3,159 top_p=0.7,160 top_k=20,161 spk=-1,162 infer_seed=-1,163 use_decoder=True,164 prompt1="",165 prompt2="",166 prefix="",167 style="",168 disable_normalize=False,169 batch_size=4,170 enable_enhance=False,171 enable_denoise=False,172 spk_file=None,173 spliter_thr: int = 100,174 eos: str = "[uv_break]",175 pitch: float = 0,176 speed_rate: float = 1,177 volume_gain_db: float = 0,178 normalize: bool = True,179 headroom: float = 1,180 progress=gr.Progress(track_tqdm=True),181):182 try:183 batch_size = int(batch_size)184 except Exception:185 batch_size = 4186 187 max_len = webui_config.tts_max188 text = text.strip()[0:max_len]189 190 if text == "":191 raise gr.Error("Text is empty, please input some text")192 193 if style == "*auto":194 style = ""195 196 if isinstance(top_k, float):197 top_k = int(top_k)198 199 params = calc_spk_style(spk=spk, style=style)200 spk = params.get("spk", spk)201 202 infer_seed = infer_seed or params.get("seed", infer_seed)203 temperature = temperature or params.get("temperature", temperature)204 prefix = prefix or params.get("prefix", prefix)205 prompt1 = prompt1 or params.get("prompt1", "")206 prompt2 = prompt2 or params.get("prompt2", "")207 208 infer_seed = np.clip(infer_seed, -1, 2**32 - 1, out=None, dtype=np.float64)209 infer_seed = int(infer_seed)210 211 if isinstance(spk, int):212 spk = Speaker.from_seed(spk)213 214 if spk_file:215 try:216 spk: Speaker = Speaker.from_file(spk_file)217 except Exception:218 raise gr.Error("Failed to load speaker file")219 220 if not isinstance(spk.emb, torch.Tensor):221 raise gr.Error("Speaker file is not supported")222 223 tts_config = ChatTTSConfig(224 style=style,225 temperature=temperature,226 top_k=top_k,227 top_p=top_p,228 prefix=prefix,229 prompt1=prompt1,230 prompt2=prompt2,231 )232 infer_config = InferConfig(233 batch_size=batch_size,234 spliter_threshold=spliter_thr,235 eos=eos,236 seed=infer_seed,237 )238 adjust_config = AdjustConfig(239 pitch=pitch,240 speed_rate=speed_rate,241 volume_gain_db=volume_gain_db,242 normalize=normalize,243 headroom=headroom,244 )245 enhancer_config = EnhancerConfig(246 enabled=enable_denoise or enable_enhance or False,247 lambd=0.9 if enable_denoise else 0.1,248 )249 250 handler = TTSHandler(251 text_content=text,252 spk=spk,253 tts_config=tts_config,254 infer_config=infer_config,255 adjust_config=adjust_config,256 enhancer_config=enhancer_config,257 )258 259 audio_data, sample_rate = handler.enqueue()260 261 # NOTE: 这里必须要加,不然 gradio 没法解析成 mp3 格式262 audio_data = audio.audio_to_int16(audio_data)263 return sample_rate, audio_data264 265 266@torch.inference_mode()267@spaces.GPU(duration=120)268def refine_text(269 text: str,270 prompt: str,271 progress=gr.Progress(track_tqdm=True),272):273 text = text_normalize(text)274 return refiner.refine_text(text, prompt=prompt)275 276 277@torch.inference_mode()278@spaces.GPU(duration=120)279def split_long_text(long_text_input, spliter_threshold=100, eos=""):280 spliter = SentenceSplitter(threshold=spliter_threshold)281 sentences = spliter.parse(long_text_input)282 sentences = [text_normalize(s) + eos for s in sentences]283 data = []284 for i, text in enumerate(sentences):285 token_length = spliter.count_tokens(text)286 data.append([i, text, token_length])287 return data288 