#!/usr/bin/env python3 """ Extract sample records for each keyword category with audio files. For each category: 1. Pick 5 representative keywords 2. Collect 5 sample records containing each keyword 3. Extract vocal stems and full songs as MP3 4. Organize in structured directories """ import argparse import json import logging import os import shutil import subprocess from collections import defaultdict from datetime import datetime from pathlib import Path from typing import Dict, List, Optional, Tuple from tqdm import tqdm def setup_logging(output_dir: Path) -> logging.Logger: """Set up logging configuration.""" log_file = output_dir / f"category_samples_{datetime.now().strftime('%Y%m%d_%H%M%S')}.log" logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s", handlers=[logging.FileHandler(log_file), logging.StreamHandler()], ) return logging.getLogger(__name__) class CategorySampleExtractor: """Extract audio samples for each keyword category.""" def __init__( self, categories_file: str, dataset_file: str, output_dir: str, keywords_per_category: int = 5, samples_per_keyword: int = 5, ): """ Initialize the extractor. Args: categories_file: Path to categorized keywords JSON dataset_file: Path to captioned dataset JSONL output_dir: Base output directory keywords_per_category: Number of keywords to sample per category samples_per_keyword: Number of records to find per keyword """ self.categories_file = Path(categories_file) self.dataset_file = Path(dataset_file) self.output_dir = Path(output_dir) self.output_dir.mkdir(parents=True, exist_ok=True) self.keywords_per_category = keywords_per_category self.samples_per_keyword = samples_per_keyword self.logger = setup_logging(self.output_dir) # Track statistics self.stats = { "categories_processed": 0, "keywords_processed": 0, "records_found": 0, "audio_files_converted": 0, "errors": [], } def load_categories(self) -> Dict: """Load categorized keywords from JSON.""" self.logger.info(f"Loading categories from {self.categories_file}") with open(self.categories_file, "r") as f: data = json.load(f) return data def select_keywords(self, categories_data: Dict) -> Dict[str, List[Tuple[str, int]]]: """ Select top keywords from each category. Returns: Dict mapping category to list of (keyword, count) tuples """ selected_keywords = {} detailed_categories = categories_data.get("detailed_categories", {}) for main_cat, subcats in detailed_categories.items(): # Collect all keywords from all subcategories all_keywords = [] for subcat, subcat_data in subcats.items(): keywords = subcat_data.get("keywords", {}) for keyword, count in keywords.items(): all_keywords.append((keyword, count)) # Sort by count and select top keywords all_keywords.sort(key=lambda x: x[1], reverse=True) selected = all_keywords[: self.keywords_per_category] if selected: selected_keywords[main_cat] = selected self.logger.info(f"{main_cat}: Selected {len(selected)} keywords") for keyword, count in selected: self.logger.info(f" - {keyword} ({count:,})") return selected_keywords def find_records_with_keyword(self, keyword: str, max_records: int = 5) -> List[Dict]: """ Find records containing a specific keyword in voice_description_keywords. Args: keyword: Keyword to search for max_records: Maximum records to return Returns: List of matching records """ found_records = [] keyword_lower = keyword.lower() self.logger.debug(f"Searching for keyword: {keyword}") with open(self.dataset_file, "r") as f: for line_num, line in enumerate(f, 1): if len(found_records) >= max_records: break if not line.strip(): continue try: record = json.loads(line.strip()) # Check stems_captions for keyword stems_captions = record.get("stems_captions", {}) for stem_name, caption_list in stems_captions.items(): if len(found_records) >= max_records: break for caption_entry in caption_list: if caption_entry.get("prompt_type") == "voice_description_keywords": caption = caption_entry.get("caption", "").lower() # Check if keyword appears in caption # Split by comma to check individual keywords caption_keywords = [k.strip() for k in caption.split(",")] if keyword_lower in caption_keywords: # Add stem info to record for later processing record["_matched_stem"] = stem_name record["_matched_keyword"] = keyword found_records.append(record) self.logger.debug( f" Found in record {record.get('id')} stem {stem_name}" ) break except json.JSONDecodeError: continue except Exception as e: self.logger.warning(f"Error processing line {line_num}: {e}") self.logger.info(f"Found {len(found_records)} records for keyword '{keyword}'") return found_records def convert_audio_to_mp3( self, input_path: str, output_path: str, skip_if_exists: bool = True ) -> bool: """ Convert audio file to MP3. Args: input_path: Input audio file path output_path: Output MP3 file path skip_if_exists: Skip conversion if output exists Returns: True if successful, False otherwise """ output_path = Path(output_path) if skip_if_exists and output_path.exists(): self.logger.debug(f"MP3 already exists: {output_path.name}") return True if not os.path.exists(input_path): self.logger.warning(f"Input file not found: {input_path}") return False try: cmd = [ "ffmpeg", "-i", input_path, "-acodec", "libmp3lame", "-b:a", "192k", str(output_path), "-y", "-loglevel", "error", ] subprocess.run(cmd, check=True, capture_output=True) self.stats["audio_files_converted"] += 1 return True except subprocess.CalledProcessError as e: self.logger.error(f"Failed to convert {input_path}: {e}") self.stats["errors"].append(f"Audio conversion failed: {input_path}") return False def process_category(self, category: str, keywords: List[Tuple[str, int]]) -> Dict: """ Process a single category to extract samples. Args: category: Category name keywords: List of (keyword, count) tuples Returns: Dictionary with processing results """ self.logger.info(f"\nProcessing category: {category}") self.logger.info("=" * 60) category_dir = self.output_dir / category.lower().replace(" ", "_") category_dir.mkdir(exist_ok=True) category_results = { "category": category, "keywords_processed": [], "total_records": 0, "total_audio_files": 0, } for keyword, count in keywords: self.logger.info(f"\nProcessing keyword: {keyword} (count: {count:,})") # Create keyword directory keyword_safe = keyword.replace(" ", "_").replace("/", "_") keyword_dir = category_dir / keyword_safe keyword_dir.mkdir(exist_ok=True) # Find records with this keyword records = self.find_records_with_keyword(keyword, self.samples_per_keyword) if not records: self.logger.warning(f"No records found for keyword: {keyword}") continue keyword_results = { "keyword": keyword, "count": count, "records_found": len(records), "samples": [], } # Process each record for idx, record in enumerate(records, 1): record_id = record.get("id", f"unknown_{idx}") matched_stem = record.get("_matched_stem", "Vocals") self.logger.info(f" Processing record {idx}/{len(records)}: {record_id}") # Create record directory record_dir = keyword_dir / f"{idx:02d}_{record_id}" record_dir.mkdir(exist_ok=True) sample_info = {"id": record_id, "matched_stem": matched_stem, "audio_files": []} # Save record metadata metadata_file = record_dir / "metadata.json" with open(metadata_file, "w") as f: json.dump( { "id": record_id, "keyword": keyword, "category": category, "matched_stem": matched_stem, "text": record.get("text", ""), "tags": record.get("tags", []), "stems_captions": record.get("stems_captions", {}), }, f, indent=2, ) # Extract and convert full song local_filepath = record.get("local_filepath") if local_filepath: full_mp3 = record_dir / f"full_{record_id}.mp3" if self.convert_audio_to_mp3(local_filepath, full_mp3): sample_info["audio_files"].append("full_song") self.logger.info(f" ✓ Converted full song") else: self.logger.warning(f" No local_filepath for {record_id}") # Extract and convert vocal stem stems = record.get("stems", {}) if matched_stem in stems: stem_path = stems[matched_stem] stem_mp3 = record_dir / f"{matched_stem.lower()}_{record_id}.mp3" if self.convert_audio_to_mp3(stem_path, stem_mp3): sample_info["audio_files"].append(matched_stem) self.logger.info(f" ✓ Converted {matched_stem} stem") else: self.logger.warning(f" No {matched_stem} stem for {record_id}") # Also extract main Vocals stem if different if matched_stem != "Vocals" and "Vocals" in stems: vocals_path = stems["Vocals"] vocals_mp3 = record_dir / f"vocals_{record_id}.mp3" if self.convert_audio_to_mp3(vocals_path, vocals_mp3): sample_info["audio_files"].append("Vocals") self.logger.info(f" ✓ Converted Vocals stem") keyword_results["samples"].append(sample_info) category_results["keywords_processed"].append(keyword_results) category_results["total_records"] += len(records) category_results["total_audio_files"] += sum( len(s["audio_files"]) for s in keyword_results["samples"] ) self.stats["keywords_processed"] += 1 self.stats["records_found"] += len(records) self.stats["categories_processed"] += 1 return category_results def generate_summary(self, results: List[Dict]) -> Dict: """Generate summary of extraction results.""" summary = { "timestamp": datetime.now().isoformat(), "configuration": { "categories_file": str(self.categories_file), "dataset_file": str(self.dataset_file), "output_dir": str(self.output_dir), "keywords_per_category": self.keywords_per_category, "samples_per_keyword": self.samples_per_keyword, }, "statistics": self.stats, "categories": results, } return summary def run(self): """Run the extraction process.""" self.logger.info("=" * 80) self.logger.info("CATEGORY SAMPLE EXTRACTION") self.logger.info("=" * 80) # Load categories categories_data = self.load_categories() # Select keywords selected_keywords = self.select_keywords(categories_data) if not selected_keywords: self.logger.error("No keywords selected!") return self.logger.info(f"\nSelected keywords from {len(selected_keywords)} categories") # Process each category results = [] for category, keywords in tqdm(selected_keywords.items(), desc="Processing categories"): category_results = self.process_category(category, keywords) results.append(category_results) # Generate and save summary summary = self.generate_summary(results) summary_file = self.output_dir / "extraction_summary.json" with open(summary_file, "w") as f: json.dump(summary, f, indent=2) self.logger.info("\n" + "=" * 80) self.logger.info("EXTRACTION COMPLETE") self.logger.info("=" * 80) self.logger.info(f"Categories processed: {self.stats['categories_processed']}") self.logger.info(f"Keywords processed: {self.stats['keywords_processed']}") self.logger.info(f"Records found: {self.stats['records_found']}") self.logger.info(f"Audio files converted: {self.stats['audio_files_converted']}") self.logger.info(f"Errors: {len(self.stats['errors'])}") self.logger.info(f"\nOutput directory: {self.output_dir}") self.logger.info(f"Summary saved to: {summary_file}") # Create README self.create_readme() def create_readme(self): """Create README file explaining the directory structure.""" readme_file = self.output_dir / "README.md" with open(readme_file, "w") as f: f.write("# Category Sample Extraction Results\n\n") f.write(f"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n\n") f.write("## Directory Structure\n\n") f.write("```\n") f.write("output_dir/\n") f.write("├── README.md # This file\n") f.write("├── extraction_summary.json # Complete extraction summary\n") f.write("├── category_samples_*.log # Detailed log file\n") f.write("└── [category_name]/ # One directory per category\n") f.write(" └── [keyword]/ # One directory per keyword\n") f.write(" └── [NN_record_id]/ # Sample record directory\n") f.write(" ├── metadata.json # Record metadata\n") f.write(" ├── full_*.mp3 # Full song MP3\n") f.write(" └── vocals_*.mp3 # Vocal stem MP3\n") f.write("```\n\n") f.write("## Categories Processed\n\n") # List categories for category_dir in sorted(self.output_dir.iterdir()): if category_dir.is_dir() and not category_dir.name.startswith("."): f.write(f"- **{category_dir.name}**\n") # List keywords for keyword_dir in sorted(category_dir.iterdir())[:5]: if keyword_dir.is_dir(): sample_count = len(list(keyword_dir.iterdir())) f.write(f" - {keyword_dir.name} ({sample_count} samples)\n") f.write("\n## Usage\n\n") f.write("Each sample directory contains:\n") f.write("- `metadata.json`: Full record metadata including lyrics, tags, and captions\n") f.write("- `full_*.mp3`: The complete song audio\n") f.write("- `vocals_*.mp3` or `[stem]_*.mp3`: The vocal stem where the keyword was found\n") f.write("\n") f.write("Use the `extraction_summary.json` file for programmatic access to all results.\n") self.logger.info(f"README created: {readme_file}") def main(): """Main execution function.""" parser = argparse.ArgumentParser(description="Extract audio samples for keyword categories") parser.add_argument( "--categories-file", default="/home/vibert/tmp/voice_keywords_analysis/top_1000_detailed_categories.json", help="Path to categorized keywords JSON", ) parser.add_argument( "--dataset-file", default="/home/vibert/data/voice_designer/metas_v6_tr_vocal_stems_captioned_w30.jsonl", help="Path to captioned dataset JSONL", ) parser.add_argument( "--output-dir", default="/home/vibert/tmp/category_audio_samples", help="Output directory for samples", ) parser.add_argument( "--keywords-per-category", type=int, default=5, help="Number of keywords to sample per category (default: 5)", ) parser.add_argument( "--samples-per-keyword", type=int, default=5, help="Number of records to find per keyword (default: 5)", ) args = parser.parse_args() # Run extractor extractor = CategorySampleExtractor( categories_file=args.categories_file, dataset_file=args.dataset_file, output_dir=args.output_dir, keywords_per_category=args.keywords_per_category, samples_per_keyword=args.samples_per_keyword, ) extractor.run() if __name__ == "__main__": main()