import os
import glob
import tqdm
import argparse
import numpy as np
from suno_utils.audio import Audio


# get file list
filelist = glob.glob("/app/suno/data/audio_mono_24khz/msd/audio/*.wav")
print("%d files exsist" % len(filelist))

# parse
parser = argparse.ArgumentParser()
parser.add_argument("--index", required=True, help="partial index")
parser.add_argument("--num_partial", required=True, help="number of partial data")
args = parser.parse_args()
index = int(args.index)
num_partial = int(args.num_partial)

# split data
hop = len(filelist) // num_partial
if index == num_partial:
    filelist = filelist[hop * index - 1:]
else:
    filelist = filelist[hop * index : (index + 1) * hop]

# resample audio
save_path = "/app/suno/data/audio_mono_24khz/msd/splits"

dur = np.zeros(len(filelist))
for ix, fn in tqdm.tqdm(enumerate(filelist)):
    audio = Audio.from_file(fn, sample_rate=24000)
    dur[ix] = audio.duration_s
dur_fn = os.path.join(save_path, "dur_%s.npy" % index)
id_fn = os.path.join(save_path, "id_%s.npy" % index)
np.save(open(dur_fn, "wb"), dur)
np.save(open(id_fn, "wb"), filelist)
