Add Python API examples for PocketTTS - #3088
Conversation
Summary of ChangesHello @csukuangfj, I'm Gemini Code Assist1! I'm currently reviewing this pull request and will post my feedback shortly. In the meantime, here's a summary to help you and other reviewers quickly get up to speed! This pull request introduces Python API examples for the PocketTTS module, allowing users to generate speech from text. It provides two distinct examples: one for basic audio generation and saving, and another for real-time audio playback as the speech is being synthesized. Additionally, the changes include necessary configuration updates and crucial input validation within the C++ backend to ensure stable operation of the PocketTTS functionality. Highlights
Using Gemini Code AssistThe full guide for Gemini Code Assist can be found on our documentation page, here are some quick tips. Invoking Gemini You can request assistance from Gemini at any point by creating a comment using either
Customization To customize Gemini Code Assist for GitHub experience, repository maintainers can create a configuration file and/or provide a custom code review style guide (such as PEP-8 for Python) by creating and adding files to a Limitations & Feedback Gemini Code Assist may make mistakes. Please leave feedback on any instances where its feedback is incorrect or counter productive. You can react with 👍 and 👎 on @gemini-code-assist comments. If you're interested in giving your feedback about your experience with Gemini Code Assist for Github and other Google products, sign up here. You can also get AI-powered code generation, chat, as well as code reviews directly in the IDE at no cost with the Gemini Code Assist IDE Extension. Footnotes
|
|
Caution Review failedThe pull request is closed. 📝 WalkthroughWalkthroughThis PR introduces PocketTTS voice cloning support for sherpa-onnx by adding example scripts demonstrating TTS usage, enhancing the C++ implementation with voice embedding validation, and expanding Python API exposure to include GenerationConfig and OfflineTtsPocketModelConfig types. Changes
Sequence DiagramsequenceDiagram
participant Main as Main Thread
participant Gen as TTS Generator<br/>(Callback)
participant Buffer as Audio Buffer
participant Stream as Audio Stream<br/>(Playback)
Main->>Gen: Start generation with callback
activate Gen
Gen->>Buffer: generated_audio_callback(samples)
activate Buffer
Buffer-->>Buffer: Append samples to buffer
Buffer-->>Main: Signal playback ready (started=True)
deactivate Buffer
Main->>Stream: Start playback thread
activate Stream
Stream->>Buffer: Request audio frames (play_audio_callback)
Buffer-->>Stream: Return buffered samples
Stream-->>Stream: Feed to OutputStream
Gen->>Buffer: Continue: generated_audio_callback(samples)
Buffer-->>Stream: Frames available
Stream-->>Stream: Continue playback
Gen->>Buffer: Final: generated_audio_callback(last samples)
deactivate Gen
Buffer-->>Stream: End of buffer signal
Stream-->>Stream: Fill silence/close stream
deactivate Stream
Main->>Main: Wait for playback completion
Main->>Main: Write generated.wav + log metrics
Estimated Code Review Effort🎯 3 (Moderate) | ⏱️ ~25 minutes Poem
✨ Finishing touches
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
There was a problem hiding this comment.
Code Review
This pull request adds Python API examples for PocketTTS, including one for simple generation and another for streaming playback. The changes also include necessary C++ backend updates to support this, such as adding input validation for reference audio and exposing the GenerationConfig to Python.
My review focuses on the new Python example for streaming playback. I've identified a critical race condition in the audio playback callback that could lead to application crashes. I've provided a detailed explanation and a code suggestion to fix this by using a thread-safe approach. The rest of the changes look good and are well-implemented.
| def play_audio_callback( | ||
| outdata: np.ndarray, frames: int, time, status: sd.CallbackFlags | ||
| ): | ||
| if killed or (started and buffer.empty() and stopped): | ||
| event.set() | ||
|
|
||
| # outdata is of shape (frames, num_channels) | ||
| if buffer.empty(): | ||
| outdata.fill(0) | ||
| return | ||
|
|
||
| n = 0 | ||
| while n < frames and not buffer.empty(): | ||
| remaining = frames - n | ||
| k = buffer.queue[0].shape[0] | ||
|
|
||
| if remaining <= k: | ||
| outdata[n:, 0] = buffer.queue[0][:remaining] | ||
| buffer.queue[0] = buffer.queue[0][remaining:] | ||
| n = frames | ||
| if buffer.queue[0].shape[0] == 0: | ||
| buffer.get() | ||
|
|
||
| break | ||
|
|
||
| outdata[n : n + k, 0] = buffer.get() | ||
| n += k | ||
|
|
||
| if n < frames: | ||
| outdata[n:, 0] = 0 | ||
|
|
There was a problem hiding this comment.
Accessing buffer.queue directly is not thread-safe. The queue.Queue is designed for thread-safe communication via its public methods (put, get, etc.), but its underlying queue attribute (a collections.deque) is not protected for direct access. This can lead to race conditions between the audio generation thread (calling put) and the audio playback thread (this callback), potentially causing an IndexError or other unpredictable behavior if the queue is modified by another thread concurrently.
A safer approach is to only use the public API of queue.Queue and manage any leftover audio samples from a chunk within the callback's state. This can be achieved by attaching a state variable to the callback function itself.
def play_audio_callback(
outdata: np.ndarray, frames: int, time, status: sd.CallbackFlags
):
if not hasattr(play_audio_callback, "leftover"):
play_audio_callback.leftover = np.array([], dtype=np.float32)
if killed or (
started
and buffer.empty()
and stopped
and not play_audio_callback.leftover.any()
):
event.set()
outdata.fill(0)
return
outdata.fill(0)
num_copied = 0
if play_audio_callback.leftover.any():
n = min(frames, play_audio_callback.leftover.shape[0])
outdata[:n, 0] = play_audio_callback.leftover[:n]
play_audio_callback.leftover = play_audio_callback.leftover[n:]
num_copied = n
while num_copied < frames:
try:
samples = buffer.get_nowait()
except queue.Empty:
break
needed = frames - num_copied
k = min(needed, samples.shape[0])
outdata[num_copied : num_copied + k, 0] = samples[:k]
num_copied += k
if k < samples.shape[0]:
play_audio_callback.leftover = samples[k:]
break # outdata is full
Summary by CodeRabbit
New Features
Bug Fixes
✏️ Tip: You can customize this high-level summary in your review settings.