Merge pull request #166 from pipecat-ai/fix-llm-response-wake-check
fix llm response wake check
This commit is contained in:
@@ -19,6 +19,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
|
- Fixed an issue in `LLMUserResponseAggregator` and `UserResponseAggregator`
|
||||||
|
that would cause frames after a brief pause to not be pushed to the LLM.
|
||||||
|
|
||||||
- Clear the audio output buffer if we are interrupted.
|
- Clear the audio output buffer if we are interrupted.
|
||||||
|
|
||||||
- Re-add exponential smoothing after volume calculation. This makes sure the
|
- Re-add exponential smoothing after volume calculation. This makes sure the
|
||||||
|
|||||||
@@ -84,11 +84,6 @@ async def main(room_url: str, token):
|
|||||||
transport.capture_participant_transcription(participant["id"])
|
transport.capture_participant_transcription(participant["id"])
|
||||||
await tts.say("Hi! If you want to talk to me, just say 'Hey Robot'.")
|
await tts.say("Hi! If you want to talk to me, just say 'Hey Robot'.")
|
||||||
|
|
||||||
# Kick off the conversation.
|
|
||||||
# messages.append(
|
|
||||||
# {"role": "system", "content": "Please introduce yourself to the user."})
|
|
||||||
# await task.queue_frames([LLMMessagesFrame(messages)])
|
|
||||||
|
|
||||||
runner = PipelineRunner()
|
runner = PipelineRunner()
|
||||||
|
|
||||||
await runner.run(task)
|
await runner.run(task)
|
||||||
|
|||||||
@@ -119,7 +119,7 @@ class TextFrame(DataFrame):
|
|||||||
text: str
|
text: str
|
||||||
|
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
return f"{self.name}(text: {self.text})"
|
return f"{self.name}(text: [{self.text}])"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -132,7 +132,7 @@ class TranscriptionFrame(TextFrame):
|
|||||||
timestamp: str
|
timestamp: str
|
||||||
|
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
return f"{self.name}(user_id: {self.user_id}, text: {self.text}, timestamp: {self.timestamp})"
|
return f"{self.name}(user_id: {self.user_id}, text: [{self.text}], timestamp: {self.timestamp})"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -143,7 +143,7 @@ class InterimTranscriptionFrame(TextFrame):
|
|||||||
timestamp: str
|
timestamp: str
|
||||||
|
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
return f"{self.name}(user: {self.user_id}, text: {self.text}, timestamp: {self.timestamp})"
|
return f"{self.name}(user: {self.user_id}, text: [{self.text}], timestamp: {self.timestamp})"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -78,10 +78,14 @@ class LLMResponseAggregator(FrameProcessor):
|
|||||||
send_aggregation = False
|
send_aggregation = False
|
||||||
|
|
||||||
if isinstance(frame, self._start_frame):
|
if isinstance(frame, self._start_frame):
|
||||||
self._seen_start_frame = True
|
self._aggregation = ""
|
||||||
self._aggregating = True
|
self._aggregating = True
|
||||||
|
self._seen_start_frame = True
|
||||||
|
self._seen_end_frame = False
|
||||||
|
self._seen_interim_results = False
|
||||||
elif isinstance(frame, self._end_frame):
|
elif isinstance(frame, self._end_frame):
|
||||||
self._seen_end_frame = True
|
self._seen_end_frame = True
|
||||||
|
self._seen_start_frame = False
|
||||||
|
|
||||||
# We might have received the end frame but we might still be
|
# We might have received the end frame but we might still be
|
||||||
# aggregating (i.e. we have seen interim results but not the final
|
# aggregating (i.e. we have seen interim results but not the final
|
||||||
@@ -118,10 +122,9 @@ class LLMResponseAggregator(FrameProcessor):
|
|||||||
if len(self._aggregation) > 0:
|
if len(self._aggregation) > 0:
|
||||||
self._messages.append({"role": self._role, "content": self._aggregation})
|
self._messages.append({"role": self._role, "content": self._aggregation})
|
||||||
|
|
||||||
# Reset our accumulator state. Reset it before pushing it down,
|
# Reset the aggregation. Reset it before pushing it down, otherwise
|
||||||
# otherwise if the tasks gets cancelled we won't be able to clear
|
# if the tasks gets cancelled we won't be able to clear things up.
|
||||||
# things up.
|
self._aggregation = ""
|
||||||
self._reset()
|
|
||||||
|
|
||||||
frame = LLMMessagesFrame(self._messages)
|
frame = LLMMessagesFrame(self._messages)
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|||||||
@@ -85,10 +85,13 @@ class ResponseAggregator(FrameProcessor):
|
|||||||
send_aggregation = False
|
send_aggregation = False
|
||||||
|
|
||||||
if isinstance(frame, self._start_frame):
|
if isinstance(frame, self._start_frame):
|
||||||
self._seen_start_frame = True
|
|
||||||
self._aggregating = True
|
self._aggregating = True
|
||||||
|
self._seen_start_frame = True
|
||||||
|
self._seen_end_frame = False
|
||||||
|
self._seen_interim_results = False
|
||||||
elif isinstance(frame, self._end_frame):
|
elif isinstance(frame, self._end_frame):
|
||||||
self._seen_end_frame = True
|
self._seen_end_frame = True
|
||||||
|
self._seen_start_frame = False
|
||||||
|
|
||||||
# We might have received the end frame but we might still be
|
# We might have received the end frame but we might still be
|
||||||
# aggregating (i.e. we have seen interim results but not the final
|
# aggregating (i.e. we have seen interim results but not the final
|
||||||
@@ -120,10 +123,9 @@ class ResponseAggregator(FrameProcessor):
|
|||||||
if len(self._aggregation) > 0:
|
if len(self._aggregation) > 0:
|
||||||
frame = TextFrame(self._aggregation.strip())
|
frame = TextFrame(self._aggregation.strip())
|
||||||
|
|
||||||
# Reset our accumulator state. Reset it before pushing it down,
|
# Reset the aggregation. Reset it before pushing it down, otherwise
|
||||||
# otherwise if the tasks gets cancelled we won't be able to clear
|
# if the tasks gets cancelled we won't be able to clear things up.
|
||||||
# things up.
|
self._aggregation = ""
|
||||||
self._reset()
|
|
||||||
|
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ class WakeCheckFilter(FrameProcessor):
|
|||||||
self.wake_timer = 0.0
|
self.wake_timer = 0.0
|
||||||
self.accumulator = ""
|
self.accumulator = ""
|
||||||
|
|
||||||
def __init__(self, wake_phrases: list[str], keepalive_timeout: float = 2):
|
def __init__(self, wake_phrases: list[str], keepalive_timeout: float = 3):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self._participant_states = {}
|
self._participant_states = {}
|
||||||
self._keepalive_timeout = keepalive_timeout
|
self._keepalive_timeout = keepalive_timeout
|
||||||
@@ -55,7 +55,7 @@ class WakeCheckFilter(FrameProcessor):
|
|||||||
if p.state == WakeCheckFilter.WakeState.AWAKE:
|
if p.state == WakeCheckFilter.WakeState.AWAKE:
|
||||||
if time.time() - p.wake_timer < self._keepalive_timeout:
|
if time.time() - p.wake_timer < self._keepalive_timeout:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Wake phrase keepalive timeout has not expired. Passing frame through.")
|
f"Wake phrase keepalive timeout has not expired. Pushing {frame}")
|
||||||
p.wake_timer = time.time()
|
p.wake_timer = time.time()
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -66,7 +66,7 @@ class TTSService(AIService):
|
|||||||
else:
|
else:
|
||||||
self._current_sentence += frame.text
|
self._current_sentence += frame.text
|
||||||
if self._current_sentence.strip().endswith((".", "?", "!")):
|
if self._current_sentence.strip().endswith((".", "?", "!")):
|
||||||
text = self._current_sentence
|
text = self._current_sentence.strip()
|
||||||
self._current_sentence = ""
|
self._current_sentence = ""
|
||||||
|
|
||||||
if text:
|
if text:
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ class ElevenLabsTTSService(TTSService):
|
|||||||
self._model = model
|
self._model = model
|
||||||
|
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
logger.debug(f"Transcribing text: {text}")
|
logger.debug(f"Transcribing text: [{text}]")
|
||||||
|
|
||||||
url = f"https://api.elevenlabs.io/v1/text-to-speech/{self._voice_id}/stream"
|
url = f"https://api.elevenlabs.io/v1/text-to-speech/{self._voice_id}/stream"
|
||||||
|
|
||||||
|
|||||||
@@ -137,10 +137,12 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
if self._allow_interruptions:
|
if self._allow_interruptions:
|
||||||
# Make sure we notify about interruptions quickly out-of-band
|
# Make sure we notify about interruptions quickly out-of-band
|
||||||
if isinstance(frame, UserStartedSpeakingFrame):
|
if isinstance(frame, UserStartedSpeakingFrame):
|
||||||
|
logger.debug("User started speaking")
|
||||||
self._push_frame_task.cancel()
|
self._push_frame_task.cancel()
|
||||||
self._create_push_task()
|
self._create_push_task()
|
||||||
await self.push_frame(StartInterruptionFrame())
|
await self.push_frame(StartInterruptionFrame())
|
||||||
elif isinstance(frame, UserStoppedSpeakingFrame):
|
elif isinstance(frame, UserStoppedSpeakingFrame):
|
||||||
|
logger.debug("User stopped speaking")
|
||||||
await self.push_frame(StopInterruptionFrame())
|
await self.push_frame(StopInterruptionFrame())
|
||||||
await self._internal_push_frame(frame)
|
await self._internal_push_frame(frame)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user