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

@@ -10,8 +10,8 @@ from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService
from dailyai.services.open_ai_services import OpenAILLMService from dailyai.services.open_ai_services import OpenAILLMService
from dailyai.services.ai_services import FrameLogger from dailyai.services.ai_services import FrameLogger
from dailyai.pipeline.aggregators import ( from dailyai.pipeline.aggregators import (
LLMAssistantContextAggregator, LLMAssistantResponseAggregator,
LLMUserContextAggregator, LLMUserResponseAggregator,
) )
from runner import configure from runner import configure
@@ -55,11 +55,9 @@ async def main(room_url: str, token):
}, },
] ]
tma_in = LLMUserContextAggregator( tma_in = LLMUserResponseAggregator(messages)
messages, transport._my_participant_id) tma_out = LLMAssistantResponseAggregator(messages)
tma_out = LLMAssistantContextAggregator(
messages, transport._my_participant_id
)
pipeline = Pipeline( pipeline = Pipeline(
processors=[ processors=[
fl, fl,
@@ -78,8 +76,6 @@ async def main(room_url: str, token):
{"role": "system", "content": "Please introduce yourself to the user."}) {"role": "system", "content": "Please introduce yourself to the user."})
await pipeline.queue_frames([LLMMessagesFrame(messages)]) await pipeline.queue_frames([LLMMessagesFrame(messages)])
transport.transcription_settings["extra"]["endpointing"] = True
transport.transcription_settings["extra"]["punctuate"] = True
await transport.run(pipeline) await transport.run(pipeline)

View File

@@ -47,13 +47,12 @@ async def main(room_url: str, token):
token, token,
"Respond bot", "Respond bot",
5, 5,
camera_enabled=True,
camera_width=1024,
camera_height=1024,
mic_enabled=True,
mic_sample_rate=16000,
) )
transport._camera_enabled = True
transport._camera_width = 1024
transport._camera_height = 1024
transport._mic_enabled = True
transport._mic_sample_rate = 16000
transport.transcription_settings["extra"]["punctuate"] = True
tts = ElevenLabsTTSService( tts = ElevenLabsTTSService(
aiohttp_session=session, aiohttp_session=session,

View File

@@ -67,7 +67,6 @@ async def main(room_url: str, token):
pre_processor=LLMUserResponseAggregator(messages), pre_processor=LLMUserResponseAggregator(messages),
) )
transport.transcription_settings["extra"]["punctuate"] = False
await asyncio.gather(transport.run(), run_conversation()) await asyncio.gather(transport.run(), run_conversation())

View File

@@ -129,12 +129,6 @@ async def main(room_url: str, token):
camera_width=720, camera_width=720,
camera_height=1280, camera_height=1280,
) )
transport._mic_enabled = True
transport._mic_sample_rate = 16000
transport._camera_enabled = True
transport._camera_width = 720
transport._camera_height = 1280
transport.transcription_settings["extra"]["punctuate"] = True
llm = OpenAILLMService( llm = OpenAILLMService(
api_key=os.getenv("OPENAI_API_KEY"), api_key=os.getenv("OPENAI_API_KEY"),

View File

@@ -82,7 +82,6 @@ async def main(room_url: str, token):
mic_sample_rate=16000, mic_sample_rate=16000,
camera_enabled=False, camera_enabled=False,
) )
transport.transcription_settings["extra"]["punctuate"] = True
llm = OpenAILLMService( llm = OpenAILLMService(
api_key=os.getenv("OPENAI_API_KEY"), api_key=os.getenv("OPENAI_API_KEY"),

View File

@@ -77,8 +77,6 @@ async def main(room_url: str, token):
async for audio in audio_generator: async for audio in audio_generator:
transport.output_queue.put(Frame(FrameType.AUDIO_FRAME, audio)) transport.output_queue.put(Frame(FrameType.AUDIO_FRAME, audio))
transport.transcription_settings["extra"]["punctuate"] = False
transport.transcription_settings["extra"]["endpointing"] = False
await asyncio.gather(transport.run(), handle_transcriptions()) await asyncio.gather(transport.run(), handle_transcriptions())

View File

@@ -127,8 +127,6 @@ async def main(room_url: str, token, phone):
transport.start_recording() transport.start_recording()
transport.dialout(phone) transport.dialout(phone)
transport.transcription_settings["extra"]["punctuate"] = True
await asyncio.gather(transport.run(), handle_transcriptions()) await asyncio.gather(transport.run(), handle_transcriptions())

View File

@@ -139,8 +139,6 @@ async def main(room_url: str, token):
pre_processor=LLMUserResponseAggregator(messages), pre_processor=LLMUserResponseAggregator(messages),
) )
transport.transcription_settings["extra"]["endpointing"] = True
transport.transcription_settings["extra"]["punctuate"] = True
await asyncio.gather(transport.run(), run_conversation()) await asyncio.gather(transport.run(), run_conversation())

View File

@@ -340,8 +340,6 @@ async def main(room_url: str, token):
pre_processor=OpenAIUserContextAggregator(context), pre_processor=OpenAIUserContextAggregator(context),
) )
transport.transcription_settings["extra"]["endpointing"] = True
transport.transcription_settings["extra"]["punctuate"] = True
try: try:
await asyncio.gather(transport.run(), handle_intake()) await asyncio.gather(transport.run(), handle_intake())
except (asyncio.CancelledError, KeyboardInterrupt): except (asyncio.CancelledError, KeyboardInterrupt):

View File

@@ -278,8 +278,6 @@ async def main(room_url: str, token):
pipeline, pipeline,
) )
transport.transcription_settings["extra"]["endpointing"] = True
transport.transcription_settings["extra"]["punctuate"] = True
try: try:
await asyncio.gather(transport.run(), storytime()) await asyncio.gather(transport.run(), storytime())
except (asyncio.CancelledError, KeyboardInterrupt): except (asyncio.CancelledError, KeyboardInterrupt):

View File

@@ -99,8 +99,6 @@ async def main(room_url: str, token):
ts = TranslationSubtitles("spanish") ts = TranslationSubtitles("spanish")
pipeline = Pipeline([sa, tp, llm, lfra, ts, tts]) pipeline = Pipeline([sa, tp, llm, lfra, ts, tts])
transport.transcription_settings["extra"]["endpointing"] = True
transport.transcription_settings["extra"]["punctuate"] = True
await transport.run(pipeline) await transport.run(pipeline)

View File

@@ -1,5 +1,6 @@
import asyncio import asyncio
import re import re
import time
from dailyai.pipeline.frame_processor import FrameProcessor from dailyai.pipeline.frame_processor import FrameProcessor
@@ -8,6 +9,7 @@ from dailyai.pipeline.frames import (
EndPipeFrame, EndPipeFrame,
Frame, Frame,
ImageFrame, ImageFrame,
InterimTranscriptionFrame,
LLMMessagesFrame, LLMMessagesFrame,
LLMResponseEndFrame, LLMResponseEndFrame,
LLMResponseStartFrame, LLMResponseStartFrame,
@@ -106,6 +108,7 @@ class LLMResponseAggregator(FrameProcessor):
start_frame, start_frame,
end_frame, end_frame,
accumulator_frame, accumulator_frame,
interim_accumulator_frame=None,
pass_through=True, pass_through=True,
): ):
self.aggregation = "" self.aggregation = ""
@@ -115,31 +118,75 @@ class LLMResponseAggregator(FrameProcessor):
self._start_frame = start_frame self._start_frame = start_frame
self._end_frame = end_frame self._end_frame = end_frame
self._accumulator_frame = accumulator_frame self._accumulator_frame = accumulator_frame
self._interim_accumulator_frame = interim_accumulator_frame
self._pass_through = pass_through 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]: async def process_frame(self, frame: Frame) -> AsyncGenerator[Frame, None]:
if not self.messages: if not self.messages:
return return
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
elif isinstance(frame, self._end_frame): elif isinstance(frame, self._end_frame):
self.aggregating = False self._seen_end_frame = True
# Sometimes VAD triggers quickly on and off. If we don't get any transcription,
# it creates empty LLM message queue frames # We might have received the end frame but we might still be
if len(self.aggregation) > 0: # aggregating (i.e. we have seen interim results but not the final
self.messages.append( # text).
{"role": self._role, "content": self.aggregation}) self.aggregating = self._seen_interim_results
self.aggregation = ""
yield self._end_frame() # Send the aggregation if we are not aggregating anymore (i.e. no
yield LLMMessagesFrame(self.messages) # more interim results received).
elif isinstance(frame, self._accumulator_frame) and self.aggregating: send_aggregation = not self.aggregating
self.aggregation += f" {frame.text}" 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: if self._pass_through:
yield frame 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: else:
yield frame 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): class LLMAssistantResponseAggregator(LLMResponseAggregator):
def __init__(self, messages: list[dict]): def __init__(self, messages: list[dict]):
@@ -160,6 +207,7 @@ class LLMUserResponseAggregator(LLMResponseAggregator):
start_frame=UserStartedSpeakingFrame, start_frame=UserStartedSpeakingFrame,
end_frame=UserStoppedSpeakingFrame, end_frame=UserStoppedSpeakingFrame,
accumulator_frame=TranscriptionFrame, accumulator_frame=TranscriptionFrame,
interim_accumulator_frame=InterimTranscriptionFrame,
pass_through=False, 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}" 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): class TTSStartFrame(ControlFrame):
"""Used to indicate the beginning of a TTS response. Following AudioFrames """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 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 typing import Any
from dailyai.pipeline.frames import ( from dailyai.pipeline.frames import (
InterimTranscriptionFrame,
ReceivedAppMessageFrame, ReceivedAppMessageFrame,
TranscriptionFrame, TranscriptionFrame,
UserImageFrame, UserImageFrame,
@@ -88,9 +89,11 @@ class DailyTransport(ThreadedTransport, EventHandler):
"model": "2-conversationalai", "model": "2-conversationalai",
"profanity_filter": True, "profanity_filter": True,
"redact": False, "redact": False,
"endpointing": True,
"punctuate": True,
"includeRawResponse": True,
"extra": { "extra": {
"endpointing": True, "interim_results": True,
"punctuate": False,
}, },
} }
@@ -368,8 +371,12 @@ class DailyTransport(ThreadedTransport, EventHandler):
elif "session_id" in message: elif "session_id" in message:
participantId = message["session_id"] participantId = message["session_id"]
if self._my_participant_id and participantId != self._my_participant_id: if self._my_participant_id and participantId != self._my_participant_id:
frame = TranscriptionFrame( is_final = message["rawResponse"]["is_final"]
message["text"], participantId, message["timestamp"]) if is_final:
frame = TranscriptionFrame(message["text"], participantId, message["timestamp"])
else:
frame = InterimTranscriptionFrame(
message["text"], participantId, message["timestamp"])
asyncio.run_coroutine_threadsafe( asyncio.run_coroutine_threadsafe(
self.receive_queue.put(frame), self._loop) self.receive_queue.put(frame), self._loop)