@@ -6,6 +6,7 @@ from ml_model.services.deepl import Deepl from ml_model.services.epicphotogasm import Epicphotogasm from ml_model.services.djourney import Djourney from ml_model.services.flux import Flux +from ml_model.services.granite import Granite from ml_model.services.iconic import Iconic from ml_model.services.kandinsky import Kandinsky from ml_model.services.llama import Llama @@ -0,0 +1,120 @@ +import time + +import requests +from datetime import timedelta +from decimal import Decimal + +from django.conf import settings + +from messages.models import Message, BaseStore +from ml_model.exceptions.external_api import ExternalAPIException +from ml_model.models import ModelCategory, ModelInput, ModelParameter +from ml_model.services.base import SimpleService + + +class Granite(SimpleService): + """ + Granite-3.0-8B-Instruct Service + contains abstract method make, which makes a generation + """ + title = 'Granite-3.0-8B-Instruct' + description = 'Нейросеть, способная генерировать качественный текст из вашего промпта' + category = ModelCategory(title='Чат-боты', slug='chat-bots') + versions = [] + inputs = [ModelInput(type=ModelInput.TypeChoices.TEXT, required=True)] + parameters = [ + ModelParameter( + name='Системный промпт', + key='system_prompt', + type=ModelParameter.TypeChoices.STR, + hidden=True + ), + ModelParameter( + name='Лучший процент', + key='top_p', + type=ModelParameter.TypeChoices.FLOATRANGE, + values={'start': 1.0, 'end': 2.0, 'step': 0.1, 'default': 0.9}, + ), + ModelParameter( + name='Температура', + key='temperature', + type=ModelParameter.TypeChoices.FLOATRANGE, + values={'start': 0.0, 'end': 1.0, 'step': 0.1, 'default': 0.6}, + ), + ] + + TOKEN_PAYMENT_RULES = { + 'granite-input': Decimal('27.5'), # 1M tokens + 'granite-output': Decimal('137.5') # 1M tokens + } + + def __init__(self, store: BaseStore) -> None: + super().__init__(store) + self.urls = { + 'generate': 'https://api.replicate.com/v1/models/ibm-granite/granite-3.0-8b-instruct/', + 'get': 'https://api.replicate.com/v1/predictions/', + } + + def _call_api(self, payload: dict) -> list: + headers = { + 'Authorization': f'Bearer {settings.REPLICATE_API_KEY}', + 'Prefer': 'wait', + } + data = {'input': payload} + response = requests.post( + url=f'{self.urls['generate']}predictions', + headers=headers, + json=data, + ) + if response.status_code != 201: + raise Exception(response.json()) + result = requests.get( + url=f'{self.urls['get']}{response.json().get('id')}', + headers=headers + ) + while result.json()['status'] not in ('succeeded', 'failed', 'canceled'): + result = requests.get( + url=f'{self.urls['get']}{response.json().get('id')}', + headers=headers + ) + return result.json() + + def calculate_price(self, result: str, input_message: Message) -> Decimal: + return Decimal( + sum( + [self.TOKEN_PAYMENT_RULES['granite-output'] / 1_000_000 * len(result.split(' '))] + + [self.TOKEN_PAYMENT_RULES['granite-input'] / 1_000_000 * len(input_message.content.split(' '))] + ) + ) + + def save_results( + self, result: str, time: timedelta, save: bool = True + ) -> list[Message]: + msgs: list[Message] = [ + Message( + content_object=self.store, + content=result, + elapsed_time=time, + ) + ] + if save: + return Message.objects.bulk_create(msgs) + return msgs + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + start_time = time.time() + callback_data = dict( + { + 'prompt': input_message.content, + 'system_prompt': f'Use only Russian language', + **input_message.info, + } + ) + try: + result = ''.join(self._call_api(payload=callback_data)['output']) + except Exception: + raise ExternalAPIException + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, result, input_message) + msgs = self.save_results(result, process_time, save) + return msgs