lenML/ChatTTS-Forge
301
1import copy2import io3import json4import logging5import re6from typing import List, Union7 8import numpy as np9from box import Box10from pydub import AudioSegment11from scipy.io import wavfile12 13from modules import generate_audio14from modules.api.utils import calc_spk_style15from modules.normalization import text_normalize16from modules.SentenceSplitter import SentenceSplitter17from modules.speaker import Speaker18from modules.ssml_parser.SSMLParser import SSMLBreak, SSMLContext, SSMLSegment19from modules.utils import rng20from modules.utils.audio import apply_prosody_to_audio_segment21 22logger = logging.getLogger(__name__)23 24 25def audio_data_to_segment_slow(audio_data, sr):26 byte_io = io.BytesIO()27 wavfile.write(byte_io, rate=sr, data=audio_data)28 byte_io.seek(0)29 30 return AudioSegment.from_file(byte_io, format="wav")31 32 33def clip_audio(audio_data: np.ndarray, threshold: float = 0.99):34 audio_data = np.clip(audio_data, -threshold, threshold)35 return audio_data36 37 38def normalize_audio(audio_data: np.ndarray, norm_factor: float = 0.8):39 max_amplitude = np.max(np.abs(audio_data))40 if max_amplitude > 0:41 audio_data = audio_data / max_amplitude * norm_factor42 return audio_data43 44 45def audio_data_to_segment(audio_data: np.ndarray, sr: int):46 """47 optimize: https://github.com/lenML/ChatTTS-Forge/issues/5748 """49 50 audio_data = normalize_audio(audio_data)51 audio_data = clip_audio(audio_data)52 53 audio_data = (audio_data * 32767).astype(np.int16)54 audio_segment = AudioSegment(55 audio_data.tobytes(),56 frame_rate=sr,57 sample_width=audio_data.dtype.itemsize,58 channels=1,59 )60 return audio_segment61 62 63def combine_audio_segments(audio_segments: list[AudioSegment]) -> AudioSegment:64 combined_audio = AudioSegment.empty()65 for segment in audio_segments:66 combined_audio += segment67 return combined_audio68 69 70def to_number(value, t, default=0):71 try:72 number = t(value)73 return number74 except (ValueError, TypeError) as e:75 return default76 77 78class TTSAudioSegment(Box):79 def __init__(self, *args, **kwargs):80 super().__init__(*args, **kwargs)81 self._type = kwargs.get("_type", "voice")82 self.text = kwargs.get("text", "")83 self.temperature = kwargs.get("temperature", 0.3)84 self.top_P = kwargs.get("top_P", 0.5)85 self.top_K = kwargs.get("top_K", 20)86 self.spk = kwargs.get("spk", -1)87 self.infer_seed = kwargs.get("infer_seed", -1)88 self.prompt1 = kwargs.get("prompt1", "")89 self.prompt2 = kwargs.get("prompt2", "")90 self.prefix = kwargs.get("prefix", "")91 92 93class SynthesizeSegments:94 def __init__(self, batch_size: int = 8, eos="", spliter_thr=100):95 self.batch_size = batch_size96 self.batch_default_spk_seed = rng.np_rng()97 self.batch_default_infer_seed = rng.np_rng()98 self.eos = eos99 self.spliter_thr = spliter_thr100 101 def segment_to_generate_params(102 self, segment: Union[SSMLSegment, SSMLBreak]103 ) -> TTSAudioSegment:104 if isinstance(segment, SSMLBreak):105 return TTSAudioSegment(_type="break")106 107 if segment.get("params", None) is not None:108 params = segment.get("params")109 text = segment.get("text", None) or segment.text or ""110 return TTSAudioSegment(**params, text=text)111 112 text = segment.get("text", None) or segment.text or ""113 is_end = segment.get("is_end", False)114 115 text = str(text).strip()116 117 attrs = segment.attrs118 spk = attrs.spk119 style = attrs.style120 121 ss_params = calc_spk_style(spk, style)122 123 if "spk" in ss_params:124 spk = ss_params["spk"]125 126 seed = to_number(attrs.seed, int, ss_params.get("seed") or -1)127 top_k = to_number(attrs.top_k, int, None)128 top_p = to_number(attrs.top_p, float, None)129 temp = to_number(attrs.temp, float, None)130 131 prompt1 = attrs.prompt1 or ss_params.get("prompt1")132 prompt2 = attrs.prompt2 or ss_params.get("prompt2")133 prefix = attrs.prefix or ss_params.get("prefix")134 disable_normalize = attrs.get("normalize", "") == "False"135 136 seg = TTSAudioSegment(137 _type="voice",138 text=text,139 temperature=temp if temp is not None else 0.3,140 top_P=top_p if top_p is not None else 0.5,141 top_K=top_k if top_k is not None else 20,142 spk=spk if spk else -1,143 infer_seed=seed if seed else -1,144 prompt1=prompt1 if prompt1 else "",145 prompt2=prompt2 if prompt2 else "",146 prefix=prefix if prefix else "",147 )148 149 if not disable_normalize:150 seg.text = text_normalize(text, is_end=is_end)151 152 # NOTE 每个batch的默认seed保证前后一致即使是没设置spk的情况153 if seg.spk == -1:154 seg.spk = self.batch_default_spk_seed155 if seg.infer_seed == -1:156 seg.infer_seed = self.batch_default_infer_seed157 158 return seg159 160 def process_break_segments(161 self,162 src_segments: List[SSMLBreak],163 bucket_segments: List[SSMLBreak],164 audio_segments: List[AudioSegment],165 ):166 for segment in bucket_segments:167 index = src_segments.index(segment)168 audio_segments[index] = AudioSegment.silent(169 duration=int(segment.attrs.duration)170 )171 172 def process_voice_segments(173 self,174 src_segments: List[SSMLSegment],175 bucket: List[SSMLSegment],176 audio_segments: List[AudioSegment],177 ):178 for i in range(0, len(bucket), self.batch_size):179 batch = bucket[i : i + self.batch_size]180 param_arr = [self.segment_to_generate_params(segment) for segment in batch]181 182 def append_eos(text: str):183 text = text.strip()184 eos_arr = ["[uv_break]", "[v_break]", "[lbreak]", "[llbreak]"]185 has_eos = False186 for eos in eos_arr:187 if eos in text:188 has_eos = True189 break190 if not has_eos:191 text += self.eos192 return text193 194 # 这里会添加 end_of_text 到 text 之后195 texts = [append_eos(params.text) for params in param_arr]196 197 params = param_arr[0]198 audio_datas = generate_audio.generate_audio_batch(199 texts=texts,200 temperature=params.temperature,201 top_P=params.top_P,202 top_K=params.top_K,203 spk=params.spk,204 infer_seed=params.infer_seed,205 prompt1=params.prompt1,206 prompt2=params.prompt2,207 prefix=params.prefix,208 )209 for idx, segment in enumerate(batch):210 sr, audio_data = audio_datas[idx]211 rate = float(segment.get("rate", "1.0"))212 volume = float(segment.get("volume", "0"))213 pitch = float(segment.get("pitch", "0"))214 215 audio_segment = audio_data_to_segment(audio_data, sr)216 audio_segment = apply_prosody_to_audio_segment(217 audio_segment, rate=rate, volume=volume, pitch=pitch218 )219 # compare by Box object220 original_index = src_segments.index(segment)221 audio_segments[original_index] = audio_segment222 223 def bucket_segments(224 self, segments: List[Union[SSMLSegment, SSMLBreak]]225 ) -> List[List[Union[SSMLSegment, SSMLBreak]]]:226 buckets = {"<break>": []}227 for segment in segments:228 if isinstance(segment, SSMLBreak):229 buckets["<break>"].append(segment)230 continue231 232 params = self.segment_to_generate_params(segment)233 234 if isinstance(params.spk, Speaker):235 params.spk = str(params.spk.id)236 237 key = json.dumps(238 {k: v for k, v in params.items() if k != "text"}, sort_keys=True239 )240 if key not in buckets:241 buckets[key] = []242 buckets[key].append(segment)243 244 return buckets245 246 def split_segments(self, segments: List[Union[SSMLSegment, SSMLBreak]]):247 """248 将 segments 中的 text 经过 spliter 处理成多个 segments249 """250 spliter = SentenceSplitter(threshold=self.spliter_thr)251 ret_segments: List[Union[SSMLSegment, SSMLBreak]] = []252 253 for segment in segments:254 if isinstance(segment, SSMLBreak):255 ret_segments.append(segment)256 continue257 258 text = segment.text259 if not text:260 continue261 262 sentences = spliter.parse(text)263 for sentence in sentences:264 seg = SSMLSegment(265 text=sentence,266 attrs=segment.attrs.copy(),267 params=copy.copy(segment.params),268 )269 ret_segments.append(seg)270 setattr(seg, "_idx", len(ret_segments) - 1)271 272 def is_none_speak_segment(segment: SSMLSegment):273 text = segment.text.strip()274 regexp = r"\[[^\]]+?\]"275 text = re.sub(regexp, "", text)276 text = text.strip()277 if not text:278 return True279 return False280 281 # 将 none_speak 合并到前一个 speak segment282 for i in range(1, len(ret_segments)):283 if is_none_speak_segment(ret_segments[i]):284 ret_segments[i - 1].text += ret_segments[i].text285 ret_segments[i].text = ""286 # 移除空的 segment287 ret_segments = [seg for seg in ret_segments if seg.text.strip()]288 289 return ret_segments290 291 def synthesize_segments(292 self, segments: List[Union[SSMLSegment, SSMLBreak]]293 ) -> List[AudioSegment]:294 segments = self.split_segments(segments)295 audio_segments = [None] * len(segments)296 buckets = self.bucket_segments(segments)297 298 break_segments = buckets.pop("<break>")299 self.process_break_segments(segments, break_segments, audio_segments)300 301 buckets = list(buckets.values())302 303 for bucket in buckets:304 self.process_voice_segments(segments, bucket, audio_segments)305 306 return audio_segments307 308 309# 示例使用310if __name__ == "__main__":311 ctx1 = SSMLContext()312 ctx1.spk = 1313 ctx1.seed = 42314 ctx1.temp = 0.1315 ctx2 = SSMLContext()316 ctx2.spk = 2317 ctx2.seed = 42318 ctx2.temp = 0.1319 ssml_segments = [320 SSMLSegment(text="大🍌,一条大🍌,嘿,你的感觉真的很奇妙", attrs=ctx1.copy()),321 SSMLBreak(duration_ms=1000),322 SSMLSegment(text="大🍉,一个大🍉,嘿,你的感觉真的很奇妙", attrs=ctx1.copy()),323 SSMLSegment(text="大🍊,一个大🍊,嘿,你的感觉真的很奇妙", attrs=ctx2.copy()),324 ]325 326 synthesizer = SynthesizeSegments(batch_size=2)327 audio_segments = synthesizer.synthesize_segments(ssml_segments)328 print(audio_segments)329 combined_audio = combine_audio_segments(audio_segments)330 combined_audio.export("output.wav", format="wav")331 