Merge pull request #117 from daily-co/llm-use-aggregator-pass-through-fix

aggregators: fix LLMUserResponseAggregator passs-through
This commit is contained in:
Aleix Conchillo Flaqué
2024-04-12 04:24:56 +08:00
committed by GitHub
14 changed files with 91 additions and 50 deletions

View File

@@ -1,5 +1,6 @@
import asyncio
import re
import time
from dailyai.pipeline.frame_processor import FrameProcessor
@@ -8,6 +9,7 @@ from dailyai.pipeline.frames import (
EndPipeFrame,
Frame,
ImageFrame,
InterimTranscriptionFrame,
LLMMessagesFrame,
LLMResponseEndFrame,
LLMResponseStartFrame,
@@ -106,6 +108,7 @@ class LLMResponseAggregator(FrameProcessor):
start_frame,
end_frame,
accumulator_frame,
interim_accumulator_frame=None,
pass_through=True,
):
self.aggregation = ""
@@ -115,31 +118,75 @@ class LLMResponseAggregator(FrameProcessor):
self._start_frame = start_frame
self._end_frame = end_frame
self._accumulator_frame = accumulator_frame
self._interim_accumulator_frame = interim_accumulator_frame
self._pass_through = pass_through
self._seen_start_frame = False
self._seen_end_frame = False
self._seen_interim_results = False
# Use cases implemented:
#
# S: Start, E: End, T: Transcription, I: Interim, X: Text
#
# S E -> None
# S T E -> X
# S I T E -> X
# S I E T -> X
# S I E I T -> X
#
# The following case would not be supported:
#
# S I E T1 I T2 -> X
#
# and T2 would be dropped.
async def process_frame(self, frame: Frame) -> AsyncGenerator[Frame, None]:
if not self.messages:
return
send_aggregation = False
if isinstance(frame, self._start_frame):
self._seen_start_frame = True
self.aggregating = True
elif isinstance(frame, self._end_frame):
self.aggregating = False
# Sometimes VAD triggers quickly on and off. If we don't get any transcription,
# it creates empty LLM message queue frames
if len(self.aggregation) > 0:
self.messages.append(
{"role": self._role, "content": self.aggregation})
self.aggregation = ""
yield self._end_frame()
yield LLMMessagesFrame(self.messages)
elif isinstance(frame, self._accumulator_frame) and self.aggregating:
self.aggregation += f" {frame.text}"
self._seen_end_frame = True
# We might have received the end frame but we might still be
# aggregating (i.e. we have seen interim results but not the final
# text).
self.aggregating = self._seen_interim_results
# Send the aggregation if we are not aggregating anymore (i.e. no
# more interim results received).
send_aggregation = not self.aggregating
elif isinstance(frame, self._accumulator_frame):
if self.aggregating:
self.aggregation += f" {frame.text}"
# We have receied a complete sentence, so if we have seen the
# end frame and we were still aggregating, it means we should
# send the aggregation.
send_aggregation = self._seen_end_frame
if self._pass_through:
yield frame
# We just got our final result, so let's reset interim results.
self._seen_interim_results = False
elif self._interim_accumulator_frame and isinstance(frame, self._interim_accumulator_frame):
self._seen_interim_results = True
else:
yield frame
if send_aggregation and len(self.aggregation) > 0:
self.messages.append({"role": self._role, "content": self.aggregation})
yield self._end_frame()
yield LLMMessagesFrame(self.messages)
# Reset
self.aggregation = ""
self._seen_start_frame = False
self._seen_end_frame = False
self._seen_interim_results = False
class LLMAssistantResponseAggregator(LLMResponseAggregator):
def __init__(self, messages: list[dict]):
@@ -160,6 +207,7 @@ class LLMUserResponseAggregator(LLMResponseAggregator):
start_frame=UserStartedSpeakingFrame,
end_frame=UserStoppedSpeakingFrame,
accumulator_frame=TranscriptionFrame,
interim_accumulator_frame=InterimTranscriptionFrame,
pass_through=False,
)

View File

@@ -164,6 +164,17 @@ class TranscriptionFrame(TextFrame):
return f"{self.__class__.__name__}, text: '{self.text}' participantId: {self.participantId}, timestamp: {self.timestamp}"
@dataclass()
class InterimTranscriptionFrame(TextFrame):
"""A text frame with interim transcription-specific data. Will be placed in
the transport's receive queue when a participant speaks."""
participantId: str
timestamp: str
def __str__(self):
return f"{self.__class__.__name__}, text: '{self.text}' participantId: {self.participantId}, timestamp: {self.timestamp}"
class TTSStartFrame(ControlFrame):
"""Used to indicate the beginning of a TTS response. Following AudioFrames
are part of the TTS response until an TTEndFrame. These frames can be used

View File

@@ -10,6 +10,7 @@ from functools import partial
from typing import Any
from dailyai.pipeline.frames import (
InterimTranscriptionFrame,
ReceivedAppMessageFrame,
TranscriptionFrame,
UserImageFrame,
@@ -88,9 +89,11 @@ class DailyTransport(ThreadedTransport, EventHandler):
"model": "2-conversationalai",
"profanity_filter": True,
"redact": False,
"endpointing": True,
"punctuate": True,
"includeRawResponse": True,
"extra": {
"endpointing": True,
"punctuate": False,
"interim_results": True,
},
}
@@ -368,8 +371,12 @@ class DailyTransport(ThreadedTransport, EventHandler):
elif "session_id" in message:
participantId = message["session_id"]
if self._my_participant_id and participantId != self._my_participant_id:
frame = TranscriptionFrame(
message["text"], participantId, message["timestamp"])
is_final = message["rawResponse"]["is_final"]
if is_final:
frame = TranscriptionFrame(message["text"], participantId, message["timestamp"])
else:
frame = InterimTranscriptionFrame(
message["text"], participantId, message["timestamp"])
asyncio.run_coroutine_threadsafe(
self.receive_queue.put(frame), self._loop)