diff --git a/indextts/infer_v2.py b/indextts/infer_v2.py index cefc875..acd13b6 100644 --- a/indextts/infer_v2.py +++ b/indextts/infer_v2.py @@ -265,6 +265,20 @@ class IndexTTS2: code_lens = torch.tensor(code_lens, dtype=torch.long, device=device) return codes, code_lens + def interval_silence(self, wavs, sampling_rate=22050, interval_silence=200): + """ + Silences to be insert between generated segments. + """ + + if not wavs or interval_silence <= 0: + return wavs + + # get channel_size + channel_size = wavs[0].size(0) + # get silence tensor + sil_dur = int(sampling_rate * interval_silence / 1000.0) + return torch.zeros(channel_size, sil_dur) + def insert_interval_silence(self, wavs, sampling_rate=22050, interval_silence=200): """ Insert silences between generated segments. @@ -327,7 +341,29 @@ class IndexTTS2: emo_audio_prompt=None, emo_alpha=1.0, emo_vector=None, use_emo_text=False, emo_text=None, use_random=False, interval_silence=200, - verbose=False, max_text_tokens_per_segment=120, **generation_kwargs): + verbose=False, max_text_tokens_per_segment=120, stream_return=False, more_segment_before=0, **generation_kwargs): + if stream_return: + return self.infer_generator( + spk_audio_prompt, text, output_path, + emo_audio_prompt, emo_alpha, + emo_vector, + use_emo_text, emo_text, use_random, interval_silence, + verbose, max_text_tokens_per_segment, stream_return, more_segment_before, **generation_kwargs + ) + else: + return list(self.infer_generator( + spk_audio_prompt, text, output_path, + emo_audio_prompt, emo_alpha, + emo_vector, + use_emo_text, emo_text, use_random, interval_silence, + verbose, max_text_tokens_per_segment, stream_return, more_segment_before, **generation_kwargs + ))[0] + + def infer_generator(self, spk_audio_prompt, text, output_path, + emo_audio_prompt=None, emo_alpha=1.0, + emo_vector=None, + use_emo_text=False, emo_text=None, use_random=False, interval_silence=200, + verbose=False, max_text_tokens_per_segment=120, stream_return=False, quick_streaming_tokens=0, **generation_kwargs): print(">> starting inference...") self._set_gr_progress(0, "starting inference...") if verbose: @@ -445,7 +481,7 @@ class IndexTTS2: self._set_gr_progress(0.1, "text processing...") text_tokens_list = self.tokenizer.tokenize(text) - segments = self.tokenizer.split_segments(text_tokens_list, max_text_tokens_per_segment) + segments = self.tokenizer.split_segments(text_tokens_list, max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens) segments_count = len(segments) text_token_ids = self.tokenizer.convert_tokens_to_ids(text_tokens_list) @@ -476,6 +512,7 @@ class IndexTTS2: s2mel_time = 0 bigvgan_time = 0 has_warned = False + silence = None # for stream_return for seg_idx, sent in enumerate(segments): self._set_gr_progress(0.2 + 0.7 * seg_idx / segments_count, f"speech synthesis {seg_idx + 1}/{segments_count}...") @@ -608,6 +645,11 @@ class IndexTTS2: print(f"wav shape: {wav.shape}", "min:", wav.min(), "max:", wav.max()) # wavs.append(wav[:, :-512]) wavs.append(wav.cpu()) # to cpu before saving + if stream_return: + yield wav.cpu() + if silence == None: + silence = self.interval_silence(wavs, sampling_rate=sampling_rate, interval_silence=interval_silence) + yield silence end_time = time.perf_counter() self._set_gr_progress(0.9, "saving audio...") @@ -633,8 +675,12 @@ class IndexTTS2: os.makedirs(os.path.dirname(output_path), exist_ok=True) torchaudio.save(output_path, wav.type(torch.int16), sampling_rate) print(">> wav file saved to:", output_path) + if stream_return: + return None return output_path else: + if stream_return: + return None # 返回以符合Gradio的格式要求 wav_data = wav.type(torch.int16) wav_data = wav_data.numpy().T diff --git a/indextts/utils/front.py b/indextts/utils/front.py index 7f69fac..5864073 100644 --- a/indextts/utils/front.py +++ b/indextts/utils/front.py @@ -343,7 +343,10 @@ class TextTokenizer: @staticmethod def split_segments_by_token( - tokenized_str: List[str], split_tokens: List[str], max_text_tokens_per_segment: int + tokenized_str: List[str], + split_tokens: List[str], + max_text_tokens_per_segment: int, + quick_streaming_tokens: int = 0 ) -> List[List[str]]: """ 将tokenize后的结果按特定token进一步分割 @@ -358,7 +361,17 @@ class TextTokenizer: token = tokenized_str[i] current_segment.append(token) current_segment_tokens_len += 1 - if current_segment_tokens_len <= max_text_tokens_per_segment: + if not ("," in split_tokens or "▁," in split_tokens ) and ("," in current_segment or "▁," in current_segment): + # 如果当前tokens中有,,则按,分割 + sub_segments = TextTokenizer.split_segments_by_token( + current_segment, [",", "▁,"], max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens + ) + elif "-" not in split_tokens and "-" in current_segment: + # 没有,,则按-分割 + sub_segments = TextTokenizer.split_segments_by_token( + current_segment, ["-"], max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens + ) + elif current_segment_tokens_len <= max_text_tokens_per_segment: if token in split_tokens and current_segment_tokens_len > 2: if i < len(tokenized_str) - 1: if tokenized_str[i + 1] in ["'", "▁'"]: @@ -370,16 +383,6 @@ class TextTokenizer: current_segment_tokens_len = 0 continue # 如果当前tokens的长度超过最大限制 - if not ("," in split_tokens or "▁," in split_tokens ) and ("," in current_segment or "▁," in current_segment): - # 如果当前tokens中有,,则按,分割 - sub_segments = TextTokenizer.split_segments_by_token( - current_segment, [",", "▁,"], max_text_tokens_per_segment=max_text_tokens_per_segment - ) - elif "-" not in split_tokens and "-" in current_segment: - # 没有,,则按-分割 - sub_segments = TextTokenizer.split_segments_by_token( - current_segment, ["-"], max_text_tokens_per_segment=max_text_tokens_per_segment - ) else: # 按照长度分割 sub_segments = [] @@ -400,14 +403,19 @@ class TextTokenizer: if current_segment_tokens_len > 0: assert current_segment_tokens_len <= max_text_tokens_per_segment segments.append(current_segment) - # 如果相邻的句子加起来长度小于最大限制,则合并 + # 如果相邻的句子加起来长度小于最大限制,且此前token总数超过quick_streaming_tokens,则合并 merged_segments = [] + total_token = 0 for segment in segments: + total_token += len(segment) if len(segment) == 0: continue if len(merged_segments) == 0: merged_segments.append(segment) - elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment: + elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment and total_token > quick_streaming_tokens: + merged_segments[-1] = merged_segments[-1] + segment + # 或小于最大长度限制的一半,则合并 + elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment / 2: merged_segments[-1] = merged_segments[-1] + segment else: merged_segments.append(segment) @@ -422,9 +430,9 @@ class TextTokenizer: "▁?", "▁...", # ellipsis ] - def split_segments(self, tokenized: List[str], max_text_tokens_per_segment=120) -> List[List[str]]: + def split_segments(self, tokenized: List[str], max_text_tokens_per_segment=120, quick_streaming_tokens = 0) -> List[List[str]]: return TextTokenizer.split_segments_by_token( - tokenized, self.punctuation_marks_tokens, max_text_tokens_per_segment=max_text_tokens_per_segment + tokenized, self.punctuation_marks_tokens, max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens )