import os import random import time from collections import defaultdict from datetime import datetime from typing import Optional, Literal, Iterator import openai from openai.types.chat import ChatCompletion, ChatCompletionChunk from suno_utils.utils.bedrock import BedrockAdapter try: import together except ImportError: print("Warning: together not installed locally, will not use together client") together = None MINUTES = 60 # seconds / minute FAST_TIER_FRACTION = 0.01 MODAL_BASE_URL = "https://suno-ai--lyrics-gen-serve.modal.run" def _get_openai_compatible_url_from_base_url(base_url: str) -> str: return os.path.join(base_url, "v1") REMI = "ft:gpt-4o-2024-08-06:suno:recent-billboard3:ASqCxlGH" DEFAULT_GPT_MODEL = "gpt-4o-2024-11-20" DUSAN_BOT2 = "ft:gpt-4o-2024-08-06:suno:dusan-ascii:APwxFPs4" # 4o FT with ascii chars fixed LORENZO_V1_LYRICS_MODEL = "ft:gpt-4o-2024-08-06:suno:lorenzo-gym:AXXM0K5u" CLAUDE_SONNET_MODEL = "us.anthropic.claude-3-7-sonnet-20250219-v1:0" CLAUDE_SONNET_4_MODEL = "us.anthropic.claude-sonnet-4-20250514-v1:0" # Provider types as constants PROVIDER_ANTHROPIC = "anthropic" PROVIDER_OPENAI = "openai" PROVIDER_TOGETHER = "together" PROVIDER_MODAL = "modal" PROVIDER_VERTEX = "vertex" PROVIDER_UNKNOWN = "unknown" ProviderType = Literal["anthropic", "openai", "together", "modal", "vertex", "unknown"] LYRICS_PARAMETER_CONFIGS = defaultdict( dict, # if not found, assume no special param values should be passed in { "my-remi-8b-1": {"top_p": 0.95, "temperature": 0.7}, "remi-v1-hps1": {"top_p": 0.95, "temperature": 1}, "remi-v1-hps2": {"top_p": 0.9, "temperature": 1}, "remi-v1-hps3": {"top_p": 0.8, "temperature": 1}, "remi-v1-hps4": {"top_p": 0.95, "temperature": 0.9}, "remi-v1-hps5": {"top_p": 0.9, "temperature": 0.9}, "remi-v1-hps6": {"top_p": 0.8, "temperature": 0.9}, "remi-v1-hps7": {"top_p": 0.95, "temperature": 1.1}, "remi-v1-hps8": {"top_p": 0.9, "temperature": 1.1}, "remi-v1-hps9": {"top_p": 0.8, "temperature": 1.1}, "remi-v1-hps10": {"top_p": 0.8, "temperature": 1.2}, "remi-v1": {"top_p": 0.8, "temperature": 1.1}, "claude-sonnet": {"topP": 0.8, "temperature": 0.9}, "claude-sonnet-4": {"topP": 0.8, "temperature": 0.9}, "default-hps1": {"top_p": 0.95, "temperature": 1}, "default-hps2": {"top_p": 0.9, "temperature": 1}, "default-hps3": {"top_p": 0.8, "temperature": 1}, "default-hps4": {"top_p": 0.95, "temperature": 0.9}, "default-hps5": {"top_p": 0.9, "temperature": 0.9}, "default-hps6": {"top_p": 0.8, "temperature": 0.9}, "default-hps7": {"top_p": 0.95, "temperature": 1.1}, "default-hps8": {"top_p": 0.9, "temperature": 1.1}, "default-hps9": {"top_p": 0.8, "temperature": 1.1}, "default-hps10": {"top_p": 0.8, "temperature": 1.2}, # A/B hyperparam sweep winner, corresponding to default-hps6 "default": {"top_p": 0.8, "temperature": 0.9}, }, ) def _get_backend_lyrics_model(frontend_lyrics_model: str | None) -> str: """Translate the lyrics model from the frontend-visible name to the corresponding backend name.""" # map lyrics model options from FE to real names print("getting backend lyrics model for:", frontend_lyrics_model) LYRICS_MODEL_LOOKUP = { "dusanbot2": DUSAN_BOT2, "remi-v1": REMI, "lorenzo-v1": LORENZO_V1_LYRICS_MODEL, "my-remi-8b-1": "my-remi-8b-1", "remi-v1-hps1": REMI, "remi-v1-hps2": REMI, "remi-v1-hps3": REMI, "remi-v1-hps4": REMI, "remi-v1-hps5": REMI, "remi-v1-hps6": REMI, "remi-v1-hps7": REMI, "remi-v1-hps8": REMI, "remi-v1-hps9": REMI, "claude-sonnet": CLAUDE_SONNET_MODEL, "claude-sonnet-4": CLAUDE_SONNET_4_MODEL, "default-hps1": DEFAULT_GPT_MODEL, "default-hps2": DEFAULT_GPT_MODEL, "default-hps3": DEFAULT_GPT_MODEL, "default-hps4": DEFAULT_GPT_MODEL, "default-hps5": DEFAULT_GPT_MODEL, "default-hps6": DEFAULT_GPT_MODEL, "default-hps7": DEFAULT_GPT_MODEL, "default-hps8": DEFAULT_GPT_MODEL, "default-hps9": DEFAULT_GPT_MODEL, "default-hps10": DEFAULT_GPT_MODEL, None: DEFAULT_GPT_MODEL, "default": DEFAULT_GPT_MODEL, DEFAULT_GPT_MODEL: DEFAULT_GPT_MODEL, "gemini-2.0": "google/gemini-2.0-flash-001", } backend_lyrics_model = LYRICS_MODEL_LOOKUP.get(frontend_lyrics_model) if backend_lyrics_model is None: valid_choices = list(LYRICS_MODEL_LOOKUP.keys()) err_msg = f"Lyrics model: {frontend_lyrics_model} unrecognized, provide one of {valid_choices}" raise ValueError(err_msg) return backend_lyrics_model class LyricsClient: def __init__(self, openai_client, together_client, vertex_client, bedrock_client): print("initializing lyrics llm client") self.openai_client = openai_client self.together_client = together_client self.modal_client = openai.OpenAI( base_url=_get_openai_compatible_url_from_base_url(MODAL_BASE_URL), api_key="sunosunosuno" ) self.anthropic_client = BedrockAdapter(bedrock_client) self.vertex_client = vertex_client def _get_provider_type(self, model: str | None = None) -> ProviderType: """Determine the provider type based on the model name. Args: model: Model name to check Returns: Provider type as a string constant """ if model is None: return PROVIDER_OPENAI if ( model == "lorenzo-v1" or model == "dusanbot2" or model.startswith("gpt-4") or model.startswith("remi-v1") or model.startswith("default") ): return PROVIDER_OPENAI elif model.startswith("patsuno/Meta-Llama"): return PROVIDER_TOGETHER elif model.startswith("my-remi-8b-1"): return PROVIDER_MODAL elif model.startswith("gemini"): return PROVIDER_VERTEX elif model.startswith("anthropic") or model == "claude-sonnet" or model == "claude-sonnet-4": return PROVIDER_ANTHROPIC else: return PROVIDER_UNKNOWN def get_provider_name(self, model: str | None = None) -> str: """Return provider name based on model. Args: model: Model name to check. If None, returns "openai". Returns: Provider name: "openai", "together", "modal", or "unknown" """ return self._get_provider_type(model) def create( self, model: str, system_prompt: str, user_prompt: str, max_tokens: Optional[int] = None, n: int = 1, stream: bool = False, ) -> ChatCompletion | Iterator[ChatCompletionChunk]: print("creating with LyricsLLMClient for:", model, user_prompt) client = self._get_client_for_model(model) lyrics_param_config = LYRICS_PARAMETER_CONFIGS[model] backend_model = _get_backend_lyrics_model(model) print("lyrics model:", model, "lyrics_param_config:", lyrics_param_config) if model.startswith("default"): lyrics_param_config = { **lyrics_param_config, "timeout": 30, } fast_tier_fraction = _interpolate_time(time.ctime()) if random.random() < fast_tier_fraction: print( f"upgrading {model} to Fast Tier at: {time.ctime()}, fraction: {fast_tier_fraction}" ) lyrics_param_config["service_tier"] = "fast_tier_temp_pilot" try: resp = client.chat.completions.create( model=backend_model, messages=[ {"role": "system", "content": system_prompt}, {"role": "user", "content": user_prompt}, ], max_tokens=max_tokens, n=n, stream=stream, **lyrics_param_config, ) except Exception as e: print("got exception in LyricsClient.create:", e) raise e return resp def _get_client_for_model(self, model: str): print("getting model for:", model) provider_type = self._get_provider_type(model) if provider_type == PROVIDER_OPENAI: return self.openai_client elif provider_type == PROVIDER_TOGETHER: return self.together_client elif provider_type == PROVIDER_MODAL: print("getting modal client") return self.modal_client elif provider_type == PROVIDER_VERTEX: return self.vertex_client elif provider_type == PROVIDER_ANTHROPIC: return self.anthropic_client else: raise ValueError(f"Couldn't find a client for model: {model}") def _is_openai_client(client): """Determine whether client is the openai client.""" is_openai = "OpenAI" in str(client) print(f"checking if {client} is openai: {is_openai}") return is_openai def _interpolate_time(current_time): """ Interpolates between start_time and end_time based on current_time. Args: current_time (str): Current time in ctime format (e.g., 'Thu Apr 3 12:34:56 2025') start_time (str): Start time in ctime format end_time (str): End time in ctime format Returns: float: 0 if current_time < start_time, 1 if current_time > end_time, linear interpolation between 0 and 1 otherwise """ # Convert ctime strings to datetime objects START_TIME = "Thu Apr 3 22:00:00 2025" # NB UTC time in modal land END_TIME = "Fri Apr 4 02:00:00 2025" TIME_FORMAT = "%a %b %d %H:%M:%S %Y" current_dt = datetime.strptime(current_time, TIME_FORMAT) start_dt = datetime.strptime(START_TIME, TIME_FORMAT) end_dt = datetime.strptime(END_TIME, TIME_FORMAT) # Convert to timestamps (seconds since epoch) current_ts = current_dt.timestamp() start_ts = start_dt.timestamp() end_ts = end_dt.timestamp() # Check boundary conditions if current_ts <= start_ts: return 0.0 if current_ts >= end_ts: return 1.0 # Linear interpolation progress = (current_ts - start_ts) / (end_ts - start_ts) return progress