@@ -1,3 +1,5 @@ +import logging + from django.conf import settings from django.conf.urls.static import static from django.contrib import admin @@ -16,9 +18,12 @@ api.add_router('users/', 'authentication.routes.v1.router') api.add_router('chats/', 'tools.chats.routes.v1.router') api.add_router('media/', 'tools.media.routes.v1.router') +logger = logging.getLogger(__name__) + @api.exception_handler(ObjectDoesNotExist) def object_does_not_exists_error_handler(request, exc: ObjectDoesNotExist): + logger.exception(exc) return api.create_response( request, {'message': _('Requested object does not exists')}, status=404 ) @@ -26,16 +31,19 @@ def object_does_not_exists_error_handler(request, exc: ObjectDoesNotExist): @api.exception_handler(InvalidToken) def invalid_token_error_handler(request, exc: InvalidToken): + logger.exception(exc) return api.create_response(request, {'message': _('Token is invalid')}, status=401) @api.exception_handler(InvalidPassword) def invalid_password_error_handler(request, exc: InvalidPassword): + logger.exception(exc) return api.create_response(request, {'message': _('Wrong password')}, status=401) @api.exception_handler(InvalidUsername) def invalid_username_error_handler(request, exc: InvalidUsername): + logger.exception(exc) return api.create_response(request, {'message': _('Wrong username')}, status=401) @@ -4,13 +4,12 @@ from datetime import timedelta from io import BytesIO import requests -from celery.result import AsyncResult from django.core.files import File from messages.models import Message from ml_model.models import ModelCategory, ModelInput, ModelParameter, ModelVersion from ml_model.services.base import SimpleService -from ml_model.tasks import create_d_image +from ml_model.tasks import replicate_run class Dalle(SimpleService): @@ -19,113 +18,107 @@ class Dalle(SimpleService): contains abstract method make, which makes a generation """ - TOKEN_PAYMENT_RULES = { - 'dall-e-2': { - '256x256': Decimal('10.56'), - '512x512': Decimal('11.88'), - '1024x1024': Decimal('13.2'), - }, - 'dall-e-3': { - '1024x1024': Decimal('26.4'), - '1792x1024': Decimal('35.2'), - '1024x1792': Decimal('35.2'), - }, - 'dall-e-3-hd': { - '1024x1024': Decimal('35.2'), - '1792x1024': Decimal('52.8'), - '1024x1792': Decimal('52.8'), - }, - } - title = 'Dalle' description = 'Нейросеть, способная генерировать фотографии из вашего текста' category = ModelCategory(title='Изображения', slug='images') versions = [ - ModelVersion(name='Dalle 3', slug='dall-e-3', default=True), - ModelVersion(name='Dalle 2', slug='dall-e-2'), + ModelVersion(name='Dalle 3', slug='sdxl-lightning-4step', default=True), ] inputs = [ModelInput(type=ModelInput.TypeChoices.TEXT, required=True)] parameters = [ ModelParameter( - name='Размер', - key='size', - type=ModelParameter.TypeChoices.LIST, - values={ - 'availables': ['256x256', '512x512', '1024x1024'], - 'default': '1024x1024', - }, + name='Негативный промпт', + key='negative_prompt', + type=ModelParameter.TypeChoices.STR, ), ModelParameter( - name='Размер', - key='size', - type=ModelParameter.TypeChoices.LIST, - values={ - 'availables': ['1024x1024', '1024x1792', '1792x1024'], - 'default': '1024x1024', - }, + name='Ширина', + key='width', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 1024, 'end': 1280, 'step': 256, 'default': 1024}, + ), + ModelParameter( + name='Высота', + key='height', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 1024, 'end': 1280, 'step': 256, 'default': 1024}, ), ModelParameter( - name='Качество', - key='quality', + name='Планировщик', + key='scheduler', type=ModelParameter.TypeChoices.LIST, values={ - 'availables': ['default', 'hd'], - 'default': 'hd', + 'availables': [ + "DDIM", + "DPMSolverMultistep", + "HeunDiscrete", + "KarrasDPM", + "K_EULER_ANCESTRAL", + "K_EULER", + "PNDM", + "DPM++2MSDE" + ], + 'default': 'K_EULER', }, ), ModelParameter( name='Количество изображений', - key='n', + key='num_outputs', type=ModelParameter.TypeChoices.INTRANGE, - values={'start': 1, 'end': 10, 'step': 1, 'default': 1}, + values={'start': 1, 'end': 4, 'step': 1, 'default': 1}, + ), + ModelParameter( + name='Точность запроса', + key='guidance_scale', + type=ModelParameter.TypeChoices.FLOATRANGE, + values={'start': 0, 'end': 50, 'step': 1, 'default': 0}, + ), + ModelParameter( + name='Шаги предобработки', + key='num_inference_steps', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 1, 'end': 10, 'step': 1, 'default': 4}, ), ] + PRICE = Decimal('1') + + _CALLBACK = ( + 'bytedance/sdxl-lightning-4step:5599ed30703defd1d160a25a63321b4dec97101d98b4674bcc56e41f62f35637' + ) + def calculate_price(self, input_message: Message) -> Decimal: - price = self.TOKEN_PAYMENT_RULES[input_message.info.get('version', 'dall-e-2')][ - input_message.info.get('size', '1024x1024') - ] * input_message.info.get('n', 1) + price = input_message.info.get('num_outputs', 1) * self.PRICE return price.quantize(Decimal('0.1'), rounding='ROUND_UP') def save_results( - self, input_prompt: str, r: list[str], t: timedelta, save: bool = True + self, prompt: str, images: list, time: timedelta, save: bool = True ) -> list[Message]: - out: list[Message] = [] - for obj in r: - out.append( + messages: list[Message] = [] + for image in images: + messages.append( Message( content_object=self.store, - elapsed_time=t, - content=input_prompt, - file=File( - BytesIO(requests.get(obj['url']).content), - '.png', - ), + elapsed_time=time, + content=prompt, + file=File(BytesIO(requests.get(image).content), '.png'), ) ) if save: - return Message.objects.bulk_create(out) - return out + return Message.objects.bulk_create(messages) + return messages def make(self, input_message: Message, save: bool = True) -> list[Message]: - info = input_message.info.copy() - info['model'] = info.pop('version') - if info.get('quality') == 'default': - del info['quality'] + start_time = time.time() + translated_prompt = self.translate_prompt(input_message.content) callback_data = dict( { - 'prompt': input_message.content, - **info, + 'prompt': translated_prompt, + **input_message.info, } ) - if input_message.file: - callback_data.update({'image': BytesIO(input_message.file.read())}) - start_time = time.time() - results: AsyncResult = create_d_image.delay(callback_data) - data = results.get()['data'] + images = replicate_run(self._CALLBACK, callback_data) process_time = timedelta(seconds=(time.time() - start_time)) - self.handle_invoice( - input_message.content_object.model, input_message=input_message - ) - msgs = self.save_results(input_message.content, data, process_time, save) + self.handle_invoice(input_message.content_object.model, input_message=input_message) + msgs = self.save_results(input_message.content, images, process_time, save) return msgs @@ -31,7 +31,8 @@ class Flux(SimpleService): description = 'Нейросеть, способная генерировать картинки из вашего текста' category = ModelCategory(title='Изображения', slug='images') versions = [ - ModelVersion(name='Flux-Pro1.1', slug='flux-pro-1.1', default=True), + ModelVersion(name='Flux-Schnell', slug='flux-schnell', default=True), + ModelVersion(name='Flux-Pro1.1', slug='flux-pro-1.1'), ModelVersion(name='Flux-Dev', slug='flux-dev'), ModelVersion(name='Ultra', slug='flux-1.1-pro-ultra'), ] @@ -88,6 +89,12 @@ class Flux(SimpleService): type=ModelParameter.TypeChoices.INTRANGE, values={'start': 1, 'end': 50, 'step': 1, 'default': 28}, ), + ModelParameter( + name='Количество шагов вывода', + key='num_inference_steps', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 1, 'end': 4, 'step': 1, 'default': 4}, + ), ModelParameter( name='Соотношение сторон', key='aspect_ratio', @@ -132,20 +139,26 @@ class Flux(SimpleService): versions[0].slug: ModelPaymentRule( strategy=ModelPaymentRule.StrategyChoices.FIXED, interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, - cost=4.4, + cost=0.3, coefficient=5.00, ), versions[1].slug: ModelPaymentRule( strategy=ModelPaymentRule.StrategyChoices.FIXED, interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, - cost=2.75, - coefficient=5.00, + cost=4.00, + coefficient=2.00, ), versions[2].slug: ModelPaymentRule( strategy=ModelPaymentRule.StrategyChoices.FIXED, interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, - cost=6.6, - coefficient=5.00, + cost=2.5, + coefficient=2.00, + ), + versions[3].slug: ModelPaymentRule( + strategy=ModelPaymentRule.StrategyChoices.FIXED, + interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, + cost=6.00, + coefficient=2.00, ), } @@ -135,7 +135,7 @@ class Stablediffusion(SimpleService): ), ] - TOKEN_PRICE = Decimal('11') + TOKEN_PRICE = Decimal('2') _API_KEY = settings.STABLE_DIFFUSION_API_KEY @@ -143,11 +143,11 @@ class Stablediffusion(SimpleService): if ( input_message.info.get('version') or input_message.info.get('engine', 'sd3') ) == 'sd3': - return Decimal('71.5') + return Decimal('13') elif ( input_message.info.get('version') or input_message.info.get('engine', 'sd3') ) == 'sd3-turbo': - return Decimal('44') + return Decimal('8') steps = input_message.info.get('steps', 30) default_price = Decimal('0.9') * self.TOKEN_PRICE return default_price if steps <= 30 else default_price * Decimal((steps / 30)) @@ -81,7 +81,8 @@ services: networks: default: - name: "air" + name: infrastructure + external: true volumes: static: