CoolFace
Apppublic

lenML/ChatTTS-Forge

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