Merge pull request #114 from daily-co/jpt/fal-updates

Updated Fal.ai service to take a params model and allow for model string param
This commit is contained in:
Aleix Conchillo Flaqué
2024-04-11 00:47:33 +08:00
committed by GitHub
10 changed files with 60 additions and 29 deletions

View File

@@ -31,7 +31,9 @@ async def main(room_url):
) )
imagegen = FalImageGenService( imagegen = FalImageGenService(
image_size="square_hd", params=FalImageGenService.InputParams(
image_size="square_hd"
),
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"),

View File

@@ -35,7 +35,9 @@ async def main():
) )
imagegen = FalImageGenService( imagegen = FalImageGenService(
image_size="square_hd", params=FalImageGenService.InputParams(
image_size="square_hd"
),
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"),

View File

@@ -85,7 +85,9 @@ async def main(room_url):
model="gpt-4-turbo-preview") model="gpt-4-turbo-preview")
imagegen = FalImageGenService( imagegen = FalImageGenService(
image_size="square_hd", params=FalImageGenService.InputParams(
image_size="square_hd"
),
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"),

View File

@@ -45,7 +45,9 @@ async def main():
model="gpt-4-turbo-preview") model="gpt-4-turbo-preview")
imagegen = FalImageGenService( imagegen = FalImageGenService(
image_size="1024x1024", params=FalImageGenService.InputParams(
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"),

View File

@@ -51,7 +51,9 @@ async def main(room_url: str):
voice_id="jBpfuIE2acCO8z3wKNLl", voice_id="jBpfuIE2acCO8z3wKNLl",
) )
dalle = FalImageGenService( dalle = FalImageGenService(
image_size="1024x1024", params=FalImageGenService.InputParams(
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"),

View File

@@ -204,7 +204,9 @@ async def main(room_url: str, token):
voice_id="Xb7hH8MSUJpSbSDYk0k2", voice_id="Xb7hH8MSUJpSbSDYk0k2",
) # matilda ) # matilda
img = FalImageGenService( img = FalImageGenService(
image_size="1024x1024", params={
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"),

View File

@@ -83,13 +83,12 @@ class TTSService(AIService):
class ImageGenService(AIService): class ImageGenService(AIService):
def __init__(self, image_size, **kwargs): def __init__(self, **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
self.image_size = image_size
# 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, tuple[int, int]]: async def run_image_gen(self, prompt: 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]:

View File

@@ -97,23 +97,24 @@ class AzureImageGenServiceREST(ImageGenService):
endpoint, endpoint,
model, model,
): ):
super().__init__(image_size=image_size) super().__init__()
self._api_key = api_key self._api_key = api_key
self._azure_endpoint = endpoint self._azure_endpoint = endpoint
self._api_version = api_version self._api_version = api_version
self._model = model self._model = model
self._aiohttp_session = aiohttp_session self._aiohttp_session = aiohttp_session
self._image_size = image_size
async def run_image_gen(self, sentence) -> tuple[str, bytes, tuple[int, int]]: async def run_image_gen(self, prompt: str) -> 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,
"Content-Type": "application/json"} "Content-Type": "application/json"}
body = { body = {
# Enter your prompt text here # Enter your prompt text here
"prompt": sentence, "prompt": prompt,
"size": self.image_size, "size": self._image_size,
"n": 1, "n": 1,
} }
async with self._aiohttp_session.post( async with self._aiohttp_session.post(

View File

@@ -3,6 +3,8 @@ import asyncio
import io import io
import os import os
from PIL import Image from PIL import Image
from pydantic import BaseModel
from typing import Optional, Union, Dict
from dailyai.services.ai_services import ImageGenService from dailyai.services.ai_services import ImageGenService
@@ -16,30 +18,44 @@ except ModuleNotFoundError as e:
class FalImageGenService(ImageGenService): class FalImageGenService(ImageGenService):
class InputParams(BaseModel):
seed: Optional[int] = None
num_inference_steps: int = 4
num_images: int = 1
image_size: Union[str, Dict[str, int]] = "square_hd"
expand_prompt: bool = False
enable_safety_checker: bool = True
format: str = "png"
def __init__( def __init__(
self, self,
*, *,
image_size,
aiohttp_session: aiohttp.ClientSession, aiohttp_session: aiohttp.ClientSession,
params: InputParams,
model="fal-ai/fast-sdxl",
key_id=None, key_id=None,
key_secret=None key_secret=None
): ):
super().__init__(image_size) super().__init__()
self._model = model
self._params = params
self._aiohttp_session = aiohttp_session self._aiohttp_session = aiohttp_session
if key_id: if key_id:
os.environ["FAL_KEY_ID"] = key_id os.environ["FAL_KEY_ID"] = key_id
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, tuple[int, int]]: async def run_image_gen(self, prompt: str) -> tuple[str, bytes, tuple[int, int]]:
def get_image_url(sentence, size): def get_image_url(prompt):
handler = fal.apps.submit( handler = fal.apps.submit( # type: ignore
"110602490-fast-sdxl", self._model,
# "fal-ai/fast-sdxl", arguments={
arguments={"prompt": sentence}, "prompt": prompt,
**self._params.dict(),
},
) )
for event in handler.iter_events(): for event in handler.iter_events():
if isinstance(event, fal.apps.InProgress): if isinstance(event, fal.apps.InProgress): # type: ignore
pass pass
result = handler.get() result = handler.get()
@@ -50,7 +66,8 @@ class FalImageGenService(ImageGenService):
return image_url return image_url
image_url = await asyncio.to_thread(get_image_url, sentence, self.image_size) image_url = await asyncio.to_thread(get_image_url, prompt)
# Load the image from the url # Load the image from the url
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())

View File

@@ -1,3 +1,4 @@
from typing import Literal
import aiohttp import aiohttp
from PIL import Image from PIL import Image
import io import io
@@ -26,24 +27,25 @@ class OpenAIImageGenService(ImageGenService):
def __init__( def __init__(
self, self,
*, *,
image_size: str, image_size: Literal["256x256", "512x512", "1024x1024", "1792x1024", "1024x1792"],
aiohttp_session: aiohttp.ClientSession, aiohttp_session: aiohttp.ClientSession,
api_key, api_key,
model="dall-e-3", model="dall-e-3",
): ):
super().__init__(image_size=image_size) super().__init__()
self._model = model self._model = model
self._image_size = image_size
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, tuple[int, int]]: async def run_image_gen(self, prompt: str) -> tuple[str, bytes, tuple[int, int]]:
self.logger.info("Generating OpenAI image", sentence) self.logger.info("Generating OpenAI image", prompt)
image = await self._client.images.generate( image = await self._client.images.generate(
prompt=sentence, prompt=prompt,
model=self._model, model=self._model,
n=1, n=1,
size=self.image_size size=self._image_size
) )
image_url = image.data[0].url image_url = image.data[0].url
if not image_url: if not image_url: