@@ -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