added params and model attribute to fal service
This commit is contained in:
@@ -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())
|
||||||
|
|||||||
Reference in New Issue
Block a user