from suno_utils.worker.settings import s3_client from pathlib import Path from concurrent.futures import ThreadPoolExecutor, as_completed from tqdm import tqdm from zipfile import ZipFile, BadZipFile import numpy as np import io import json import random import bisect def _write_beats(f, data): downbeats = list(sorted(data["downbeat_times"].tolist())) beats = list(sorted(data["beat_times"].tolist())) bars = [] for left_downbeat, right_downbeat in zip( [-float("inf")] + downbeats, downbeats + [float("inf")] ): start_idx = bisect.bisect_left(beats, left_downbeat) end_idx = bisect.bisect_left(beats, right_downbeat) beats_in_bar = end_idx - start_idx if beats_in_bar == 0: continue bars.append(beats[start_idx:end_idx]) result = [] rest_idx = 0 if len(bars) > 1: # fix up the first bar offset = len(bars[1]) - len(bars[0]) if offset < 0: offset = 0 result.extend([b, i + 1 + offset] for i, b in enumerate(bars[0])) rest_idx = 1 for bar in bars[rest_idx:]: result.extend([b, i + 1] for i, b in enumerate(bar)) for b, d in result: f.write(f"{b} {d}\n") def main(suffix: str): paginator = s3_client.get_paginator("list_objects_v2") pages = paginator.paginate( Bucket="suno-data", Prefix=f"m4burns/downbeats_processed/spectrograms_round5_{suffix}/", ) data_dir = Path(__file__).parent.parent / "data" # temp_dir = data_dir / "suno_download" temp_dir = Path(f"/mnt/localdisk/tmp_m4burns/beat_this/{suffix}") temp_dir.mkdir(parents=True, exist_ok=True) spect_uuids = [] def download_file(obj): if not obj["Key"].endswith(".npz"): return None npz_name = obj["Key"].split("/")[-1] npz_out_path = temp_dir / npz_name spect_uuids.append(npz_name.split(".")[0]) if npz_out_path.exists(): return None try: with open(npz_out_path, "wb") as f: s3_client.download_fileobj( Bucket="suno-data", Key=obj["Key"], Fileobj=f ) except Exception: npz_out_path.unlink() raise return npz_name with ThreadPoolExecutor(max_workers=16) as executor: futures = [ executor.submit(download_file, obj) for page in pages for obj in page.get("Contents", []) ] for future in tqdm( as_completed(futures), total=len(futures), desc="Downloading files" ): future.result() suno_ann_dir = data_dir / "annotations" / f"suno_synth_{suffix}" suno_ann_dir.mkdir(parents=True, exist_ok=True) with open(suno_ann_dir / "info.json", "w") as f: json.dump({"has_downbeats": True}, f) with open(suno_ann_dir / "single.split", "w") as f: for u in spect_uuids: split = "train" if random.random() < 0.9 else "val" f.write(f"{u}\t{split}\n") beats_dir = suno_ann_dir / "annotations" / "beats" beats_dir.mkdir(parents=True, exist_ok=True) with ZipFile( data_dir / "audio" / "spectrograms" / f"suno_synth_{suffix}.npz", "w" ) as z: for u in tqdm(spect_uuids, desc="Writing bundle+annotations"): try: data = np.load(temp_dir / f"{u}.npz") except (BadZipFile, EOFError): print(f"Bad zip file: {u}") continue buf = io.BytesIO() spec_array = data["spec_array"] if spec_array.ndim == 3 and spec_array.shape[-1] == 1: spec_array = spec_array.squeeze(-1) np.save(buf, spec_array.astype(np.float16)) z.writestr(f"{u}/track.npy", buf.getvalue()) with open(beats_dir / f"{u}.beats", "w") as f: _write_beats(f, data) if __name__ == "__main__": import sys assert len(sys.argv) == 2, "Usage: python download_suno.py " main(sys.argv[1])