Simple streaming return implementation, lower latency for the first sound. (#417)

* Add stream_return switch to get wavs from yield

* Add more_segment_before arg for more segmenting.

more_segment_before is a int, for token_index < more_segment_before, more segmenting will be applied.
0: no effect; 80 is recommended for better first-wav-latency

* Uncomment silence insertion

* fix: rename quick streaming tokens argument

* fix: rename quick streaming tokens argument

* fix: Add a wrapper for the yield function. It will not return a generator in normal condition.
This commit is contained in:
Yt Zhong 2025-09-30 14:05:39 +08:00 committed by GitHub
parent 2ca41d738f
commit b0c6ab8a93
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 72 additions and 18 deletions

View File

@ -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

View File

@ -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
)