import os os.environ["OPENBLAS_NUM_THREADS"] = "1" import json import numpy as np import argparse from tqdm.contrib.concurrent import process_map, thread_map from multiprocessing import cpu_count import librosa from suno_utils.utils.text import read_jsonl from suno_utils.audio import Audio genius_data = read_jsonl("/app/suno/data/high_freq/genius_meta.jsonl") print(len(genius_data)) n_worker = cpu_count() - 2 n_batch_size = 1000 def get_high_frequency_cutoff(job): try: job_id = list(job.keys())[0] job_file_path = job[job_id] audio = Audio.from_s3(job_file_path) # spectral roll-off roll_off = librosa.feature.spectral_rolloff( y=audio.array_float.astype("float32"), sr=48000, n_fft=4096, hop_length=480, roll_percent=0.99, )[0] # max pooling with 30s window window_step = 100 pooled_roll_off = np.array( [ np.max(roll_off[ix : ix + 3000]) for ix in range(0, len(roll_off) - 3000 + window_step, window_step) ] ) roll_off_median = np.median(pooled_roll_off) # get pass / fail # pass_fail = "pass" if roll_off_median > 17000 else "fail" return {job_id: roll_off_median} except: return {job_id: 0} def get_high_frequency_cutoff_within_range(start_index=0, length=n_batch_size): output_json_path = f"/app/suno/data/high_freq/{start_index}.json" if os.path.exists(output_json_path): return # create a fake place holder with open(output_json_path, "w") as fp: json.dump({}, fp) output_result = {} jobs = [] for index in range(start_index, min(start_index + length, len(genius_data))): jobs.append({genius_data[index]["id"]: genius_data[index]["audio_filepath"]}) results = process_map( get_high_frequency_cutoff, jobs, max_workers=n_worker, chunksize=1, ) for result in results: output_result.update(result) with open(output_json_path, "w") as fp: json.dump(output_result, fp) def parse_args(): parser = argparse.ArgumentParser() parser.add_argument("--start_index", type=int, default=0) parser.add_argument("--end_index", type=int, default=len(genius_data)) args = parser.parse_args() return args if __name__ == "__main__": input_args = parse_args() print(f"Start!!! {input_args.start_index, input_args.end_index}") for i in range(input_args.start_index, input_args.end_index, n_batch_size): get_high_frequency_cutoff_within_range(i, n_batch_size) print(f"DONE!!! {input_args.start_index, input_args.end_index}")