#!/usr/bin/env python3 """ Distributed semantic encoding pipeline for audio dataset. Encodes audio files to semantic codes using MERT-25 model with: - Multi-node, multi-GPU support via torch.distributed - Efficient batching with padding - Resume support (skip already encoded files) - Direct mono 24kHz audio loading """ import argparse import json import os import time from typing import Dict, List, Optional, Tuple import numpy as np import torch import torch.distributed as dist import torchaudio from suno_utils.audio import Audio try: from tqdm import tqdm except ImportError: def tqdm(iterable, *args, **kwargs): return iterable def setup_distributed() -> Tuple[int, int, int]: """ Initialize distributed process group for SLURM. Returns: Tuple of (rank, world_size, local_rank) """ # SLURM environment variables rank = int(os.environ.get("SLURM_PROCID", 0)) world_size = int(os.environ.get("SLURM_NTASKS", 1)) local_rank = int(os.environ.get("SLURM_LOCALID", 0)) # Initialize process group if world_size > 1: dist.init_process_group( backend="nccl", init_method="env://", rank=rank, world_size=world_size ) # Set device for this process torch.cuda.set_device(local_rank) return rank, world_size, local_rank def load_metadata(metadata_path: str) -> List[Dict]: """Load metadata from JSONL file.""" metadata = [] with open(metadata_path, "r") as f: for line in f: line = line.strip() if line: try: metadata.append(json.loads(line)) except json.JSONDecodeError: continue return metadata def load_audio_list(audio_list_path: str) -> List[Dict]: """ Load list of audio file paths and convert to metadata format. Supports multiple formats: - Plain text file with one path per line - JSON file with list of paths - JSON file with list of dicts containing 'path' or 'filepath' keys Args: audio_list_path: Path to file containing audio paths Returns: List of metadata dicts with 'id' and 'local_filepath' fields """ metadata = [] # Try loading as JSON first try: with open(audio_list_path, "r") as f: # Try to load as single JSON object first f.seek(0) first_line = f.readline().strip() f.seek(0) # Check if it's JSONL format (one JSON per line) if first_line and first_line.startswith("{"): # JSONL format for line in f: line = line.strip() if not line: continue item = json.loads(line) if isinstance(item, dict): audio_path = ( item.get("audio_path") or item.get("path") or item.get("filepath") or item.get("local_filepath") ) audio_id = ( item.get("id") or os.path.splitext(os.path.basename(audio_path))[0] ) metadata.append({"id": audio_id, "local_filepath": audio_path}) return metadata else: # Regular JSON file data = json.load(f) if isinstance(data, list): for item in data: if isinstance(item, str): # Simple list of paths audio_path = item audio_id = os.path.splitext(os.path.basename(audio_path))[0] metadata.append({"id": audio_id, "local_filepath": audio_path}) elif isinstance(item, dict): # List of dicts - extract path and id audio_path = ( item.get("audio_path") or item.get("path") or item.get("filepath") or item.get("local_filepath") ) audio_id = ( item.get("id") or os.path.splitext(os.path.basename(audio_path))[0] ) metadata.append({"id": audio_id, "local_filepath": audio_path}) return metadata except (json.JSONDecodeError, ValueError): pass # Fall back to plain text file (one path per line) with open(audio_list_path, "r") as f: for line in f: line = line.strip() if line and not line.startswith("#"): # Skip comments audio_path = line audio_id = os.path.splitext(os.path.basename(audio_path))[0] metadata.append({"id": audio_id, "local_filepath": audio_path}) return metadata def load_audio_mono_24k(filepath: str) -> Optional[Audio]: """ Load audio file directly as mono 24kHz for semantic encoding. Args: filepath: Path to audio file Returns: Mono 24kHz Audio object or None if loading fails """ try: # Load with torchaudio waveform, sample_rate = torchaudio.load(filepath) # Convert to mono by averaging channels if waveform.shape[0] > 1: waveform = waveform.mean(dim=0, keepdim=True) # Resample to 24kHz if needed if sample_rate != 24000: waveform = torchaudio.functional.resample(waveform, sample_rate, 24000) # Convert to Audio object (mono, 24kHz) audio = Audio.from_array_float( waveform.numpy().squeeze(), sample_rate=24000, max_allowed_val=100 ) return audio except Exception as e: return None def encode_audio_batch( audio_list: List[Audio], device: str, encode_semantic_fn, n_codebooks: Optional[int] = None, ) -> List[np.ndarray]: """ Encode a batch of audio files to semantic codes. Args: audio_list: List of Audio objects (mono 24kHz) device: Device to use for encoding encode_semantic_fn: Encoding function from suno_utils n_codebooks: Number of codebooks to use (None = use all available) Returns: List of semantic code arrays """ if len(audio_list) == 0: return [] # Encode each audio individually (encode_semantic expects single audio) codes_list = [] for audio in audio_list: try: codes = encode_semantic_fn(audio, device=device, n_codebooks=n_codebooks) if isinstance(codes, torch.Tensor): codes = codes.cpu().numpy() codes_list.append(codes.astype(np.int64)) except Exception as e: # If encoding fails, append None and handle in caller codes_list.append(None) return codes_list def process_worker( rank: int, world_size: int, local_rank: int, metadata: List[Dict], output_dir: str, batch_size: int, overwrite: bool, encode_semantic_fn, preload_semantic_models_fn, n_codebooks: Optional[int] = None, ) -> Dict[str, int]: """ Process assigned subset of dataset on this worker. Args: rank: Global rank of this worker world_size: Total number of workers local_rank: Local rank on this node metadata: Full metadata list output_dir: Output directory for .npz files batch_size: Batch size for encoding overwrite: Whether to overwrite existing files encode_semantic_fn: Encoding function preload_semantic_models_fn: Model loading function n_codebooks: Number of codebooks to use (None = use all available) Returns: Dictionary with processing statistics """ device = f"cuda:{local_rank}" # Load model on this GPU if rank == 0: print(f"Loading semantic models on {world_size} workers...") # IMPORTANT: Clear any cached models first to ensure we load with correct centroids # (suno_utils caches models by device only, not by centroids filepath) from suno_utils.tasks.mert_25 import clean_models clean_models() preload_semantic_models_fn( checkpoint_filepath="/app/suno/data/dpo/models/mert_25.pt", # this is the 2 codebook model # centroids_filepath="/app/suno/data/dpo/models/mert_25_2x4k.npy", # this is the 768d model, 64 codebooks centroids_filepath="/home/minz/temp/mert_768d_centroids_4000_50.npy", device=device, ) # Create output directory os.makedirs(output_dir, exist_ok=True) error_log_path = os.path.join(output_dir, f"errors_rank{rank}.log") # Statistics stats = { "processed": 0, "skipped": 0, "errors": 0, } # Get subset for this worker worker_metadata = [metadata[i] for i in range(rank, len(metadata), world_size)] if rank == 0: print(f"Rank {rank}: Processing {len(worker_metadata)} samples") # Process in batches batch_audio = [] batch_ids = [] batch_filepaths = [] start_time = time.time() processed_count = 0 for idx, meta in enumerate(worker_metadata): audio_id = meta.get("id") local_filepath = meta.get("local_filepath") if not audio_id or not local_filepath: stats["errors"] += 1 continue output_path = os.path.join(output_dir, f"{audio_id}.npz") # Check if already exists (resume support) if not overwrite and os.path.exists(output_path): stats["skipped"] += 1 continue # Load audio as mono 24kHz audio_mono = load_audio_mono_24k(local_filepath) if audio_mono is None: stats["errors"] += 1 with open(error_log_path, "a") as f: f.write(f"Failed to load: {audio_id} - {local_filepath}\n") continue # Add to batch batch_audio.append(audio_mono) batch_ids.append(audio_id) batch_filepaths.append(local_filepath) # Process batch when full or at end if len(batch_audio) >= batch_size or idx == len(worker_metadata) - 1: # Encode batch try: codes_list = encode_audio_batch( batch_audio, device, encode_semantic_fn, n_codebooks=n_codebooks, ) # Save individual files for bid, codes in zip(batch_ids, codes_list): if codes is not None: output_path = os.path.join(output_dir, f"{bid}.npz") np.savez(output_path, codes=codes) stats["processed"] += 1 processed_count += 1 else: stats["errors"] += 1 with open(error_log_path, "a") as f: f.write(f"Failed to encode: {bid}\n") # Log progress every 100 samples if processed_count % 100 == 0 and processed_count > 0: elapsed = time.time() - start_time rate = processed_count / elapsed if elapsed > 0 else 0 if rank == 0: print( f"Rank {rank}: Processed {processed_count}/{len(worker_metadata)} " f"({rate:.2f} samples/sec)" ) except Exception as e: stats["errors"] += len(batch_ids) with open(error_log_path, "a") as f: f.write(f"Batch encoding failed: {batch_ids}\n") f.write(f"Error: {str(e)}\n") # Clear batch batch_audio = [] batch_ids = [] batch_filepaths = [] # Final statistics elapsed = time.time() - start_time rate = stats["processed"] / elapsed if elapsed > 0 else 0 if rank == 0 or stats["processed"] > 0: print(f"\nRank {rank} Summary:") print(f" Processed: {stats['processed']}") print(f" Skipped: {stats['skipped']}") print(f" Errors: {stats['errors']}") print(f" Rate: {rate:.2f} samples/sec") print(f" Time: {elapsed:.1f}s") return stats def main(): """Main entry point for semantic encoding.""" parser = argparse.ArgumentParser( description="Distributed semantic encoding pipeline" ) parser.add_argument( "--metadata_path", type=str, default=None, help="Path to JSONL metadata file (format: one JSON object per line with 'id' and 'local_filepath' fields)", ) parser.add_argument( "--audio_list", type=str, default=None, help="Path to audio list file (plalsin text with one path per line, or JSON list of paths)", ) parser.add_argument( "--output_dir", type=str, default="/app2/suno/data/semantic_code/sft", help="Output directory for .npz files", ) parser.add_argument( "--batch_size", type=int, default=16, help="Batch size for encoding" ) parser.add_argument( "--overwrite", action="store_true", help="Overwrite existing files" ) parser.add_argument( "--n_codebooks", type=int, default=None, help="Number of codebooks to use (default: use all available from centroids file)", ) args = parser.parse_args() # Setup distributed processing first rank, world_size, local_rank = setup_distributed() # Validate input arguments if args.metadata_path is None and args.audio_list is None: # Set default if neither provided args.metadata_path = ( "/home/tony/Data/Preference/RealGen/metas_v5_val_filtered_clean.jsonl" ) if args.metadata_path and args.audio_list: if rank == 0: print("Error: Cannot specify both --metadata_path and --audio_list") return if rank == 0: print("=" * 80) print("SEMANTIC ENCODING PIPELINE") print("=" * 80) if args.metadata_path: print(f"Input mode: JSONL metadata") print(f"Metadata: {args.metadata_path}") else: print(f"Input mode: Audio list") print(f"Audio list: {args.audio_list}") print(f"Output dir: {args.output_dir}") print(f"Batch size: {args.batch_size}") print(f"Overwrite: {args.overwrite}") print( f"N codebooks: {args.n_codebooks if args.n_codebooks else 'all available'}" ) print(f"World size: {world_size}") print(f"Nodes: {world_size // 8}") print(f"GPUs/node: 8") print("=" * 80) print() # Load metadata if rank == 0: print("Loading input data...") if args.metadata_path: metadata = load_metadata(args.metadata_path) else: metadata = load_audio_list(args.audio_list) if rank == 0: print(f"Loaded {len(metadata):,} metadata entries") # Count existing files if os.path.exists(args.output_dir): existing_count = len( [f for f in os.listdir(args.output_dir) if f.endswith(".npz")] ) print(f"Found {existing_count:,} existing .npz files") if not args.overwrite and existing_count > 0: print("Will skip existing files (use --overwrite to re-encode)") print() # Import encoding functions try: from suno_utils.tasks.mert_25 import ( preload_models as preload_semantic_models_, encode as encode_semantic, ) except ImportError as e: if rank == 0: print(f"Error: Failed to import suno_utils: {e}") print("Make sure suno_utils is installed and accessible") return # Process worker's subset stats = process_worker( rank=rank, world_size=world_size, local_rank=local_rank, metadata=metadata, output_dir=args.output_dir, batch_size=args.batch_size, overwrite=args.overwrite, encode_semantic_fn=encode_semantic, preload_semantic_models_fn=preload_semantic_models_, n_codebooks=args.n_codebooks, ) # Gather statistics from all workers if world_size > 1: # Synchronize dist.barrier() # Collect stats on rank 0 if rank == 0: total_stats = { "processed": stats["processed"], "skipped": stats["skipped"], "errors": stats["errors"], } for i in range(1, world_size): # Simple gathering (could use dist.gather for more efficiency) pass print("\n" + "=" * 80) print("FINAL SUMMARY (Rank 0 only, see individual logs for full stats)") print("=" * 80) print(f"Total processed (rank 0): {total_stats['processed']:,}") print(f"Total skipped (rank 0): {total_stats['skipped']:,}") print(f"Total errors (rank 0): {total_stats['errors']:,}") print("=" * 80) else: print("\n" + "=" * 80) print("FINAL SUMMARY") print("=" * 80) print(f"Total processed: {stats['processed']:,}") print(f"Total skipped: {stats['skipped']:,}") print(f"Total errors: {stats['errors']:,}") print("=" * 80) # Cleanup if world_size > 1: dist.destroy_process_group() if rank == 0: print("\nEncoding complete!") if __name__ == "__main__": main()