added params and model attribute to fal service

This commit is contained in:
Jon Taylor
2024-04-09 17:43:27 -07:00
parent 4bd29b0080
commit 7b44a79a5b

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) -> tuple[str, bytes]:
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,9 +66,10 @@ 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())
image = Image.open(image_stream) image = Image.open(image_stream)
return (image_url, image.tobytes(), image.size) return (image_url, image.tobytes())