probably will undo most of this

This commit is contained in:
Moishe Lettvin
2024-02-27 18:52:01 -05:00
parent fc19a55f04
commit d03dd62941
3 changed files with 120 additions and 28 deletions

View File

@@ -18,11 +18,19 @@ from dailyai.queue_frame import (
from abc import abstractmethod from abc import abstractmethod
from typing import AsyncGenerator, AsyncIterable, BinaryIO, Iterable from typing import AsyncGenerator, AsyncIterable, BinaryIO, Iterable
class AbstractPipeService(): class AbstractPipeService:
def __init__(self, source:asyncio.Queue[QueueFrame], sink:asyncio.Queue[QueueFrame]):
self.source: asyncio.Queue[QueueFrame] = source def __init__(
self.sink: asyncio.Queue[QueueFrame] = sink self,
pass ):
self.source_queue: asyncio.Queue[QueueFrame] = asyncio.Queue()
self.sink_queue: asyncio.Queue[QueueFrame] = asyncio.Queue()
async def get(self) -> QueueFrame:
return await self.sink_queue.get()
async def put(self, frame: QueueFrame) -> None:
await self.source_queue.put(frame)
@abstractmethod @abstractmethod
async def process_queue(self): async def process_queue(self):
@@ -32,34 +40,20 @@ class PipeService(AbstractPipeService):
def __init__( def __init__(
self, self,
source: asyncio.Queue[QueueFrame] | AbstractPipeService | None=None, source: AbstractPipeService | None = None,
sink: asyncio.Queue[QueueFrame] | AbstractPipeService | None=None, sink: AbstractPipeService | None = None,
): ):
super().__init__()
self.logger: logging.Logger = logging.getLogger("dailyai") self.logger: logging.Logger = logging.getLogger("dailyai")
if not source:
source = asyncio.Queue()
elif isinstance(source, AbstractPipeService):
source = source.sink
if not sink:
sink = asyncio.Queue()
elif isinstance(sink, AbstractPipeService):
sink = sink.source
self.source = source self.source = source
self.sink = sink self.sink = sink
def chain(self, service: AbstractPipeService) -> AbstractPipeService:
self.source = service.sink
return self
async def process_queue(self): async def process_queue(self):
if not self.source: if not self.source:
return return
while True: while True:
frame = await self.source.get() frame: QueueFrame = await self.source.get()
async for output_frame in self.process_frame(frame): async for output_frame in self.process_frame(frame):
await self.sink.put(output_frame) await self.sink.put(output_frame)
if isinstance(frame, EndStreamQueueFrame): if isinstance(frame, EndStreamQueueFrame):

View File

@@ -0,0 +1,10 @@
import unittest
from unittest.mock import MagicMock, patch
from dailyai.queue_frame import AudioQueueFrame, ImageQueueFrame
from dailyai.services.ai_services import PipeService
class TestDailyTransport(unittest.IsolatedAsyncioTestCase):
def test_pipe_chain(self):
pipe1 = PipeService()

View File

@@ -1,8 +1,10 @@
import asyncio import asyncio
from typing import Any, AsyncGenerator, Callable, Tuple
import aiohttp import aiohttp
import os import os
from dailyai.queue_frame import AudioQueueFrame, ImageQueueFrame from dailyai.queue_frame import AudioQueueFrame, EndStreamQueueFrame, ImageQueueFrame, LLMResponseEndQueueFrame, QueueFrame, TextQueueFrame
from dailyai.services.ai_services import PipeService
from dailyai.services.azure_ai_services import AzureLLMService, AzureImageGenServiceREST, AzureTTSService from dailyai.services.azure_ai_services import AzureLLMService, AzureImageGenServiceREST, AzureTTSService
from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService
from dailyai.services.daily_transport_service import DailyTransportService from dailyai.services.daily_transport_service import DailyTransportService
@@ -27,23 +29,109 @@ async def main(room_url):
camera_height=1024 camera_height=1024
) )
"""
/ TTS \
Month prompt -> LLM -> Fork -> -> Gate -> Transport
\ ImageGen /
"""
class QueueFork(PipeService):
def __init__(self, source:PipeService, sinks: list[asyncio.Queue[QueueFrame]]):
self.source = source
self.sinks = sinks
async def process_queue(self):
while True:
frame = await self.source.get()
for sink in self.sinks:
await sink.put(frame)
if isinstance(frame, EndStreamQueueFrame):
break
class QueueGateOnFrame(PipeService):
def __init__(
self,
source: asyncio.Queue[QueueFrame],
sink: asyncio.Queue[QueueFrame],
aggregator: Callable[[Any, QueueFrame], Tuple[Any, QueueFrame | None]]
):
self.source = source
self.sink = sink
self.aggregator = aggregator
self.accumulation = None
async def process_frame(
self, frame: QueueFrame
) -> AsyncGenerator[QueueFrame, None]:
output_frame: QueueFrame | None = None
(self.aggregation, output_frame) = self.aggregator(
self.aggregation, frame
)
if output_frame:
yield output_frame
class QueueMergeGateOnFirst(PipeService):
def __init__(
self, queue_fork_service: QueueFork, sink: asyncio.Queue[QueueFrame]
):
self.queue_fork_service = queue_fork_service
self.sink = sink
async def process_queue(self):
(frames): list[QueueFrame] = await asyncio.gather(
*[source.get() for source in self.queue_fork_service.get_end_sinks()]
)
for idx, frame in enumerate(frames):
await self.sink.put(frame)
# if the first frame we got from a source is an EndStreamQueueFrame, remove that source
if isinstance(frame, EndStreamQueueFrame):
self.sources.pop(idx)
async def pass_through(sink, source):
while True:
frame = await source.get()
await sink.put(frame)
if isinstance(frame, EndStreamQueueFrame):
break
await asyncio.gather(*[pass_through(self.sink, source) for source in self.sources])
llm = AzureLLMService( llm = AzureLLMService(
api_key=os.getenv("AZURE_CHATGPT_API_KEY"), api_key=os.getenv("AZURE_CHATGPT_API_KEY"),
endpoint=os.getenv("AZURE_CHATGPT_ENDPOINT"), endpoint=os.getenv("AZURE_CHATGPT_ENDPOINT"),
model=os.getenv("AZURE_CHATGPT_MODEL")) model=os.getenv("AZURE_CHATGPT_MODEL"))
tts = ElevenLabsTTSService( tts = ElevenLabsTTSService(
aiohttp_session=session, aiohttp_session=session,
api_key=os.getenv("ELEVENLABS_API_KEY"), api_key=os.getenv("ELEVENLABS_API_KEY"),
voice_id="ErXwobaYiN019PkySvjV") voice_id="ErXwobaYiN019PkySvjV")
# tts = AzureTTSService(api_key=os.getenv("AZURE_SPEECH_API_KEY"), region=os.getenv("AZURE_SPEECH_REGION"))
dalle = FalImageGenService( dalle = FalImageGenService(
image_size="1024x1024", image_size="1024x1024",
aiohttp_session=session, aiohttp_session=session,
key_id=os.getenv("FAL_KEY_ID"), key_id=os.getenv("FAL_KEY_ID"),
key_secret=os.getenv("FAL_KEY_SECRET")) key_secret=os.getenv("FAL_KEY_SECRET"),
# dalle = OpenAIImageGenService(aiohttp_session=session, api_key=os.getenv("OPENAI_DALLE_API_KEY"), image_size="1024x1024") )
# dalle = AzureImageGenServiceREST(image_size="1024x1024", aiohttp_session=session, api_key=os.getenv("AZURE_DALLE_API_KEY"), endpoint=os.getenv("AZURE_DALLE_ENDPOINT"), model=os.getenv("AZURE_DALLE_MODEL"))
def aggregator(
accumulation, frame: QueueFrame
) -> tuple[Any, QueueFrame | None]:
if isinstance(frame, TextQueueFrame):
accumulation += frame.text
return (accumulation, None)
elif isinstance(frame, LLMResponseEndQueueFrame):
return ("", TextQueueFrame(accumulation))
else:
return (accumulation, frame)
llm_image_gate = QueueGateOnFrame(llm.sink, dalle.source, aggregator)
fork_audio_image = QueueFork(llm.sink, [tts.source, llm_image_gate.source])
audio_image_gate = QueueMergeGateOnFirst([tts.sink, dalle.sink], transport.send_queue)
# Get a complete audio chunk from the given text. Splitting this into its own # Get a complete audio chunk from the given text. Splitting this into its own
# coroutine lets us ensure proper ordering of the audio chunks on the send queue. # coroutine lets us ensure proper ordering of the audio chunks on the send queue.