@@ -325,6 +325,7 @@ GOOGLE_API_KEY = env.str('GOOGLE_API_KEY', default='defaultapikey') SERPER_API_KEY = env.str('SERPER_API_KEY', 'defaultapikey') FLUX_API_KEY = env.str('FLUX_API_KEY', 'defaultapikey') OPENROUTER_API_KEY = env.str('OPENROUTER_API_KEY', 'defaultapikey') +FAL_API_KEY = env.str('FAL_API_KEY', 'defaultapikey') OPENAI_PROXY_HOST = env.str('OPENAI_PROXY_HOST', 'neuron-proxy:8080') UPSCALE_MULTIPLIER_HOST = env.str('UPSCALE_MULTIPLIER_HOST', 'packet:8080') @@ -9,8 +9,8 @@ def migrate_nm_type(apps, schema_editor): NeuronModel = apps.get_model('ml_model', 'NeuronModel') models = NeuronModel.objects.all() for model in models: - model.type = model.category.slug - NeuronModel.objects.bulk_update(models, fields=['type']) + model.types = [model.category.slug] + NeuronModel.objects.bulk_update(models, fields=['types']) def migrate_versions_to_deployments(apps, schema_editor): Deployment = apps.get_model('ml_model', 'Deployment') @@ -0,0 +1,18 @@ +# Generated by Django 5.0.11 on 2025-05-27 07:29 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ('ml_model', '0055_remove_neuronmodel_category_and_more'), + ] + + operations = [ + migrations.AlterField( + model_name='deployment', + name='runner_import_path', + field=models.CharField(choices=[('ml_model.runners:FalAIRunner', 'FalAI'), ('ml_model.runners:OpenAIGPTRunner', 'OpenAIGPT'), ('ml_model.runners:OpenrouterRunner', 'Openrouter'), ('ml_model.runners:ReplicateAudioRunner', 'ReplicateAudio'), ('ml_model.runners:ReplicateImageRunner', 'ReplicateImage'), ('ml_model.runners:ReplicateTextRunner', 'ReplicateText'), ('ml_model.runners:ReplicateVideoRunner', 'ReplicateVideo')], max_length=100, verbose_name='Runner'), + ), + ] @@ -6,6 +6,7 @@ from ml_model.runners.replicate import ( ReplicateTextRunner, ReplicateVideoRunner, ) +from ml_model.runners.falai import FalAIRunner __all__ = [ 'OpenAIGPTRunner', @@ -14,4 +15,5 @@ __all__ = [ 'ReplicateAudioRunner', 'ReplicateImageRunner', 'ReplicateVideoRunner', + 'FalAIRunner' ] @@ -0,0 +1,47 @@ +import logging +from io import BytesIO, StringIO +from typing import Any, Iterable, Generator + +import httpx +from django.conf import settings + +from ml_model.runners.base import BaseRunner + +logger = logging.getLogger(__name__) + + +class FalAIRunner(BaseRunner): + @classmethod + def generate( + cls, + content: str | None = None, + file: BytesIO | StringIO | None = None, + parameters: dict[str, Any] = {}, + history: Iterable['Message'] = [], + scrape_results: list[StringIO] | list[BytesIO] = [], + ) -> Generator[str, Any, None]: + try: + version = parameters.pop('model') + model_owner = parameters.pop('model_owner') + official = version and model_owner + payload = {'prompt': content, **parameters} + if not official: + payload.update({'model_version': version}) + with httpx.Client( + base_url="https://fal.run/", + headers={"Authorization": f"Key {settings.FAL_API_KEY}"}, + timeout=600, + ) as client: + with client.stream( + 'POST', f'{model_owner}/{version}/stream', json=payload, headers={'Accept': 'text/event-stream'} + ) as stream: + for chunk in stream.iter_text(): + if 'images' in chunk: + yield f'\n{chunk.split('"')[-1]}' + elif '"' in chunk: + yield chunk.split('\"')[0] + else: + yield chunk + except httpx.TimeoutException as exc: + logger.exception(exc) + raise Exception('Timeout happened') @@ -128,20 +128,26 @@ class MediaAPIView(APIView): serializer = MessageSerializer(data=request.data) try: serializer.is_valid(raise_exception=True) - chat = self.manager.objects.prefetch_related('model', 'model__inferences').get( - user=request.user, model__slug=model + gallery, _ = self.manager.objects.get_or_create( + user=request.user, + model__slug=model, + defaults={ + 'user': request.user, + 'model': NeuronModel.objects.get(slug=model), + }, ) + gallery = self.manager.objects.prefetch_related('model', 'model__inferences').get(pk=gallery.pk) info = serializer.validated_data.pop('info') inference_slug = info.pop('inference') if not any( [ inference.slug == inference_slug and inference.enabled - for inference in chat.model.inferences.all() + for inference in gallery.model.inferences.all() ] ): raise Exception(_('Enabled inference not found in model')) input_message = Message.objects.create( - **serializer.validated_data, info=info, content_object=chat, from_model=False + **serializer.validated_data, info=info, content_object=gallery, from_model=False ) except Exception as exc: return Response({'detail': str(exc)}, status=400) @@ -154,7 +160,7 @@ class MediaAPIView(APIView): history_ids=[ id async for id in Message.objects.filter( - object_id=chat.uid, content_type__model='chat', is_deleted=False, is_sent=True + object_id=gallery.uid, content_type__model='chat', is_deleted=False, is_sent=True )[:10].values_list('uid', flat=True) ], ) @@ -171,7 +177,8 @@ class MediaAPIView(APIView): if chunk: content += chunk yield f'id: {output_slot_id}\nevent: output\ndata: {"".join(chunk).replace("\n", "\\n")}\n\n' - + else: + await asyncio.sleep(0.1) yield f'id: {output_slot_id}\nevent: done\ndata: [DONE]\n\n' except Exception as exc: input_message.is_sent = False