import os import sys import time import modal import torch import pathlib import funcy import json import pandas as pd import numpy as np import tempfile import torchaudio import pyloudnorm as pyln from suno_utils.audio import Audio from suno_utils.worker.settings import s3_client # from suno_utils.worker.modal_base import MODAL_MOUNTS from suno_utils.utils.text import read_jsonl from suno_utils.utils.s3 import list_s3_dir, read_from_s3 from suno_utils.tasks.hoot import ( encode_filepaths, encode, clean_text, encode_and_align, _get_alignable_tokens, _assign_word_timings, get_aligned_lyrics, ctc_align, load_model_list, get_word_timing_from_audio_and_lyrics, _merge_into_lines, decode_logits, load_model, ) from suno_utils.utils.metrics import get_cer from suno_utils.tasks.shimmerscore import shimmerscore from suno_utils.tasks.ear import load_model as load_ear_model from suno_utils.diffusion.generation import ( preload_dit_model, preload_tokenizer, TOKENIZER_FILEPATH, SEMANTIC_MODEL_FILEPATH, SEMANTIC_CLUSTERS_FILEPATH, _retrieve_models, ) from suno_utils.tasks.mert_25 import ( preload_models as preload_semantic_models, encode as encode_semantic, ) from suno_utils.diffusion import generation as diffusion_gen from suno_utils.tasks.upsample_engine import UpsampleEngine, Request, Job MOUNT_PATH = "/suno/models" aws_secret = modal.Secret.from_name("studio-aws") SECRETS = [ aws_secret, modal.Secret.from_dict( { "SUNO_ASSETS_PATH": "/suno/models/assets", "XDG_CACHE_HOME": "/suno/models/", } ), modal.Secret.from_name("api-callback-token"), modal.Secret.from_name("datadog-metrics"), ] base_image = ( modal.Image.from_registry("nvidia/cuda:12.4.0-devel-ubuntu22.04", add_python="3.10") .apt_install("curl", "ffmpeg", "sox", "unzip", "libsox-fmt-mp3", "zlib1g-dev", "git", "clang") .run_commands( [ 'curl "https://awscli.amazonaws.com/awscli-exe-linux-x86_64.zip" -o "awscliv2.zip"', "unzip -q awscliv2.zip", "./aws/install", ] ) .dockerfile_commands( [ "COPY --from=datadog/serverless-init:1.2.1 /datadog-init /app/datadog-init", 'ENTRYPOINT ["/app/datadog-init"]', ] ) .pip_install("torch==2.5.1", "torchaudio==2.5.1") .pip_install("flashinfer-python", index_url="https://flashinfer.ai/whl/cu124/torch2.5/") .pip_install_private_repos( "github.com/suno-ai/glockenspiel.git@a5dba4e50#subdirectory=descript-audio-codec&egg=descript-audio-codec", git_user="mcamac", secrets=[modal.Secret.from_name("victor-modal-github-token")], ) .pip_install_private_repos( "github.com/suno-ai/neon.git@1c83548#subdirectory=hoot", git_user="mcamac", secrets=[modal.Secret.from_name("victor-modal-github-token")], ) .pip_install( "boto3", "transformers", "tokenizers", "encodec", "ctc_segmentation", "psutil", "redis", "pydantic", "nnAudio", "rpyc", "biopython>=1.81", # TODO: don't love this depdendency, for hoot "pynvml", # for torch cuda utilization "torchsde", "ninja", "wheel", ) .pip_install_from_pyproject( "/home/christian/code/glockenspiel/suno_utils/pyproject.toml", ) .run_commands( # This is really slow "git clone https://github.com/Dao-AILab/flash-attention.git", "cd flash-attention/hopper && python setup.py install", gpu="h100", ) .apt_install( "libogg0", "libopus0", "opus-tools", ) .pip_install("transformers==4.44.0", "wandb") ) def calculate_stereo_width(waveform): # Split into left and right channels left = waveform[0] right = waveform[1] # Compute mid/side representation mid = (left + right) / 2 side = (left - right) / 2 # Compute RMS energy of mid and side channels mid_energy = torch.sqrt(torch.mean(mid**2)) side_energy = torch.sqrt(torch.mean(side**2)) # Compute stereo width based on mid/side ratio # Normalize to range 0-1 using sigmoid-like function width_ratio = (side_energy / (mid_energy + 1e-8)).item() stereo_width = 2 * (1 / (1 + np.exp(-width_ratio)) - 0.5) return stereo_width def calculate_loudness(waveform, sr): meter = pyln.Meter(sr) lufs_db = meter.integrated_loudness(waveform) return lufs_db def calculate_loudness_factor(waveform, sr): meter = pyln.Meter(sr) normalized_waveform = waveform / np.clip(np.max(np.abs(waveform)), 1e-10, None) lufs_db = meter.integrated_loudness(normalized_waveform) return lufs_db def calculate_possible_clipped_samples(waveform): # Convert to numpy if it's a torch tensor if isinstance(waveform, torch.Tensor): waveform_np = waveform.numpy() else: waveform_np = waveform # Count samples that are at or above the clipping threshold return ( np.sum(np.abs(waveform_np) >= 1.0).item() if isinstance(np.sum(np.abs(waveform_np) >= 1.0), torch.Tensor) else np.sum(np.abs(waveform_np) >= 1.0) ) def calculate_average_spectrum_db(waveform, n_fft=16384, hop_length=8192): # Keep as torch tensor or convert to torch tensor if it's numpy if not isinstance(waveform, torch.Tensor): waveform = torch.from_numpy(waveform) # Calculate the average spectrum using STFT for efficiency n_fft = 2048 # Choose an appropriate FFT size hop_length = n_fft // 4 # Standard hop length # Compute STFT using torch if waveform.dim() > 1: # For stereo, compute STFT for each channel stft_results = [] for channel in range(waveform.shape[0]): stft = torch.stft( waveform[channel], n_fft=n_fft, hop_length=hop_length, window=torch.hann_window(n_fft), return_complex=True, ) # Get magnitude stft_magnitude = torch.abs(stft) stft_results.append(stft_magnitude) # Average across time frames for each channel magnitude_spectrum = torch.stack([torch.mean(stft, dim=1) for stft in stft_results]) else: # For mono stft = torch.stft( waveform, n_fft=n_fft, hop_length=hop_length, window=torch.hann_window(n_fft), return_complex=True, ) # Get magnitude stft_magnitude = torch.abs(stft) magnitude_spectrum = torch.mean(stft_magnitude, dim=1) # Convert to dB scale spectrum_db = 20 * torch.log10(magnitude_spectrum + 1e-10) # Adding small value to avoid log(0) # Convert to numpy for consistency with the rest of the code return spectrum_db.numpy() def calculate_average_stereo_spectrum(waveform, sr): assert waveform.dim() == 2 and waveform.shape[0] == 2 # split into left and right channels left = waveform[0] right = waveform[1] # compute mid and side channels mid = (left + right) / 2 side = (left - right) / 2 # calculate spectrum for mid and side channels spectrum_mid = calculate_average_spectrum_db(mid, sr) spectrum_side = calculate_average_spectrum_db(side, sr) return spectrum_mid, spectrum_side def calculate_spectrum_evolution(waveform, sr, n_fft=16384, hop_length=8192): # compute spectrum for first 30s waveform_first = waveform[:, : 30 * sr] spectrum_first = calculate_average_spectrum_db(waveform_first, n_fft, hop_length) # compute spectrum for last 30s waveform_last = waveform[:, -30 * sr :] spectrum_last = calculate_average_spectrum_db(waveform_last, n_fft, hop_length) return spectrum_first, spectrum_last def analyze_audio(filepath): audio, sr = torchaudio.load(filepath) if sr != 48000: audio = torchaudio.functional.resample(audio, sr, 48000) # calculate loudness lufs_db = calculate_loudness(audio.permute(1, 0).numpy(), sr) lufs_db_factor = calculate_loudness_factor(audio.permute(1, 0).numpy(), sr) # calculate stereo width stereo_width = calculate_stereo_width(audio) # clipped samples clipped_samples = calculate_possible_clipped_samples(audio) # average spectrum average_spectrum_db = calculate_average_spectrum_db(audio) # average stereo spectrum average_stereo_spectrum_mid, average_stereo_spectrum_side = calculate_average_stereo_spectrum( audio, sr ) # spectrum evolution spectrum_first, spectrum_last = calculate_spectrum_evolution(audio, sr) return { "lufs_db": lufs_db, "lufs_db_factor": lufs_db_factor, "stereo_width": stereo_width, "clipped_samples": clipped_samples, "average_spectrum_db": average_spectrum_db, "average_spectrum_db_first": spectrum_first, "average_spectrum_db_last": spectrum_last, "average_stereo_spectrum_mid": average_stereo_spectrum_mid, "average_stereo_spectrum_side": average_stereo_spectrum_side, } def _reload_models_if_needed(hoot_filepath, hoot_tokenizer_filepath): model_list = load_model_list( checkpoint_filepath=hoot_filepath, tokenizer_filepath=hoot_tokenizer_filepath, n_gpus=None, ) model = model_list[0]["model"] tokenizer = model_list[0]["tokenizer"] return model_list, model, tokenizer class GenerateWorker: def __init__( self, dit_model_filepath: str, output_path: str, test_type: str, checkpoint_name: str, objective: str, ): self.output_path = output_path self.test_type = test_type self.checkpoint_name = checkpoint_name start_time = time.time() print("Start loading models") # load diffusion model num_gpus = torch.cuda.device_count() cuda_device = torch.cuda.current_device() print(f"Found {num_gpus} GPUs. Using GPU {cuda_device}.") tokenizer_filepath = "s3://suno-data/georg/models/tokenizers/tokenizer_60k.json" semantic_model_filepath = "s3://suno-data/georg/models/semantic/mert_25.pt" semantic_clusters_filepath = "s3://suno-data/georg/models/semantic/mert_25_2x4k.npy" codec_filepath = CODEC_FILEPATH _ = diffusion_gen.preload_dit_model( dit_model_filepath=dit_model_filepath, use_ema_if_exists=True, compile=False, weights_precision=torch.bfloat16, ) _ = preload_tokenizer(tokenizer_filepath) _ = preload_semantic_models(semantic_model_filepath, semantic_clusters_filepath) _ = preload_codec_models(codec_filepath) self.diffusion_engine = UpsampleEngine(min_chunk_size=30 * 25) self.objective = objective # "v" or "rectified_flow" assert self.objective in ["v", "rectified_flow"] # load hoot model print("Loading hoot model") hoot_filepath = "s3://suno-data/christian/checkpoints/hoot_v3.pt" hoot_tokenizer_filepath = "s3://suno-data/christian/checkpoints/hoot_v3_tokenizer_v3.pt" self.hoot_model_list, self.hoot_model, self.hoot_tokenizer = _reload_models_if_needed( hoot_filepath, hoot_tokenizer_filepath ) # load ear model ear_model_filepath = "s3://suno-data/christian/checkpoints/ear/ear_v2_s3080.pt" self.ear_model = load_ear_model(ear_model_filepath, compile=True) print(f"Finish loading models. Took {round(time.time() - start_time, 2)} seconds") @staticmethod def download_models(dit_model_filepath, dir_path=MOUNT_PATH): """Download diffusion models.""" print("Start downloading models") _ = diffusion_gen.get_model_if_needed( CODEC_FILEPATH, cache_dir=dir_path, ) _ = diffusion_gen.get_model_if_needed(diffusion_gen.SEMANTIC_MODEL_FILEPATH, cache_dir=dir_path) _ = diffusion_gen.get_model_if_needed( diffusion_gen.SEMANTIC_CLUSTERS_FILEPATH, cache_dir=dir_path ) _ = diffusion_gen.get_model_if_needed(diffusion_gen.CODEC_FILEPATH, cache_dir=dir_path) _ = diffusion_gen.get_model_if_needed(dit_model_filepath, cache_dir=dir_path) # _ = chirp_v2._get_model_if_needed(gpt_ckpt, cache_dir=dir_path) # _ = chirp_v2._get_model_if_needed(chirp_v2.TOKENIZER_PATH, cache_dir=dir_path) print("Finish downloading models") def generate(self, work_item): global models """Generate audio from a work item.""" # when saving to s3 we will use a structure like this: # self.output_path/ # item_id/ # original_semantic.npz # 0_generated_audio.mp3 # 0_generated_semantic.npz # 0_metadata.json # 1_generated_audio.mp3 # 1_generated_semantic.npz # 1_metadata.json # ... # metadata.json for item in work_item: item_id = item["id"] if "s3_filepath" in item: s3_filepath = item["s3_filepath"] # read audio from s3 and then semantic encode audio = Audio.from_s3(s3_filepath, n_channels=2) if SAVE_CODES: # encode the vae latents vae_latents_cycled = codec_encode( audio.convert(sample_rate=48_000, byte_width=2, n_channels=2).normalize_volume( target_db=-16 ) ) with tempfile.TemporaryDirectory() as td: vae_latents_path = os.path.join(td, f"{item_id}_cycled_vae.npz") np.savez(vae_latents_path, vae_latents=vae_latents_cycled) s3_filepath = os.path.join( self.output_path, f"{item_id}", f"{item_id}_vae.npz", ) s3_client.upload_file( vae_latents_path, "suno-data", s3_filepath, ExtraArgs={"ContentType": "application/octet-stream"}, ) if CYCLE_ONLY: continue # also encode semantic codes codes = encode_semantic( audio.convert(sample_rate=24_000, byte_width=2, n_channels=1) ).astype(np.int64) if SAVE_CODES: with tempfile.TemporaryDirectory() as td: codes_path = os.path.join(td, f"{item_id}_semantic.npz") np.savez(codes_path, codes=codes) s3_filepath = os.path.join( self.output_path, f"{item_id}", f"{item_id}_semantic.npz", ) s3_client.upload_file( codes_path, "suno-data", s3_filepath, ExtraArgs={"ContentType": "application/octet-stream"}, ) else: # load the semantic codes from s3 s3_filepath = f"s3://suno-data-uploads/studio/uploads/{item_id}.npz" try: data = read_from_s3(s3_filepath, read_f=np.load) except Exception as e: print(f"Error loading {s3_filepath}: {e}") continue if "v3.0_raw" in data: codes = data["v3.0_raw"] elif "v3.5_raw" in data: codes = data["v3.5_raw"] elif "v4.0_raw" in data: codes = data["v4.0_raw"] elif "v4.5_raw" in data: codes = data["v4.5_raw"] elif "v5.0_raw" in data: codes = data["v5.0_raw"] else: raise ValueError("No codes found") semantic_codes = torch.from_numpy(codes[:, 0]).long() # .cuda() # semantic_codes = semantic_codes[:3000] # if tags is a list, join it into a string if isinstance(item["tags"], list): tags_str = ", ".join(item["tags"][:5]) else: tags_str = item["tags"] if tags_str is None: tags_str = "" lyrics = item["text"] # create a list of configs to run gen_cfgs = [] if self.test_type == "steps": diffusion_steps = [8, 10, 12, 14, 16, 18, 20] diffusion_seed = 42 diffusion_text_cfg_coef = 2.0 noise_ctx_level = NOISE_CTX_LEVEL noise_ctx_pad_len = NOISE_CTX_PAD_LEN rho = RHO sigma_min = SIGMA_MIN sigma_max = SIGMA_MAX for diffusion_steps in diffusion_steps: gen_cfgs.append( ( f"steps_{diffusion_steps}", diffusion_gen.DiffusionGenerationConfig( steps=diffusion_steps, lyrics=lyrics, tags=tags_str, text_cfg_coef=diffusion_text_cfg_coef, codec_scale_factor=CODEC_SCALE_FACTOR, scale_ctx_vector=SCALE_CTX_VECTOR, noise_ctx_level=noise_ctx_level, noise_ctx_pad_len=noise_ctx_pad_len, # sigma_min=sigma_min, sigma_max=sigma_max, seed=diffusion_seed, # noise_schedule=NOISE_SCHEDULE, objective=self.objective, ), ) ) elif self.test_type == "standard": # diffusion parameters if RANDOMIZE_DIFFUSION_PARAMS: diffusion_seed = np.random.randint(0, 1000000) diffusion_steps = np.random.choice([8, 10, 12, 14, 16, 18, 20]) diffusion_text_cfg_coef = np.random.choice([1.0, 1.5, 1.75, 2.0, 2.5, 3.0]) noise_ctx_level = np.random.choice([0.0, 0.25, 0.5, 0.75, 1.0]) rho = np.random.choice([0.9, 1.0, 1.1]) sigma_min = np.random.choice([0.05, 0.1, 0.25, 0.5]) sigma_max = np.random.choice([40.0, 50.0, 60.0, 70.0, 80.0, 100.0]) noise_ctx_pad_len = np.random.choice([0, 10, 20, 30]) else: if DIFFUSION_SEED is None: diffusion_seed = 0 for char in item_id: diffusion_seed = (diffusion_seed * 31 + ord(char)) % 1000000 else: diffusion_seed = DIFFUSION_SEED diffusion_steps = DIFFUSION_STEPS diffusion_text_cfg_coef = 2.0 noise_ctx_level = NOISE_CTX_LEVEL noise_ctx_pad_len = NOISE_CTX_PAD_LEN rho = RHO sigma_min = SIGMA_MIN sigma_max = SIGMA_MAX gen_cfgs.append( ( "", diffusion_gen.DiffusionGenerationConfig( steps=diffusion_steps, lyrics=lyrics, tags=tags_str, text_cfg_coef=diffusion_text_cfg_coef, codec_scale_factor=CODEC_SCALE_FACTOR, scale_ctx_vector=SCALE_CTX_VECTOR, noise_ctx_level=noise_ctx_level, noise_ctx_pad_len=noise_ctx_pad_len, semantic_skip_factor=SEMANTIC_SKIP_FACTOR, rho=rho, sigma_min=sigma_min, sigma_max=sigma_max, seed=diffusion_seed, objective=self.objective, # noise_schedule=NOISE_SCHEDULE, ), ) ) else: raise ValueError(f"Invalid test type: {self.test_type}") # now run the requests for name, gen_cfg in gen_cfgs: request = Request( id="dummy", generation_config=gen_cfg, tokens=semantic_codes, input_tokens_finished=True, ) result = self.diffusion_engine.run_request(request) vae_latents = torch.concat(result.vae_latents) upsampled_audio = decode_stream_to_full_audio(vae_latents) metadata = { # "original_audio": item["s3_filepath"], "id": item_id, "text": lyrics, "tags": tags_str, "diffusion": { "steps": int(gen_cfg.steps), "seed": int(gen_cfg.seed), "text_cfg_coef": float(gen_cfg.text_cfg_coef), "noise_ctx_level": float(gen_cfg.noise_ctx_level), "codec_scale_factor": float(gen_cfg.codec_scale_factor), "scale_ctx_vector": gen_cfg.scale_ctx_vector, # "rho": float(gen_cfg.rho), }, } # we want to save out # audio file of the final audio # npz of the estimated semantics # metadata json with the original prompt and tags # copy some stuff # we want to copy the original audio and the npz of the original semantics with tempfile.TemporaryDirectory() as td: # save audio to s3 upsampled_audio_path = os.path.join(td, f"{item_id}_{name}.mp3") upsampled_audio.write_hq_mp3(upsampled_audio_path) # run hoot evaluation (CER) # check if the lyrics are not empty if lyrics != "": out = encode_filepaths([upsampled_audio_path], return_logits=True) basic_cleaned_lyrics = lyrics decoded_preds = self.hoot_tokenizer.decode_logits( out[0], prior_text=basic_cleaned_lyrics ) true_text_norm = clean_text(basic_cleaned_lyrics) cer_val = round(get_cer(true_text_norm, decoded_preds), 3) metadata["hoot_cer"] = cer_val # add cer to metadata else: metadata["hoot_cer"] = None # run ear evaluation (quality score) # check if the length of the audio is greater than 5 seconds if upsampled_audio.duration_s > 5.0: quality_score = self.ear_model.get_score(upsampled_audio_path) metadata["ear_score"] = quality_score # add quality score to metadata else: metadata["ear_score"] = None # run shimmer score evaluation shimmer_score = shimmerscore(upsampled_audio_path) metadata["shimmer_score"] = shimmer_score # run the other evals audio_analysis = analyze_audio(upsampled_audio_path) # add all the keys in audio_analysis to metadata for key, value in audio_analysis.items(): metadata[key] = value s3_filepath = os.path.join( self.output_path, f"{item_id}", ( f"{item_id}_{self.checkpoint_name}_{name}.mp3" if name != "" else f"{item_id}_{self.checkpoint_name}.mp3" ), ) s3_client.upload_file( upsampled_audio_path, "suno-data", s3_filepath, ExtraArgs={ "ContentType": "audio/mpeg", }, ) if SAVE_CODES: # save the vae latents vae_latents_path = os.path.join(td, f"{item_id}_upsampled_vae.npz") np.savez(vae_latents_path, vae_latents=vae_latents.cpu().numpy()) s3_filepath = os.path.join( self.output_path, f"{item_id}", f"{item_id}_upsampled_vae.npz", ) s3_client.upload_file( vae_latents_path, "suno-data", s3_filepath, ExtraArgs={ "ContentType": "application/octet-stream", }, ) # save the metadata metadata_path = os.path.join(td, f"{item_id}_{name}_metadata.npz") np.savez(metadata_path, **metadata) # move the metadata to s3 s3_filepath = os.path.join( self.output_path, f"{item_id}", ( f"{item_id}_{self.checkpoint_name}_{name}__metadata.npz" if name != "" else f"{item_id}_{self.checkpoint_name}__metadata.npz" ), ) s3_client.upload_file( metadata_path, "suno-data", s3_filepath, ExtraArgs={ "ContentType": "application/octet-stream", }, ) DIT_MODEL_FILEPATH = "s3://suno-data/christian/checkpoints/diffusion/modal_eval_ckpt.pt" OUTPUT_STR = "auk-clips-up-u-2" TEST_TYPE = "standard" # standard, steps SAVE_CODES = False CYCLE_ONLY = False GPU_TYPE = "H100" # H100, A10G # extra params RHO = 1.0 SIGMA_MIN = 0.5 SIGMA_MAX = 50.0 DIFFUSION_SEED = None # 42 is default RANDOMIZE_DIFFUSION_PARAMS = False if "dit_v6_dpo_t11_9k_5e6_b100" in DIT_MODEL_FILEPATH: CODEC_SCALE_FACTOR = 2.5 SCALE_CTX_VECTOR = False NOISE_CTX_LEVEL = 0.0 NOISE_CTX_PAD_LEN = 0 DIFFUSION_STEPS = 10 SEMANTIC_SKIP_FACTOR = 1 CODEC_FILEPATH = "s3://suno-data/christian/25hz_vae_peaq_kl_0.005.pth" from suno_utils.tasks.dac_vae_100hz_peaq import ( preload_models as preload_codec_models, decode as codec_decode, encode as codec_encode, decode_stream_to_full_audio, ) else: CODEC_SCALE_FACTOR = 0.4 SCALE_CTX_VECTOR = True NOISE_CTX_LEVEL = 0.75 NOISE_CTX_PAD_LEN = 0 DIFFUSION_STEPS = 10 SEMANTIC_SKIP_FACTOR = 1 NOISE_SCHEDULE = "polyexponential" # NOISE_SCHEDULE = "cosine" # CHUNK_SIZE_SCHEDULE = [10 * 25, 20 * 25, 30 * 25] # 10s, 20s, 30s chunks # CHUNK_SIZE_SCHEDULE = [30 * 25, 30 * 25, 30 * 25] # 30s chunks CODEC_FILEPATH = "s3://suno-data/minz/models/dac_vae_tuned_25hz.pth" from suno_utils.tasks.dac_vae_fixed_25hz import ( preload_models as preload_codec_models, decode as codec_decode, encode as codec_encode, decode_stream_to_full_audio, ) def download_model_wrapper_d(): # this print is necessary to have modal rerun this when MODEL changes # Modal tracks referenced global variables # Change the name of the function to force a rerun print("Downloading model for history encoder") GenerateWorker.download_models(DIT_MODEL_FILEPATH) image = base_image.run_function(download_model_wrapper_d, secrets=SECRETS) APP_NAME = f"batch-generate-gpu" app = modal.App(APP_NAME, image=image, secrets=SECRETS) N_MAX_REPLICAS = 8 @app.cls( # gpu=modal.gpu.H100(count=1), gpu=modal.gpu.A10G(count=1) if GPU_TYPE == "A10G" else modal.gpu.H100(count=1), cpu=4, secrets=SECRETS, timeout=2 * 60 * 60, container_idle_timeout=240, # mounts=MODAL_MOUNTS, memory=15000, concurrency_limit=N_MAX_REPLICAS, ) class GenerateStub: def __init__( self, dit_ckpt: str, output_path: str, test_type: str, checkpoint_name: str, objective: str, ): import torch num_gpus = torch.cuda.device_count() print(f"Found {num_gpus} GPUs.") self.worker = GenerateWorker( dit_ckpt, output_path, test_type, checkpoint_name, objective, ) @modal.method() def generate(self, work_item: list[dict]): return self.worker.generate(work_item) @app.local_entrypoint() def main(run_name: str, objective: str = "v"): # checkpoints in s3://suno-data/christian/checkpoints # prompts in s3://suno-data/christian/prompts # outputs in s3://suno-data/christian/outputs # load the positive prompts # work_items = read_from_s3( # "s3://suno-data/christian/sft/pos_interesting_clips_up_u_1_20241201_full.jsonl", # read_f=read_jsonl, # ) work_items = read_jsonl( "/home/christian/code/christian/metadata/sft/auk_clips_up_u_1_20241201_pos.jsonl", ) print(f"Total work items: {len(work_items)}") chunksize = 16 # num of prompts per worker worker = GenerateStub( DIT_MODEL_FILEPATH, f"christian/outputs/{OUTPUT_STR}", test_type=TEST_TYPE, checkpoint_name=run_name, objective=objective, ) work_items = list(funcy.chunks(chunksize, work_items)) print(f"Chunksize: {chunksize}, total chunks: {len(work_items)}") print("Running batch inference...") t0 = time.time() _ = list(worker.generate.map(work_items)) print(round((time.time() - t0) / 60 / 60), "h total for batch generation")