import os import json import random import time import pandas as pd from tqdm import tqdm from suno_utils.utils.text import write_jsonl from joblib import Parallel, delayed from typing import Dict, List, Callable, Any, Optional, Tuple from dataclasses import dataclass, field # ============================================================================ # Utility Functions # ============================================================================ def read_jsonl_lazy(filepath): """Lazy generator to read JSONL file line by line without loading everything into memory""" with open(filepath, "r", encoding="utf-8") as f: for line in f: line = line.strip() if line: try: yield json.loads(line) except json.JSONDecodeError: continue def count_jsonl_lines(filepath): """Count lines in a JSONL file efficiently""" count = 0 with open(filepath, "r", encoding="utf-8") as f: for _ in f: count += 1 return count # ============================================================================ # Metadata Loader Functions # ============================================================================ def load_json_dict(filepath: str, transform: Optional[Callable] = None) -> Dict: """Load a JSON file as a dictionary""" with open(filepath, "r") as f: data = json.load(f) if transform: data = transform(data) return data def load_jsonl_dict(filepath: str, key_field: str = "id") -> Dict: """Load a JSONL file as a dictionary indexed by key_field""" return {meta[key_field]: meta for meta in read_jsonl_lazy(filepath)} def load_csv_dict(filepath: str, key_field: str = "id") -> Dict: """Load a CSV file as a dictionary indexed by key_field""" df = pd.read_csv(filepath) return df.set_index(key_field).to_dict(orient="index") def load_alignment_map(filepath: str) -> Dict: """Load alignment data from JSONL file""" result = {} for item in tqdm(read_jsonl_lazy(filepath), desc="Building alignment map"): k, v, cer = item meta = v[0] lines, starts, ends = ( meta["line_text"], meta["line_start_s"], meta["line_end_s"], ) line_entries = [] for text, start, end in zip(lines, starts, ends): if start is None or end is None: continue line_entries.append((start, end, text)) result[k] = { "lines": line_entries, "cer": cer, "text": meta.get("text"), "start_s": meta.get("start_s"), "end_s": meta.get("end_s"), "vocal_start_s": meta.get("vocal_start_s"), "vocal_end_s": meta.get("vocal_end_s"), } return result # ============================================================================ # Metadata Configuration # ============================================================================ @dataclass class MetadataSource: """Configuration for a metadata source""" name: str # Name for this metadata source filepath: str # Path to the metadata file loader: Callable # Function to load the metadata loader_kwargs: Dict[str, Any] = field(default_factory=dict) # Arguments for loader merge_key: str = "id" # Key to use for merging (default: "id") merge_function: Optional[Callable] = None # Custom merge function if needed output_key: Optional[str] = None # Key to store in output (default: same as name) priority: int = 0 # Priority for merging (lower = higher priority, for conflicts) enabled: bool = True # Whether to enable this metadata source @dataclass class MetadataRegistry: """Registry of all metadata sources""" sources: Dict[str, MetadataSource] = field(default_factory=dict) loaded_data: Dict[str, Any] = field(default_factory=dict) def register(self, source: MetadataSource): """Register a metadata source""" self.sources[source.name] = source def load_all(self): """Load all enabled metadata sources""" print("Loading external metadata...") for name, source in self.sources.items(): if not source.enabled: continue print(f" Loading {name}...") try: self.loaded_data[name] = source.loader( source.filepath, **source.loader_kwargs ) size = ( len(self.loaded_data[name]) if isinstance(self.loaded_data[name], dict) else "N/A" ) print(f" Loaded {name}: {size:,} entries") except Exception as e: print(f" WARNING: Failed to load {name}: {e}") self.loaded_data[name] = {} def merge_into_meta(self, meta: Dict, meta_id: str) -> Dict: """Merge all registered metadata into a meta dictionary""" new_meta = meta.copy() # Sort sources by priority (lower priority number = higher priority) sorted_sources = sorted( self.sources.items(), key=lambda x: (x[1].priority, x[0]) ) for name, source in sorted_sources: if not source.enabled or name not in self.loaded_data: continue data = self.loaded_data[name] if meta_id not in data: continue value = data[meta_id] # Use custom merge function if provided if source.merge_function: source.merge_function(new_meta, value, name) else: # Default: merge directly using output_key or name output_key = source.output_key or name new_meta[output_key] = value return new_meta # ============================================================================ # Custom Merge Functions # ============================================================================ def merge_alignment_data(meta: Dict, alignment_data: Dict, source_name: str): """Custom merge function for alignment data""" meta["text_aligned"] = alignment_data["lines"] meta["cer"] = alignment_data["cer"] meta["vocal_start_s"] = alignment_data["vocal_start_s"] meta["vocal_end_s"] = alignment_data["vocal_end_s"] def merge_tags(meta: Dict, tags: Any, source_name: str): """Custom merge function for tags (combines into a list) Accepts either a list of tags or a dict with a 'tags' key """ if "tags" not in meta: meta["tags"] = [] # If tags is a dict with a 'tags' key, extract it if isinstance(tags, dict) and "tags" in tags: tags = tags["tags"] # Now handle tags as a list or single value if isinstance(tags, list): meta["tags"].extend(tags) elif tags: meta["tags"].append(tags) def merge_stem_captions(meta: Dict, stem_captions_dict: Dict, source_name: str): """Custom merge function for stem captions""" if "stem_captions" in stem_captions_dict: meta["stem_captions"] = stem_captions_dict["stem_captions"] if "vocal_pitch_range" in stem_captions_dict: meta["vocal_pitch_range"] = stem_captions_dict["vocal_pitch_range"] def merge_all_keys( meta: Dict, extra_metadata: Dict, source_name: str, ignore_keys: List[str] = ["artists"], ): """Merge all keys from extra_metadata directly into meta, handling tags carefully.""" if not isinstance(extra_metadata, dict): return for key, value in extra_metadata.items(): if key in ignore_keys: continue if key == "tags": if "tags" not in meta: # If meta doesn't already have a tags field, create it as a list or single item meta["tags"] = ( list(value) if isinstance(value, list) else ([value] if value else []) ) else: # Merge tags into existing tags list if isinstance(value, list): meta["tags"].extend(value) elif value: meta["tags"].append(value) else: meta[key] = value def merge_priority_alignments(meta: Dict, alignment_data: Dict, source_name: str): """Merge alignment data only if not already present (priority-based)""" if "text_aligned" not in meta: merge_alignment_data(meta, alignment_data, source_name) # ============================================================================ # Processing Functions # ============================================================================ def filter_meta(meta: Dict) -> Tuple[Optional[Dict], Optional[str], float]: """Filter a single meta entry""" # Duration filter if meta["duration_s"] < 30 or meta["duration_s"] > 720: return None, "duration", meta["duration_s"] # check if we have audio_stats, if we dont, skip if "audio_stats" not in meta: return None, "audio_stats", meta["duration_s"] # check if the path exists in the local_filepath_dir local_filepath = meta["local_filepath"] if not os.path.exists(local_filepath): return None, "local_filepath", meta["duration_s"] # Loudness filter if meta["audio_stats"]["loudness"] < -20 or meta["audio_stats"]["loudness"] > -6: return None, "loudness", meta["duration_s"] # Silence filter if meta["audio_stats"]["silence_percentage"] > 4: return None, "silence", meta["duration_s"] # CER filter (if alignments exist) if meta.get("text_aligned"): if meta["cer"] > 0.8: return None, "cer", meta["duration_s"] # passed all checks, keep meta return meta, None, meta["duration_s"] def normalize_artist_field(meta: Dict) -> Dict: """Normalize artist field: rename 'artist' to 'artist_ids' and ensure it's a list of strings""" if "artist" in meta: artist_value = meta.pop("artist") # Convert to list of strings if it's not already if isinstance(artist_value, list): # Ensure all elements are strings artist_ids = [str(item) for item in artist_value if item is not None] elif artist_value is not None: # Convert single value to list of strings artist_ids = [str(artist_value)] else: # Handle None case artist_ids = [] meta["artist_ids"] = artist_ids return meta def merge_and_filter_meta( meta: Dict, registry: MetadataRegistry, local_filepath_dir: str, ) -> Tuple[Optional[Dict], Optional[str], float, str]: """Merge metadata and filter in a single pass""" meta_id = meta["id"] new_meta = meta.copy() # Merge all registered metadata new_meta = registry.merge_into_meta(new_meta, meta_id) # Normalize artist field (rename to artist_ids and ensure it's a list of strings) new_meta = normalize_artist_field(new_meta) # Construct local filepath local_filepath = os.path.join(local_filepath_dir, f"{meta_id}.opus") new_meta["local_filepath"] = local_filepath # Filter immediately after merging filtered_meta, reason, dur = filter_meta(new_meta) return filtered_meta, reason, dur, meta_id def process_meta_chunk( meta_chunk: List[Dict], registry: MetadataRegistry, local_filepath_dir: str, ) -> Tuple[List[Tuple], set]: """Process a chunk of metas in parallel""" results = [] chunk_ids_processed = set() for meta in meta_chunk: meta_id = meta["id"] # Skip duplicates within this chunk if meta_id in chunk_ids_processed: continue chunk_ids_processed.add(meta_id) filtered_meta, reason, dur, meta_id = merge_and_filter_meta( meta, registry, local_filepath_dir ) results.append((filtered_meta, reason, dur, meta.get("duration_s", 0.0))) return results, chunk_ids_processed # ============================================================================ # Main Script # ============================================================================ def setup_metadata_registry(version: str, alignments_version: str) -> MetadataRegistry: """Setup and configure all metadata sources""" registry = MetadataRegistry() # Base directories metadata_base = "/app2/suno/data/christian/metadata/" codebase_metadata = "/home/christian/code/christian/metadata/" alignments_base = "/home/tony/Work/tony/hoot/tmp/" # Stem metadata registry.register( MetadataSource( name="stems", filepath=f"{metadata_base}stems_metadata_v9.json", loader=load_json_dict, output_key="stems", ) ) # Vox stem metadata registry.register( MetadataSource( name="vox_stem", filepath=f"{metadata_base}trimmed_vocals_stem_map_v9.json", loader=load_json_dict, output_key="vox_stem", ) ) # Stem captions registry.register( MetadataSource( name="stem_captions", filepath=f"{metadata_base}stem_captions_v6.json", loader=load_json_dict, merge_function=merge_stem_captions, ) ) # Audio production metadata registry.register( MetadataSource( name="audio_production", filepath=f"{codebase_metadata}v4/combined_audio_production_features_v2.csv", loader=load_csv_dict, output_key="audio_stats", ) ) # Alignments (with priority: genius > discogs > deezer) registry.register( MetadataSource( name="alignments_genius", filepath=f"{alignments_base}genius_hq_alignments_h5_t480_v{alignments_version}.jsonl", loader=load_alignment_map, merge_function=merge_priority_alignments, priority=1, # Highest priority ) ) registry.register( MetadataSource( name="alignments_discogs", filepath=f"{alignments_base}discogs_hq_alignments_h5_t480_v{alignments_version}.jsonl", loader=load_alignment_map, merge_function=merge_priority_alignments, priority=2, ) ) registry.register( MetadataSource( name="alignments_deezer", filepath=f"{alignments_base}deezer_hq_alignments_h5_t480_v{alignments_version}.jsonl", loader=load_alignment_map, merge_function=merge_priority_alignments, priority=3, # Lowest priority ) ) # Tags - discogs titles registry.register( MetadataSource( name="tags_discogs_titles", filepath=f"{codebase_metadata}organized/titles/discogs_subset_title_terms_map.json", loader=load_json_dict, merge_function=merge_tags, ) ) # Tags - genius titles registry.register( MetadataSource( name="tags_genius_titles", filepath=f"{codebase_metadata}organized/titles/genius_title_terms_map.json", loader=load_json_dict, merge_function=merge_tags, ) ) # Tags - GPT tags (with transform to extract gpt_tag) def transform_gpt_tags(data): return {id: tags["gpt_tag"] for id, tags in data.items()} registry.register( MetadataSource( name="tags_discogs_gpt", filepath=f"{codebase_metadata}organized/llm/gpt-4_1-tagging/gpt-4_1-tagging_results_discogs_subset.json", loader=load_json_dict, loader_kwargs={"transform": transform_gpt_tags}, merge_function=merge_tags, ) ) registry.register( MetadataSource( name="tags_genius_gpt", filepath=f"{codebase_metadata}organized/llm/gpt-4_1-tagging/gpt-4_1-tagging_results_genius.json", loader=load_json_dict, loader_kwargs={"transform": transform_gpt_tags}, merge_function=merge_tags, ) ) registry.register( MetadataSource( name="tags_imslp_gpt", filepath=f"{codebase_metadata}organized/llm/gpt-4_1-tagging/gpt-4_1-tagging_results_imslp.json", loader=load_json_dict, loader_kwargs={"transform": transform_gpt_tags}, merge_function=merge_tags, ) ) # Chart metadata - merge into tags registry.register( MetadataSource( name="chart_metas", filepath=f"{codebase_metadata}organized/charts/discogs_subset_chart_metas.jsonl", loader=load_jsonl_dict, merge_function=merge_tags, ) ) # Grammy metadata - merge into tags registry.register( MetadataSource( name="grammy_metas", filepath=f"{codebase_metadata}organized/charts/discogs_subset_grammy_metas.jsonl", loader=load_jsonl_dict, merge_function=merge_tags, ) ) registry.register( MetadataSource( name="extra_metadata", filepath=f"{metadata_base}merged_metadata_map.json", loader=load_json_dict, merge_function=merge_all_keys, ) ) # add more metadata sources here # ... return registry if __name__ == "__main__": version = "4" alignments_version = "5" local_filepath_dir = "/app2/suno/data/raw_audio_opus_v0" # Setup metadata registry registry = setup_metadata_registry(version, alignments_version) # Load all metadata registry.load_all() # ------------------------------------------------------------ # Source metadata files # ------------------------------------------------------------ print("\nLoading source metadata files...") metas_dir = "/app2/suno/data/christian/metadata/clean_metas/" source_metas_generators = [ ( "discogs_subset", read_jsonl_lazy( os.path.join(metas_dir, "clean_discogs_subset_v0_metas.jsonl") ), ), ( "genius", read_jsonl_lazy(os.path.join(metas_dir, "clean_genius_v0_metas.jsonl")), ), ( "imslp", read_jsonl_lazy(os.path.join(metas_dir, "clean_imslp_v0_metas.jsonl")), ), ( "deezer", read_jsonl_lazy(os.path.join(metas_dir, "clean_deezer_v0_metas.jsonl")), ), ( "pond5", read_jsonl_lazy(os.path.join(metas_dir, "clean_pond5_v0_metas.jsonl")), ), ( "karaoke", read_jsonl_lazy(os.path.join(metas_dir, "clean_karaoke_v0_metas.jsonl")), ), ] # ------------------------------------------------------------ # Merge and filter all metadata in parallel chunks, streaming to file # ------------------------------------------------------------ print("\nMerging and filtering metadata in parallel (streaming)...") # Start timing start_time = time.time() # Track all the ids that we have merged (for deduplication across sources) all_ids_merged = set() # Statistics tracking total_duration_unfiltered = 0.0 total_duration_filtered = 0.0 total_unfiltered_count = 0 # Per-dataset statistics dataset_stats = {} # Filtering statistics filter_counts = { "duration": 0, "audio_stats": 0, "loudness": 0, "silence": 0, "cer": 0, "local_filepath": 0, } # Stream filtered results directly to file instead of keeping in memory filtered_filepath = ( f"/app2/suno/data/christian/metadata/filtered_metas_v{version}.jsonl" ) filtered_count = 0 # Configuration for parallel processing chunk_size = 2000 # Process metas in chunks of this size n_jobs = min(32, os.cpu_count() or 1) print(f"Using {n_jobs} parallel workers with chunk size {chunk_size}") # Process each source metas generator with open(filtered_filepath, "w", encoding="utf-8") as filtered_file: for source_name, metas_gen in source_metas_generators: print(f"\nProcessing {source_name} in parallel...") # Initialize stats for this dataset (must be before processing) dataset_stats[source_name] = { "unfiltered_count": 0, "filtered_count": 0, "unfiltered_duration": 0.0, "filtered_duration": 0.0, } # Collect chunks from generator and process in batches meta_chunk = [] chunk_batch = [] chunk_idx = 0 batch_size = n_jobs * 2 # Process this many chunks in parallel at once def process_batch(batch_chunks): """Process a batch of chunks in parallel""" return Parallel(n_jobs=n_jobs, backend="threading")( delayed(process_meta_chunk)( chunk, registry, local_filepath_dir, ) for chunk in batch_chunks ) for meta in metas_gen: meta_chunk.append(meta) # When chunk is full, add to batch if len(meta_chunk) >= chunk_size: chunk_batch.append(meta_chunk) meta_chunk = [] # When batch is full, process it in parallel if len(chunk_batch) >= batch_size: chunk_results_list = process_batch(chunk_batch) # Write results sequentially (thread-safe) for chunk_results, _ in chunk_results_list: for ( filtered_meta, reason, dur, unfiltered_dur, ) in chunk_results: total_duration_unfiltered += unfiltered_dur total_unfiltered_count += 1 dataset_stats[source_name]["unfiltered_count"] += 1 dataset_stats[source_name][ "unfiltered_duration" ] += unfiltered_dur meta_id = filtered_meta["id"] if filtered_meta else None if meta_id and meta_id in all_ids_merged: continue if filtered_meta is not None: all_ids_merged.add(meta_id) filtered_file.write( json.dumps(filtered_meta) + "\n" ) filtered_count += 1 total_duration_filtered += dur dataset_stats[source_name]["filtered_count"] += 1 dataset_stats[source_name][ "filtered_duration" ] += dur else: if reason and reason in filter_counts: filter_counts[reason] += 1 chunk_batch = [] chunk_idx += len(chunk_results_list) if chunk_idx % (batch_size * 5) == 0: print( f" Processed {chunk_idx * chunk_size:,} metas, {filtered_count:,} passed filter" ) # Process remaining chunk if meta_chunk: chunk_batch.append(meta_chunk) # Process remaining batch if chunk_batch: chunk_results_list = process_batch(chunk_batch) for chunk_results, _ in chunk_results_list: for filtered_meta, reason, dur, unfiltered_dur in chunk_results: total_duration_unfiltered += unfiltered_dur total_unfiltered_count += 1 dataset_stats[source_name]["unfiltered_count"] += 1 dataset_stats[source_name][ "unfiltered_duration" ] += unfiltered_dur meta_id = filtered_meta["id"] if filtered_meta else None if meta_id and meta_id in all_ids_merged: continue if filtered_meta is not None: all_ids_merged.add(meta_id) filtered_file.write(json.dumps(filtered_meta) + "\n") filtered_count += 1 total_duration_filtered += dur dataset_stats[source_name]["filtered_count"] += 1 dataset_stats[source_name]["filtered_duration"] += dur else: if reason and reason in filter_counts: filter_counts[reason] += 1 print( f" Completed {source_name}: {filtered_count:,} filtered metas so far" ) # End timing build_time = time.time() - start_time print("\n" + "=" * 80) print("BUILD SUMMARY") print("=" * 80) print(f"\nTotal build time: {build_time:.2f} seconds ({build_time/60:.2f} minutes)") print("\nPer-Dataset Statistics:") print("-" * 80) for dataset_name, stats in sorted(dataset_stats.items()): unfiltered_hrs = stats["unfiltered_duration"] / 3600 filtered_hrs = stats["filtered_duration"] / 3600 filter_pct = ( (stats["filtered_duration"] / stats["unfiltered_duration"] * 100) if stats["unfiltered_duration"] > 0 else 0 ) print(f" {dataset_name:20s}:") print( f" Unfiltered: {stats['unfiltered_count']:8,} entries ({unfiltered_hrs:8.2f} hours)" ) print( f" Filtered: {stats['filtered_count']:8,} entries ({filtered_hrs:8.2f} hours, {filter_pct:5.1f}%)" ) print("\nOverall Statistics:") print("-" * 80) print(f"Total unfiltered entries: {total_unfiltered_count:,}") print(f"Total filtered entries: {filtered_count:,}") print(f"Total unfiltered duration: {total_duration_unfiltered/3600:.2f} hours") print(f"Total filtered duration: {total_duration_filtered/3600:.2f} hours") print( f"Filtered percentage: {(total_duration_filtered/total_duration_unfiltered*100):.1f}%" ) print("\nFiltering breakdown:") print("-" * 80) for reason, count in filter_counts.items(): print(f" {reason.capitalize():20s} filtered: {count:8,} entries") print(f"\nWrote {filtered_count:,} filtered metas to {filtered_filepath}") print("=" * 80) # ------------------------------------------------------------ # Shuffle and split (read from file, shuffle, write splits) # ------------------------------------------------------------ print(f"\nShuffling and splitting {filtered_count:,} filtered metas...") # Read all filtered metas for shuffling (we need them all for proper shuffle) # This is the only time we load all filtered metas into memory filtered_metas = [] for meta in read_jsonl_lazy(filtered_filepath): filtered_metas.append(meta) # Set random seed for reproducibility random.seed(42) random.shuffle(filtered_metas) # Calculate split sizes (1% for validation) val_size = int(len(filtered_metas) * 0.01) train_size = len(filtered_metas) - val_size print(f"Train size: {train_size:,} ({train_size/len(filtered_metas)*100:.1f}%)") print(f"Validation size: {val_size:,} ({val_size/len(filtered_metas)*100:.1f}%)") # Split the data train_metas = filtered_metas[:train_size] val_metas = filtered_metas[train_size:] # Free memory del filtered_metas # Calculate durations for each split train_duration = sum(meta["duration_s"] for meta in train_metas) val_duration = sum(meta["duration_s"] for meta in val_metas) total_duration = train_duration + val_duration print("\nDuration Summary:") print( f"Train duration: {train_duration/3600:.2f} hours ({train_duration/total_duration*100:.1f}%)" ) print( f"Validation duration: {val_duration/3600:.2f} hours ({val_duration/total_duration*100:.1f}%)" ) # Define output directory and filenames output_dir = "/app2/suno/data/diffusion/v1/" train_filename = f"metas_diff_v{version}_tr.jsonl" val_filename = f"metas_diff_v{version}_val.jsonl" # Ensure output directory exists os.makedirs(output_dir, exist_ok=True) # Write train and validation files train_filepath = os.path.join(output_dir, train_filename) val_filepath = os.path.join(output_dir, val_filename) print(f"\nWriting train metas to: {train_filepath}") write_jsonl(train_metas, train_filepath) print(f"Writing validation metas to: {val_filepath}") write_jsonl(val_metas, val_filepath) print("\n✅ Successfully split and saved:") print(f" Train: {len(train_metas):,} entries ({train_duration/3600:.2f} hours)") print(f" Validation: {len(val_metas):,} entries ({val_duration/3600:.2f} hours)") print(f" Total: {filtered_count:,} entries ({total_duration/3600:.2f} hours)")