#!/usr/bin/env python3 """ Unified script for downloading NPZ datasets from S3. Supports different dataset types and configurations. # For auk_t0 dataset python download_npz_dataset.py \ --preset auk_t0 \ --data_path "/home/tony/Data/Preference/auk_t0/interesting_clips_auk_t0_20250527.pkl" # For auk_t1 dataset python download_npz_dataset.py \ --preset auk_t1 \ --data_path "/home/tony/Data/Preference/auk_t1/interesting_clips_auk_t1_20250527.pkl" # For vae_diffv2_d3 dataset python download_npz_dataset.py \ --preset vae_diffv2_d3 \ --data_path "/home/tony/Data/Preference/up_v2_d3/interesting_clips_ahi_d3_20250527.pkl" """ import argparse import os import pandas as pd from suno_utils.utils.s3 import download_s3_files def get_preset_config(preset_name, data_path): """Get predefined configurations for common use cases.""" presets = { "auk_t0": { "data_path": data_path, "npz_dir": "/app2/suno/data/dpo/auk_t0_npz/", "json_dir": "/app2/suno/data/dpo/auk_t0_json/", "dataset_type": "standard", "model_filter": "auk", "download_hoot": True, "download_vae": False, }, "auk_t1": { "data_path": data_path, "npz_dir": "/app2/suno/data/dpo/auk_t1_npz/", "json_dir": "/app2/suno/data/dpo/auk_t1_json/", "dataset_type": "standard", "model_filter": "auk", "download_hoot": True, "download_vae": False, }, "vae_diffv2_d3": { "data_path": data_path, "npz_dir": "/app2/suno/data/dpo/diff2_v2_d3", "json_dir": None, "dataset_type": "upsample", "model_filter": "up", "download_hoot": False, "download_vae": True, }, } return presets.get(preset_name) def parse_args(): parser = argparse.ArgumentParser(description="Download NPZ datasets from S3") # Preset mode - simplified usage parser.add_argument( "--preset", type=str, choices=["auk_t0", "auk_t1", "vae_diffv2_d3"], help="Use predefined configuration preset", ) # Required arguments parser.add_argument( "--data_path", type=str, required=True, help="Path to the input data file (.pkl or .csv)", ) # Optional arguments (for manual configuration) parser.add_argument( "--npz_dir", type=str, help="Directory to store downloaded NPZ files", ) parser.add_argument( "--json_dir", type=str, default=None, help="Directory to store downloaded JSON files (hoot files)", ) parser.add_argument( "--dataset_type", type=str, choices=["standard", "upsample", "vae"], default="standard", help="Type of dataset processing", ) parser.add_argument( "--model_filter", type=str, default=None, help="Filter for model names (e.g., 'auk', 'up')", ) parser.add_argument( "--n_cores", type=int, default=32, help="Number of cores for parallel downloading", ) parser.add_argument( "--download_vae", action="store_true", help="Download VAE files (_vae.npz)" ) parser.add_argument( "--download_hoot", action="store_true", help="Download hoot JSON files" ) parser.add_argument( "--skip_deleted", action="store_true", help="Skip trying to download from deleted folder", ) args = parser.parse_args() # Apply preset configuration if specified if args.preset: preset_config = get_preset_config(args.preset, args.data_path) if not preset_config: raise ValueError(f"Unknown preset: {args.preset}") print(f"Using preset configuration: {args.preset}") # Apply preset values, but allow command line overrides for key, value in preset_config.items(): if key == "data_path": continue # Always use the provided data_path # Only override if the argument wasn't explicitly set current_value = getattr(args, key) if key in ["download_vae", "download_hoot"]: # For boolean flags, only override if not set if not current_value: setattr(args, key, value) else: # For other arguments, override if None or not set if current_value is None: setattr(args, key, value) # Validate required arguments if not args.npz_dir: parser.error("--npz_dir is required (or use --preset)") return args def load_data(data_path): """Load data from pickle or CSV file.""" print(f"Loading data from: {data_path}") if data_path.endswith(".pkl"): df = pd.read_pickle(data_path) elif data_path.endswith(".csv"): df = pd.read_csv(data_path) else: raise ValueError("Data file must be .pkl or .csv") print(f"Loaded data shape: {df.shape}") return df def filter_dataframe(df, model_filter): """Filter dataframe by model name if specified.""" if model_filter: print(f"Filtering by model name containing: {model_filter}") df = df[df["model_name"].str.contains(model_filter)] print(f"Filtered data shape: {df.shape}") print("Model name counts:") print(df["model_name"].value_counts()) return df def extract_s3_ids(df, dataset_type): """Extract S3 IDs based on dataset type.""" print(f"Extracting S3 IDs for dataset type: {dataset_type}") if dataset_type == "upsample": # For upsample datasets, extract from metadata edit_clip_ids = df["metadata"].apply(lambda x: x.get("upsample_clip_id", "")) s3_ids = [s3_id for s3_id in edit_clip_ids if s3_id and len(s3_id) > 0] print(f"All upsample clip IDs: {len(s3_ids)}") unique_s3_ids = sorted(set(s3_ids)) print(f"Unique upsample clip IDs: {len(unique_s3_ids)}") return unique_s3_ids elif dataset_type == "vae": # For VAE datasets, use original s3_id s3_ids = df["s3_id"].unique() print(f"Unique S3 IDs for VAE: {len(s3_ids)}") return s3_ids else: # standard # For standard datasets, use s3_id values s3_ids = df["s3_id"].values print(f"Total S3 IDs: {len(s3_ids)}") return s3_ids def create_s3_paths(s3_ids, file_type="npz"): """Create S3 and local paths for downloading.""" if file_type == "npz": s3_paths = [ f"s3://suno-data-uploads/studio/uploads/{s3_id}.npz" for s3_id in s3_ids ] elif file_type == "vae": s3_paths = [ f"s3://suno-data-uploads/studio/uploads/{s3_id}_vae.npz" for s3_id in s3_ids ] elif file_type == "hoot": s3_paths = [ f"s3://suno-data-uploads/studio/uploads/{s3_id}_hoot.json" for s3_id in s3_ids ] else: raise ValueError(f"Unknown file type: {file_type}") return s3_paths def create_local_paths(s3_ids, local_dir, file_type="npz"): """Create local paths for downloaded files.""" if file_type == "npz": local_paths = [f"{local_dir}/{s3_id}.npz" for s3_id in s3_ids] elif file_type == "vae": local_paths = [f"{local_dir}/{s3_id}_vae.npz" for s3_id in s3_ids] elif file_type == "hoot": local_paths = [f"{local_dir}/{s3_id}_hoot.json" for s3_id in s3_ids] else: raise ValueError(f"Unknown file type: {file_type}") return local_paths def get_unfinished_downloads(s3_paths, local_paths, local_dir): """Get list of files that haven't been downloaded yet.""" os.makedirs(local_dir, exist_ok=True) finished_paths = os.listdir(local_dir) finished_paths_set = set(finished_paths) unfinished_s3_paths = [ path for path in s3_paths if os.path.basename(path) not in finished_paths_set ] unfinished_local_paths = [ path for path in local_paths if os.path.basename(path) not in finished_paths_set ] return unfinished_s3_paths, unfinished_local_paths def download_files(s3_paths, local_paths, n_cores, skip_deleted=False): """Download files from S3.""" if not s3_paths: print("No files to download") return [] print(f"Downloading {len(s3_paths)} files with {n_cores} cores...") # Try downloading from uploads folder first result = download_s3_files(s3_paths, local_paths, n_cores=n_cores) if not skip_deleted: # Check for remaining files and try deleted folder remaining_s3_paths = [] remaining_local_paths = [] for s3_path, local_path in zip(s3_paths, local_paths): if not os.path.exists(local_path): remaining_s3_paths.append(s3_path) remaining_local_paths.append(local_path) if remaining_s3_paths: print( f"Trying to download {len(remaining_s3_paths)} files from deleted folder..." ) deleted_s3_paths = [ path.replace("/uploads/", "/deleted/") for path in remaining_s3_paths ] download_s3_files(deleted_s3_paths, remaining_local_paths, n_cores=n_cores) return result def main(): args = parse_args() # Load and filter data df = load_data(args.data_path) df = filter_dataframe(df, args.model_filter) # Extract S3 IDs s3_ids = extract_s3_ids(df, args.dataset_type) # Download NPZ files print("\n=== Downloading NPZ files ===") file_type = "vae" if args.download_vae else "npz" s3_paths = create_s3_paths(s3_ids, file_type) local_paths = create_local_paths(s3_ids, args.npz_dir, file_type) unfinished_s3_paths, unfinished_local_paths = get_unfinished_downloads( s3_paths, local_paths, args.npz_dir ) print(f"Jobs to be done: {len(unfinished_s3_paths)}") if unfinished_s3_paths: download_files( unfinished_s3_paths, unfinished_local_paths, args.n_cores, args.skip_deleted ) print("Finished downloading NPZ files") # Download VAE files if requested and not already done if args.download_vae and file_type != "vae": print("\n=== Downloading VAE files ===") # For VAE files, we need the original s3_id, not upsample_clip_id if args.dataset_type == "upsample": vae_s3_ids = df["s3_id"].unique() else: vae_s3_ids = s3_ids vae_s3_paths = create_s3_paths(vae_s3_ids, "vae") vae_local_paths = create_local_paths(vae_s3_ids, args.npz_dir, "vae") unfinished_vae_s3_paths, unfinished_vae_local_paths = get_unfinished_downloads( vae_s3_paths, vae_local_paths, args.npz_dir ) print(f"VAE jobs to be done: {len(unfinished_vae_s3_paths)}") if unfinished_vae_s3_paths: download_files( unfinished_vae_s3_paths, unfinished_vae_local_paths, args.n_cores, args.skip_deleted, ) print("Finished downloading VAE files") # Download hoot JSON files if requested if args.download_hoot and args.json_dir: print("\n=== Downloading Hoot JSON files ===") hoot_s3_paths = create_s3_paths(s3_ids, "hoot") hoot_local_paths = create_local_paths(s3_ids, args.json_dir, "hoot") unfinished_hoot_s3_paths, unfinished_hoot_local_paths = ( get_unfinished_downloads(hoot_s3_paths, hoot_local_paths, args.json_dir) ) print(f"Hoot jobs to be done: {len(unfinished_hoot_s3_paths)}") if unfinished_hoot_s3_paths: download_files( unfinished_hoot_s3_paths, unfinished_hoot_local_paths, args.n_cores, args.skip_deleted, ) print("Finished downloading hoot files") print("\n=== All downloads completed ===") if __name__ == "__main__": main()