CoolFace
Apppublic

lenML/ChatTTS-Forge

sourceHugging Faceagpl-3.0updated 2y agoView on Hugging Face
301likes
generate_audio.py234 linesDownload Raw Back to modules
1import gc2import logging3from typing import Generator, Union4 5import numpy as np6import torch7 8from modules import config, models9from modules.ChatTTS import ChatTTS10from modules.devices import devices11from modules.speaker import Speaker12from modules.utils.cache import conditional_cache13from modules.utils.SeedContext import SeedContext14 15logger = logging.getLogger(__name__)16 17SAMPLE_RATE = 2400018 19 20def generate_audio(21    text: str,22    temperature: float = 0.3,23    top_P: float = 0.7,24    top_K: float = 20,25    spk: Union[int, Speaker] = -1,26    infer_seed: int = -1,27    use_decoder: bool = True,28    prompt1: str = "",29    prompt2: str = "",30    prefix: str = "",31):32    (sample_rate, wav) = generate_audio_batch(33        [text],34        temperature=temperature,35        top_P=top_P,36        top_K=top_K,37        spk=spk,38        infer_seed=infer_seed,39        use_decoder=use_decoder,40        prompt1=prompt1,41        prompt2=prompt2,42        prefix=prefix,43    )[0]44 45    return (sample_rate, wav)46 47 48def parse_infer_params(49    texts: list[str],50    chat_tts: ChatTTS.Chat,51    temperature: float = 0.3,52    top_P: float = 0.7,53    top_K: float = 20,54    spk: Union[int, Speaker] = -1,55    infer_seed: int = -1,56    prompt1: str = "",57    prompt2: str = "",58    prefix: str = "",59):60    params_infer_code = {61        "spk_emb": None,62        "temperature": temperature,63        "top_P": top_P,64        "top_K": top_K,65        "prompt1": prompt1 or "",66        "prompt2": prompt2 or "",67        "prefix": prefix or "",68        "repetition_penalty": 1.0,69        "disable_tqdm": config.runtime_env_vars.off_tqdm,70    }71 72    if isinstance(spk, int):73        with SeedContext(spk, True):74            params_infer_code["spk_emb"] = chat_tts.sample_random_speaker()75        logger.debug(("spk", spk))76    elif isinstance(spk, Speaker):77        if not isinstance(spk.emb, torch.Tensor):78            raise ValueError("spk.pt is broken, please retrain the model.")79        params_infer_code["spk_emb"] = spk.emb80        logger.debug(("spk", spk.name))81    else:82        logger.warn(83            f"spk must be int or Speaker, but: <{type(spk)}> {spk}, wiil set to default voice"84        )85        with SeedContext(2, True):86            params_infer_code["spk_emb"] = chat_tts.sample_random_speaker()87 88    logger.debug(89        {90            "text": texts,91            "infer_seed": infer_seed,92            "temperature": temperature,93            "top_P": top_P,94            "top_K": top_K,95            "prompt1": prompt1 or "",96            "prompt2": prompt2 or "",97            "prefix": prefix or "",98        }99    )100 101    return params_infer_code102 103 104@torch.inference_mode()105def generate_audio_batch(106    texts: list[str],107    temperature: float = 0.3,108    top_P: float = 0.7,109    top_K: float = 20,110    spk: Union[int, Speaker] = -1,111    infer_seed: int = -1,112    use_decoder: bool = True,113    prompt1: str = "",114    prompt2: str = "",115    prefix: str = "",116):117    chat_tts = models.load_chat_tts()118    params_infer_code = parse_infer_params(119        texts=texts,120        chat_tts=chat_tts,121        temperature=temperature,122        top_P=top_P,123        top_K=top_K,124        spk=spk,125        infer_seed=infer_seed,126        prompt1=prompt1,127        prompt2=prompt2,128        prefix=prefix,129    )130 131    with SeedContext(infer_seed, True):132        wavs = chat_tts.generate_audio(133            prompt=texts, params_infer_code=params_infer_code, use_decoder=use_decoder134        )135 136    if config.auto_gc:137        devices.torch_gc()138        gc.collect()139 140    return [(SAMPLE_RATE, np.array(wav).flatten().astype(np.float32)) for wav in wavs]141 142 143# TODO: generate_audio_stream 也应该支持 lru cache144@torch.inference_mode()145def generate_audio_stream(146    text: str,147    temperature: float = 0.3,148    top_P: float = 0.7,149    top_K: float = 20,150    spk: Union[int, Speaker] = -1,151    infer_seed: int = -1,152    use_decoder: bool = True,153    prompt1: str = "",154    prompt2: str = "",155    prefix: str = "",156) -> Generator[tuple[int, np.ndarray], None, None]:157    chat_tts = models.load_chat_tts()158    texts = [text]159    params_infer_code = parse_infer_params(160        texts=texts,161        chat_tts=chat_tts,162        temperature=temperature,163        top_P=top_P,164        top_K=top_K,165        spk=spk,166        infer_seed=infer_seed,167        prompt1=prompt1,168        prompt2=prompt2,169        prefix=prefix,170    )171 172    with SeedContext(infer_seed, True):173        wavs_gen = chat_tts.generate_audio(174            prompt=texts,175            params_infer_code=params_infer_code,176            use_decoder=use_decoder,177            stream=True,178        )179 180        for wav in wavs_gen:181            yield [SAMPLE_RATE, np.array(wav).flatten().astype(np.float32)]182 183    if config.auto_gc:184        devices.torch_gc()185        gc.collect()186 187    return188 189 190lru_cache_enabled = False191 192 193def setup_lru_cache():194    global generate_audio_batch195    global lru_cache_enabled196 197    if lru_cache_enabled:198        return199    lru_cache_enabled = True200 201    def should_cache(*args, **kwargs):202        spk_seed = kwargs.get("spk", -1)203        infer_seed = kwargs.get("infer_seed", -1)204        return spk_seed != -1 and infer_seed != -1205 206    lru_size = config.runtime_env_vars.lru_size207    if isinstance(lru_size, int):208        generate_audio_batch = conditional_cache(lru_size, should_cache)(209            generate_audio_batch210        )211        logger.info(f"LRU cache enabled with size {lru_size}")212    else:213        logger.debug(f"LRU cache failed to enable, invalid size {lru_size}")214 215 216if __name__ == "__main__":217    import soundfile as sf218 219    # 测试batch生成220    inputs = ["你好[lbreak]", "再见[lbreak]", "长度不同的文本片段[lbreak]"]221    outputs = generate_audio_batch(inputs, spk=5, infer_seed=42)222 223    for i, (sample_rate, wav) in enumerate(outputs):224        print(i, sample_rate, wav.shape)225 226        sf.write(f"batch_{i}.wav", wav, sample_rate, format="wav")227 228    # 单独生成229    for i, text in enumerate(inputs):230        sample_rate, wav = generate_audio(text, spk=5, infer_seed=42)231        print(i, sample_rate, wav.shape)232 233        sf.write(f"one_{i}.wav", wav, sample_rate, format="wav")234