@@ -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,16 +139,22 @@ class Flux(SimpleService): versions[0].slug: ModelPaymentRule( strategy=ModelPaymentRule.StrategyChoices.FIXED, interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, - cost=4.4, + cost=0.33, coefficient=5.00, ), versions[1].slug: ModelPaymentRule( strategy=ModelPaymentRule.StrategyChoices.FIXED, interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, - cost=2.75, + cost=4.4, coefficient=5.00, ), versions[2].slug: ModelPaymentRule( + strategy=ModelPaymentRule.StrategyChoices.FIXED, + interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, + cost=2.75, + coefficient=5.00, + ), + versions[3].slug: ModelPaymentRule( strategy=ModelPaymentRule.StrategyChoices.FIXED, interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, cost=6.6, @@ -18,43 +18,52 @@ class Openjourney(SimpleService): contains abstract method make, which makes a generation """ - title = 'OpenJourney' + title = 'Midjourney' description = 'Нейросеть, способная генерировать фотографии из вашего текста' - price = Decimal('1.518') + price = Decimal('2') category = ModelCategory(title='Изображения', slug='images') - versions = [] inputs = [ModelInput(type=ModelInput.TypeChoices.TEXT, required=True)] parameters = [ ModelParameter( - name='Количество изображений', - key='num_outputs', - type=ModelParameter.TypeChoices.INTRANGE, - values={'start': 1, 'end': 10, 'step': 1, 'default': 1}, + name='Соотношение сторон', + key='aspect_ratio', + type=ModelParameter.TypeChoices.LIST, + values={ + 'availables': [ + "1:1", + "16:9", + "4:3", + "3:2", + "2:3", + "3:4", + "9:16", + "21:9" + ], + 'default': '1:1', + }, ), ModelParameter( - name='Количество шагов', - key='num_inference_steps', + name='Количество изображений', + key='number_of_images', type=ModelParameter.TypeChoices.INTRANGE, - values={'start': 0, 'end': 100, 'step': 1, 'default': 10}, + values={'start': 1, 'end': 9, 'step': 1, 'default': 1}, ), ModelParameter( - name='Негативный промпт', - key='negative_prompt', - type=ModelParameter.TypeChoices.STR, + name='Оптимизация промпта', + key='prompt_optimizer', + type=ModelParameter.TypeChoices.BOOL, + values={'default': True}, ), ] - _CALLBACK = ( - 'prompthero/openjourney' - ':ad59ca21177f9e217b9075e7300cf6e14f7e5b4505b87b9689dbd866e9768969' - ) + _CALLBACK = 'minimax/image-01' def __init__(self, store): super().__init__(store) - def calculate_price(self, process_time: timedelta) -> Decimal: - price = Decimal(process_time.total_seconds()) * self.price - return price.quantize(Decimal('0.1'), rounding='ROUND_UP') + def calculate_price(self, input_message: Message) -> Decimal: + price = input_message.info.get('number_of_images', 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 @@ -78,10 +87,8 @@ class Openjourney(SimpleService): translated_prompt = self.translate_prompt(input_message.content) activation_prompt = f'mdjrny-v4 style a highly detailed {translated_prompt}' callback_data = dict(prompt=activation_prompt, **input_message.info) - if input_message.file: - callback_data.update({'image': BytesIO(input_message.file.read())}) results = replicate_run(self._CALLBACK, callback_data) process_time = timedelta(seconds=(time.time() - start_time)) - self.handle_invoice(input_message.content_object.model, process_time=process_time) + self.handle_invoice(input_message.content_object.model, input_message=input_message) msgs = self.save_results(input_message.content, results, process_time, save) return msgs