lenML/ChatTTS-Forge
301
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 