| """ |
| Overlap task generator for temporal reasoning dataset. |
| |
| Generates audio samples where two sound events partially overlap in time, |
| and asks questions about which sounds overlap, identification of overlapping |
| pairs, and yes/no overlap verification. |
| """ |
|
|
| import csv |
| import json |
| import random |
| from pathlib import Path |
| from typing import Dict, List, Optional, Tuple |
|
|
| from utils import ( |
| AudioProcessor, |
| QuestionGenerator, |
| setup_logger, |
| set_random_seed, |
| generate_sample_durations_for_task, |
| generate_single_clip_duration, |
| concatenate_to_target_duration, |
| build_overlap_task_audio, |
| create_dataset |
| ) |
|
|
|
|
| class OverlapTaskGenerator: |
| """Generates overlap task dataset samples.""" |
| |
| def __init__(self, config: dict, logger=None): |
| """ |
| Initialize OverlapTaskGenerator. |
| |
| Args: |
| config: Full pipeline configuration dictionary |
| logger: Logger instance |
| """ |
| self.config = config |
| self.logger = logger or setup_logger(__name__) |
| |
| self.task_config = config['tasks']['overlap'] |
| |
| |
| audio_config = config['audio'] |
| self.min_clip_duration = audio_config['min_clip_duration'] |
| self.max_clip_duration = audio_config['max_clip_duration'] |
| self.source_clip_duration = audio_config.get('source_clip_duration', 5.0) |
| |
| |
| self.task_duration_hours = self.task_config['task_duration_size'] |
| self.num_sounds = self.task_config.get('num_sounds', 2) |
| self.min_overlap_ratio = self.task_config.get('min_overlap_ratio', 0.2) |
| self.max_overlap_ratio = self.task_config.get('max_overlap_ratio', 0.5) |
| |
| |
| self.dataset = create_dataset(config) |
| |
| |
| self.audio_processor = AudioProcessor( |
| crossfade_duration=audio_config.get('crossfade_duration', 500), |
| silence_duration=audio_config.get('silence_duration', 1000), |
| with_silence=False, |
| normalize=audio_config.get('normalize', False), |
| normalize_target_dBFS=audio_config.get('normalize_target_dBFS', -20.0) |
| ) |
| |
| |
| mcq_config = config.get('mcq', {}) |
| self.question_generator = QuestionGenerator( |
| num_options=mcq_config.get('num_options', 4), |
| option_labels=mcq_config.get('option_labels', ['A', 'B', 'C', 'D']), |
| distractor_strategy=mcq_config.get('distractor_strategy', 'balanced') |
| ) |
| |
| |
| self.output_base = Path(config['output']['base_path']) / 'overlap' |
| self.audio_output = self.output_base / 'audio' |
| self.output_base.mkdir(parents=True, exist_ok=True) |
| self.audio_output.mkdir(parents=True, exist_ok=True) |
| |
| def generate_sample( |
| self, |
| sample_id: int, |
| target_duration_seconds: float = None |
| ) -> Optional[Dict]: |
| """ |
| Generate a single overlap task sample. |
| |
| Pipeline: |
| 1. Sample 2 categories |
| 2. Load audio for each |
| 3. Extend clips to reasonable duration |
| 4. Build overlapping audio using overlay() |
| 5. Generate questions about the overlap |
| |
| Args: |
| sample_id: Sample ID |
| target_duration_seconds: Pre-generated target duration |
| |
| Returns: |
| Metadata dictionary, or None if failed |
| """ |
| |
| try: |
| categories = self.dataset.sample_categories(self.num_sounds) |
| except ValueError: |
| self.logger.warning(f"Sample {sample_id}: Cannot sample {self.num_sounds} categories") |
| return None |
| |
| category_a, category_b = categories[0], categories[1] |
| |
| |
| filename_a, filepath_a = self.dataset.sample_file_from_category(category_a) |
| filename_b, filepath_b = self.dataset.sample_file_from_category(category_b) |
| |
| audio_a = self.audio_processor.load_audio(filepath_a) |
| audio_b = self.audio_processor.load_audio(filepath_b) |
| |
| |
| |
| clip_target_s = max(self.source_clip_duration, 5.0) |
| if target_duration_seconds: |
| |
| clip_target_s = max(clip_target_s, target_duration_seconds * 0.6) |
| |
| try: |
| final_audio, build_metadata = build_overlap_task_audio( |
| audio_a, audio_b, |
| category_a, category_b, |
| min_overlap_ratio=self.min_overlap_ratio, |
| max_overlap_ratio=self.max_overlap_ratio |
| ) |
| except Exception as e: |
| self.logger.warning(f"Sample {sample_id}: Failed to build overlap audio: {e}") |
| return None |
| |
| |
| output_audio_path = self.audio_output / f"{sample_id}.wav" |
| final_audio.export(str(output_audio_path), format="wav") |
| |
| |
| question_type = random.choice(self.task_config['question_types']) |
| |
| mcq_data, open_text_data = self._generate_question( |
| question_type, category_a, category_b, build_metadata |
| ) |
| |
| if mcq_data is None: |
| self.logger.warning(f"Sample {sample_id}: Failed to generate question") |
| return None |
| |
| metadata = { |
| 'id': sample_id, |
| 'audio_path': str(output_audio_path.relative_to(self.output_base.parent)), |
| 'num_sounds': self.num_sounds, |
| 'source_files': [filename_a, filename_b], |
| 'category_a': category_a, |
| 'category_b': category_b, |
| 'categories': categories, |
| 'question_type': question_type, |
| 'a_start_ms': build_metadata['a_start_ms'], |
| 'a_end_ms': build_metadata['a_end_ms'], |
| 'b_start_ms': build_metadata['b_start_ms'], |
| 'b_end_ms': build_metadata['b_end_ms'], |
| 'overlap_duration_ms': build_metadata['overlap_duration_ms'], |
| 'overlap_ratio': build_metadata['overlap_ratio'], |
| 'total_duration_ms': build_metadata['total_duration_ms'], |
| 'actual_duration_s': len(final_audio) / 1000.0, |
| 'mcq_question': mcq_data['question'], |
| 'mcq_options': mcq_data['options'], |
| 'mcq_correct_answer': mcq_data['correct_answer'], |
| 'open_text_question': open_text_data['question'], |
| 'open_text_answer': open_text_data['correct_answer'] |
| } |
| |
| self.logger.info( |
| f"Generated overlap sample {sample_id}: {category_a} + {category_b}, " |
| f"overlap={build_metadata['overlap_duration_ms']}ms " |
| f"({build_metadata['overlap_ratio']*100:.1f}%), type={question_type}" |
| ) |
| |
| return metadata |
| |
| def _generate_question( |
| self, |
| question_type: str, |
| category_a: str, |
| category_b: str, |
| build_metadata: Dict |
| ) -> Tuple[Optional[Dict], Optional[Dict]]: |
| """Generate MCQ and open-text question based on type.""" |
| |
| mcq_template = self.task_config['mcq_questions'].get(question_type, '') |
| open_template = self.task_config['open_text_questions'].get(question_type, '') |
| |
| present_categories = [category_a, category_b] |
| |
| if question_type == 'identify_overlap': |
| |
| anchor = random.choice([category_a, category_b]) |
| correct = category_b if anchor == category_a else category_a |
| |
| formatted_mcq = mcq_template.format(anchor_sound=anchor) |
| formatted_open = open_template.format(anchor_sound=anchor) |
| |
| mcq_data = self.question_generator.generate_category_mcq( |
| formatted_mcq, correct, present_categories, self.dataset.CATEGORIES |
| ) |
| open_data = self.question_generator.generate_category_open_text( |
| formatted_open, correct |
| ) |
| |
| elif question_type == 'overlap_pair': |
| |
| correct_pair = (category_a, category_b) |
| mcq_data = self.question_generator.generate_pair_mcq( |
| mcq_template, correct_pair, present_categories, self.dataset.CATEGORIES |
| ) |
| open_data = self.question_generator.generate_pair_open_text( |
| open_template, correct_pair |
| ) |
| |
| elif question_type == 'yes_no_overlap': |
| |
| |
| formatted_mcq = mcq_template.format(sound1=category_a, sound2=category_b) |
| formatted_open = open_template.format(sound1=category_a, sound2=category_b) |
| |
| mcq_data = self.question_generator.generate_yes_no_mcq( |
| formatted_mcq, correct_answer=True |
| ) |
| open_data = self.question_generator.generate_yes_no_open_text( |
| formatted_open, correct_answer=True |
| ) |
| |
| elif question_type == 'starts_before_end': |
| |
| |
| anchor = category_a |
| correct = category_b |
| |
| formatted_mcq = mcq_template.format(anchor_sound=anchor) |
| formatted_open = open_template.format(anchor_sound=anchor) |
| |
| mcq_data = self.question_generator.generate_category_mcq( |
| formatted_mcq, correct, present_categories, self.dataset.CATEGORIES |
| ) |
| open_data = self.question_generator.generate_category_open_text( |
| formatted_open, correct |
| ) |
| |
| else: |
| return None, None |
| |
| return mcq_data, open_data |
| |
| def generate_dataset(self) -> tuple: |
| """Generate the complete overlap dataset.""" |
| sample_durations = generate_sample_durations_for_task( |
| self.task_duration_hours, |
| self.min_clip_duration, |
| self.max_clip_duration |
| ) |
| num_samples = len(sample_durations) |
| |
| self.logger.info( |
| f"Generating {num_samples} overlap samples " |
| f"(target: {self.task_duration_hours}h)..." |
| ) |
| |
| all_metadata = [] |
| for i, duration in enumerate(sample_durations): |
| metadata = self.generate_sample(i, target_duration_seconds=duration) |
| if metadata is not None: |
| all_metadata.append(metadata) |
| |
| self.logger.info(f"Generated {len(all_metadata)}/{num_samples} samples successfully") |
| |
| |
| mcq_csv_path = self.output_base / 'overlap_mcq.csv' |
| self._save_mcq_csv(all_metadata, mcq_csv_path) |
| |
| open_text_csv_path = self.output_base / 'overlap_open_text.csv' |
| self._save_open_text_csv(all_metadata, open_text_csv_path) |
| |
| metadata_csv_path = self.output_base / 'overlap_metadata.csv' |
| self._save_metadata_csv(all_metadata, metadata_csv_path) |
| |
| self.logger.info(f"Overlap task complete!") |
| self.logger.info(f" - MCQ CSV: {mcq_csv_path}") |
| self.logger.info(f" - Open-text CSV: {open_text_csv_path}") |
| self.logger.info(f" - Metadata CSV: {metadata_csv_path}") |
| |
| return mcq_csv_path, open_text_csv_path |
| |
| def _save_mcq_csv(self, metadata_list: List[Dict], output_path: Path): |
| """Save MCQ format CSV.""" |
| with open(output_path, 'w', newline='') as f: |
| writer = csv.writer(f) |
| writer.writerow([ |
| 'question', 'id', 'audio_path', |
| 'optionA', 'optionB', 'optionC', 'optionD', |
| 'correct', 'question_type', 'source_wavs', 'source_categories', |
| 'overlap_duration_ms', 'overlap_ratio' |
| ]) |
| for meta in metadata_list: |
| writer.writerow([ |
| meta['mcq_question'], |
| meta['id'], |
| meta['audio_path'], |
| meta['mcq_options']['A'], |
| meta['mcq_options']['B'], |
| meta['mcq_options']['C'], |
| meta['mcq_options']['D'], |
| meta['mcq_correct_answer'], |
| meta['question_type'], |
| str(meta['source_files']), |
| str(meta['categories']), |
| meta['overlap_duration_ms'], |
| round(meta['overlap_ratio'], 3) |
| ]) |
| |
| def _save_open_text_csv(self, metadata_list: List[Dict], output_path: Path): |
| """Save open-text format CSV.""" |
| with open(output_path, 'w', newline='') as f: |
| writer = csv.writer(f) |
| writer.writerow([ |
| 'question', 'id', 'audio_path', 'answer', |
| 'question_type', 'source_wavs', 'source_categories', |
| 'overlap_duration_ms' |
| ]) |
| for meta in metadata_list: |
| writer.writerow([ |
| meta['open_text_question'], |
| meta['id'], |
| meta['audio_path'], |
| meta['open_text_answer'], |
| meta['question_type'], |
| str(meta['source_files']), |
| str(meta['categories']), |
| meta['overlap_duration_ms'] |
| ]) |
| |
| def _save_metadata_csv(self, metadata_list: List[Dict], output_path: Path): |
| """Save detailed metadata CSV.""" |
| with open(output_path, 'w', newline='') as f: |
| writer = csv.writer(f) |
| writer.writerow([ |
| 'id', 'audio_path', 'num_sounds', |
| 'source_files', 'source_categories', |
| 'category_a', 'category_b', |
| 'a_start_ms', 'a_end_ms', 'b_start_ms', 'b_end_ms', |
| 'overlap_duration_ms', 'overlap_ratio', 'total_duration_ms', |
| 'actual_duration_s', 'question_type' |
| ]) |
| |
| for meta in metadata_list: |
| writer.writerow([ |
| meta['id'], |
| meta['audio_path'], |
| meta['num_sounds'], |
| str(meta['source_files']), |
| str(meta['categories']), |
| meta['category_a'], |
| meta['category_b'], |
| meta['a_start_ms'], |
| meta['a_end_ms'], |
| meta['b_start_ms'], |
| meta['b_end_ms'], |
| meta['overlap_duration_ms'], |
| round(meta['overlap_ratio'], 3), |
| meta['total_duration_ms'], |
| meta['actual_duration_s'], |
| meta['question_type'] |
| ]) |
|
|
|
|
| def main(config_path: str = None): |
| """Main entry point for overlap task generation.""" |
| import yaml |
| |
| if config_path is None: |
| config_path = Path(__file__).parent.parent / 'config.yaml' |
| |
| with open(config_path, 'r') as f: |
| config = yaml.safe_load(f) |
| |
| set_random_seed(config['random_seed']) |
| |
| logger = setup_logger( |
| 'overlap_task', |
| log_file=str(Path(config['output']['base_path']) / config['logging']['log_file']), |
| level=config['logging']['level'], |
| console_output=config['logging']['console_output'] |
| ) |
| |
| generator = OverlapTaskGenerator(config, logger) |
| generator.generate_dataset() |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|