[流式回复播放]:完成LLM到TTS流式播报,包含句子切分和可取消播放队列
This commit is contained in:
@@ -54,6 +54,8 @@ from .full_duplex_speech import (
|
||||
VadProvider,
|
||||
)
|
||||
from .full_duplex_response import (
|
||||
AudioSegmentStreamingTtsProvider,
|
||||
AudioSegmentStreamingTtsSession,
|
||||
FakeStreamingLlmProvider,
|
||||
FakeStreamingTtsProvider,
|
||||
InterruptiblePlaybackQueue,
|
||||
@@ -138,6 +140,8 @@ __all__ = [
|
||||
"webrtc_apm_probe",
|
||||
"FullDuplexAgentRuntime",
|
||||
"FullDuplexRuntimeHealth",
|
||||
"AudioSegmentStreamingTtsProvider",
|
||||
"AudioSegmentStreamingTtsSession",
|
||||
"FakeStreamingSttProvider",
|
||||
"FakeVadProvider",
|
||||
"InterruptController",
|
||||
|
||||
@@ -4,9 +4,9 @@ from collections import deque
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Iterable, Literal, Protocol
|
||||
|
||||
from .full_duplex_audio import RenderReferenceRingBuffer
|
||||
from .full_duplex_audio import AudioHub, RenderReferenceRingBuffer
|
||||
from .full_duplex_control import CancellationToken
|
||||
from .models import AudioFrame, Message, ProviderError
|
||||
from .models import AudioFrame, AudioSegment, Message, ProviderError
|
||||
from .tts import sanitize_tts_text
|
||||
|
||||
|
||||
@@ -162,6 +162,94 @@ class FakeStreamingTtsSession:
|
||||
self.cancelled = True
|
||||
|
||||
|
||||
class AudioSegmentStreamingTtsProvider:
|
||||
name = "audio_segment_streaming_tts"
|
||||
|
||||
def __init__(self, tts_provider, *, chunk_ms: int = 30) -> None:
|
||||
if chunk_ms <= 0:
|
||||
raise ValueError("chunk_ms must be positive")
|
||||
self.tts_provider = tts_provider
|
||||
self.chunk_ms = chunk_ms
|
||||
self.started: list[tuple[str, int]] = []
|
||||
|
||||
def start_stream(self, *, voice: str, sample_rate: int) -> "AudioSegmentStreamingTtsSession":
|
||||
self.started.append((voice, sample_rate))
|
||||
return AudioSegmentStreamingTtsSession(
|
||||
tts_provider=self.tts_provider,
|
||||
sample_rate=sample_rate,
|
||||
chunk_ms=self.chunk_ms,
|
||||
)
|
||||
|
||||
|
||||
class AudioSegmentStreamingTtsSession:
|
||||
def __init__(self, *, tts_provider, sample_rate: int, chunk_ms: int) -> None:
|
||||
self.tts_provider = tts_provider
|
||||
self.sample_rate = sample_rate
|
||||
self.chunk_ms = chunk_ms
|
||||
self.cancelled = False
|
||||
self.cancel_reason = ""
|
||||
self.accepted_text: list[str] = []
|
||||
self._next_frame_id = 0
|
||||
if hasattr(self.tts_provider, "load"):
|
||||
self.tts_provider.load()
|
||||
|
||||
def accept_text(self, text: str) -> list[AudioFrame]:
|
||||
if self.cancelled:
|
||||
return []
|
||||
spoken = prepare_tts_sentence(text)
|
||||
if not spoken:
|
||||
return []
|
||||
segment = self.tts_provider.synthesize(spoken)
|
||||
self.accepted_text.append(spoken)
|
||||
return self._segment_to_frames(segment, spoken)
|
||||
|
||||
def flush(self) -> list[AudioFrame]:
|
||||
return []
|
||||
|
||||
def cancel(self, reason: str) -> None:
|
||||
self.cancelled = True
|
||||
self.cancel_reason = reason
|
||||
|
||||
def _segment_to_frames(self, segment: AudioSegment, text: str) -> list[AudioFrame]:
|
||||
metadata = dict(segment.metadata)
|
||||
if metadata.get("format"):
|
||||
self._next_frame_id += 1
|
||||
metadata.update({"duration_ms": segment.duration_ms, "tts_text": text})
|
||||
return [
|
||||
AudioFrame(
|
||||
segment.pcm,
|
||||
segment.sample_rate,
|
||||
segment.channels,
|
||||
segment.start_time_ms,
|
||||
self._next_frame_id,
|
||||
metadata,
|
||||
)
|
||||
]
|
||||
bytes_per_ms = max(1, int(segment.sample_rate * segment.channels * 2 / 1000))
|
||||
chunk_bytes = max(2 * segment.channels, bytes_per_ms * self.chunk_ms)
|
||||
chunk_bytes -= chunk_bytes % (2 * segment.channels)
|
||||
frames: list[AudioFrame] = []
|
||||
offset = 0
|
||||
timestamp_ms = segment.start_time_ms
|
||||
while offset < len(segment.pcm):
|
||||
data = segment.pcm[offset : offset + chunk_bytes]
|
||||
duration_ms = max(1, int(len(data) / bytes_per_ms))
|
||||
self._next_frame_id += 1
|
||||
frames.append(
|
||||
AudioFrame(
|
||||
data,
|
||||
segment.sample_rate,
|
||||
segment.channels,
|
||||
timestamp_ms,
|
||||
self._next_frame_id,
|
||||
{"duration_ms": duration_ms, "tts_text": text},
|
||||
)
|
||||
)
|
||||
offset += chunk_bytes
|
||||
timestamp_ms += duration_ms
|
||||
return frames
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class PlaybackItem:
|
||||
text: str
|
||||
@@ -208,9 +296,12 @@ class InterruptiblePlaybackQueue:
|
||||
def play_next(
|
||||
self,
|
||||
*,
|
||||
render_reference: RenderReferenceRingBuffer,
|
||||
render_reference: RenderReferenceRingBuffer | None = None,
|
||||
audio_hub: AudioHub | None = None,
|
||||
cancellation: CancellationToken,
|
||||
) -> PlaybackChunkResult:
|
||||
if render_reference is None and audio_hub is None:
|
||||
raise ValueError("render_reference or audio_hub is required")
|
||||
if not self._items:
|
||||
return PlaybackChunkResult(0, interrupted=False)
|
||||
item = self._items.popleft()
|
||||
@@ -219,7 +310,10 @@ class InterruptiblePlaybackQueue:
|
||||
if cancellation.cancelled:
|
||||
self.clear_unplayed()
|
||||
return PlaybackChunkResult(played, interrupted=True)
|
||||
render_reference.write(frame)
|
||||
if audio_hub is not None:
|
||||
audio_hub.accept_render(frame)
|
||||
elif render_reference is not None:
|
||||
render_reference.write(frame)
|
||||
played += 1
|
||||
self.spoken.commit(item.text)
|
||||
return PlaybackChunkResult(played, interrupted=False, committed_text=item.text)
|
||||
|
||||
@@ -5,8 +5,17 @@ from dataclasses import dataclass
|
||||
from .config import AppConfig
|
||||
from .full_duplex_audio import AudioHub, AudioProcessingProvider, build_audio_processing_provider
|
||||
from .full_duplex_control import CancellationGraph, FullDuplexStateMachine
|
||||
from .full_duplex_response import (
|
||||
FakeStreamingLlmProvider,
|
||||
FakeStreamingTtsProvider,
|
||||
InterruptiblePlaybackQueue,
|
||||
LlmStreamEvent,
|
||||
SentenceSegmenter,
|
||||
StreamingLlmProvider,
|
||||
StreamingTtsProvider,
|
||||
)
|
||||
from .full_duplex_speech import FakeVadProvider, InterruptController, InterruptionDetector
|
||||
from .models import AudioFrame, PipelineState
|
||||
from .models import AudioFrame, Message, PipelineState
|
||||
from .runtime import RuntimeSummary
|
||||
|
||||
|
||||
@@ -32,14 +41,19 @@ class FullDuplexAgentRuntime:
|
||||
config: AppConfig,
|
||||
processor: AudioProcessingProvider | None = None,
|
||||
audio_hub: AudioHub | None = None,
|
||||
llm_provider: StreamingLlmProvider | None = None,
|
||||
tts_provider: StreamingTtsProvider | None = None,
|
||||
) -> None:
|
||||
self.config = config
|
||||
self.processor = processor
|
||||
self.audio_hub = audio_hub
|
||||
self.llm_provider = llm_provider
|
||||
self.tts_provider = tts_provider
|
||||
self.health: FullDuplexRuntimeHealth | None = None
|
||||
self.state_machine = FullDuplexStateMachine()
|
||||
self.cancellation_graph = CancellationGraph("turn")
|
||||
self.interrupt_controller: InterruptController | None = None
|
||||
self.playback_queue = InterruptiblePlaybackQueue()
|
||||
|
||||
def load_audio(self) -> FullDuplexRuntimeHealth:
|
||||
if self.audio_hub is None:
|
||||
@@ -116,6 +130,45 @@ class FullDuplexAgentRuntime:
|
||||
return RuntimeSummary(completed_turns=0, failed_turns=0, interrupted=True)
|
||||
return RuntimeSummary(completed_turns=1, failed_turns=0)
|
||||
|
||||
def run_streaming_response_fixture(
|
||||
self,
|
||||
messages: list[Message],
|
||||
*,
|
||||
llm_events: list[LlmStreamEvent] | None = None,
|
||||
) -> str:
|
||||
self.load_audio()
|
||||
if self.audio_hub is None:
|
||||
raise RuntimeError("audio hub is not loaded")
|
||||
self.cancellation_graph = CancellationGraph("turn")
|
||||
llm = self.llm_provider or FakeStreamingLlmProvider(
|
||||
llm_events or [LlmStreamEvent("delta", "你好。"), LlmStreamEvent("finish", finish_reason="stop")]
|
||||
)
|
||||
tts = self.tts_provider or FakeStreamingTtsProvider()
|
||||
tts_session = tts.start_stream(voice=self.config.tts_voice, sample_rate=self.config.sample_rate)
|
||||
segmenter = SentenceSegmenter()
|
||||
self.playback_queue = InterruptiblePlaybackQueue()
|
||||
|
||||
for event in llm.stream(messages, cancellation=self.cancellation_graph.root):
|
||||
if event.kind == "delta" and event.text_delta:
|
||||
for sentence in segmenter.accept_delta(event.text_delta):
|
||||
self._synthesize_and_play_sentence(sentence, tts_session)
|
||||
elif event.kind == "finish":
|
||||
break
|
||||
tail = segmenter.flush()
|
||||
if tail:
|
||||
self._synthesize_and_play_sentence(tail, tts_session)
|
||||
for frame in tts_session.flush():
|
||||
self.playback_queue.enqueue(str(frame.metadata.get("tts_text", "")), [frame])
|
||||
self.playback_queue.play_next(audio_hub=self.audio_hub, cancellation=self.cancellation_graph.root)
|
||||
return self.playback_queue.spoken.text
|
||||
|
||||
def _synthesize_and_play_sentence(self, sentence: str, tts_session) -> None:
|
||||
if self.audio_hub is None:
|
||||
raise RuntimeError("audio hub is not loaded")
|
||||
frames = tts_session.accept_text(sentence)
|
||||
self.playback_queue.enqueue(sentence, frames)
|
||||
self.playback_queue.play_next(audio_hub=self.audio_hub, cancellation=self.cancellation_graph.root)
|
||||
|
||||
def _run_audio_smoke_once(self) -> None:
|
||||
if self.audio_hub is None:
|
||||
raise RuntimeError("audio hub is not loaded")
|
||||
|
||||
Reference in New Issue
Block a user