@@ -8,9 +8,11 @@ from ml_model.services.djourney import Djourney from ml_model.services.flux import Flux from ml_model.services.kandinsky import Kandinsky from ml_model.services.llama import Llama +from ml_model.services.lightning import Lightning from ml_model.services.mistral import Mistral from ml_model.services.musicgen import Musicgen from ml_model.services.openjourney import Openjourney +from ml_model.services.recraft import Recraft from ml_model.services.sdxlemoji import Sdxlemoji from ml_model.services.stablediffusion import Stablediffusion from ml_model.services.upscaleai import Upscaleai @@ -0,0 +1,135 @@ +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO + +import requests +from django.core.files import File + +from messages.models import Message +from ml_model.exceptions.external_api import ExternalAPIException +from ml_model.models import ( + ModelCategory, + ModelInput, + ModelParameter, +) +from ml_model.services.base import SimpleService +from ml_model.tasks import replicate_run + + +class Lightning(SimpleService): + """ + Lightning Service + contains abstract method make, which makes a generation + """ + title = 'Lightning' + description = 'Нейросеть, способная генерировать картинки из вашего текста' + category = ModelCategory(title='Изображения', slug='images') + versions = [] + inputs = [ModelInput(type=ModelInput.TypeChoices.TEXT, required=True)] + parameters = [ + ModelParameter( + name='Негативный промпт', + key='negative_prompt', + type=ModelParameter.TypeChoices.STR, + ), + ModelParameter( + name='Ширина', + key='width', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 256, 'end': 1280, 'step': 128, 'default': 1024}, + ), + ModelParameter( + name='Высота', + key='height', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 256, 'end': 1280, 'step': 128, 'default': 1024}, + ), + ModelParameter( + name='Планировщик', + key='scheduler', + type=ModelParameter.TypeChoices.LIST, + values={ + 'availables': [ + 'DDIM', + 'DPMSolverMultistep', + 'HeunDiscrete', + 'KarrasDPM', + 'K_EULER_ANCESTRAL', + 'K_EULER', + 'PNDM', + 'DPM++2MSDE' + ], + 'default': 'K_EULER', + }, + ), + ModelParameter( + name='Количество изображений', + key='num_outputs', + type=ModelParameter.TypeChoices.INTRANGE, + 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}, + ), + ModelParameter( + name='Отключить проверку безопасности', + key='disable_safety_checker', + type=ModelParameter.TypeChoices.BOOL, + values={'default': False}, + ), + ] + + PRICE = Decimal('0.750') + + _CALLBACK = ( + 'bytedance/sdxl-lightning-4step' + ':5599ed30703defd1d160a25a63321b4dec97101d98b4674bcc56e41f62f35637' + ) + + def calculate_price(self, process_time: timedelta) -> Decimal: + total_seconds = Decimal(process_time.total_seconds()) + return self.PRICE * total_seconds + + def save_results( + self, prompt: str, images: list, time: timedelta, save: bool = True + ) -> list[Message]: + messages: list[Message] = [] + for image in images: + messages.append( + Message( + content_object=self.store, + elapsed_time=time, + content=prompt, + file=File(BytesIO(requests.get(image).content), '.png') + ) + ) + if save: + return Message.objects.bulk_create(messages) + return messages + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + start_time = time.time() + callback_data = dict( + { + 'prompt': self.translate_prompt(input_message.content), + **input_message.info, + } + ) + try: + images = replicate_run(self._CALLBACK, callback_data) + except Exception: + raise ExternalAPIException() + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, process_time=process_time) + msgs = self.save_results(input_message.content, images, process_time, save) + return msgs @@ -0,0 +1,180 @@ +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO + +import requests +from django.core.files import File + +from backend import settings +from messages.models import BaseStore, Message +from ml_model.exceptions.external_api import ExternalAPIException +from ml_model.models import ( + ModelCategory, + ModelInput, + ModelParameter, + ModelVersion, +) +from ml_model.services.base import SimpleService + + +class Recraft(SimpleService): + """ + Recraft Service + contains abstract method make, which makes a generation + """ + + title = 'Recraft' + description = 'Нейросеть, способная генерировать картинки из вашего текста' + category = ModelCategory(title='Изображения', slug='images') + versions = [ + ModelVersion(name='Recraft V3', slug='recraft-v3', default=True), + ModelVersion(name='Recraft V3 SVG', slug='recraft-v3-svg'), + ] + inputs = [ + ModelInput(type=ModelInput.TypeChoices.TEXT, required=True), + ] + parameters = [ + ModelParameter( + name='Ширина', + key='width', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 1024, 'end': 2048, 'step': 128, 'default': 1024}, + ), + ModelParameter( + name='Высота', + key='height', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 1024, 'end': 2048, 'step': 128, 'default': 1024}, + ), + ModelParameter( + name='Стиль', + key='style', + type=ModelParameter.TypeChoices.LIST, + values={ + 'availables': [ + 'any', + 'realistic_image', + 'digital_illustration', + 'digital_illustration/pixel_art', + 'digital_illustration/hand_drawn', + 'digital_illustration/grain', + 'digital_illustration/infantile_sketch', + 'digital_illustration/2d_art_poster', + 'digital_illustration/handmade_3d', + 'digital_illustration/hand_drawn_outline', + 'digital_illustration/engraving_color', + 'digital_illustration/2d_art_poster_2', + 'realistic_image/b_and_w', + 'realistic_image/hard_flash', + 'realistic_image/hdr', + 'realistic_image/natural_light', + 'realistic_image/studio_portrait', + 'realistic_image/enterprise', + 'realistic_image/motion_blur' + ], + 'default': 'any', + }, + ), + ModelParameter( + name='Стиль', + key='style', + type=ModelParameter.TypeChoices.LIST, + values={ + 'availables': [ + 'any', + 'engraving', + 'line_art', + 'line_circuit', + 'linocut', + ], + 'default': 'any', + }, + ), + ] + + payment_rules = { + versions[0].slug: Decimal('20'), + versions[1].slug: Decimal('40'), + } + + def __init__(self, store: BaseStore) -> None: + super().__init__(store) + self.generate_url = 'https://api.replicate.com/v1/models/recraft-ai/' + self.get_url = 'https://api.replicate.com/v1/predictions/' + + def _get_size(self, width: int, height: int) -> str: + available_sizes = ( + (1024, 1024), (1365, 1024), (1024, 1365), (1536, 1024), (1024, 1536), + (1820, 1024), (1024, 1820), (1024, 2048), (2048, 1024), (1434, 1024), + (1024, 1434), (1024, 1280), (1280, 1024), (1024, 1707), (1707, 1024) + ) + if width >= height: + size = min(available_sizes, key=lambda size: abs(width - size[0])) + else: + size = min(available_sizes, key=lambda size: abs(height - size[1])) + return f'{size[0]}x{size[1]}' + + def _call_api(self, payload: dict) -> list: + try: + headers = { + 'Authorization': f'Bearer {settings.REPLICATE_API_KEY}', + 'Content-Type': 'application/json', + 'Prefer': 'wait', + } + data = {'input': payload} + response = requests.post( + url=f'{self.generate_url}{payload['version']}/predictions', + headers=headers, + json=data, + ) + if response.status_code != 201: + raise Exception(response.json()) + + result = requests.get(url=f'{self.get_url}{response.json().get('id')}', headers=headers) + while result.json()['status'] not in ('succeeded', 'failed', 'canceled'): + result = requests.get(url=f'{self.get_url}{response.json().get('id')}', headers=headers) + return result.json()['output'] + except Exception: + raise ExternalAPIException + + def calculate_price(self, input_message: Message) -> Decimal: + return self.payment_rules[input_message.info.get('version', 'recraft-v3')] + + def save_results( + self, + prompt: str, + image: str, + extension: str, + time: timedelta, + save: bool = True, + ) -> list[Message]: + messages: list[Message] = [] + messages.append( + Message( + content_object=self.store, + elapsed_time=time, + content=prompt, + file=File(BytesIO(requests.get(image).content), extension), + ) + ) + if save: + return Message.objects.bulk_create(messages) + return messages + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + start_time = time.time() + extension = '.svg' if input_message.info['version'] == self.versions[1].slug else '.png' + size = self._get_size(input_message.info.pop('width', 1024), input_message.info.pop('height', 1024)) + callback_data = dict( + { + 'prompt': self.translate_prompt(input_message.content), + 'size': size, + **input_message.info + } + ) + image = self._call_api(payload=callback_data) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, input_message) + message = self.save_results(input_message.content, image, extension, process_time, save) + return message