@@ -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') @@ -129,9 +129,15 @@ class MediaAPIView(APIView): serializer = MessageSerializer(data=request.data) try: serializer.is_valid(raise_exception=True) - gallery = 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( @@ -165,7 +171,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