import os import json import funcy import boto3 import shutil import numpy as np from tqdm import tqdm from concurrent.futures import ThreadPoolExecutor, as_completed from suno_utils.utils.text import read_jsonl, write_jsonl from suno_utils.utils.s3 import read_from_s3, _verify_s3_filepath def get_s3_files(bucket_name, prefix, max_keys: int = 100000): all_files = [] continuation_token = None while True: # Prepare the arguments for the request list_kwargs = { "Bucket": bucket_name, "Prefix": prefix, # List objects under this prefix, or leave blank for all objects } if continuation_token: list_kwargs["ContinuationToken"] = continuation_token # Make the request to list objects response = s3.list_objects_v2(**list_kwargs) # Collect the file keys all_files += [obj["Key"] for obj in response.get("Contents", [])] # Check if more results are available if response.get("IsTruncated"): # True if there are more results to fetch continuation_token = response["NextContinuationToken"] else: break # No more results to fetch return all_files def list_s3_directories(bucket_name: str, prefix: str = ""): """ List all directories (prefixes) in an S3 bucket. Args: bucket_name (str): Name of the S3 bucket prefix (str): Optional prefix to filter results (like a directory path) Returns: List[str]: List of directory paths (prefixes) """ s3_client = boto3.client("s3") directories = set() # Use paginator to handle buckets with many objects paginator = s3_client.get_paginator("list_objects_v2") page_iterator = paginator.paginate(Bucket=bucket_name, Prefix=prefix, Delimiter="/") # Collect all prefixes (directories) for page in page_iterator: # Get common prefixes (directories) if "CommonPrefixes" in page: for prefix_obj in page["CommonPrefixes"]: directories.add(prefix_obj["Prefix"]) # Also check Contents for any directory-like objects if "Contents" in page: for obj in page["Contents"]: key = obj["Key"] # If the key contains a slash, add the directory part if "/" in key: directory = key.rsplit("/", 1)[0] + "/" directories.add(directory) return sorted(list(directories)) if __name__ == "__main__": # load the base metas VAE_DIM = 128 VAE_RATE_HZ = 25 VAE_MEMMAP_SIZE = 750 SEMANTIC_VOCAB_SIZE = 4000 SEMANTIC_MEMMAP_SIZE = 750 VAL_SIZE = 1000 USE_PAIRS = True # use pairs instead of single examples USE_AUDIO_QUALITY = False OUT_DATA_DIR = "/app/suno/data/diffusion_ft/genius_hq_filtered_20k_25hz_20241115_v1" if os.path.exists(OUT_DATA_DIR): print(f"out data dir {OUT_DATA_DIR} already exists, deleting") shutil.rmtree(OUT_DATA_DIR) os.makedirs(OUT_DATA_DIR, exist_ok=True) # christian/data/upsample_100hz_v4_t_5_20241018 bucket_name = "suno-data" base_dir = "christian/data/genius_hq_filtered_20k" output_name = "25hz_20241115_v2/" base_metas_path = os.path.join("s3://", bucket_name, base_dir, "metas.jsonl") base_metas = read_from_s3(base_metas_path, read_f=read_jsonl) print(len(base_metas)) base_metas_map = {meta["id"]: meta for meta in base_metas} # s3 client s3 = boto3.client("s3") # load quality scores # quality_scores_path = "/home/christian/code/christian/notebooks/diff_dpo/upsample_v4_t_5_20241018_25hz_20241031_v1_quality_scores.json" # with open(quality_scores_path, "r") as f: # quality_scores = json.load(f) # get all dir paths on s3 dir_paths = list_s3_directories(bucket_name, f"{base_dir}/{output_name}") print("total dirs: ", len(dir_paths)) valid_dir_paths = dir_paths # for dir_path in tqdm(dir_paths): # if dir_path in quality_scores: # valid_dir_paths.append(dir_path) # print( # f"total dir paths with quality scores: {len(valid_dir_paths)}/{len(dir_paths)}" # ) # split dir paths into train and val train_dir_paths = valid_dir_paths[:-VAL_SIZE] val_dir_paths = valid_dir_paths[-VAL_SIZE:] # now iterate over the val, then train metas for dset_type in ["val", "train"]: dset_dir_paths = val_dir_paths if dset_type == "val" else train_dir_paths out_mm_vae_filepath = os.path.join(OUT_DATA_DIR, f"data_vae_{dset_type}.bin") out_metas_filepath = os.path.join(OUT_DATA_DIR, f"metas_{dset_type}.jsonl") out_mm_semantic_filepath = os.path.join( OUT_DATA_DIR, f"data_semantic_{dset_type}.bin" ) # initial write out_mm_semantic = np.memmap( out_mm_semantic_filepath, dtype=np.uint16, mode="w+", shape=(1), ) out_mm_vae = np.memmap( out_mm_vae_filepath, dtype=np.float16, mode="w+", shape=(1), ) n_offs_s = 0 n_offs_v = 0 CHUNK_SIZE = 100 # split valid metas into chunks of CHUNK_SIZE valid_dir_paths_chunks = [ valid_dir_paths[i : i + CHUNK_SIZE] for i in range(0, len(dset_dir_paths), CHUNK_SIZE) ] print("total chunks: ", len(valid_dir_paths_chunks)) for chunk_idx, dir_path_chunk in enumerate(tqdm(valid_dir_paths_chunks)): arr_s_list = [] arr_v_list = [] new_metas = [] def process_dir_path(dir_path): # get quality scores # dir_path_quality_scores = quality_scores[dir_path] meta_id = dir_path.strip("/").split("/")[-1] # get meta meta = base_metas_map[meta_id] # select the paths with the highest and lowest quality scores # sorted_s3_paths = sorted( # dir_path_quality_scores, # key=lambda x: dir_path_quality_scores[x], # reverse=True, # ) # positive_s3_path = sorted_s3_paths[0] # negative_s3_path = sorted_s3_paths[-1] # get files in dir files = get_s3_files(bucket_name, dir_path) npz_filepath = [f for f in files if f.endswith(".npz")][0] # load the npz file data = read_from_s3( f"s3://{bucket_name}/{npz_filepath}", read_f=np.load ) arr_s = data["semantic_codes"] arr_v_neg = data["upsampled_latents"] arr_v_pos = data["original_latents"] try: result = [] for arr_v in [arr_v_neg, arr_v_pos]: if arr_s.size < SEMANTIC_MEMMAP_SIZE: return None if arr_v.size < VAE_MEMMAP_SIZE * VAE_DIM: return None result.append((arr_s, arr_v, meta)) return result except Exception as e: print(f"error loading {dir_path}: {e}") return None with ThreadPoolExecutor(max_workers=16) as executor: futures = [ executor.submit(process_dir_path, dir_path) for dir_path in dir_path_chunk ] for future in as_completed(futures): result = future.result() # list of tuples (arr_s, arr_v, meta) if result: for arr_s, arr_v, meta in result: arr_s_list.append(arr_s) arr_v_list.append(arr_v) new_meta = meta.copy() new_meta["text"] = meta["lyrics"] new_meta["tags"] = meta["tags_text"] new_meta["n_vae_tokens"] = VAE_MEMMAP_SIZE new_metas.append(new_meta) print(len(arr_s_list), len(arr_v_list), len(new_metas)) assert len(arr_s_list) == len(arr_v_list) == len(new_metas) # now write to the memmap # get a list of all the ids in the id_to_s3_paths to_write_len_s = SEMANTIC_MEMMAP_SIZE * len(arr_v_list) to_write_len_v = VAE_MEMMAP_SIZE * VAE_DIM * len(arr_v_list) out_mm_semantic = np.memmap( out_mm_semantic_filepath, dtype=np.uint16, mode="r+", shape=(n_offs_s + to_write_len_s,), ) out_mm_vae = np.memmap( out_mm_vae_filepath, dtype=np.float16, mode="r+", shape=(n_offs_v + to_write_len_v,), ) # write to memmap (has to happen sequentially) for new_meta, arr_s, arr_v in zip(new_metas, arr_s_list, arr_v_list): out_mm_semantic[n_offs_s : n_offs_s + arr_s.size] = arr_s.reshape( -1, ) out_mm_vae[n_offs_v : n_offs_v + arr_v.size] = arr_v.reshape( -1, ) n_offs_s += arr_s.size n_offs_v += arr_v.size # write it once out_mm_semantic.flush() out_mm_vae.flush() del out_mm_semantic, out_mm_vae write_jsonl(new_metas, out_metas_filepath, do_append=True)