Skip to content

Commit 1a63dda

Browse files
committed
Several fixes for utils_infer.py; separate streaming and non-streaming func and add back parallelism
1 parent 2414e3d commit 1a63dda

1 file changed

Lines changed: 40 additions & 21 deletions

File tree

src/f5_tts/infer/utils_infer.py

Lines changed: 40 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
# Make adjustments inside functions, and consider both gradio and cli scripts if need to change func output format
33
import os
44
import sys
5+
from concurrent.futures import ThreadPoolExecutor
56

67

78
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1" # for MPS device compatibility
@@ -86,6 +87,8 @@ def chunk_text(text, max_chars=135):
8687
sentences = re.split(r"(?<=[;:,.!?])\s+|(?<=[;:,。!?])", text)
8788

8889
for sentence in sentences:
90+
if not sentence:
91+
continue
8992
if len(current_chunk.encode("utf-8")) + len(sentence.encode("utf-8")) <= max_chars:
9093
current_chunk += sentence + " " if sentence and len(sentence[-1].encode("utf-8")) == 1 else sentence
9194
else:
@@ -279,12 +282,12 @@ def remove_silence_edges(audio, silence_threshold=-42):
279282
audio = audio[non_silent_start_idx:]
280283

281284
# Remove silence from the end
282-
non_silent_end_duration = audio.duration_seconds
283-
for ms in reversed(audio):
284-
if ms.dBFS > silence_threshold:
285-
break
286-
non_silent_end_duration -= 0.001
287-
trimmed_audio = audio[: int(non_silent_end_duration * 1000)]
285+
reversed_audio = audio.reverse()
286+
non_silent_end_idx = silence.detect_leading_silence(reversed_audio, silence_threshold=silence_threshold)
287+
if non_silent_end_idx > 0:
288+
trimmed_audio = audio[: len(audio) - non_silent_end_idx]
289+
else:
290+
trimmed_audio = audio
288291

289292
return trimmed_audio
290293

@@ -400,11 +403,16 @@ def infer_process(
400403
audio, sr = torchaudio.load(ref_audio)
401404
max_chars = int(len(ref_text.encode("utf-8")) / (audio.shape[-1] / sr) * (22 - audio.shape[-1] / sr) * speed)
402405
gen_text_batches = chunk_text(gen_text, max_chars=max_chars)
403-
for i, gen_text in enumerate(gen_text_batches):
404-
print(f"gen_text {i}", gen_text)
406+
for i, gen_text_i in enumerate(gen_text_batches):
407+
print(f"gen_text {i}", gen_text_i)
405408
print("\n")
406409

407410
show_info(f"Generating audio in {len(gen_text_batches)} batches...")
411+
412+
if not gen_text_batches:
413+
show_info("No text batches to generate.")
414+
return None, target_sample_rate, None
415+
408416
return next(
409417
infer_batch_process(
410418
(audio, sr),
@@ -466,7 +474,7 @@ def infer_batch_process(
466474
if len(ref_text[-1].encode("utf-8")) == 1:
467475
ref_text = ref_text + " "
468476

469-
def process_batch(gen_text):
477+
def _infer_basic(gen_text):
470478
local_speed = speed
471479
if len(gen_text.encode("utf-8")) < 10:
472480
local_speed = 0.3
@@ -509,23 +517,34 @@ def process_batch(gen_text):
509517
# wav -> numpy
510518
generated_wave = generated_wave.squeeze().cpu().numpy()
511519

512-
if streaming:
513-
for j in range(0, len(generated_wave), chunk_size):
514-
yield generated_wave[j : j + chunk_size], target_sample_rate
515-
else:
516-
generated_cpu = generated[0].cpu().numpy()
517-
del generated
518-
yield generated_wave, generated_cpu
520+
return generated_wave, generated
521+
522+
def infer_single_process(gen_text):
523+
generated_wave, generated = _infer_basic(gen_text)
524+
generated_cpu = generated[0].cpu().numpy()
525+
del generated
526+
return generated_wave, generated_cpu
527+
528+
def infer_single_process_streaming(gen_text):
529+
# for src/f5_tts/socket_server.py
530+
generated_wave, generated = _infer_basic(gen_text)
531+
del generated
532+
for j in range(0, len(generated_wave), chunk_size):
533+
yield generated_wave[j : j + chunk_size], target_sample_rate
519534

520535
if streaming:
521536
for gen_text in progress.tqdm(gen_text_batches) if progress is not None else gen_text_batches:
522-
for chunk in process_batch(gen_text):
537+
for chunk in infer_single_process_streaming(gen_text):
523538
yield chunk
524539
else:
525-
for gen_text in progress.tqdm(gen_text_batches) if progress is not None else gen_text_batches:
526-
generated_wave, generated_mel_spec = next(process_batch(gen_text))
527-
generated_waves.append(generated_wave)
528-
spectrograms.append(generated_mel_spec)
540+
with ThreadPoolExecutor() as executor:
541+
futures = [executor.submit(infer_single_process, gen_text) for gen_text in gen_text_batches]
542+
for future in progress.tqdm(futures) if progress is not None else futures:
543+
result = future.result()
544+
if result:
545+
generated_wave, generated_mel_spec = result
546+
generated_waves.append(generated_wave)
547+
spectrograms.append(generated_mel_spec)
529548

530549
if generated_waves:
531550
if cross_fade_duration <= 0:

0 commit comments

Comments
 (0)