import os import torch import argparse import torchaudio import numpy as np import pandas as pd import pyloudnorm as pyln import multiprocessing as mp from tqdm import tqdm from suno_utils.audio import Audio from ear.system import EarSystem from suno_boost.utils import apply_normalization if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument( "--csv_path", default="/home/tony/Data/Preference/7b_v0/interesting_clips_v3_processed.csv", type=str, ) parser.add_argument( "--output_dir", default="/app/suno/christian/data/interesting_clips_v3_processed", ) parser.add_argument("--num_examples", type=int, default=-1) args = parser.parse_args() if not os.path.isdir(args.output_dir): os.makedirs(args.output_dir) # create func def download_from_s3_and_save(s3_id: str): filepath = os.path.join(args.output_dir, s3_id + ".mp3") if os.path.isfile(filepath): return else: try: pos_audio = Audio.from_s3( f"s3://suno-data-uploads/studio/uploads/{s3_id}.mp3" ) except: return pos_audio.write_mp3(filepath) # load csv df = pd.read_csv(args.csv_path) print(df.shape) # get pairs df = df.sort_values(by=["request_id", "preference"]) print(df.head()) s3_ids = [] for idx in tqdm(np.arange(0, len(df), 1)): # idx = np.random.randint(0, len(df)) # idx += idx % 2 row = df.iloc[idx] # pos_row = df.iloc[idx + 1] s3_ids.append(row["s3_id"]) print(len(s3_ids)) with mp.Pool(32) as pool: pool.map(download_from_s3_and_save, s3_ids)