CoolFace
Apppublic

lenML/ChatTTS-Forge

sourceHugging Faceagpl-3.0updated 2y agoView on Hugging Face
301likes
webui_utils.py288 linesDownload Raw Back to webui
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