@@ -3,6 +3,7 @@ from ml_model.runners.openrouter import OpenrouterRunner from ml_model.runners.replicate import ( ReplicateAudioRunner, ReplicateImageRunner, + ReplicateIconicRunner, ReplicateTextRunner, ReplicateVideoRunner, ) @@ -14,6 +15,7 @@ __all__ = [ 'ReplicateTextRunner', 'ReplicateAudioRunner', 'ReplicateImageRunner', + 'ReplicateIconicRunner', 'ReplicateVideoRunner', 'FalAIRunner' ] @@ -21,10 +21,14 @@ class FalAIRunner(BaseRunner): scrape_results: list[StringIO] | list[BytesIO] = [], ) -> Generator[str, Any, None]: try: + image_size = { + "width": parameters.pop('width'), + "height": parameters.pop('height') + } version = parameters.pop('model') model_owner = parameters.pop('model_owner') official = version and model_owner - payload = {'prompt': content, **parameters} + payload = {'prompt': content, 'image_size': image_size, **parameters} if not official: payload.update({'model_version': version}) with httpx.Client( @@ -1,4 +1,5 @@ import base64 +import logging from abc import ABC, abstractmethod from io import BytesIO, StringIO from typing import TYPE_CHECKING, Any, Iterable @@ -12,6 +13,7 @@ from ml_model.runners.base import BaseRunner if TYPE_CHECKING: from messages.models import Message +logger = logging.getLogger(__name__) class ReplicateBaseRunner(BaseRunner, ABC): @classmethod @@ -131,6 +133,7 @@ class ReplicateImageRunner(BaseRunner): model_owner = parameters.pop('model_owner') official = version and model_owner payload = {'input': {parameters.pop('prompt_key', None) or 'prompt': content, **parameters}} + logger.info(official) if not official: payload.update({'version': version}) if file: @@ -152,9 +155,12 @@ class ReplicateImageRunner(BaseRunner): resp = client.post( f'/models/{model_owner}/{version}/predictions' if official else '/predictions', json=payload ) + logger.info(resp) data = resp.json() + logger.info(data) stream_url = data['urls']['stream'] get_url = data['urls']['get'] + logger.info(stream_url, get_url) with client.stream( 'GET', stream_url, headers={'Accept': 'text/event-stream', 'Cache-Control': 'no-store'} ) as stream: @@ -168,6 +174,7 @@ class ReplicateImageRunner(BaseRunner): if content_length == 0: resp = client.get(get_url) data = resp.json() + logger.info(data) if data['status'] == 'failed': raise Exception('Model failed generation') @@ -176,7 +183,59 @@ class ReplicateVideoRunner(BaseRunner): @classmethod def generate(cls, content=None, file=None, parameters={}, history=[], scrape_results=[]): ... - class ReplicateAudioRunner(BaseRunner): @classmethod def generate(cls, content=None, file=None, parameters={}, history=[], scrape_results=[]): ... + + +class ReplicateIconicRunner(BaseRunner): + @classmethod + def generate(cls, content=None, file=None, parameters={}, history=[], scrape_results=[]): + version = parameters.pop('version') + model_owner = parameters.pop('model_owner') + official = version and model_owner + payload = {'input': {parameters.pop('prompt_key', None) or 'prompt': content, **parameters}} + logger.info(official) + if not official: + payload.update({'version': version}) + if file: + kind = filetype.guess(file.read(20)) + format = 'jpeg' if kind.extension == 'jpg' else kind.extension + if format in ('jpeg', 'png'): + mime = kind.mime if kind else 'application/octet-stream' + payload['input'][parameters.pop('file_key', None) or 'file'] = ( + f'data:{mime};base64,{base64.b64encode(file.read()).decode("utf-8")}' + ) + with httpx.Client( + base_url='https://api.replicate.com/v1', + headers={ + 'Authorization': f'Bearer {settings.REPLICATE_API_KEY}', + 'Content-Type': 'application/json', + }, + timeout=None, + ) as client: + resp = client.post( + f'/models/{model_owner}/{version}/predictions' if official else '/predictions', json=payload + ) + logger.info(resp) + data = resp.json() + logger.info(data) + stream_url = data['urls']['stream'] + get_url = data['urls']['get'] + logger.info(stream_url, get_url) + with client.stream( + 'GET', stream_url, headers={'Accept': 'text/event-stream', 'Cache-Control': 'no-store'} + ) as stream: + content_length = 0 + for chunk in stream.iter_text(): + chunk = chunk.strip().split('\n') + chunk = chunk[0] if content_length > 0 else chunk[-1][len('data:') :].strip() + content_length += len(chunk) + yield chunk + + if content_length == 0: + resp = client.get(get_url) + data = resp.json() + logger.info(data) + if data['status'] == 'failed': + raise Exception('Model failed generation') @@ -166,10 +166,10 @@ class InferenceService: if ( parameter.type in ( - Parameter.TypeChoices.FLOAT - or Parameter.TypeChoices.INT - or Parameter.TypeChoices.STR - or Parameter.TypeChoices.BOOL + Parameter.TypeChoices.FLOAT, + Parameter.TypeChoices.INT, + Parameter.TypeChoices.STR, + Parameter.TypeChoices.BOOL, ) or ( parameter.type