|
2 | 2 | # Make adjustments inside functions, and consider both gradio and cli scripts if need to change func output format |
3 | 3 | import os |
4 | 4 | import sys |
| 5 | +from concurrent.futures import ThreadPoolExecutor |
5 | 6 |
|
6 | 7 |
|
7 | 8 | os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1" # for MPS device compatibility |
@@ -86,6 +87,8 @@ def chunk_text(text, max_chars=135): |
86 | 87 | sentences = re.split(r"(?<=[;:,.!?])\s+|(?<=[;:,。!?])", text) |
87 | 88 |
|
88 | 89 | for sentence in sentences: |
| 90 | + if not sentence: |
| 91 | + continue |
89 | 92 | if len(current_chunk.encode("utf-8")) + len(sentence.encode("utf-8")) <= max_chars: |
90 | 93 | current_chunk += sentence + " " if sentence and len(sentence[-1].encode("utf-8")) == 1 else sentence |
91 | 94 | else: |
@@ -279,12 +282,12 @@ def remove_silence_edges(audio, silence_threshold=-42): |
279 | 282 | audio = audio[non_silent_start_idx:] |
280 | 283 |
|
281 | 284 | # 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 |
288 | 291 |
|
289 | 292 | return trimmed_audio |
290 | 293 |
|
@@ -400,11 +403,16 @@ def infer_process( |
400 | 403 | audio, sr = torchaudio.load(ref_audio) |
401 | 404 | max_chars = int(len(ref_text.encode("utf-8")) / (audio.shape[-1] / sr) * (22 - audio.shape[-1] / sr) * speed) |
402 | 405 | 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) |
405 | 408 | print("\n") |
406 | 409 |
|
407 | 410 | 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 | + |
408 | 416 | return next( |
409 | 417 | infer_batch_process( |
410 | 418 | (audio, sr), |
@@ -466,7 +474,7 @@ def infer_batch_process( |
466 | 474 | if len(ref_text[-1].encode("utf-8")) == 1: |
467 | 475 | ref_text = ref_text + " " |
468 | 476 |
|
469 | | - def process_batch(gen_text): |
| 477 | + def _infer_basic(gen_text): |
470 | 478 | local_speed = speed |
471 | 479 | if len(gen_text.encode("utf-8")) < 10: |
472 | 480 | local_speed = 0.3 |
@@ -509,23 +517,34 @@ def process_batch(gen_text): |
509 | 517 | # wav -> numpy |
510 | 518 | generated_wave = generated_wave.squeeze().cpu().numpy() |
511 | 519 |
|
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 |
519 | 534 |
|
520 | 535 | if streaming: |
521 | 536 | 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): |
523 | 538 | yield chunk |
524 | 539 | 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) |
529 | 548 |
|
530 | 549 | if generated_waves: |
531 | 550 | if cross_fade_duration <= 0: |
|
0 commit comments