import librosa import numpy as np import pickle import requests import os import time import io import pandas as pd import json import tqdm from multiprocessing import Pool from suno_utils.audio import Audio # git+https://github.com/alphacsc/alphacsc.git@843990555914f2ebbcaef5fc40647d112ed3a51e try: import warnings from numba.core.errors import NumbaPerformanceWarning warnings.filterwarnings("ignore", category=NumbaPerformanceWarning) except ImportError: pass # s3.download_file("suno-data", "m4burns/csc_model.pkl", csc_model_path) csc_model_path = os.path.join("/app/suno/data/dpo/models", "csc_model.pkl") def _load_csc_model(): return pickle.load(open(csc_model_path, "rb")) model_dict = _load_csc_model() csc_model = model_dict["csc"] shimmer_component_idx = model_dict["shimmer_component_idx"] sr = model_dict["sr"] low_cutoff = model_dict["low_cutoff"] high_cutoff = model_dict["high_cutoff"] window_size = model_dict["window_size"] hop_length = model_dict["hop_length"] def load_audio(src): y, _ = librosa.load(src, sr=sr, mono=True) return y def spectrogram(signal): window = np.hanning(window_size) low_bin = int(low_cutoff * window_size / sr) high_bin = int(high_cutoff * window_size / sr) hops = signal.shape[0] // hop_length real_spec = np.zeros((hops, high_bin - low_bin + 1)) for i in range(hops): start = i * hop_length if start + window_size > signal.shape[0]: break real_spec[i, :] = 10 * np.log10( np.abs( np.fft.fft(signal[start : start + window_size] * window)[ low_bin : high_bin + 1 ] ) + 1e-6 ) return real_spec def eval_csc(sol, real_specs): return sol.transform(real_specs.transpose(0, 2, 1)) def find_peaks(activation): threshold = np.max(activation) * 0.5 peaks = [] pos = 0 while pos < len(activation): if activation[pos] > threshold: start = pos while pos < len(activation) and activation[pos] > threshold: pos += 1 peaks.append(np.argmax(activation[start:pos]) + start) pos += 1 return peaks def get_shimmer_score_from_audio_array(audio_array): # start_time = time.time() real_spec = spectrogram(audio_array) # print(f"Time taken spectrogram: {round(time.time() - start_time, 2)} seconds") activations = eval_csc(csc_model, real_spec[None, :, :]) # print(f"Time taken eval_csc: {round(time.time() - start_time, 2)} seconds") shimmer_activation = activations[0][shimmer_component_idx] # print( # f"Time taken shimmer_activation: {round(time.time() - start_time, 2)} seconds" # ) peaks = find_peaks(shimmer_activation) # print(f"Time taken find_peaks: {round(time.time() - start_time, 2)} seconds") score = np.sum(shimmer_activation[peaks]) / (len(shimmer_activation) + 0.00001) # print(f"Time taken score: {round(time.time() - start_time, 2)} seconds") return score def shimmer_score(path_or_io): signal = load_audio(path_or_io) return get_shimmer_score_from_audio_array(signal) def shimmer_score_url(url): return shimmer_score(io.BytesIO(requests.get(url).content)) def shimmer_core_from_s3_id(s3_id): try: # start_time = time.time() mp3_filepath = f"s3://suno-data-uploads/studio/uploads/{s3_id}.mp3" audio = Audio.from_s3(mp3_filepath, n_channels=2) audio = audio.get_segment(0, 30) audio_array = audio.convert( sample_rate=sr, byte_width=2, n_channels=1 ).array_float # print(f"Time taken: {round(time.time() - start_time, 2)} seconds") score = get_shimmer_score_from_audio_array(audio_array) # print(f"Time taken: {round(time.time() - start_time, 2)} seconds") # print(f"Shimmer score: {score}") return score except Exception as e: # likely the clip has been deleted from s3 print(f"Error processing {s3_id}: {e}") return 100 # if __name__ == "__main__": # # # debug # arg = sys.argv[1] # print(shimmer_core_from_s3_id(arg)) if __name__ == "__main__": # # debug # arg = sys.argv[1] # print(shimmer_core_from_s3_id(arg)) start_time = time.time() input_folder_path = "/home/tony/Data/Preference/up_v2" input_file_name = "interesting_clips_up_u_2_20241210_full.pkl" df = pd.read_pickle(f"{input_folder_path}/{input_file_name}") valid_s3_ids = df["s3_id"].tolist() shimmer_results = {} with open(f"{input_folder_path}/total_shimmer_scores.json", "r") as fp: known_shimmer_results = json.load(fp) print(f"Total known shimmer scores: {len(known_shimmer_results)}") need_to_process_s3_ids = [ s3_id for s3_id in valid_s3_ids if s3_id not in known_shimmer_results ] # debug # for s3_id in tqdm.tqdm(need_to_process_s3_ids[:4]): # shimmer_results[s3_id] = shimmer_core_from_s3_id(s3_id) # from tqdm.contrib.concurrent import process_map from multiprocessing.pool import ThreadPool with ThreadPool(processes=64) as pool: results = list( tqdm.tqdm( pool.imap( shimmer_core_from_s3_id, need_to_process_s3_ids, chunksize=1, ), total=len(need_to_process_s3_ids), desc="Processing audio files", ) ) shimmer_results.update(dict(zip(need_to_process_s3_ids, results))) known_shimmer_results.update(shimmer_results) with open(f"{input_folder_path}/total_shimmer_scores.json", "w") as fp: json.dump(known_shimmer_results, fp, indent=4) print( f"Done! {len(known_shimmer_results)} shimmer scores saved. Total new {len(shimmer_results)}, total time {round(time.time() - start_time, 2)} seconds" )