#!/usr/bin/env python3 """ Parallel audio processing script for SFX beat alignment. Converts notebook workflow to parallelized script using joblib. """ import os import logging from pathlib import Path from typing import Dict, List, Optional, Tuple from joblib import Parallel, delayed from tqdm import tqdm from suno_utils.utils.text import read_jsonl, read_json from suno_utils.audio import Audio # Configuration OUTPUT_DIR = "/app2/suno/data/sara/sfx_get_beats/aligned_audio" BEAT_DATA_DIR = "/app2/suno/data/sara/sfx_get_beats/sfx_full_beat_data/" META_FILE = "/app2/suno/data/sara/sfx_get_beats/combined_v3_w_extreme_metas_v0.jsonl" # Parallelization settings N_JOBS = 16 # Use all CPUs, or set to specific number (e.g., 16) BACKEND = "threading" # Options: "threading", "loky", "multiprocessing" BATCH_SIZE = 'auto' # Number of files per batch # Setup logging logging.basicConfig( level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s', handlers=[ logging.FileHandler('sfx_processing.log'), logging.StreamHandler() ] ) logger = logging.getLogger(__name__) def bpm_std_to_time_difference(tempo_bpm: float, std_bpm: float) -> Tuple[float, float]: """ Convert BPM standard deviation to time difference. Args: tempo_bpm: Tempo in beats per minute std_bpm: Standard deviation of BPM Returns: Tuple of (time_diff_slower, time_diff_faster) in seconds """ # Beat periods in seconds mean_period = 60 / tempo_bpm slower_period = 60 / (tempo_bpm - 2 * std_bpm) # tempo - 1 std dev faster_period = 60 / (tempo_bpm + 2 * std_bpm) # tempo + 1 std dev # Time differences from mean time_diff_slower = slower_period - mean_period time_diff_faster = mean_period - faster_period return time_diff_slower, time_diff_faster def load_and_prepare_metadata() -> List[Dict]: """ Load metadata and beat detection data, filter for valid entries. Returns: List of valid metadata dictionaries with beat_start_time_s """ logger.info("Loading metadata...") meta = read_jsonl(META_FILE) metas_map = {m["id"]: m for m in meta} logger.info(f"Loaded {len(metas_map)} metadata entries") # Load beat detection data logger.info("Loading beat detection data...") kept = 0 total_found = 0 for filename in os.listdir(BEAT_DATA_DIR): if filename.endswith('.json'): file_path = os.path.join(BEAT_DATA_DIR, filename) data = read_json(file_path) for id, row in data.items(): if id not in metas_map: continue meta_row = metas_map[id] meta_row["beat_start_time_s"] = 0 if not row["processing_failed"]: detected_bpm = row["inferred_tempo"] beats_detected = row["beats_detected"] beat_times = row["beat_times"] tempo_std = row["tempo_std"] if detected_bpm is not None and beats_detected is not None and tempo_std is not None: std_time = max(bpm_std_to_time_difference(detected_bpm, tempo_std)) total_found += 1 if beats_detected >= 4 and std_time < 0.01: meta_row["beat_start_time_s"] = float(beat_times[0]) kept += 1 logger.info(f"Found {total_found} files with beat data, kept {kept} ({kept/total_found*100:.1f}%)") # Filter for valid entries valid = [] for id, data in metas_map.items(): if "beat_start_time_s" in data and data["beat_start_time_s"] > 0: valid.append(data) logger.info(f"Total valid files to process: {len(valid)}") return valid def process_single_file(v: Dict) -> Tuple[str, bool, Optional[str]]: """ Process a single audio file: trim, align, and save. Args: v: Metadata dictionary for the file Returns: Tuple of (file_id, success, error_message) """ file_id = v["id"] try: # Skip extreme dataset if v["dataset"] == "extreme": return file_id, True, "skipped:extreme_dataset" # Check if output already exists output_path = os.path.join(OUTPUT_DIR, f"{file_id}.opus") if os.path.exists(output_path): return file_id, True, "skipped:already_exists" s3_path = v["s3_filepath"] beat_start_time_s = v['beat_start_time_s'] # Load and process audio audio = Audio.from_s3(s3_path) trim_audio = audio.get_slice(from_s=beat_start_time_s) # Calculate target length aligned to token boundaries token_len_samples = int(trim_audio.sample_rate * 0.04) audio_len_samples = int(trim_audio.sample_rate * trim_audio.duration_s) leftover_samples = audio_len_samples % token_len_samples target_len_samples = audio_len_samples - leftover_samples if leftover_samples > 0: target_len_samples += token_len_samples target_len_seconds = target_len_samples / trim_audio.sample_rate # Pad if necessary if target_len_samples > audio_len_samples: trim_audio = trim_audio.pad_to_length(target_len_seconds) # Write output if modified or mp3 source if abs(trim_audio.duration_s - audio.duration_s) > 0.005 or s3_path.endswith(".mp3"): # Create output directory if it doesn't exist os.makedirs(OUTPUT_DIR, exist_ok=True) trim_audio.write_opus(output_path) return file_id, True, "processed" else: return file_id, True, "skipped:no_modification_needed" except Exception as e: logger.error(f"Error processing {file_id}: {str(e)}") return file_id, False, str(e) def main(): """Main execution function.""" logger.info("Starting parallel audio processing") logger.info(f"Configuration: n_jobs={N_JOBS}, backend={BACKEND}, batch_size={BATCH_SIZE}") # Ensure output directory exists os.makedirs(OUTPUT_DIR, exist_ok=True) # Load metadata valid = load_and_prepare_metadata() if not valid: logger.warning("No valid files to process") return # Process files in parallel logger.info(f"Processing {len(valid)} files...") results = Parallel( n_jobs=N_JOBS, backend=BACKEND, batch_size=BATCH_SIZE, verbose=0 )(delayed(process_single_file)(v) for v in tqdm(valid, desc="Processing files")) # Summarize results success_count = sum(1 for _, success, _ in results if success) failed_count = len(results) - success_count # Count by status status_counts = {} for _, success, msg in results: if success: status = msg.split(':')[0] if msg else "processed" status_counts[status] = status_counts.get(status, 0) + 1 logger.info("=" * 60) logger.info("Processing complete!") logger.info(f"Total files: {len(results)}") logger.info(f"Successful: {success_count}") logger.info(f"Failed: {failed_count}") logger.info(f"Status breakdown:") for status, count in sorted(status_counts.items()): logger.info(f" {status}: {count}") # Log failures if failed_count > 0: logger.warning("Failed files:") for file_id, success, error in results: if not success: logger.warning(f" {file_id}: {error}") logger.info("=" * 60) if __name__ == "__main__": main()