@@ -25,13 +25,22 @@ class NeuronModelSelector: return NeuronModelSerializer(models, many=True) return models - def get_models_by_output_content_type(self, serialize: bool = False, hidden: bool = False): + def get_models_by_output_content_type( + self, output_content_type: str = None, serialize: bool = False, hidden: bool = False + ): + categories = { + 'text': 'chat-bots', + 'image': 'images', + 'audio': 'audio', + 'video': 'videos', + 'code': 'code' + } models = NeuronModel.objects.prefetch_related( Prefetch( 'model_modelparameters', queryset=ModelParameter.objects.filter(hidden=hidden), ) - ).all() + ).filter(category__slug=categories[output_content_type]) if serialize: return NeuronModelSerializer(models, many=True) return models @@ -65,6 +65,30 @@ class NeuronModelSerializer(serializers.ModelSerializer): exclude = ('created_at', 'updated_at', 'order', 'category', 'alternative_titles') +class PublicModelParameterSerializer(ModelParameterSerializer): + class Meta: + model = ModelParameter + exclude = ('id', 'model', 'hidden', 'order') + + +class PublicModelVersionSerializer(ModelVersionSerializer): + class Meta: + model = ModelVersion + exclude = ('id', 'model', 'order') + + +class PublicNeuronModelSerializer(NeuronModelSerializer): + parameters = PublicModelParameterSerializer(many=True) + versions = PublicModelVersionSerializer(many=True) + + class Meta: + model = NeuronModel + fields = ( + 'title', 'description', 'slug', 'blocked', 'versions', + 'inputs', 'parameters' + ) + + class NeuronModelsSerializer(serializers.ModelSerializer): blocked = serializers.BooleanField() tags = ModelTagSerializer(many=True) @@ -16,6 +16,7 @@ from messages.serializers import MessageSerializer from ml_model.choices import ContentTypes from ml_model.models import NeuronModel from ml_model.selectors.ml_models_selector import NeuronModelSelector +from ml_model.serializers import PublicNeuronModelSerializer from ml_model.services.base import SimpleService from payments.exceptions.insufficient_balance import InsufficientBalance from tools.public_api.models import APIKey, APIStore @@ -42,12 +43,14 @@ class BaseGenerationView(APIView): def get(self, request, *args, **kwargs): """List available generative Models for output content type: text, image, audio, video, or code.""" - return Response( - model['slug'] - for model in NeuronModelSelector(request.user) - .get_models_by_output_content_type(serialize=True) - .data + models = ( + NeuronModelSelector(request.user) + .get_models_by_output_content_type( + serialize=False, + output_content_type=self.output_content_type, + ) ) + return Response(PublicNeuronModelSerializer(models, many=True).data) def post(self, request, model_slug, *args, **kwargs): """Create new content. Type of content depends on model output content type: