@@ -1,50 +1,62 @@ import base64 import logging +import time import re import subprocess -import time import zipfile -from datetime import timedelta -from decimal import Decimal -from io import BytesIO, StringIO -from math import ceil -from typing import Any, Callable, Iterable, Literal -from uuid import UUID +from dataclasses import dataclass, field import docx2txt -import openpyxl import filetype -from django.core.cache import cache -from django.core.files.base import ContentFile -from django.db.models import Prefetch +import openpyxl + +from datetime import timedelta +from io import StringIO, BytesIO +from typing import Iterable, Any, Callable, Literal, Optional, Union from django.utils.translation import gettext_lazy as _ +from uuid import UUID from PIL import Image as ImageModule from PyPDF2 import PdfReader - -from authentication.models.user import CustomUserModel -from messages.models import Message from ml_model.exceptions import ( InferenceDisabled, - PaymentRuleNotImplemented, ScraperDoesNotExists, UnknownFileException, ) + + +from django.core.cache import cache +from django.core.files.base import ContentFile +from django.db.models import Prefetch + +from authentication.models.user import CustomUserModel +from messages.models import Message + from ml_model.models import ( Deployment, Inference, OverridenParameter, Parameter, - PaymentBias, - PaymentRule, TrackingRecord, ) -from ml_model.tools.tokenizer import TokenizerTool +from ml_model.services.pricing import PricingService from payments.exceptions.insufficient_balance import InsufficientBalance from payments.services.payment_plan_service import PaymentPlanService logger = logging.getLogger(__name__) +@dataclass +class ProcessedInput: + inference: 'Inference' + content: Optional[str] = None + file: Optional[Union[StringIO, BytesIO]] = None + normalized_image: Optional[Any] = None + file_extension: Optional[str] = None + parameters: dict = field(default_factory=dict) + history: Iterable[Any] = field(default_factory=list) + scrape_results: list = field(default_factory=list) + + class InferenceService: def __init__(self, user: CustomUserModel): self.user = user @@ -79,273 +91,179 @@ class InferenceService: 'deployment__deployment_payment_rules', ).get(slug=slug) - def run( - self, - slug: str, - input_message: Message, - output_slot: Message, - history: Iterable[Message] = Message.objects.none(), - ): - cache_key = f'{input_message.content_object._meta.model_name}s:{input_message.content_object.uid}' - try: - inference = Inference.objects.prefetch_related( - 'inference_parameters', - 'inference_parameters__parameter', - 'inference_payment_biases', - Prefetch( - 'inference_tracking_records', queryset=TrackingRecord.objects.order_by('-created_at') - ), - 'deployment', - 'deployment__scraper_config', - 'deployment__deployment_inputs', - 'deployment__deployment_parameters', - 'deployment__deployment_payment_rules', - ).get(slug=slug) + @staticmethod + def _load_inference(slug: str) -> Inference: + return Inference.objects.prefetch_related( + 'inference_parameters', + 'inference_parameters__parameter', + 'inference_payment_biases', + Prefetch( + 'inference_tracking_records', queryset=TrackingRecord.objects.order_by('-created_at') + ), + 'deployment', + 'deployment__scraper_config', + 'deployment__deployment_inputs', + 'deployment__deployment_parameters', + 'deployment__deployment_payment_rules', + ).get(slug=slug) - if not inference.enabled or not inference.deployment.enabled: - raise InferenceDisabled - if not inference.deployment.payment_rules: - raise Exception(_('Payment rules are missing; Inference: %s' % (inference.name))) - if not inference.payment_biases: - logger.warning('Payment biases are missing; Inference: %s' % (inference.name)) + def preprocess(self, slug: str, input_message: Message, history: Iterable[Message]) -> ProcessedInput: + inference = self._load_inference(slug) - file = None - raw_file = input_message.file - raw_info = input_message.info.copy() - scrape_results = [] + if not inference.enabled or not inference.deployment.enabled: + raise InferenceDisabled + if not inference.deployment.payment_rules: + raise Exception(_('Payment rules are missing; Inference: %s' % (inference.name))) + file: StringIO | BytesIO | None = None + normalized_image = None + file_extension: str | None = None + raw_file = input_message.file + raw_info = input_message.info.copy() + scrape_results: list[StringIO] | list[BytesIO] = [] + + if raw_file: try: file_header = raw_file.read(50) - file_extension = filetype.guess(file_header).extension + ft = filetype.guess(file_header) raw_file.seek(0) file_buf = BytesIO(raw_file.read()) - - if file_extension == 'zip': - file_extension = None - signatures = { - 'xlsx': 'xl/workbook.xml', - 'docx': 'word/document.xml' - } - with zipfile.ZipFile(file_buf, 'r') as zip_file: - namelist = zip_file.namelist() - for format_name, required_file in signatures.items(): - if required_file in namelist: - file_extension = format_name - break - if not file_extension: - raise UnknownFileException - if file_extension in ('png', 'jpg', 'jpeg'): - file = BytesIO() - normalized_image = ImageModule.open(file_buf) - normalized_image.save(file, format='jpeg' if file_extension == 'jpg' else file_extension) - file.seek(0) - elif file_extension in ('pdf',): - file = StringIO() - reader = PdfReader(file_buf) - for page in reader.pages: - file.write(page.extract_text()) - elif file_extension in ('doc', 'docx'): - extractors: dict[Literal['doc', 'docx'], Callable[[], str]] = { - 'doc': lambda: subprocess.Popen( - ['antiword', '-w', '0', '-'], - stdin=subprocess.PIPE, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - ) - .communicate(file_buf.getvalue()) - .decode(), - 'docx': lambda: docx2txt.process(file_buf), - } - file = StringIO() - file.write('Remember this Document included in request:') - file.write('[DOCUMENT-START]\n') - file.write(extractors[file_extension]()) - file.write('\n[DOCUMENT-END]') - elif file_extension in ('xlsx',): - file = StringIO() - xlsx_file = openpyxl.load_workbook(file_buf) - for sheet_name in xlsx_file.sheetnames: - sheet = xlsx_file[sheet_name] - for row in sheet.iter_rows(values_only=True): - file.write(f'Данные ряда: {row}') - + if ft: + file_extension = ft.extension + if file_extension == 'zip': + file_extension = None + signatures = { + 'xlsx': 'xl/workbook.xml', + 'docx': 'word/document.xml' + } + with zipfile.ZipFile(file_buf, 'r') as zip_file: + namelist = zip_file.namelist() + for format_name, required_file in signatures.items(): + if required_file in namelist: + file_extension = format_name + break + if not file_extension: + raise UnknownFileException + if file_extension in ('png', 'jpg', 'jpeg'): + file = BytesIO() + normalized_image = ImageModule.open(file_buf) + normalized_image.save(file, format='jpeg' if file_extension == 'jpg' else file_extension) + file.seek(0) + elif file_extension in ('pdf',): + file = StringIO() + reader = PdfReader(file_buf) + for page in reader.pages: + text = page.extract_text() or '' + file.write(text) + elif file_extension in ('doc', 'docx'): + extractors: dict[Literal['doc', 'docx'], Callable[[], str]] = { + 'doc': lambda: subprocess.Popen( + ['antiword', '-w', '0', '-'], + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + .communicate(file_buf.getvalue())[0] + .decode(), + 'docx': lambda: docx2txt.process(file_buf), + } + file = StringIO() + file.write('Remember this Document included in request:') + file.write('[DOCUMENT-START]\n') + file.write(extractors[file_extension]()) + file.write('\n[DOCUMENT-END]') + elif file_extension in ('xlsx',): + file = StringIO() + xlsx_file = openpyxl.load_workbook(file_buf) + for sheet_name in xlsx_file.sheetnames: + sheet = xlsx_file[sheet_name] + for row in sheet.iter_rows(values_only=True): + file.write(f'Данные ряда: {row}') except Exception: file = None - logger.info(f'File: {file}') - - if isinstance(file, StringIO): - file_data = file.getvalue() - cleaned_data = re.sub(r'\n{2,}', '\n', file_data) - file.seek(0) - file.truncate() - file.write(cleaned_data) - file.seek(0) - - # first round: create initial params - parameters: dict[str, Any] = {} + if isinstance(file, StringIO): + file_data = file.getvalue() + cleaned_data = re.sub(r'\n{2,}', '\n', file_data) + file.seek(0) + file.truncate() + file.write(cleaned_data) + file.seek(0) + + parameters: dict[str, Any] = {} + for parameter in inference.deployment.parameters: + parameters.update({parameter.key: parameter.values['default']}) + for overriden_parameter in inference.overriden_parameters: + parameters.update({overriden_parameter.parameter.key: overriden_parameter.value}) + for intersected_param_key in parameters.keys() & raw_info.keys(): for parameter in inference.deployment.parameters: - parameters.update({parameter.key: parameter.values['default']}) - - # second round: override initial - for overriden_parameter in inference.overriden_parameters: - parameters.update({overriden_parameter.parameter.key: overriden_parameter.value}) - - # third round: override with incoming params - for intersected_param_key in parameters.keys() & raw_info.keys(): - for parameter in inference.deployment.parameters: - if intersected_param_key == parameter.key and not parameter.hidden: - if ( + if intersected_param_key == parameter.key and not parameter.hidden: + if ( parameter.type in ( - Parameter.TypeChoices.FLOAT, - Parameter.TypeChoices.INT, - Parameter.TypeChoices.STR, - Parameter.TypeChoices.BOOL, - ) + Parameter.TypeChoices.FLOAT, + Parameter.TypeChoices.INT, + Parameter.TypeChoices.STR, + Parameter.TypeChoices.BOOL, + ) or ( - parameter.type - in (Parameter.TypeChoices.FLOATRANGE, Parameter.TypeChoices.INTRANGE) - and parameter.values['start'] - <= raw_info[parameter.key] - <= parameter.values['stop'] - ) + parameter.type + in (Parameter.TypeChoices.FLOATRANGE, Parameter.TypeChoices.INTRANGE) + and parameter.values['start'] + <= raw_info[parameter.key] + <= parameter.values['stop'] + ) or parameter.type in (Parameter.TypeChoices.CHOICES,) - and len( - [ - real - for real, _ in parameter.values['availables'] - if real == raw_info[parameter.key] - ] - ) + and len([ + real for real, _ in parameter.values['availables'] if real == raw_info[parameter.key] + ]) > 0 - ): - parameters.update({parameter.key: raw_info[parameter.key]}) - - if parameters.pop('use_scraping', None): - if not inference.deployment.scraper_config: - raise ScraperDoesNotExists - scrape_results = inference.deployment.scraper_config.scraper.scrape(input_message) - - # TODO: caching three rounds calculation - - calculated_price, expected_additional_price = Decimal('0'), Decimal('0') - content = input_message.content + ): + parameters.update({parameter.key: raw_info[parameter.key]}) + + if parameters.pop('use_scraping', None): + if not inference.deployment.scraper_config: + raise ScraperDoesNotExists + scrape_results = inference.deployment.scraper_config.scraper.scrape(input_message) + + return ProcessedInput( + inference=inference, + content=input_message.content, + file=file, + normalized_image=normalized_image, + file_extension=file_extension, + parameters=parameters, + history=history, + scrape_results=scrape_results, + ) - # first round: predict main price - for payment_rule in inference.deployment.payment_rules: - if ( - payment_rule.strategy == PaymentRule.StrategyChoices.FIXED - and payment_rule.interaction_type is None - and payment_rule.content_type is None - ): - calculated_price += payment_rule.cost - elif ( - payment_rule.strategy == PaymentRule.StrategyChoices.PER_SECOND - and payment_rule.interaction_type is None - and payment_rule.content_type is None - ): - if len(inference.tracking_records) < 1: - raise Exception(_('No tracking records found')) - last_tracking_record = inference.tracking_records[0] - average_generation_time = last_tracking_record.generation_time - expected_additional_price += payment_rule.cost * Decimal( - average_generation_time.total_seconds() - ) - elif ( - payment_rule.strategy == PaymentRule.StrategyChoices.FIXED - and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.INPUT - and payment_rule.content_type == PaymentRule.ContentTypeChoices.FILE - ): - if input_message.file: - calculated_price += payment_rule.cost - elif ( - payment_rule.strategy == PaymentRule.StrategyChoices.PER_TEXT_TOKEN - and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.INPUT - and payment_rule.content_type == PaymentRule.ContentTypeChoices.EMBEDDINGS - ): - if isinstance(file, StringIO) and len(file_data := file.getvalue()) >= 20_000: - calculated_price += TokenizerTool.token_count( - file_data - ) * payment_rule.cost - elif ( - content - and payment_rule.strategy == PaymentRule.StrategyChoices.PER_TEXT_TOKEN - and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.INPUT - ): - calculated_price += TokenizerTool.token_count(content) * payment_rule.cost - if isinstance(file, StringIO) and len(file_data := file.getvalue()) < 20_000: - calculated_price += TokenizerTool.token_count( - file_data - ) * payment_rule.cost - elif isinstance(file, StringIO): - calculated_price += Decimal(f'{(110 + 100 + 10 * 4000) // 3}') * payment_rule.cost - if history: - calculated_price += ( - sum([TokenizerTool.token_count(message.content) for message in history if message.content]) - * payment_rule.cost - ) - if scrape_results: - calculated_price += ( - sum([TokenizerTool.token_count(result.getvalue()) for result in scrape_results]) - ) * payment_rule.cost - elif ( - payment_rule.strategy == PaymentRule.StrategyChoices.PER_TEXT_TOKEN - and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.OUTPUT - ): - max_tokens_size = Decimal(parameters.get('max_tokens', '128000')) - expected_additional_price += max_tokens_size * payment_rule.cost - elif ( - payment_rule.strategy == PaymentRule.StrategyChoices.PER_PIXEL - and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.INPUT - and payment_rule.content_type == PaymentRule.ContentTypeChoices.FILE - ): - if file and file_extension in ('png', 'jpg', 'jpeg') and normalized_image: - width, height = normalized_image.width, normalized_image.height - if max(width, height) > 2048: - a_ratio = width / height - width, height = ( - (2048, int(2048 / a_ratio)) if a_ratio > 1 else (int(2048 * a_ratio), 2048) - ) - if width >= height and height > 768: - width, height = int((768 / height) * width), 768 - elif height > width and width > 768: - width, height = 768, int((768 / width) * height) - tiles_size = ceil(width / 512) * ceil(height / 512) - calculated_price += tiles_size * payment_rule.cost - else: - raise PaymentRuleNotImplemented - - expected_price = calculated_price + expected_additional_price - if self.user.balance < expected_price: - raise InsufficientBalance(self.user.balance, expected_price) - - logger.debug( - f'Predicted price before biasing: {expected_price}; Calculated: {calculated_price}; Additional: {expected_additional_price}; Message: {output_slot.pk}' + def run( + self, + slug: str, + input_message: Message, + output_slot: Message, + history: Iterable[Message] = Message.objects.none(), + ): + cache_key = f'{input_message.content_object._meta.model_name}s:{input_message.content_object.uid}' + try: + processed = self.preprocess(slug, input_message, history) + calculated_price, additional_price = PricingService.compute_pre_cost(processed) + predicted_price = PricingService.apply_biases( + processed.inference, + calculated_price + additional_price ) + if self.user.balance < predicted_price: + raise InsufficientBalance(self.user.balance, predicted_price) - # second round: predict bias price - for payment_bias in inference.payment_biases: - if payment_bias.type == PaymentBias.TypeChoices.ADDITION: - expected_price += payment_bias.coefficient - elif payment_bias.type == PaymentBias.TypeChoices.MULTIPLICATION: - expected_price *= payment_bias.coefficient - - if self.user.balance < expected_price: - raise InsufficientBalance(self.user.balance, expected_price) - - logger.debug(f'Predicted price after biasing: {expected_price}; Message: {output_slot.pk}') - # if sufficient - reserve tokens - - runner_cls = inference.deployment.runner + runner_cls = processed.inference.deployment.runner start = time.time() activated_generation = runner_cls.generate( - content=content, - file=file, - parameters=parameters, - history=history, - scrape_results=scrape_results, + content=processed.content, + file=processed.file, + parameters=processed.parameters, + history=processed.history, + scrape_results=processed.scrape_results, ) output_content = '' @@ -357,36 +275,28 @@ class InferenceService: ) end = time.time() - process_time = timedelta(seconds=end - start) + total_price = PricingService.apply_biases( + processed.inference, + PricingService.finalize_price( + calculated_price, processed.inference, output_content, process_time) + ) + output_content = output_content.split('base64,')[-1] # cutoff b64-prefix if exists output_slot.elapsed_time = process_time - match inference.deployment.output_type: + match processed.inference.deployment.output_type: case Deployment.OutputTypeChoices.TEXT: output_slot.content = output_content case Deployment.OutputTypeChoices.FILE: raw = base64.b64decode(output_content) - extension = filetype.guess_extension(raw[:100]) + ft_out = filetype.guess(raw) + extension = ft_out.extension if ft_out else 'bin' output_slot.file = ContentFile(raw, name=f'.{extension}') output_slot.content = input_message.content - for payment_rule in inference.deployment.payment_rules: - if payment_rule.strategy == PaymentRule.StrategyChoices.PER_SECOND: - calculated_price += Decimal(process_time.total_seconds()) * payment_rule.cost - elif ( - payment_rule.strategy == PaymentRule.StrategyChoices.PER_TEXT_TOKEN - and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.OUTPUT - ): - calculated_price += TokenizerTool.token_count(output_content) * payment_rule.cost - for payment_bias in inference.payment_biases: - if payment_bias.type == PaymentBias.TypeChoices.ADDITION: - calculated_price += payment_bias.coefficient - elif payment_bias.type == PaymentBias.TypeChoices.MULTIPLICATION: - calculated_price *= payment_bias.coefficient - PaymentPlanService(self.user).update_per_token_plan_details( - calculated_price, input_message.content_object.model + total_price, input_message.content_object.model ) output_slot.save() @@ -0,0 +1,148 @@ +from dataclasses import dataclass +from datetime import timedelta +from decimal import Decimal +from io import BytesIO, StringIO +from math import ceil +from typing import Any, Iterable + +from django.utils.translation import gettext_lazy as _ + +from messages.models import Message +from ml_model.exceptions import ( + PaymentRuleNotImplemented, +) +from ml_model.models import ( + Inference, + PaymentBias, + PaymentRule, +) +from ml_model.tools.tokenizer import TokenizerTool + + +@dataclass +class ProcessedInput: + inference: Inference + content: str | None + file: StringIO | BytesIO | None + normalized_image: Any | None + file_extension: str | None + parameters: dict[str, Any] + history: Iterable[Message] + scrape_results: list[StringIO] | list[BytesIO] + + +class PricingService: + + @staticmethod + def compute_pre_cost(processed: ProcessedInput) -> tuple[Decimal, Decimal]: + calculated_price = Decimal('0') + expected_additional_price = Decimal('0') + inference = processed.inference + + for payment_rule in inference.deployment.payment_rules: + if ( + payment_rule.strategy == PaymentRule.StrategyChoices.FIXED + and payment_rule.interaction_type is None + and payment_rule.content_type is None + ): + calculated_price += payment_rule.cost + elif ( + payment_rule.strategy == PaymentRule.StrategyChoices.PER_SECOND + and payment_rule.interaction_type is None + and payment_rule.content_type is None + ): + if len(inference.tracking_records) < 1: + raise Exception(_('No tracking records found')) + last_tracking_record = inference.tracking_records[0] + average_generation_time = last_tracking_record.generation_time + expected_additional_price += payment_rule.cost * Decimal(average_generation_time.total_seconds()) + elif ( + payment_rule.strategy == PaymentRule.StrategyChoices.FIXED + and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.INPUT + and payment_rule.content_type == PaymentRule.ContentTypeChoices.FILE + ): + if processed.file: + calculated_price += payment_rule.cost + elif ( + payment_rule.strategy == PaymentRule.StrategyChoices.PER_TEXT_TOKEN + and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.INPUT + and payment_rule.content_type == PaymentRule.ContentTypeChoices.EMBEDDINGS + ): + if isinstance(processed.file, StringIO) and len(file_data := processed.file.getvalue()) >= 20_000: + calculated_price += TokenizerTool.token_count(file_data) * payment_rule.cost + elif ( + processed.content + and payment_rule.strategy == PaymentRule.StrategyChoices.PER_TEXT_TOKEN + and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.INPUT + ): + calculated_price += TokenizerTool.token_count(processed.content) * payment_rule.cost + if isinstance(processed.file, StringIO) and len(file_data := processed.file.getvalue()) < 20_000: + calculated_price += TokenizerTool.token_count(file_data) * payment_rule.cost + elif isinstance(processed.file, StringIO): + calculated_price += Decimal(f'{(110 + 100 + 10 * 4000) // 3}') * payment_rule.cost + if processed.history: + calculated_price += ( + sum([ + TokenizerTool.token_count(message.content) + for message in processed.history if message.content + ]) + * payment_rule.cost + ) + if processed.scrape_results: + calculated_price += ( + sum([TokenizerTool.token_count(result.getvalue()) for result in processed.scrape_results]) + ) * payment_rule.cost + elif ( + payment_rule.strategy == PaymentRule.StrategyChoices.PER_TEXT_TOKEN + and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.OUTPUT + ): + max_tokens_size = Decimal(processed.parameters.get('max_tokens', '128000')) + expected_additional_price += max_tokens_size * payment_rule.cost + elif ( + payment_rule.strategy == PaymentRule.StrategyChoices.PER_PIXEL + and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.INPUT + and payment_rule.content_type == PaymentRule.ContentTypeChoices.FILE + ): + if ( + processed.file + and processed.file_extension in ('png', 'jpg', 'jpeg') + and processed.normalized_image + ): + width, height = processed.normalized_image.width, processed.normalized_image.height + if max(width, height) > 2048: + a_ratio = width / height + width, height = ( + (2048, int(2048 / a_ratio)) if a_ratio > 1 else (int(2048 * a_ratio), 2048) + ) + if width >= height and height > 768: + width, height = int((768 / height) * width), 768 + elif height > width and width > 768: + width, height = 768, int((768 / width) * height) + tiles_size = ceil(width / 512) * ceil(height / 512) + calculated_price += tiles_size * payment_rule.cost + else: + raise PaymentRuleNotImplemented + return calculated_price, expected_additional_price + + @staticmethod + def apply_biases(inference: Inference, biased: Decimal) -> Decimal: + for payment_bias in inference.payment_biases: + if payment_bias.type == PaymentBias.TypeChoices.ADDITION: + biased += payment_bias.coefficient + elif payment_bias.type == PaymentBias.TypeChoices.MULTIPLICATION: + biased *= payment_bias.coefficient + return biased + + @staticmethod + def finalize_price( + pre_cost: Decimal, inference: Inference, output_content: str,process_time: timedelta + ) -> Decimal: + for payment_rule in inference.deployment.payment_rules: + if payment_rule.strategy == PaymentRule.StrategyChoices.PER_SECOND: + pre_cost += Decimal(process_time.total_seconds()) * payment_rule.cost + elif ( + payment_rule.strategy == PaymentRule.StrategyChoices.PER_TEXT_TOKEN + and payment_rule.interaction_type == PaymentRule.InteractionTypeChoices.OUTPUT + ): + pre_cost += TokenizerTool.token_count(output_content) * payment_rule.cost + return pre_cost @@ -1,5 +1,4 @@ from django.apps import AppConfig -from django.core.signals import setting_changed from django.utils.translation import gettext_lazy as _ @@ -12,5 +11,4 @@ class MLModelConfig(AppConfig): from .signals import disable_inference_by_deployment from .utils import create_redis_search_index - setting_changed.connect(disable_inference_by_deployment) create_redis_search_index() @@ -7,7 +7,7 @@ from ml_model.models import Deployment, Inference @receiver(pre_save, sender=Deployment) def disable_inference_by_deployment(instance: Deployment, **kwargs): if not instance.enabled: - inferences = instance.inferences + inferences = list(instance.inferences) for inference in inferences: if inference.enabled: inference.enabled = False