import os import gc import json import uuid import tqdm import time import torch import funcy import random import tempfile import collections import numpy as np import pandas as pd from joblib import Parallel, delayed from matplotlib import pyplot as plt from scipy.io import wavfile from suno_utils.audio import Audio from suno_utils.tasks.data_loader import load_audio_mp from suno_utils.utils.text import write_jsonl, read_jsonl, write_json, read_json from suno_utils.utils.s3 import read_from_s3, check_s3_file_exists, open_from_s3 from joblib.externals.loky import get_reusable_executor def _convert_float_audio(sig): dtype = np.int16 dtype_info = np.iinfo(dtype) abs_max = 2 ** (dtype_info.bits - 1) offset = dtype_info.min + abs_max return (sig * abs_max + offset).clip(dtype_info.min, dtype_info.max).astype(dtype) def load_metas(filepath, simple=True): assert filepath.startswith("s3://") data = [] with open_from_s3(filepath) as f: for line in f: line = line.strip() if len(line) == 0: continue m = json.loads(line) _id = m["id"] # duration_s = m["duration_s"] filepath = m.get("s3_filepath", m.get("audio_filepath", m.get("filepath"))) assert filepath is not None m_new = { "id": _id, "filepath": filepath, # "duration_s": duration_s, } if not simple: for k in m.keys(): # if k not in ["id", "duration_s", "s3_filepath", "audio_filepath", "filepath"]: if k not in ["id", "s3_filepath", "audio_filepath", "filepath"]: m_new[k] = m[k] data.append(m_new) return data def _write_item(work_item): audio_arr, out_filepath = work_item try: wavfile.write(out_filepath, SAMPLE_RATE, audio_arr.T) except: return False return True def _mp_write(audio_arr_list, out_filepaths, num_workers=16): assert len(audio_arr_list) == len(out_filepaths) work_items = list(zip(audio_arr_list, out_filepaths)) confirmed_list = Parallel(n_jobs=num_workers, prefer="threads", batch_size=1)( delayed(_write_item)(work_item) for work_item in work_items ) get_reusable_executor().shutdown(wait=True) return confirmed_list def build_dataset( raw_metas_info, sample_rate: int = 48000, base_s3_dir: str = "s3://suno-data/datasets", is_test: bool = True, ): filepath_info = {} for dset_key, rel_fp, _, _ in tqdm.tqdm(raw_metas_info): metas = load_metas(os.path.join(base_s3_dir, rel_fp)) random.shuffle(metas) filepath_info[dset_key] = metas if is_test: metas_info = [(a, b, c, 1) for a, b, c, d in raw_metas_info] else: metas_info = [e for e in raw_metas_info] sample_metas_tr = [] sample_metas_val = [] is_finished = False # for each dset, load in chunks of 1k files, and keep doing until we have enough for dset_key, _, (req_min_s, req_max_s), req_duration_h in metas_info: dset_duration_h = 0 for n_iter, filepaths_chunk in enumerate( funcy.chunks(5000, filepath_info[dset_key]) ): is_val = n_iter == 0 if is_val: # val, so make it smaller filepaths_chunk = filepaths_chunk[:500] out_dir = os.path.join(out_data_dir, "val" if is_val else "train", dset_key) os.makedirs(out_dir, exist_ok=True) filepaths_chunk_list = [m["filepath"] for m in filepaths_chunk] t0 = time.time() audio_arr_list = load_audio_mp( filepaths_chunk_list, target_sample_rate=sample_rate, n_channels=2, # min_duration_s=req_min_s, # TODO: not here cause we want to skip max_duration_s=req_max_s, normalize_volume=True, num_workers=32, force_threads=False, # debug=False, silent=True, ) # add random offset just incase filtered_audio_arr_list = [] offset_list = [] for arr in audio_arr_list: if arr is None: filtered_audio_arr_list.append(arr) offset_list.append(0) continue offset = int( round(random.uniform(0, arr.shape[-1] // 4 / sample_rate), 1) * sample_rate ) filtered_audio_arr_list.append(arr[:, offset:]) offset_list.append(offset / sample_rate) audio_arr_list = filtered_audio_arr_list del filtered_audio_arr_list audio_arr_list = [ _convert_float_audio(arr.numpy()) if arr is not None else None for arr in audio_arr_list ] td_fetch = int(round(time.time() - t0)) time.sleep(5) # make sure things close # multicore writing t0 = time.time() n_offset = len(sample_metas_val) if is_val else len(sample_metas_tr) new_ids = [str(uuid.uuid4()) for _ in filepaths_chunk] out_filepaths = [ os.path.join(out_dir, f"{new_id}.wav") for new_id in new_ids ] confirmed_list = _mp_write(audio_arr_list, out_filepaths) tot_duration_s = 0 for is_confirmed, new_id, fp, m, offset_s, arr in zip( confirmed_list, new_ids, out_filepaths, filepaths_chunk, offset_list, audio_arr_list, ): if not is_confirmed: continue duration_s = arr.shape[-1] / sample_rate if duration_s < req_min_s: continue new_m = { "dataset": dset_key, "id": new_id, "original_id": m["id"], "filepath": fp, "offset_s": offset_s, "duration_s": round(duration_s, 2), } if is_val: sample_metas_val.append(new_m) else: sample_metas_tr.append(new_m) tot_duration_s += duration_s chunk_duration_h = round(tot_duration_s / 60 / 60, 1) del audio_arr_list gc.collect() td_write = int(round(time.time() - t0)) time.sleep(5) # make sure things close dset_type = "val" if is_val else "train" print( f"{dset_key}: {chunk_duration_h:,} hours of data fetched in {td_fetch:,}s" f" and written in {td_write:,}s as `{dset_type}`, retained {np.mean(confirmed_list)*100:.1f}%" ) if not is_val: dset_duration_h += chunk_duration_h if dset_duration_h >= req_duration_h: print( f"done with {dset_key}, collected total of {dset_duration_h:,.1f} hours for train" ) break is_finished = True # TODO: summarize fetched data amounts from collections import defaultdict val_durations_s = defaultdict(int) tr_durations_s = defaultdict(int) for m in sample_metas_val: val_durations_s[m["dataset"]] += m["duration_s"] for m in sample_metas_tr: tr_durations_s[m["dataset"]] += m["duration_s"] for k, v in tr_durations_s.items(): print(f"{v/60/60:,.1f} hours of {k} in train") print() for k, v in val_durations_s.items(): print(f"{v/60/60:,.1f} hours of {k} in val") assert is_finished write_jsonl(sample_metas_val, os.path.join(out_data_dir, "metas_val.jsonl")) write_jsonl(sample_metas_tr, os.path.join(out_data_dir, "metas_tr.jsonl")) def make_dac_manifests(out_data_dir: str): metas_tr = read_jsonl(os.path.join(out_data_dir, "metas_tr.jsonl")) metas_val = read_jsonl(os.path.join(out_data_dir, "metas_val.jsonl")) MUSIC_DATASETS = set( [ "podcasts", "genius_hq", "youtube_music", "jamendo", "imslp", "pond5_music", "spot_genres", "tency", "shutter_music", "pond5_sfx", ] ) metas_tr = [m for m in metas_tr if m["dataset"] in MUSIC_DATASETS] metas_val = [m for m in metas_val if m["dataset"] in MUSIC_DATASETS] print(f"{sum(m['duration_s'] for m in metas_tr)/60/60:,.0f} hours in train") print(f"{sum(m['duration_s'] for m in metas_val)/60/60:,.0f} hours in val") random.seed(7007) random.shuffle(metas_tr) random.shuffle(metas_val) df_tr = pd.DataFrame([{"path": m["filepath"]} for m in metas_tr]) df_val = pd.DataFrame([{"path": m["filepath"]} for m in metas_val]) df_tr.to_csv(os.path.join(out_data_dir, "music_tr.csv"), index=False) df_val.to_csv(os.path.join(out_data_dir, "music_val.csv"), index=False) if __name__ == "__main__": SAMPLE_RATE = 48_000 is_test = False out_data_dir = "/app/suno/data/audio_2ch_48khz_lg" os.makedirs(out_data_dir, exist_ok=True) base_s3_dir = "s3://suno-data/datasets" # dset_name, metas_filepath, (min_s, max_s), max_hours raw_metas_info = [ # speech ("podcasts", "bundles/v0/podcasts/metas.jsonl", (60, 2 * 60), 2_500), # music` ("genius_hq", "bundles/v1/genius_hq/metas.jsonl", (60, 2 * 60), 5_000), ("youtube_music", "bundles/v1/youtube_music/metas.jsonl", (60, 2 * 60), 10_000), ("jamendo", "bundles/v1/jamendo/metas.jsonl", (60, 2 * 60), 1_000), ("imslp", "bundles/v1/imslp/metas.jsonl", (60, 2 * 60), 1_000), ("pond5_music", "bundles/v2/pond5_music/metas.jsonl", (60, 2 * 60), 2_500), ( "spot_genres", "harvest/spotify/ytm_spotify_genres_50_simple.jsonl", (60, 2 * 60), 10_000, ), ("tency", "harvest/tency/tency_plus_fp_flat.jsonl", (20, 2 * 60), 500), ( "shutter_music", "bundles/v3/shutter_music/metas_stems_flat.jsonl", (20, 2 * 60), 500, ), # misc ("pond5_sfx", "bundles/v3/pond5_sfx/metas.jsonl", (5, 30), 500), ] random.seed(6006) build_dataset( raw_metas_info, sample_rate=SAMPLE_RATE, base_s3_dir=base_s3_dir, is_test=is_test, ) make_dac_manifests(out_data_dir)