@@ -7,6 +7,7 @@ from ml_model.services.deepseek import Deepseek from ml_model.services.djourney import Djourney from ml_model.services.epicphotogasm import Epicphotogasm from ml_model.services.flux import Flux +from ml_model.services.fluxproultra import Fluxproultra from ml_model.services.granite import Granite from ml_model.services.grok import Grok from ml_model.services.iconic import Iconic @@ -17,7 +17,7 @@ class Dalle(SimpleService): contains abstract method make, which makes a generation """ - PRICE = Decimal('1') + PRICE = Decimal('2') _CALLBACK = ( 'bytedance/sdxl-lightning-4step:5599ed30703defd1d160a25a63321b4dec97101d98b4674bcc56e41f62f35637' @@ -1,14 +1,11 @@ -import base64 import time from datetime import timedelta from decimal import Decimal from io import BytesIO -import filetype import requests from django.core.files import File -from backend import settings from messages.models import Message from ml_model.models import ( ModelCategory, @@ -33,39 +30,17 @@ class Flux(SimpleService): category = ModelCategory(title='Изображения', slug='images') versions = [ ModelVersion(name='Flux-Schnell', slug='flux-schnell'), - 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'), ] inputs = [ ModelInput(type=ModelInput.TypeChoices.TEXT, required=True), - ModelInput(type=ModelInput.TypeChoices.IMAGE), ] parameters = [ ModelParameter( - name='Ширина', - key='width', - type=ModelParameter.TypeChoices.INTRANGE, - values={'start': 256, 'end': 1440, 'step': 32, 'default': 1024}, - ), - ModelParameter( - name='Высота', - key='height', - type=ModelParameter.TypeChoices.INTRANGE, - values={'start': 256, 'end': 1440, 'step': 32, 'default': 768}, - ), - ModelParameter( - name='Ускорить', + name='Ускорение генерации', key='go_fast', type=ModelParameter.TypeChoices.BOOL, values={'default': True}, ), - ModelParameter( - name='Приближенность к запросу', - key='guidance', - type=ModelParameter.TypeChoices.INTRANGE, - values={'start': 0, 'end': 10, 'step': 1, 'default': 3}, - ), ModelParameter( name='Мегапиксели', key='megapixels', @@ -85,13 +60,7 @@ class Flux(SimpleService): values={'start': 1, 'end': 4, 'step': 1, 'default': 1}, ), ModelParameter( - name='Количество шагов вывода', - key='num_inference_steps', - type=ModelParameter.TypeChoices.INTRANGE, - values={'start': 1, 'end': 50, 'step': 1, 'default': 28}, - ), - ModelParameter( - name='Количество шагов вывода', + name='Количество шагов обработки', key='num_inference_steps', type=ModelParameter.TypeChoices.INTRANGE, values={'start': 1, 'end': 4, 'step': 1, 'default': 4}, @@ -118,94 +87,34 @@ class Flux(SimpleService): }, ), ModelParameter( - name='Качество вывода', + name='Качество вывода (в %)', key='output_quality', type=ModelParameter.TypeChoices.INTRANGE, values={'start': 0, 'end': 100, 'step': 1, 'default': 80}, ), - ModelParameter( - name='Апсемплинг', - key='prompt_upsampling', - type=ModelParameter.TypeChoices.BOOL, - values={'default': False}, - ), - ModelParameter( - name='Отключить обработку', - key='raw', - type=ModelParameter.TypeChoices.BOOL, - values={'default': False}, - ), ] payments_rules = { versions[0].slug: ModelPaymentRule( strategy=ModelPaymentRule.StrategyChoices.FIXED, interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, cost=0.3, - coefficient=5.00, - ), - versions[1].slug: ModelPaymentRule( - strategy=ModelPaymentRule.StrategyChoices.FIXED, - interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, - cost=4.00, - coefficient=2.00, - ), - versions[2].slug: ModelPaymentRule( - strategy=ModelPaymentRule.StrategyChoices.FIXED, - interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, - 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, + coefficient=10.00, ), } - _CALLBACK_BASE = 'black-forest-labs/' @property def neuron_model(self): return NeuronModel.objects.get(title='Flux') - def _call_bfl_api(self, payload: dict) -> list: - bfl_headers = { - 'Content-Type': 'application/json', - 'X-Key': settings.FLUX_API_KEY, - } - bfl_urls = { - 'generate': 'https://api.bfl.ml/v1/', - 'get': 'https://api.bfl.ml/v1/get_result?id=', - } - response = requests.post( - url=f'{bfl_urls["generate"]}{payload["version"]}', - headers=bfl_headers, - json=payload, - ) - if response.status_code != 200: - raise Exception(response.json()) - result = requests.get( - url=f'{bfl_urls["get"]}{response.json().get("id")}', - headers=bfl_headers, - ) - while result.json()['status'] not in ('Ready', 'Error'): - result = requests.get( - url=f'{bfl_urls["get"]}{response.json().get("id")}', - headers=bfl_headers, - ) - return result.json()['result']['sample'] - def calculate_price(self, input_message: Message) -> Decimal: - version_slug = input_message.info.get('version', 'flux-pro-1.1') + version_slug = input_message.info.get('version', 'flux-schnell') payment_rule = self.payments_rules[version_slug] if not payment_rule.pk: payment_rule.model = self.neuron_model payment_rule.save() - if version_slug in (self.versions[1].slug,): - return payment_rule.rate * input_message.info.get('num_outputs', 1) - else: - return payment_rule.rate + return payment_rule.rate * input_message.info.get('num_outputs', 1) + def save_results( self, @@ -236,21 +145,11 @@ class Flux(SimpleService): **input_message.info, } ) - if input_message.file: - kind = filetype.guess(input_message.file.read(20)) - mime = kind.mime if kind else 'application/octet-stream' - input_message.file.seek(0) - image = f'data:{mime};base64,{base64.b64encode(input_message.file.read()).decode("utf-8")}' - input_message.file.close() - callback_data.update({'image': image}) - if callback_data['version'] == 'flux-pro-1.1': - images = [self._call_bfl_api(payload=callback_data)] - else: - runner = replicate_run( - f'{self._CALLBACK_BASE}{callback_data["version"]}', - callback_data, - ) - images = runner if isinstance(runner, list) else [runner] + runner = replicate_run( + f'{self._CALLBACK_BASE}{callback_data.get('version', 'flux-schnell')}', + callback_data, + ) + images = runner if isinstance(runner, list) else [runner] process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice(input_message.content_object.model, input_message) msgs = self.save_results(input_message.content, images, process_time, save) @@ -0,0 +1,244 @@ +import base64 +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO + +import filetype +import requests +from django.core.files import File + +from backend import settings +from messages.models import Message +from ml_model.models import ( + ModelCategory, + ModelInput, + ModelParameter, + ModelPaymentRule, + ModelVersion, + NeuronModel, +) +from ml_model.services.base import SimpleService +from ml_model.tasks import replicate_run + + +class Fluxproultra(SimpleService): + """ + Flux Service + contains abstract method make, which makes a generation + """ + + title = 'Flux Pro Ultra' + description = 'Нейросеть, способная генерировать картинки из вашего текста' + category = ModelCategory(title='Изображения', slug='images') + versions = [ + 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'), + ] + inputs = [ + ModelInput(type=ModelInput.TypeChoices.TEXT, required=True), + ModelInput(type=ModelInput.TypeChoices.IMAGE), + ] + parameters = [ + ModelParameter( + name='Ширина', + key='width', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 256, 'end': 1440, 'step': 32, 'default': 1024}, + ), + ModelParameter( + name='Высота', + key='height', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 256, 'end': 1440, 'step': 32, 'default': 768}, + ), + ModelParameter( + name='Ускорение генерации', + key='go_fast', + type=ModelParameter.TypeChoices.BOOL, + values={'default': True}, + ), + ModelParameter( + name='Приближенность к запросу', + key='guidance', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 0, 'end': 10, 'step': 1, 'default': 3}, + ), + ModelParameter( + name='Мегапиксели', + key='megapixels', + type=ModelParameter.TypeChoices.LIST, + values={ + 'availables': [ + '1', + '0.25', + ], + 'default': '1', + }, + ), + ModelParameter( + name='Количество изображений', + key='num_outputs', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 1, 'end': 4, 'step': 1, 'default': 1}, + ), + ModelParameter( + name='Количество шагов вывода', + key='num_inference_steps', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 1, 'end': 50, 'step': 1, 'default': 28}, + ), + ModelParameter( + name='Соотношение сторон', + key='aspect_ratio', + type=ModelParameter.TypeChoices.LIST, + values={ + 'availables': [ + '1:1', + '16:9', + '21:9', + '3:2', + '2:3', + '4:5', + '5:4', + '3:4', + '4:3', + '9:16', + '9:21', + ], + 'default': '1:1', + }, + ), + ModelParameter( + name='Качество вывода (в %)', + key='output_quality', + type=ModelParameter.TypeChoices.INTRANGE, + values={'start': 0, 'end': 100, 'step': 1, 'default': 80}, + ), + ModelParameter( + name='Апсемплинг', + key='prompt_upsampling', + type=ModelParameter.TypeChoices.BOOL, + values={'default': False}, + ), + ModelParameter( + name='Отключить пост-обработку', + key='raw', + type=ModelParameter.TypeChoices.BOOL, + values={'default': False}, + ), + ] + payments_rules = { + versions[0].slug: ModelPaymentRule( + strategy=ModelPaymentRule.StrategyChoices.FIXED, + interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, + cost=4.00, + coefficient=2.00, + ), + versions[1].slug: ModelPaymentRule( + strategy=ModelPaymentRule.StrategyChoices.FIXED, + interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, + cost=2.5, + coefficient=2.00, + ), + versions[2].slug: ModelPaymentRule( + strategy=ModelPaymentRule.StrategyChoices.FIXED, + interaction_type=ModelPaymentRule.InteractionTypeChoices.INPUT, + cost=6.00, + coefficient=2.00, + ), + } + + _CALLBACK_BASE = 'black-forest-labs/' + + @property + def neuron_model(self): + return NeuronModel.objects.get(title='Flux Pro Ultra') + + def _call_bfl_api(self, payload: dict) -> list: + bfl_headers = { + 'Content-Type': 'application/json', + 'X-Key': settings.FLUX_API_KEY, + } + bfl_urls = { + 'generate': 'https://api.bfl.ml/v1/', + 'get': 'https://api.bfl.ml/v1/get_result?id=', + } + response = requests.post( + url=f'{bfl_urls["generate"]}{payload["version"]}', + headers=bfl_headers, + json=payload, + ) + if response.status_code != 200: + raise Exception(response.json()) + result = requests.get( + url=f'{bfl_urls["get"]}{response.json().get("id")}', + headers=bfl_headers, + ) + while result.json()['status'] not in ('Ready', 'Error'): + result = requests.get( + url=f'{bfl_urls["get"]}{response.json().get("id")}', + headers=bfl_headers, + ) + return result.json()['result']['sample'] + + def calculate_price(self, input_message: Message) -> Decimal: + version_slug = input_message.info.get('version', 'flux-pro-1.1') + payment_rule = self.payments_rules[version_slug] + if not payment_rule.pk: + payment_rule.model = self.neuron_model + payment_rule.save() + if version_slug in (self.versions[1].slug,): + return payment_rule.rate * input_message.info.get('num_outputs', 1) + else: + return payment_rule.rate + + 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, + } + ) + if input_message.file: + kind = filetype.guess(input_message.file.read(20)) + mime = kind.mime if kind else 'application/octet-stream' + input_message.file.seek(0) + image = f'data:{mime};base64,{base64.b64encode(input_message.file.read()).decode("utf-8")}' + input_message.file.close() + callback_data.update({'image': image}) + if callback_data['version'] == 'flux-pro-1.1': + images = [self._call_bfl_api(payload=callback_data)] + else: + runner = replicate_run( + f'{self._CALLBACK_BASE}{callback_data["version"]}', + callback_data, + ) + images = runner if isinstance(runner, list) else [runner] + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, input_message) + msgs = self.save_results(input_message.content, images, process_time, save) + return msgs