import argparse from collections import defaultdict import json import numpy as np import os import pandas as pd import pickle import time import torch import torchaudio from multiprocessing import Pool from torchaudio.transforms import MelSpectrogram from tqdm import tqdm import pyloudnorm as pyln import traceback from suno_utils.tasks.ear import load_model as load_ear_model from suno_utils.audio import Audio import torch.nn as nn import torch.nn.functional as F ## EAR # Set CUDA device device = "cuda" torch.set_num_threads(16) # neon/ear # pip install -e . # s3://suno-data/christian/ear/w5p4nhzn-epoch=27.ckpt checkpoint_filepath = "/app/suno/data/dpo/models/ear_v2_s3080.pt" NUM_FRAMES = 131072 def load_ear_system(): ear_system = load_ear_model(checkpoint_filepath) ear_system.to(device) ear_system.eval() return ear_system def process_request_batch(request_batch): try: ear_system = load_ear_system() results = defaultdict(dict) for request_id in tqdm(request_batch): try: full_audio = Audio.from_s3( f"s3://suno-data-uploads/studio/uploads/{request_id}.mp3", n_channels=2, ) ear_quality_scores, _ = ear_system.get_score( full_audio, return_scores=True ) # keep the first index for easy filtering results[request_id] = [ round(current_score, 4) for current_score in ear_quality_scores ] except Exception as e: print(f"Error processing request {request_id}: {e}") traceback.print_exc() results[request_id] = None return results except Exception as e: print(f"Error in process_request_batch: {e}") traceback.print_exc() return {} # test case def main_test(): ## Example usage test_batch = [ "df9e8d5a-720e-4382-a10a-875b9b1e488a", # negative_id "30c532ed-9422-4db3-8f4c-2b56c3aed34a", # positive_id ] test_result = process_request_batch(test_batch) print(test_result) return def main(): total_job_n_gpus = 8 parser = argparse.ArgumentParser() parser.add_argument("--input_file_name", type=str) parser.add_argument("--job_idx", type=int, default=0) args = parser.parse_args() print(f"CUDA_VISIBLE_DEVICES: {os.environ['CUDA_VISIBLE_DEVICES']}") start_time = time.time() # Load pretrained ear modelimport pandas as pd # input_folder_path = "/home/tony/Data/Preference/up_v1" # input_file_name = "interesting_clips_up_u_1_20250125_full.pkl" input_folder_path = "/app/suno/data/dpo/13b_s32_v29/quality" # input_file_name = "interesting_clips_up_u_4_20250215_full.pkl" input_file_name = "../person_info.json" output_file_name = "full_pair_quality_ld" with open(f"{input_folder_path}/{input_file_name}", "r") as f: person_info = json.load(f) request_jobs = sorted(person_info.keys()) # df = pd.read_pickle(f"{input_folder_path}/{input_file_name}") # df = df.sort_values(by=["request_id", "preference"]) # print("Preference data shape", df.shape) # request_jobs = [] # assert df["preference"].nunique() == 2 # current_requests = [] # for row_id, row in df.iterrows(): # if row["preference"] == 0 and row_id % 2 == 0: # current_requests.append(row["request_id"]) # current_requests.append(row["s3_id"]) # elif row["preference"] == 1 and row_id % 2 == 1: # current_requests.append(row["s3_id"]) # request_jobs.append(current_requests) # current_requests = [] # else: # raise ValueError(f"Invalid preference: {row['preference']}") request_jobs = sorted(request_jobs) print(f"Total jobs (requests): {len(request_jobs)}") # find the specific chunk request_jobs = request_jobs[ (len(request_jobs) // total_job_n_gpus) * args.job_idx : ( len(request_jobs) // total_job_n_gpus ) * (args.job_idx + 1) ] print(f"Chunked jobs (requests): {len(request_jobs)}") if os.path.exists(f"{input_folder_path}/{output_file_name}.json"): with open(f"{input_folder_path}/{output_file_name}.json", "r") as f: known_results = json.load(f) else: known_results = {} print( f"Pre-filtered jobs: {len(request_jobs)}, known results: {len(known_results)}" ) request_jobs = [job for job in request_jobs if job not in known_results.keys()] print(f"Total jobs: {len(request_jobs)}") # only when you debug # request_jobs = request_jobs[:20] # Using process_map from tqdm.contrib.concurrent for better parallelization # Split request_jobs into batches n_processes = 8 batch_size = len(request_jobs) // n_processes + 1 batches = [ request_jobs[i : i + batch_size] for i in range(0, len(request_jobs), batch_size) # for i in range(0, 10, batch_size) ] print(f"Total batches: {len(batches)}, batch size: {batch_size}") with Pool(processes=n_processes) as pool: results = pool.map(process_request_batch, batches) # Combine results from all batches combined_results = {} for batch_result in results: combined_results.update(batch_result) known_results.update(combined_results) # with open(f"{input_folder_path}/pair_quality_{args.job_idx}.json", "w") as f: if total_job_n_gpus > 1: with open( f"{input_folder_path}/{output_file_name}_{args.job_idx}.json", "w" ) as f: json.dump(known_results, f, indent=4) else: with open(f"{input_folder_path}/{output_file_name}.json", "w") as f: json.dump(known_results, f, indent=4) print( f"DONE!! Processed {len(combined_results)} -- to total {len(known_results)}. Total time: {round(time.time() - start_time, 2)}s" ) if __name__ == "__main__": # Make sure export is setup: export CUDA_VISIBLE_DEVICES=5 main() # main_test()