updated ImageFrame and added URLImageFrame and UserImageFrame

This commit is contained in:
Aleix Conchillo Flaqué
2024-04-08 23:10:10 -07:00
parent 96950ca6df
commit 319b8e7816
16 changed files with 97 additions and 30 deletions

View File

@@ -5,7 +5,7 @@ import tkinter as tk
import os import os
from dailyai.pipeline.aggregators import LLMFullResponseAggregator from dailyai.pipeline.aggregators import LLMFullResponseAggregator
from dailyai.pipeline.frames import AudioFrame, ImageFrame, LLMMessagesFrame, TextFrame from dailyai.pipeline.frames import AudioFrame, URLImageFrame, LLMMessagesFrame, TextFrame
from dailyai.services.open_ai_services import OpenAILLMService from dailyai.services.open_ai_services import OpenAILLMService
from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService
from dailyai.services.fal_ai_services import FalImageGenService from dailyai.services.fal_ai_services import FalImageGenService
@@ -94,6 +94,7 @@ async def main():
"text": image_description, "text": image_description,
"image_url": image_data[0], "image_url": image_data[0],
"image": image_data[1], "image": image_data[1],
"image_size": image_data[2],
"audio": audio, "audio": audio,
} }
@@ -117,7 +118,7 @@ async def main():
if data: if data:
await transport.send_queue.put( await transport.send_queue.put(
[ [
ImageFrame(data["image_url"], data["image"]), URLImageFrame(data["image_url"], data["image"], data["image_size"]),
AudioFrame(data["audio"]), AudioFrame(data["audio"]),
] ]
) )

View File

@@ -35,9 +35,9 @@ class ImageSyncAggregator(AIService):
self._waiting_image_bytes = self._waiting_image.tobytes() self._waiting_image_bytes = self._waiting_image.tobytes()
async def process_frame(self, frame: Frame) -> AsyncGenerator[Frame, None]: async def process_frame(self, frame: Frame) -> AsyncGenerator[Frame, None]:
yield ImageFrame(None, self._speaking_image_bytes) yield ImageFrame(self._speaking_image_bytes, (1024, 1024))
yield frame yield frame
yield ImageFrame(None, self._waiting_image_bytes) yield ImageFrame(self._waiting_image_bytes, (1024, 1024))
async def main(room_url: str, token): async def main(room_url: str, token):

View File

@@ -122,7 +122,7 @@ async def main(room_url: str):
) )
await transport.send_queue.put( await transport.send_queue.put(
[ [
ImageFrame(None, image_data1[1]), ImageFrame(image_data1[1], image_data1[2]),
AudioFrame(audio1), AudioFrame(audio1),
] ]
) )
@@ -134,7 +134,7 @@ async def main(room_url: str):
) )
await transport.send_queue.put( await transport.send_queue.put(
[ [
ImageFrame(None, image_data2[1]), ImageFrame(image_data2[1], image_data2[2]),
AudioFrame(audio2), AudioFrame(audio2),
] ]
) )

View File

@@ -55,7 +55,7 @@ for file in image_files:
sprites[file] = img.tobytes() sprites[file] = img.tobytes()
# When the bot isn't talking, show a static image of the cat listening # When the bot isn't talking, show a static image of the cat listening
quiet_frame = ImageFrame("", sprites["sc-listen-1.png"]) quiet_frame = ImageFrame(sprites["sc-listen-1.png"], (720, 1280))
# When the bot is talking, build an animation from two sprites # When the bot is talking, build an animation from two sprites
talking_list = [sprites["sc-default.png"], sprites["sc-talk.png"]] talking_list = [sprites["sc-default.png"], sprites["sc-talk.png"]]
talking = [random.choice(talking_list) for x in range(30)] talking = [random.choice(talking_list) for x in range(30)]

View File

@@ -48,7 +48,7 @@ for i in range(1, 26):
flipped = sprites[::-1] flipped = sprites[::-1]
sprites.extend(flipped) sprites.extend(flipped)
# When the bot isn't talking, show a static image of the cat listening # When the bot isn't talking, show a static image of the cat listening
quiet_frame = ImageFrame("", sprites[0]) quiet_frame = ImageFrame(sprites[0], (1024, 576))
talking_frame = SpriteFrame(images=sprites) talking_frame = SpriteFrame(images=sprites)

View File

@@ -99,7 +99,7 @@ class StoryProcessor(FrameProcessor):
1. Catch the frames that are generated by the LLM service 1. Catch the frames that are generated by the LLM service
""" """
if isinstance(frame, UserStoppedSpeakingFrame): if isinstance(frame, UserStoppedSpeakingFrame):
yield ImageFrame(None, images["grandma-writing.png"]) yield ImageFrame(images["grandma-writing.png"], (1024, 1024))
yield AudioFrame(sounds["talking.wav"]) yield AudioFrame(sounds["talking.wav"])
elif isinstance(frame, TextFrame): elif isinstance(frame, TextFrame):
@@ -112,7 +112,7 @@ class StoryProcessor(FrameProcessor):
self._text = self._text.replace("\n", " ") self._text = self._text.replace("\n", " ")
if len(self._text) > 2: if len(self._text) > 2:
yield ImageFrame(None, images["grandma-writing.png"]) yield ImageFrame(images["grandma-writing.png"], (1024, 1024))
yield StoryStartFrame(self._text) yield StoryStartFrame(self._text)
yield AudioFrame(sounds["ding3.wav"]) yield AudioFrame(sounds["ding3.wav"])
self._text = "" self._text = ""
@@ -146,11 +146,11 @@ class StoryProcessor(FrameProcessor):
# last bit # last bit
pass pass
elif isinstance(frame, LLMResponseEndFrame): elif isinstance(frame, LLMResponseEndFrame):
yield ImageFrame(None, images["grandma-writing.png"]) yield ImageFrame(images["grandma-writing.png"], (1024, 1024))
yield StoryPromptFrame(self._text) yield StoryPromptFrame(self._text)
self._text = "" self._text = ""
yield frame yield frame
yield ImageFrame(None, images["grandma-listening.png"]) yield ImageFrame(images["grandma-listening.png"], (1024, 1024))
yield AudioFrame(sounds["listening.wav"]) yield AudioFrame(sounds["listening.wav"])
else: else:
@@ -252,7 +252,7 @@ async def main(room_url: str, token):
[llm, lca, tts], sink=transport.send_queue) [llm, lca, tts], sink=transport.send_queue)
await local_pipeline.queue_frames( await local_pipeline.queue_frames(
[ [
ImageFrame(None, images["grandma-listening.png"]), ImageFrame(images["grandma-listening.png"], (1024, 1024)),
LLMMessagesFrame(intro_messages), LLMMessagesFrame(intro_messages),
AudioFrame(sounds["listening.wav"]), AudioFrame(sounds["listening.wav"]),
EndPipeFrame(), EndPipeFrame(),

View File

@@ -360,7 +360,7 @@ class GatedAggregator(FrameProcessor):
... start_open=False) ... start_open=False)
>>> asyncio.run(print_frames(aggregator, TextFrame("Hello"))) >>> asyncio.run(print_frames(aggregator, TextFrame("Hello")))
>>> asyncio.run(print_frames(aggregator, TextFrame("Hello again."))) >>> asyncio.run(print_frames(aggregator, TextFrame("Hello again.")))
>>> asyncio.run(print_frames(aggregator, ImageFrame(url='', image=bytes([])))) >>> asyncio.run(print_frames(aggregator, ImageFrame(image=bytes([]), size=(0, 0))))
ImageFrame ImageFrame
Hello Hello
Hello again. Hello again.

View File

@@ -70,11 +70,39 @@ class AudioFrame(Frame):
class ImageFrame(Frame): class ImageFrame(Frame):
"""An image. Will be shown by the transport if the transport's camera is """An image. Will be shown by the transport if the transport's camera is
enabled.""" enabled."""
url: str | None
image: bytes image: bytes
size: tuple[int, int]
def __str__(self): def __str__(self):
return f"{self.__class__.__name__}, url: {self.url}, image size: {len(self.image)} B" return f"{self.__class__.__name__}, image size: {self.size[0]}x{self.size[1]} buffer size: {len(self.image)} B"
@dataclass()
class URLImageFrame(ImageFrame):
"""An image. Will be shown by the transport if the transport's camera is
enabled."""
url: str | None
def __init__(self, url, image, size):
super().__init__(image, size)
self.url = url
def __str__(self):
return f"{self.__class__.__name__}, url: {self.url}, image size: {self.size[0]}x{self.size[1]}, buffer size: {len(self.image)} B"
@dataclass()
class UserImageFrame(ImageFrame):
"""An image associated to a user. Will be shown by the transport if the transport's camera is
enabled."""
user_id: str
def __init__(self, user_id, image, size):
super().__init__(image, size)
self.user_id = user_id
def __str__(self):
return f"{self.__class__.__name__}, user: {self.user_id}, image size: {self.size[0]}x{self.size[1]}, buffer size: {len(self.image)} B"
@dataclass() @dataclass()

View File

@@ -14,6 +14,7 @@ from dailyai.pipeline.frames import (
TTSStartFrame, TTSStartFrame,
TextFrame, TextFrame,
TranscriptionFrame, TranscriptionFrame,
URLImageFrame,
) )
from abc import abstractmethod from abc import abstractmethod
@@ -87,7 +88,7 @@ class ImageGenService(AIService):
# Renders the image. Returns an Image object. # Renders the image. Returns an Image object.
@abstractmethod @abstractmethod
async def run_image_gen(self, sentence: str) -> tuple[str, bytes]: async def run_image_gen(self, sentence: str) -> tuple[str, bytes, tuple[int, int]]:
pass pass
async def process_frame(self, frame: Frame) -> AsyncGenerator[Frame, None]: async def process_frame(self, frame: Frame) -> AsyncGenerator[Frame, None]:
@@ -95,8 +96,8 @@ class ImageGenService(AIService):
yield frame yield frame
return return
(url, image_data) = await self.run_image_gen(frame.text) (url, image_data, image_size) = await self.run_image_gen(frame.text)
yield ImageFrame(url, image_data) yield URLImageFrame(url, image_data, image_size)
class STTService(AIService): class STTService(AIService):

View File

@@ -105,7 +105,7 @@ class AzureImageGenServiceREST(ImageGenService):
self._model = model self._model = model
self._aiohttp_session = aiohttp_session self._aiohttp_session = aiohttp_session
async def run_image_gen(self, sentence) -> tuple[str, bytes]: async def run_image_gen(self, sentence) -> tuple[str, bytes, tuple[int, int]]:
url = f"{self._azure_endpoint}openai/images/generations:submit?api-version={self._api_version}" url = f"{self._azure_endpoint}openai/images/generations:submit?api-version={self._api_version}"
headers = { headers = {
"api-key": self._api_key, "api-key": self._api_key,
@@ -146,4 +146,4 @@ class AzureImageGenServiceREST(ImageGenService):
async with self._aiohttp_session.get(image_url) as response: async with self._aiohttp_session.get(image_url) as response:
image_stream = io.BytesIO(await response.content.read()) image_stream = io.BytesIO(await response.content.read())
image = Image.open(image_stream) image = Image.open(image_stream)
return (image_url, image.tobytes()) return (image_url, image.tobytes(), image.size)

View File

@@ -31,7 +31,7 @@ class FalImageGenService(ImageGenService):
if key_secret: if key_secret:
os.environ["FAL_KEY_SECRET"] = key_secret os.environ["FAL_KEY_SECRET"] = key_secret
async def run_image_gen(self, sentence) -> tuple[str, bytes]: async def run_image_gen(self, sentence) -> tuple[str, bytes, tuple[int, int]]:
def get_image_url(sentence, size): def get_image_url(sentence, size):
handler = fal.apps.submit( handler = fal.apps.submit(
"110602490-fast-sdxl", "110602490-fast-sdxl",
@@ -55,4 +55,4 @@ class FalImageGenService(ImageGenService):
async with self._aiohttp_session.get(image_url) as response: async with self._aiohttp_session.get(image_url) as response:
image_stream = io.BytesIO(await response.content.read()) image_stream = io.BytesIO(await response.content.read())
image = Image.open(image_stream) image = Image.open(image_stream)
return (image_url, image.tobytes()) return (image_url, image.tobytes(), image.size)

View File

@@ -36,7 +36,7 @@ class OpenAIImageGenService(ImageGenService):
self._client = AsyncOpenAI(api_key=api_key) self._client = AsyncOpenAI(api_key=api_key)
self._aiohttp_session = aiohttp_session self._aiohttp_session = aiohttp_session
async def run_image_gen(self, sentence) -> tuple[str, bytes]: async def run_image_gen(self, sentence) -> tuple[str, bytes, tuple[int, int]]:
self.logger.info("Generating OpenAI image", sentence) self.logger.info("Generating OpenAI image", sentence)
image = await self._client.images.generate( image = await self._client.images.generate(
@@ -53,4 +53,4 @@ class OpenAIImageGenService(ImageGenService):
async with self._aiohttp_session.get(image_url) as response: async with self._aiohttp_session.get(image_url) as response:
image_stream = io.BytesIO(await response.content.read()) image_stream = io.BytesIO(await response.content.read())
image = Image.open(image_stream) image = Image.open(image_stream)
return (image_url, image.tobytes()) return (image_url, image.tobytes(), image.size)

View File

@@ -19,7 +19,7 @@ class MockAIService(AIService):
image_stream = io.BytesIO(response.content) image_stream = io.BytesIO(response.content)
image = Image.open(image_stream) image = Image.open(image_stream)
time.sleep(1) time.sleep(1)
return (image_url, image) return (image_url, image.tobytes(), image.size)
def run_llm(self, messages, latest_user_message=None, stream=True): def run_llm(self, messages, latest_user_message=None, stream=True):
for i in range(5): for i in range(5):

View File

@@ -2,6 +2,7 @@ import asyncio
import inspect import inspect
import logging import logging
import signal import signal
import time
import threading import threading
import types import types
@@ -11,6 +12,7 @@ from typing import Any
from dailyai.pipeline.frames import ( from dailyai.pipeline.frames import (
ReceivedAppMessageFrame, ReceivedAppMessageFrame,
TranscriptionFrame, TranscriptionFrame,
UserImageFrame,
) )
from threading import Event from threading import Event
@@ -58,6 +60,7 @@ class DailyTransport(ThreadedTransport, EventHandler):
bot_name: str, bot_name: str,
min_others_count: int = 1, min_others_count: int = 1,
start_transcription: bool = False, start_transcription: bool = False,
video_rendering_enabled: bool = False,
**kwargs, **kwargs,
): ):
kwargs['has_webrtc_vad'] = True kwargs['has_webrtc_vad'] = True
@@ -69,6 +72,7 @@ class DailyTransport(ThreadedTransport, EventHandler):
self._token: str | None = token self._token: str | None = token
self._min_others_count = min_others_count self._min_others_count = min_others_count
self._start_transcription = start_transcription self._start_transcription = start_transcription
self._video_rendering_enabled = video_rendering_enabled
self._is_interrupted = Event() self._is_interrupted = Event()
self._stop_threads = Event() self._stop_threads = Event()
@@ -76,6 +80,8 @@ class DailyTransport(ThreadedTransport, EventHandler):
self._other_participant_has_joined = False self._other_participant_has_joined = False
self._my_participant_id = None self._my_participant_id = None
self._video_renderers = {}
self.transcription_settings = { self.transcription_settings = {
"language": "en", "language": "en",
"tier": "nova", "tier": "nova",
@@ -236,7 +242,7 @@ class DailyTransport(ThreadedTransport, EventHandler):
self.client.update_subscription_profiles({ self.client.update_subscription_profiles({
"base": { "base": {
"camera": "unsubscribed", "camera": "subscribed" if self._video_rendering_enabled else "unsubscribed",
} }
}) })
@@ -268,6 +274,37 @@ class DailyTransport(ThreadedTransport, EventHandler):
def start_recording(self): def start_recording(self):
self.client.start_recording() self.client.start_recording()
def render_participant_video(self,
participant_id,
framerate=10,
video_source="camera",
color_format="RGB") -> None:
if not self._video_rendering_enabled:
self._logger.warn("Video rendering is not enabled")
self._video_renderers[participant_id] = {
"framerate": framerate,
"timestamp": 0,
}
self.client.set_video_renderer(
participant_id,
self.on_participant_video_frame,
video_source=video_source,
color_format=color_format)
def on_participant_video_frame(self, participant_id, video_frame):
curr_time = time.time()
prev_time = self._video_renderers[participant_id]["timestamp"]
diff_time = curr_time - prev_time
period = 1 / self._video_renderers[participant_id]["framerate"]
if diff_time > period and self._loop:
self._video_renderers[participant_id]["timestamp"] = curr_time
frame = UserImageFrame(participant_id, video_frame.buffer,
(video_frame.width, video_frame.height))
asyncio.run_coroutine_threadsafe(
self.receive_queue.put(frame), self._loop
)
def on_error(self, error): def on_error(self, error):
self._logger.error(f"on_error: {error}") self._logger.error(f"on_error: {error}")

View File

@@ -54,13 +54,13 @@ class TestDailyFrameAggregators(unittest.IsolatedAsyncioTestCase):
TextFrame("Hello, "), TextFrame("Hello, "),
TextFrame("world."), TextFrame("world."),
AudioFrame(b"hello"), AudioFrame(b"hello"),
ImageFrame("image", b"image"), ImageFrame(b"image", (0, 0)),
AudioFrame(b"world"), AudioFrame(b"world"),
LLMResponseEndFrame(), LLMResponseEndFrame(),
] ]
expected_output_frames = [ expected_output_frames = [
ImageFrame("image", b"image"), ImageFrame(b"image", (0, 0)),
LLMResponseStartFrame(), LLMResponseStartFrame(),
TextFrame("Hello, "), TextFrame("Hello, "),
TextFrame("world."), TextFrame("world."),

View File

@@ -68,7 +68,7 @@ class TestDailyTransport(unittest.IsolatedAsyncioTestCase):
await transport.send_queue.put(AudioFrame(bytes([0] * 3300))) await transport.send_queue.put(AudioFrame(bytes([0] * 3300)))
async def send_video_frame(): async def send_video_frame():
await transport.send_queue.put(ImageFrame(None, b"test")) await transport.send_queue.put(ImageFrame(b"test", (0, 0)))
await asyncio.gather(transport.run(), send_audio_frame(), send_video_frame()) await asyncio.gather(transport.run(), send_audio_frame(), send_video_frame())