diff --git a/GPT_SoVITS/TTS_infer_pack/TTS.py b/GPT_SoVITS/TTS_infer_pack/TTS.py index 667f1a4a..701a51e1 100644 --- a/GPT_SoVITS/TTS_infer_pack/TTS.py +++ b/GPT_SoVITS/TTS_infer_pack/TTS.py @@ -275,15 +275,6 @@ class TTS_Config: v1_languages: list = ["auto", "en", "zh", "ja", "all_zh", "all_ja"] v2_languages: list = ["auto", "auto_yue", "en", "zh", "ja", "yue", "ko", "all_zh", "all_ja", "all_yue", "all_ko"] languages: list = v2_languages - mute_tokens: dict = { - "v1" : 486, - "v2" : 486, - "v2Pro": 486, - "v2ProPlus": 486, - "v3" : 486, - "v4" : 486, - } - mute_emb_sim_matrix: torch.Tensor = None # "all_zh",#全部按中文识别 # "en",#全部按英文识别#######不变 # "all_ja",#全部按日文识别 @@ -463,6 +454,12 @@ class TTS: self.stop_flag: bool = False self.precision: torch.dtype = torch.float16 if self.configs.is_half else torch.float32 + + # 用于存储字幕/时间轴元数据 + self.last_segments = [] + self.last_sr = None + self.last_fragment_interval = None + self.last_original_text = None def _init_models( self, @@ -499,7 +496,7 @@ class TTS: if if_lora_v3 == True and os.path.exists(path_sovits) == False: info = path_sovits + i18n("SoVITS %s 底模缺失,无法加载相应 LoRA 权重" % model_version) - raise FileNotFoundError(info) + raise FileExistsError(info) # dict_s2 = torch.load(weights_path, map_location=self.configs.device,weights_only=False) dict_s2 = load_sovits_new(weights_path) @@ -607,11 +604,6 @@ class TTS: if self.configs.is_half and str(self.configs.device) != "cpu": self.t2s_model = self.t2s_model.half() - codebook = t2s_model.model.ar_audio_embedding.weight.clone() - mute_emb = codebook[self.configs.mute_tokens[self.configs.version]].unsqueeze(0) - sim_matrix = F.cosine_similarity(mute_emb.float(), codebook.float(), dim=-1) - self.configs.mute_emb_sim_matrix = sim_matrix - def init_vocoder(self, version: str): if version == "v3": if self.vocoder is not None and self.vocoder.__class__.__name__ == "BigVGAN": @@ -1008,25 +1000,21 @@ class TTS: "aux_ref_audio_paths": [], # list.(optional) auxiliary reference audio paths for multi-speaker tone fusion "prompt_text": "", # str.(optional) prompt text for the reference audio "prompt_lang": "", # str.(required) language of the prompt text for the reference audio - "top_k": 15, # int. top k sampling + "top_k": 5, # int. top k sampling "top_p": 1, # float. top p sampling "temperature": 1, # float. temperature for sampling - "text_split_method": "cut1", # str. text split method, see text_segmentation_method.py for details. + "text_split_method": "cut0", # str. text split method, see text_segmentation_method.py for details. "batch_size": 1, # int. batch size for inference "batch_threshold": 0.75, # float. threshold for batch splitting. - "split_bucket": True, # bool. whether to split the batch into multiple buckets. + "split_bucket: True, # bool. whether to split the batch into multiple buckets. + "return_fragment": False, # bool. step by step return the audio fragment. "speed_factor":1.0, # float. control the speed of the synthesized audio. "fragment_interval":0.3, # float. to control the interval of the audio fragment. "seed": -1, # int. random seed for reproducibility. "parallel_infer": True, # bool. whether to use parallel inference. - "repetition_penalty": 1.35, # float. repetition penalty for T2S model. + "repetition_penalty": 1.35 # float. repetition penalty for T2S model. "sample_steps": 32, # int. number of sampling steps for VITS model V3. - "super_sampling": False, # bool. whether to use super-sampling for audio when using VITS model V3. - "return_fragment": False, # bool. step by step return the audio fragment. (Best Quality, Slowest response speed. old version of streaming mode) - "streaming_mode": False, # bool. return audio chunk by chunk. (Medium quality, Slow response speed) - "overlap_length": 2, # int. overlap length of semantic tokens for streaming mode. - "min_chunk_length": 16, # int. The minimum chunk length of semantic tokens for streaming mode. (affects audio chunk size) - "fixed_length_chunk": False, # bool. When turned on, it can achieve faster streaming response, but with lower quality. (lower quality, faster response speed) + "super_sampling": False, # bool. whether to use super-sampling for audio when using VITS model V3. } returns: Tuple[int, np.ndarray]: sampling rate and audio data. @@ -1034,15 +1022,16 @@ class TTS: ########## variables initialization ########### self.stop_flag: bool = False text: str = inputs.get("text", "") + self.last_original_text = text # 存储原始文本 text_lang: str = inputs.get("text_lang", "") ref_audio_path: str = inputs.get("ref_audio_path", "") aux_ref_audio_paths: list = inputs.get("aux_ref_audio_paths", []) prompt_text: str = inputs.get("prompt_text", "") prompt_lang: str = inputs.get("prompt_lang", "") - top_k: int = inputs.get("top_k", 15) + top_k: int = inputs.get("top_k", 5) top_p: float = inputs.get("top_p", 1) temperature: float = inputs.get("temperature", 1) - text_split_method: str = inputs.get("text_split_method", "cut1") + text_split_method: str = inputs.get("text_split_method", "cut0") batch_size = inputs.get("batch_size", 1) batch_threshold = inputs.get("batch_threshold", 0.75) speed_factor = inputs.get("speed_factor", 1.0) @@ -1056,43 +1045,19 @@ class TTS: repetition_penalty = inputs.get("repetition_penalty", 1.35) sample_steps = inputs.get("sample_steps", 32) super_sampling = inputs.get("super_sampling", False) - streaming_mode = inputs.get("streaming_mode", False) - overlap_length = inputs.get("overlap_length", 2) - min_chunk_length = inputs.get("min_chunk_length", 16) - fixed_length_chunk = inputs.get("fixed_length_chunk", False) - chunk_split_thershold = 0.0 # 该值代表语义token与mute token的余弦相似度阈值,若大于该阈值,则视为可切分点。 - if parallel_infer and not streaming_mode: + if parallel_infer: print(i18n("并行推理模式已开启")) self.t2s_model.model.infer_panel = self.t2s_model.model.infer_panel_batch_infer - elif not parallel_infer and streaming_mode and not self.configs.use_vocoder: - print(i18n("流式推理模式已开启")) - self.t2s_model.model.infer_panel = self.t2s_model.model.infer_panel_naive - elif streaming_mode and self.configs.use_vocoder: - print(i18n("SoVits V3/4模型不支持流式推理模式,已自动回退到分段返回模式")) - streaming_mode = False - return_fragment = True - if parallel_infer: - self.t2s_model.model.infer_panel = self.t2s_model.model.infer_panel_batch_infer - else: - self.t2s_model.model.infer_panel = self.t2s_model.model.infer_panel_naive_batched - # self.t2s_model.model.infer_panel = self.t2s_model.model.infer_panel_naive - elif parallel_infer and streaming_mode: - print(i18n("不支持同时开启并行推理和流式推理模式,已自动关闭并行推理模式")) - parallel_infer = False - self.t2s_model.model.infer_panel = self.t2s_model.model.infer_panel_naive else: - print(i18n("朴素推理模式已开启")) + print(i18n("并行推理模式已关闭")) self.t2s_model.model.infer_panel = self.t2s_model.model.infer_panel_naive_batched - if return_fragment and streaming_mode: - print(i18n("流式推理模式不支持分段返回,已自动关闭分段返回")) - return_fragment = False - - if (return_fragment or streaming_mode) and split_bucket: - print(i18n("分段返回模式/流式推理模式不支持分桶处理,已自动关闭分桶处理")) - split_bucket = False - + if return_fragment: + print(i18n("分段返回模式已开启")) + if split_bucket: + split_bucket = False + print(i18n("分段返回模式不支持分桶处理,已自动关闭分桶处理")) if split_bucket and speed_factor == 1.0 and not (self.configs.use_vocoder and parallel_infer): print(i18n("分桶处理模式已开启")) @@ -1105,9 +1070,9 @@ class TTS: else: print(i18n("分桶处理模式已关闭")) - # if fragment_interval < 0.01: - # fragment_interval = 0.01 - # print(i18n("分段间隔过小,已自动设置为0.01")) + if fragment_interval < 0.01: + fragment_interval = 0.01 + print(i18n("分段间隔过小,已自动设置为0.01")) no_prompt_text = False if prompt_text in [None, ""]: @@ -1168,7 +1133,7 @@ class TTS: ###### text preprocessing ######## t1 = time.perf_counter() data: list = None - if not (return_fragment or streaming_mode): + if not return_fragment: data = self.text_preprocessor.preprocess(text, text_lang, text_split_method, self.configs.version) if len(data) == 0: yield 16000, np.zeros(int(16000), dtype=np.int16) @@ -1228,11 +1193,11 @@ class TTS: t_34 = 0.0 t_45 = 0.0 audio = [] - is_first_package = True + segment_texts_batched = [] # 收集每段的文本(保持批次结构) output_sr = self.configs.sampling_rate if not self.configs.use_vocoder else self.vocoder_configs["sr"] for item in data: t3 = time.perf_counter() - if return_fragment or streaming_mode: + if return_fragment: item = make_batch(item) if item is None: continue @@ -1247,6 +1212,14 @@ class TTS: max_len = item["max_len"] print(i18n("前端处理后的文本(每句):"), norm_text) + + # 收集每段的文本(保持批次结构) + batch_texts = [] + if isinstance(norm_text, list): + batch_texts = norm_text + else: + batch_texts = [norm_text] + segment_texts_batched.append(batch_texts) if no_prompt_text: prompt = None else: @@ -1254,228 +1227,108 @@ class TTS: self.prompt_cache["prompt_semantic"].expand(len(all_phoneme_ids), -1).to(self.configs.device) ) + print(f"############ {i18n('预测语义Token')} ############") + pred_semantic_list, idx_list = self.t2s_model.model.infer_panel( + all_phoneme_ids, + all_phoneme_lens, + prompt, + all_bert_features, + # prompt_phone_len=ph_offset, + top_k=top_k, + top_p=top_p, + temperature=temperature, + early_stop_num=self.configs.hz * self.configs.max_sec, + max_len=max_len, + repetition_penalty=repetition_penalty, + ) + t4 = time.perf_counter() + t_34 += t4 - t3 + refer_audio_spec = [] - - sv_emb = [] if self.is_v2pro else None + if self.is_v2pro: + sv_emb = [] for spec, audio_tensor in self.prompt_cache["refer_spec"]: spec = spec.to(dtype=self.precision, device=self.configs.device) refer_audio_spec.append(spec) if self.is_v2pro: sv_emb.append(self.sv_model.compute_embedding3(audio_tensor)) - if not streaming_mode: - print(f"############ {i18n('预测语义Token')} ############") - pred_semantic_list, idx_list = self.t2s_model.model.infer_panel( - all_phoneme_ids, - all_phoneme_lens, - prompt, - all_bert_features, - # prompt_phone_len=ph_offset, - top_k=top_k, - top_p=top_p, - temperature=temperature, - early_stop_num=self.configs.hz * self.configs.max_sec, - max_len=max_len, - repetition_penalty=repetition_penalty, - ) - t4 = time.perf_counter() - t_34 += t4 - t3 - - - batch_audio_fragment = [] - - # ## vits并行推理 method 1 - # pred_semantic_list = [item[-idx:] for item, idx in zip(pred_semantic_list, idx_list)] - # pred_semantic_len = torch.LongTensor([item.shape[0] for item in pred_semantic_list]).to(self.configs.device) - # pred_semantic = self.batch_sequences(pred_semantic_list, axis=0, pad_value=0).unsqueeze(0) - # max_len = 0 - # for i in range(0, len(batch_phones)): - # max_len = max(max_len, batch_phones[i].shape[-1]) - # batch_phones = self.batch_sequences(batch_phones, axis=0, pad_value=0, max_length=max_len) - # batch_phones = batch_phones.to(self.configs.device) - # batch_audio_fragment = (self.vits_model.batched_decode( - # pred_semantic, pred_semantic_len, batch_phones, batch_phones_len,refer_audio_spec - # )) - print(f"############ {i18n('合成音频')} ############") - if not self.configs.use_vocoder: - if speed_factor == 1.0: - print(f"{i18n('并行合成中')}...") - # ## vits并行推理 method 2 - pred_semantic_list = [item[-idx:] for item, idx in zip(pred_semantic_list, idx_list)] - upsample_rate = math.prod(self.vits_model.upsample_rates) - audio_frag_idx = [ - pred_semantic_list[i].shape[0] * 2 * upsample_rate - for i in range(0, len(pred_semantic_list)) - ] - audio_frag_end_idx = [sum(audio_frag_idx[: i + 1]) for i in range(0, len(audio_frag_idx))] - all_pred_semantic = ( - torch.cat(pred_semantic_list).unsqueeze(0).unsqueeze(0).to(self.configs.device) - ) - _batch_phones = torch.cat(batch_phones).unsqueeze(0).to(self.configs.device) + batch_audio_fragment = [] + # ## vits并行推理 method 1 + # pred_semantic_list = [item[-idx:] for item, idx in zip(pred_semantic_list, idx_list)] + # pred_semantic_len = torch.LongTensor([item.shape[0] for item in pred_semantic_list]).to(self.configs.device) + # pred_semantic = self.batch_sequences(pred_semantic_list, axis=0, pad_value=0).unsqueeze(0) + # max_len = 0 + # for i in range(0, len(batch_phones)): + # max_len = max(max_len, batch_phones[i].shape[-1]) + # batch_phones = self.batch_sequences(batch_phones, axis=0, pad_value=0, max_length=max_len) + # batch_phones = batch_phones.to(self.configs.device) + # batch_audio_fragment = (self.vits_model.batched_decode( + # pred_semantic, pred_semantic_len, batch_phones, batch_phones_len,refer_audio_spec + # )) + print(f"############ {i18n('合成音频')} ############") + if not self.configs.use_vocoder: + if speed_factor == 1.0: + print(f"{i18n('并行合成中')}...") + # ## vits并行推理 method 2 + pred_semantic_list = [item[-idx:] for item, idx in zip(pred_semantic_list, idx_list)] + upsample_rate = math.prod(self.vits_model.upsample_rates) + audio_frag_idx = [ + pred_semantic_list[i].shape[0] * 2 * upsample_rate + for i in range(0, len(pred_semantic_list)) + ] + audio_frag_end_idx = [sum(audio_frag_idx[: i + 1]) for i in range(0, len(audio_frag_idx))] + all_pred_semantic = ( + torch.cat(pred_semantic_list).unsqueeze(0).unsqueeze(0).to(self.configs.device) + ) + _batch_phones = torch.cat(batch_phones).unsqueeze(0).to(self.configs.device) + if self.is_v2pro != True: _batch_audio_fragment = self.vits_model.decode( - all_pred_semantic, _batch_phones, refer_audio_spec, speed=speed_factor, sv_emb=sv_emb - ).detach()[0, 0, :] - - audio_frag_end_idx.insert(0, 0) - batch_audio_fragment = [ - _batch_audio_fragment[audio_frag_end_idx[i - 1] : audio_frag_end_idx[i]] - for i in range(1, len(audio_frag_end_idx)) - ] + all_pred_semantic, _batch_phones, refer_audio_spec, speed=speed_factor + ).detach()[0, 0, :] else: - # ## vits串行推理 - for i, idx in enumerate(tqdm(idx_list)): - phones = batch_phones[i].unsqueeze(0).to(self.configs.device) - _pred_semantic = ( - pred_semantic_list[i][-idx:].unsqueeze(0).unsqueeze(0) - ) # .unsqueeze(0)#mq要多unsqueeze一次 + _batch_audio_fragment = self.vits_model.decode( + all_pred_semantic, _batch_phones, refer_audio_spec, speed=speed_factor, sv_emb=sv_emb + ).detach()[0, 0, :] + audio_frag_end_idx.insert(0, 0) + batch_audio_fragment = [ + _batch_audio_fragment[audio_frag_end_idx[i - 1] : audio_frag_end_idx[i]] + for i in range(1, len(audio_frag_end_idx)) + ] + else: + # ## vits串行推理 + for i, idx in enumerate(tqdm(idx_list)): + phones = batch_phones[i].unsqueeze(0).to(self.configs.device) + _pred_semantic = ( + pred_semantic_list[i][-idx:].unsqueeze(0).unsqueeze(0) + ) # .unsqueeze(0)#mq要多unsqueeze一次 + if self.is_v2pro != True: audio_fragment = self.vits_model.decode( - _pred_semantic, phones, refer_audio_spec, speed=speed_factor, sv_emb=sv_emb - ).detach()[0, 0, :] - batch_audio_fragment.append(audio_fragment) ###试试重建不带上prompt部分 - else: - if parallel_infer: - print(f"{i18n('并行合成中')}...") - audio_fragments = self.using_vocoder_synthesis_batched_infer( - idx_list, pred_semantic_list, batch_phones, speed=speed_factor, sample_steps=sample_steps - ) - batch_audio_fragment.extend(audio_fragments) - else: - for i, idx in enumerate(tqdm(idx_list)): - phones = batch_phones[i].unsqueeze(0).to(self.configs.device) - _pred_semantic = ( - pred_semantic_list[i][-idx:].unsqueeze(0).unsqueeze(0) - ) # .unsqueeze(0)#mq要多unsqueeze一次 - audio_fragment = self.using_vocoder_synthesis( - _pred_semantic, phones, speed=speed_factor, sample_steps=sample_steps - ) - batch_audio_fragment.append(audio_fragment) - + _pred_semantic, phones, refer_audio_spec, speed=speed_factor + ).detach()[0, 0, :] + else: + audio_fragment = self.vits_model.decode( + _pred_semantic, phones, refer_audio_spec, speed=speed_factor, sv_emb=sv_emb + ).detach()[0, 0, :] + batch_audio_fragment.append(audio_fragment) ###试试重建不带上prompt部分 else: - # refer_audio_spec: torch.Tensor = [ - # item.to(dtype=self.precision, device=self.configs.device) - # for item in self.prompt_cache["refer_spec"] - # ] - semantic_token_generator =self.t2s_model.model.infer_panel( - all_phoneme_ids[0].unsqueeze(0), - all_phoneme_lens, - prompt, - all_bert_features[0].unsqueeze(0), - top_k=top_k, - top_p=top_p, - temperature=temperature, - early_stop_num=self.configs.hz * self.configs.max_sec, - max_len=max_len, - repetition_penalty=repetition_penalty, - streaming_mode=True, - chunk_length=min_chunk_length, - mute_emb_sim_matrix=self.configs.mute_emb_sim_matrix if not fixed_length_chunk else None, - chunk_split_thershold=chunk_split_thershold, - ) - t4 = time.perf_counter() - t_34 += t4 - t3 - phones = batch_phones[0].unsqueeze(0).to(self.configs.device) - is_first_chunk = True - - if not self.configs.use_vocoder: - # if speed_factor == 1.0: - # upsample_rate = math.prod(self.vits_model.upsample_rates)*(2 if self.vits_model.semantic_frame_rate == "25hz" else 1) - # else: - upsample_rate = math.prod(self.vits_model.upsample_rates)*((2 if self.vits_model.semantic_frame_rate == "25hz" else 1)/speed_factor) + if parallel_infer: + print(f"{i18n('并行合成中')}...") + audio_fragments = self.using_vocoder_synthesis_batched_infer( + idx_list, pred_semantic_list, batch_phones, speed=speed_factor, sample_steps=sample_steps + ) + batch_audio_fragment.extend(audio_fragments) else: - # if speed_factor == 1.0: - # upsample_rate = self.vocoder_configs["upsample_rate"]*(3.875 if self.configs.version == "v3" else 4) - # else: - upsample_rate = self.vocoder_configs["upsample_rate"]*((3.875 if self.configs.version == "v3" else 4)/speed_factor) - - last_audio_chunk = None - # last_tokens = None - last_latent = None - previous_tokens = [] - overlap_len = overlap_length - overlap_size = math.ceil(overlap_length*upsample_rate) - for semantic_tokens, is_final in semantic_token_generator: - if semantic_tokens is None and last_audio_chunk is not None: - yield self.audio_postprocess( - [[last_audio_chunk[-overlap_size:]]], - output_sr, - None, - speed_factor, - False, - 0.0, - super_sampling if self.configs.use_vocoder and self.configs.version == "v3" else False, - ) - break - - _semantic_tokens = semantic_tokens - print(f"semantic_tokens shape:{semantic_tokens.shape}") - - previous_tokens.append(semantic_tokens) - - _semantic_tokens = torch.cat(previous_tokens, dim=-1) - - if not is_first_chunk and semantic_tokens.shape[-1] < 10: - overlap_len = overlap_length+(10-semantic_tokens.shape[-1]) - else: - overlap_len = overlap_length - - - if not self.configs.use_vocoder: - token_padding_length = 0 - # token_padding_length = int(phones.shape[-1]*2)-_semantic_tokens.shape[-1] - # if token_padding_length>0: - # _semantic_tokens = F.pad(_semantic_tokens, (0, token_padding_length), "constant", 486) - # else: - # token_padding_length = 0 - - audio_chunk, latent, latent_mask = self.vits_model.decode_streaming( - _semantic_tokens.unsqueeze(0), - phones, refer_audio_spec, - speed=speed_factor, - sv_emb=sv_emb, - result_length=semantic_tokens.shape[-1]+overlap_len if not is_first_chunk else None, - overlap_frames=last_latent[:,:,-overlap_len*(2 if self.vits_model.semantic_frame_rate == "25hz" else 1):] \ - if last_latent is not None else None, - padding_length=token_padding_length - ) - audio_chunk=audio_chunk.detach()[0, 0, :] - else: - raise RuntimeError(i18n("SoVits V3/4模型不支持流式推理模式")) - - if overlap_len>overlap_length: - audio_chunk=audio_chunk[-int((overlap_length+semantic_tokens.shape[-1])*upsample_rate):] - - audio_chunk_ = audio_chunk - if is_first_chunk and not is_final: - is_first_chunk = False - audio_chunk_ = audio_chunk_[:-overlap_size] - elif is_first_chunk and is_final: - is_first_chunk = False - elif not is_first_chunk and not is_final: - audio_chunk_ = self.sola_algorithm([last_audio_chunk, audio_chunk_], overlap_size) - audio_chunk_ = ( - audio_chunk_[last_audio_chunk.shape[0]-overlap_size:-overlap_size] if not is_final \ - else audio_chunk_[last_audio_chunk.shape[0]-overlap_size:] - ) - - last_latent = latent - last_audio_chunk = audio_chunk - yield self.audio_postprocess( - [[audio_chunk_]], - output_sr, - None, - speed_factor, - False, - 0.0, - super_sampling if self.configs.use_vocoder and self.configs.version == "v3" else False, + for i, idx in enumerate(tqdm(idx_list)): + phones = batch_phones[i].unsqueeze(0).to(self.configs.device) + _pred_semantic = ( + pred_semantic_list[i][-idx:].unsqueeze(0).unsqueeze(0) + ) # .unsqueeze(0)#mq要多unsqueeze一次 + audio_fragment = self.using_vocoder_synthesis( + _pred_semantic, phones, speed=speed_factor, sample_steps=sample_steps ) - - if is_first_package: - print(f"first_package_delay: {time.perf_counter()-t0:.3f}") - is_first_package = False - - - yield output_sr, np.zeros(int(output_sr*fragment_interval), dtype=np.int16) + batch_audio_fragment.append(audio_fragment) t5 = time.perf_counter() t_45 += t5 - t4 @@ -1489,19 +1342,19 @@ class TTS: False, fragment_interval, super_sampling if self.configs.use_vocoder and self.configs.version == "v3" else False, + segment_texts_batched, ) - elif streaming_mode:... else: audio.append(batch_audio_fragment) if self.stop_flag: - yield output_sr, np.zeros(int(output_sr), dtype=np.int16) + yield 16000, np.zeros(int(16000), dtype=np.int16) return - if not (return_fragment or streaming_mode): + if not return_fragment: print("%.3f\t%.3f\t%.3f\t%.3f" % (t1 - t0, t2 - t1, t_34, t_45)) if len(audio) == 0: - yield output_sr, np.zeros(int(output_sr), dtype=np.int16) + yield 16000, np.zeros(int(16000), dtype=np.int16) return yield self.audio_postprocess( audio, @@ -1511,6 +1364,7 @@ class TTS: split_bucket, fragment_interval, super_sampling if self.configs.use_vocoder and self.configs.version == "v3" else False, + segment_texts_batched, ) except Exception as e: @@ -1547,25 +1401,62 @@ class TTS: split_bucket: bool = True, fragment_interval: float = 0.3, super_sampling: bool = False, + segment_texts: list = None, ) -> Tuple[int, np.ndarray]: - if fragment_interval>0: - zero_wav = torch.zeros( - int(self.configs.sampling_rate * fragment_interval), dtype=self.precision, device=self.configs.device - ) + zero_wav = torch.zeros( + int(self.configs.sampling_rate * fragment_interval), dtype=self.precision, device=self.configs.device + ) for i, batch in enumerate(audio): for j, audio_fragment in enumerate(batch): max_audio = torch.abs(audio_fragment).max() # 简单防止16bit爆音 if max_audio > 1: audio_fragment /= max_audio - audio_fragment: torch.Tensor = torch.cat([audio_fragment, zero_wav], dim=0) if fragment_interval>0 else audio_fragment + audio_fragment: torch.Tensor = torch.cat([audio_fragment, zero_wav], dim=0) audio[i][j] = audio_fragment + # 恢复顺序并平铺 if split_bucket: audio = self.recovery_order(audio, batch_index_list) + # 同时恢复文本顺序(segment_texts 是批次结构) + if segment_texts: + segment_texts = self.recovery_order(segment_texts, batch_index_list) else: # audio = [item for batch in audio for item in batch] audio = sum(audio, []) + + # 平铺文本列表(无论 split_bucket 是否为 True,都需要平铺) + if segment_texts and isinstance(segment_texts, list) and len(segment_texts) > 0: + # 检查第一个元素是否是列表(批次结构) + if isinstance(segment_texts[0], list): + segment_texts = sum(segment_texts, []) # 平铺成一维列表 + + # 记录分段时间轴信息(在拼接之前) + segments_info = [] + current_sample = 0 + gap_samples = int(fragment_interval * sr) + + for idx, segment_audio in enumerate(audio): + samples = len(segment_audio) + start_sample = current_sample + end_sample = current_sample + samples - gap_samples # 减去静音间隔 + + segment_info = { + "idx": idx, + "text": segment_texts[idx] if segment_texts and idx < len(segment_texts) else "", + "samples": samples - gap_samples, + "start": start_sample, + "end": end_sample, + "start_time": start_sample / sr, + "end_time": end_sample / sr, + } + segments_info.append(segment_info) + current_sample = end_sample + gap_samples # 加上静音间隔,准备下一段的起始位置 + + # 存储元数据到对象属性 + self.last_segments = segments_info + self.last_sr = sr + self.last_fragment_interval = fragment_interval audio = torch.cat(audio, dim=0) @@ -1578,16 +1469,12 @@ class TTS: max_audio = np.abs(audio).max() if max_audio > 1: audio /= max_audio - audio = (audio * 32768).astype(np.int16) - else: - audio = audio.cpu().numpy() - audio = (audio * 32768).astype(np.int16) t2 = time.perf_counter() print(f"超采样用时:{t2 - t1:.3f}s") else: audio = audio.cpu().numpy() - audio = (audio * 32768).astype(np.int16) + audio = (audio * 32768).astype(np.int16) # try: # if speed_factor != 1.0: @@ -1767,10 +1654,7 @@ class TTS: pos += chunk_len * upsample_rate audio = self.sola_algorithm(audio_fragments, overlapped_len * upsample_rate) - if padding_len > 0: - audio = audio[overlapped_len * upsample_rate : -padding_len * upsample_rate] - else: - audio = audio[overlapped_len * upsample_rate :] + audio = audio[overlapped_len * upsample_rate : -padding_len * upsample_rate] audio_fragments = [] for feat_len in feat_lens: @@ -1784,43 +1668,24 @@ class TTS: self, audio_fragments: List[torch.Tensor], overlap_len: int, - search_len:int= 320 ): - # overlap_len-=search_len - - dtype = audio_fragments[0].dtype - for i in range(len(audio_fragments) - 1): - f1 = audio_fragments[i].float() - f2 = audio_fragments[i + 1].float() + f1 = audio_fragments[i] + f2 = audio_fragments[i + 1] w1 = f1[-overlap_len:] - w2 = f2[:overlap_len+search_len] - # w2 = w2[-w2.shape[-1]//2:] - # assert w1.shape == w2.shape - corr_norm = F.conv1d(w2.view(1, 1, -1), w1.view(1, 1, -1)).view(-1) - - corr_den = F.conv1d(w2.view(1, 1, -1)**2, torch.ones_like(w1).view(1, 1, -1)).view(-1)+ 1e-8 - idx = (corr_norm/corr_den.sqrt()).argmax() - - print(f"seg_idx: {idx}") - - # idx = corr.argmax() - f1_ = f1[: -overlap_len] + w2 = f2[:overlap_len] + assert w1.shape == w2.shape + corr = F.conv1d(w1.view(1, 1, -1), w2.view(1, 1, -1), padding=w2.shape[-1] // 2).view(-1)[:-1] + idx = corr.argmax() + f1_ = f1[: -(overlap_len - idx)] audio_fragments[i] = f1_ f2_ = f2[idx:] - window = torch.hann_window((overlap_len) * 2, device=f1.device, dtype=f1.dtype) - f2_[: overlap_len] = ( - window[: overlap_len] * f2_[: overlap_len] - + window[overlap_len :] * f1[-overlap_len :] + window = torch.hann_window((overlap_len - idx) * 2, device=f1.device, dtype=f1.dtype) + f2_[: (overlap_len - idx)] = ( + window[: (overlap_len - idx)] * f2_[: (overlap_len - idx)] + + window[(overlap_len - idx) :] * f1[-(overlap_len - idx) :] ) - - # window = torch.sin(torch.arange((overlap_len - idx), device=f1.device) * np.pi / (overlap_len - idx)) - # f2_[: (overlap_len - idx)] = ( - # window * f2_[: (overlap_len - idx)] - # + (1-window) * f1[-(overlap_len - idx) :] - # ) - audio_fragments[i + 1] = f2_ - return torch.cat(audio_fragments, 0).to(dtype) + return torch.cat(audio_fragments, 0) diff --git a/api_v2.py b/api_v2.py index 21511db3..515e179a 100644 --- a/api_v2.py +++ b/api_v2.py @@ -27,23 +27,20 @@ POST: "aux_ref_audio_paths": [], # list.(optional) auxiliary reference audio paths for multi-speaker tone fusion "prompt_text": "", # str.(optional) prompt text for the reference audio "prompt_lang": "", # str.(required) language of the prompt text for the reference audio - "top_k": 15, # int. top k sampling + "top_k": 5, # int. top k sampling "top_p": 1, # float. top p sampling "temperature": 1, # float. temperature for sampling - "text_split_method": "cut5", # str. text split method, see text_segmentation_method.py for details. + "text_split_method": "cut0", # str. text split method, see text_segmentation_method.py for details. "batch_size": 1, # int. batch size for inference "batch_threshold": 0.75, # float. threshold for batch splitting. "split_bucket": True, # bool. whether to split the batch into multiple buckets. "speed_factor":1.0, # float. control the speed of the synthesized audio. - "fragment_interval":0.3, # float. to control the interval of the audio fragment. + "streaming_mode": False, # bool. whether to return a streaming response. "seed": -1, # int. random seed for reproducibility. "parallel_infer": True, # bool. whether to use parallel inference. "repetition_penalty": 1.35, # float. repetition penalty for T2S model. "sample_steps": 32, # int. number of sampling steps for VITS model V3. - "super_sampling": False, # bool. whether to use super-sampling for audio when using VITS model V3. - "streaming_mode": False, # bool or int. return audio chunk by chunk.T he available options are: 0,1,2,3 or True/False (0/False: Disabled | 1/True: Best Quality, Slowest response speed (old version streaming_mode) | 2: Medium Quality, Slow response speed | 3: Lower Quality, Faster response speed ) - "overlap_length": 2, # int. overlap length of semantic tokens for streaming mode. - "min_chunk_length": 16, # int. The minimum chunk length of semantic tokens for streaming mode. (affects audio chunk size) + "super_sampling": False # bool. whether to use super-sampling for audio when using VITS model V3. } ``` @@ -104,7 +101,7 @@ RESP: import os import sys import traceback -from typing import Generator, Union +from typing import Generator now_dir = os.getcwd() sys.path.append(now_dir) @@ -124,7 +121,6 @@ from tools.i18n.i18n import I18nAuto from GPT_SoVITS.TTS_infer_pack.TTS import TTS, TTS_Config from GPT_SoVITS.TTS_infer_pack.text_segmentation_method import get_method_names as get_cut_method_names from pydantic import BaseModel -import threading # print(sys.path) i18n = I18nAuto() @@ -158,7 +154,7 @@ class TTS_Request(BaseModel): aux_ref_audio_paths: list = None prompt_lang: str = None prompt_text: str = "" - top_k: int = 15 + top_k: int = 5 top_p: float = 1 temperature: float = 1 text_split_method: str = "cut5" @@ -169,58 +165,17 @@ class TTS_Request(BaseModel): fragment_interval: float = 0.3 seed: int = -1 media_type: str = "wav" - streaming_mode: Union[bool, int] = False + streaming_mode: bool = False parallel_infer: bool = True repetition_penalty: float = 1.35 sample_steps: int = 32 super_sampling: bool = False - overlap_length: int = 2 - min_chunk_length: int = 16 +### modify from https://github.com/RVC-Boss/GPT-SoVITS/pull/894/files def pack_ogg(io_buffer: BytesIO, data: np.ndarray, rate: int): - # Author: AkagawaTsurunaki - # Issue: - # Stack overflow probabilistically occurs - # when the function `sf_writef_short` of `libsndfile_64bit.dll` is called - # using the Python library `soundfile` - # Note: - # This is an issue related to `libsndfile`, not this project itself. - # It happens when you generate a large audio tensor (about 499804 frames in my PC) - # and try to convert it to an ogg file. - # Related: - # https://github.com/RVC-Boss/GPT-SoVITS/issues/1199 - # https://github.com/libsndfile/libsndfile/issues/1023 - # https://github.com/bastibe/python-soundfile/issues/396 - # Suggestion: - # Or split the whole audio data into smaller audio segment to avoid stack overflow? - - def handle_pack_ogg(): - with sf.SoundFile(io_buffer, mode="w", samplerate=rate, channels=1, format="ogg") as audio_file: - audio_file.write(data) - - - - # See: https://docs.python.org/3/library/threading.html - # The stack size of this thread is at least 32768 - # If stack overflow error still occurs, just modify the `stack_size`. - # stack_size = n * 4096, where n should be a positive integer. - # Here we chose n = 4096. - stack_size = 4096 * 4096 - try: - threading.stack_size(stack_size) - pack_ogg_thread = threading.Thread(target=handle_pack_ogg) - pack_ogg_thread.start() - pack_ogg_thread.join() - except RuntimeError as e: - # If changing the thread stack size is unsupported, a RuntimeError is raised. - print("RuntimeError: {}".format(e)) - print("Changing the thread stack size is unsupported.") - except ValueError as e: - # If the specified stack size is invalid, a ValueError is raised and the stack size is unmodified. - print("ValueError: {}".format(e)) - print("The specified stack size is invalid.") - + with sf.SoundFile(io_buffer, mode="w", samplerate=rate, channels=1, format="ogg") as audio_file: + audio_file.write(data) return io_buffer @@ -331,8 +286,8 @@ def check_params(req: dict): ) if media_type not in ["wav", "raw", "ogg", "aac"]: return JSONResponse(status_code=400, content={"message": f"media_type: {media_type} is not supported"}) - # elif media_type == "ogg" and not streaming_mode: - # return JSONResponse(status_code=400, content={"message": "ogg format is not supported in non-streaming mode"}) + elif media_type == "ogg" and not streaming_mode: + return JSONResponse(status_code=400, content={"message": "ogg format is not supported in non-streaming mode"}) if text_split_method not in cut_method_names: return JSONResponse( @@ -352,26 +307,25 @@ async def tts_handle(req: dict): "text": "", # str.(required) text to be synthesized "text_lang: "", # str.(required) language of the text to be synthesized "ref_audio_path": "", # str.(required) reference audio path - "aux_ref_audio_paths": [], # list.(optional) auxiliary reference audio paths for multi-speaker tone fusion + "aux_ref_audio_paths": [], # list.(optional) auxiliary reference audio paths for multi-speaker synthesis "prompt_text": "", # str.(optional) prompt text for the reference audio "prompt_lang": "", # str.(required) language of the prompt text for the reference audio - "top_k": 15, # int. top k sampling + "top_k": 5, # int. top k sampling "top_p": 1, # float. top p sampling "temperature": 1, # float. temperature for sampling "text_split_method": "cut5", # str. text split method, see text_segmentation_method.py for details. "batch_size": 1, # int. batch size for inference "batch_threshold": 0.75, # float. threshold for batch splitting. - "split_bucket": True, # bool. whether to split the batch into multiple buckets. + "split_bucket: True, # bool. whether to split the batch into multiple buckets. "speed_factor":1.0, # float. control the speed of the synthesized audio. "fragment_interval":0.3, # float. to control the interval of the audio fragment. "seed": -1, # int. random seed for reproducibility. - "parallel_infer": True, # bool. whether to use parallel inference. - "repetition_penalty": 1.35, # float. repetition penalty for T2S model. + "media_type": "wav", # str. media type of the output audio, support "wav", "raw", "ogg", "aac". + "streaming_mode": False, # bool. whether to return a streaming response. + "parallel_infer": True, # bool.(optional) whether to use parallel inference. + "repetition_penalty": 1.35 # float.(optional) repetition penalty for T2S model. "sample_steps": 32, # int. number of sampling steps for VITS model V3. - "super_sampling": False, # bool. whether to use super-sampling for audio when using VITS model V3. - "streaming_mode": False, # bool or int. return audio chunk by chunk.T he available options are: 0,1,2,3 or True/False (0/False: Disabled | 1/True: Best Quality, Slowest response speed (old version streaming_mode) | 2: Medium Quality, Slow response speed | 3: Lower Quality, Faster response speed ) - "overlap_length": 2, # int. overlap length of semantic tokens for streaming mode. - "min_chunk_length": 16, # int. The minimum chunk length of semantic tokens for streaming mode. (affects audio chunk size) + "super_sampling": False, # bool. whether to use super-sampling for audio when using VITS model V3. } returns: StreamingResponse: audio stream response. @@ -384,35 +338,9 @@ async def tts_handle(req: dict): check_res = check_params(req) if check_res is not None: return check_res - - if streaming_mode == 0: - streaming_mode = False - return_fragment = False - fixed_length_chunk = False - elif streaming_mode == 1: - streaming_mode = False - return_fragment = True - fixed_length_chunk = False - elif streaming_mode == 2: - streaming_mode = True - return_fragment = False - fixed_length_chunk = False - elif streaming_mode == 3: - streaming_mode = True - return_fragment = False - fixed_length_chunk = True - - else: - return JSONResponse(status_code=400, content={"message": f"the value of streaming_mode must be 0, 1, 2, 3(int) or true/false(bool)"}) - - req["streaming_mode"] = streaming_mode - req["return_fragment"] = return_fragment - req["fixed_length_chunk"] = fixed_length_chunk - - print(f"{streaming_mode} {return_fragment} {fixed_length_chunk}") - - streaming_mode = streaming_mode or return_fragment + if streaming_mode or return_fragment: + req["return_fragment"] = True try: tts_generator = tts_pipeline.run(req) @@ -460,10 +388,10 @@ async def tts_get_endpoint( aux_ref_audio_paths: list = None, prompt_lang: str = None, prompt_text: str = "", - top_k: int = 15, + top_k: int = 5, top_p: float = 1, temperature: float = 1, - text_split_method: str = "cut5", + text_split_method: str = "cut0", batch_size: int = 1, batch_threshold: float = 0.75, split_bucket: bool = True, @@ -471,13 +399,11 @@ async def tts_get_endpoint( fragment_interval: float = 0.3, seed: int = -1, media_type: str = "wav", + streaming_mode: bool = False, parallel_infer: bool = True, repetition_penalty: float = 1.35, sample_steps: int = 32, super_sampling: bool = False, - streaming_mode: Union[bool, int] = False, - overlap_length: int = 2, - min_chunk_length: int = 16, ): req = { "text": text, @@ -502,8 +428,6 @@ async def tts_get_endpoint( "repetition_penalty": float(repetition_penalty), "sample_steps": int(sample_steps), "super_sampling": super_sampling, - "overlap_length": int(overlap_length), - "min_chunk_length": int(min_chunk_length), } return await tts_handle(req) @@ -565,6 +489,42 @@ async def set_sovits_weights(weights_path: str = None): return JSONResponse(status_code=200, content={"message": "success"}) +@APP.get("/tts_meta") +async def get_tts_meta(): + """ + 获取最后一次 TTS 推理的字幕/时间轴元数据 + + 返回: + { + "segments": [ + { + "idx": 0, + "text": "文本内容", + "samples": 12345, + "start": 0, + "end": 12345, + "start_time": 0.0, + "end_time": 0.385 + }, + ... + ], + "sr": 32000, + "fragment_interval": 0.3, + "original_text": "原始完整文本" + } + """ + try: + meta_data = { + "segments": tts_pipeline.last_segments, + "sr": tts_pipeline.last_sr, + "fragment_interval": tts_pipeline.last_fragment_interval, + "original_text": tts_pipeline.last_original_text + } + return JSONResponse(status_code=200, content=meta_data) + except Exception as e: + return JSONResponse(status_code=400, content={"message": "get tts meta failed", "Exception": str(e)}) + + if __name__ == "__main__": try: if host == "None": # 在调用时使用 -a None 参数,可以让api监听双栈