@@ -37,6 +37,7 @@ from ml_model.services.recraft import Recraft from ml_model.services.reve import Reve from ml_model.services.sdxlemoji import Sdxlemoji from ml_model.services.seedream import Seedream +from ml_model.services.suno import Suno from ml_model.services.stablediffusion import Stablediffusion from ml_model.services.upscaleai import Upscaleai from ml_model.services.veo import Veo @@ -49,3 +50,4 @@ from ml_model.services.kling import Kling from ml_model.services.fluxkrea import Fluxkrea from ml_model.services.ideogram import Ideogram from ml_model.services.leonardo import Leonardo +from ml_model.services.lyria import Lyria @@ -0,0 +1,49 @@ +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.services.base import SimpleService +from ml_model.tasks import replicate_run +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + +class Lyria(SimpleService): + TOKENS_COST = Decimal('0.6') # per 1 sec of output audio + + def calculate_price(self, duration: int) -> Decimal: + price = self.TOKENS_COST * duration + return price.quantize(Decimal('0.1'), rounding='ROUND_UP') + + def save_results(self, content: str, t: timedelta, audio: str, save: bool = True) -> list[Message]: + msg = Message( + content=content, + content_object=self.store, + elapsed_time=t, + file=File(BytesIO(requests.get(audio).content), '.mp3'), + ) + if save: + return Message.objects.bulk_create([msg]) + return [msg] + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + if ( + (balance := PaymentPlanSelector(self.store.user).get_current_balance()) + < (cost := self.TOKENS_COST * 32) + ): + raise InsufficientBalance(balance, cost) + callback_data = { + 'prompt': self.translate_prompt(input_message.content), + **input_message.info + } + start_time = time.time() + video = replicate_run(f'google/lyria-2', callback_data) + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, duration=32) + msgs = self.save_results(input_message.content, process_time, video, save) + return msgs @@ -0,0 +1,57 @@ +import time +from _decimal import Decimal +from datetime import timedelta +from io import BytesIO + +import requests +from django.core.files import File + +from messages.models import Message +from ml_model.services.base import SimpleService +from ml_model.tasks import replicate_run +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + +class Suno(SimpleService): + + PRICE = Decimal('0.057') + + _CALLBACK = 'suno-ai/bark:b76242b40d67c76ab6742e987628a2a9ac019e11d56ab96c4e91ce03b79b2787' + + def calculate_price(self, process_time: timedelta) -> Decimal: + price = self.PRICE * Decimal(process_time.total_seconds()) + return price.quantize(Decimal('0.1'), rounding='ROUND_UP') + + def save_results(self, prompt: str, audio: 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(audio).content), '.mp3'), + ) + ) + if save: + return Message.objects.bulk_create(messages) + return messages + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + if ( + (balance := PaymentPlanSelector(self.store.user).get_current_balance()) + < (cost := Decimal('14')) + ): + raise InsufficientBalance(balance, cost) + callback_data = dict( + { + 'prompt': input_message.content, + **input_message.info, + } + ) + start_time = time.time() + result = replicate_run(self._CALLBACK, callback_data)['audio_out'] + 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, result, process_time, save) + return msgs