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:
parent
2ca41d738f
commit
b0c6ab8a93
@ -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
|
||||
|
||||
@ -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
|
||||
)
|
||||
|
||||
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user