| """ |
| During/Contains task generator for temporal reasoning dataset. |
| |
| Generates audio samples where one sound event is fully contained within |
| another (temporally). The "container" sound plays for a longer duration, |
| and the "contained" sound starts and ends entirely within the container. |
| |
| Questions test whether a model can identify containment relationships: |
| - Which sound occurs during another? |
| - Which sound contains another? |
| - Yes/No containment 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_during_contains_task_audio, |
| create_dataset |
| ) |
|
|
|
|
| class DuringContainsTaskGenerator: |
| """Generates during_contains task dataset samples.""" |
| |
| def __init__(self, config: dict, logger=None): |
| """ |
| Initialize DuringContainsTaskGenerator. |
| |
| Args: |
| config: Full pipeline configuration dictionary |
| logger: Logger instance |
| """ |
| self.config = config |
| self.logger = logger or setup_logger(__name__) |
| |
| self.task_config = config['tasks']['during_contains'] |
| |
| |
| 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_container_duration_s = self.task_config.get('min_container_duration_s', 8.0) |
| self.max_container_duration_s = self.task_config.get('max_container_duration_s', 15.0) |
| self.min_margin_s = self.task_config.get('min_margin_s', 0.5) |
| self.min_margin_ms = int(self.min_margin_s * 1000) |
| |
| |
| 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']) / 'during_contains' |
| 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 during_contains task sample. |
| |
| Pipeline: |
| 1. Sample 2 categories (container and contained) |
| 2. Load audio for each |
| 3. Extend container to min_container_duration_s..max_container_duration_s |
| 4. Keep contained at source clip duration (shorter) |
| 5. Build containment audio using overlay() |
| 6. Generate questions about the containment |
| |
| 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 |
| |
| container_category = categories[0] |
| contained_category = categories[1] |
| |
| |
| filename_container, filepath_container = self.dataset.sample_file_from_category(container_category) |
| filename_contained, filepath_contained = self.dataset.sample_file_from_category(contained_category) |
| |
| container_audio = self.audio_processor.load_audio(filepath_container) |
| contained_audio = self.audio_processor.load_audio(filepath_contained) |
| |
| |
| container_duration_s = random.uniform( |
| self.min_container_duration_s, |
| self.max_container_duration_s |
| ) |
| |
| |
| max_contained_duration_ms = len(container_audio) - 2 * self.min_margin_ms |
| if len(contained_audio) > max_contained_duration_ms: |
| contained_audio = contained_audio[:max_contained_duration_ms] |
| |
| |
| try: |
| final_audio, build_metadata = build_during_contains_task_audio( |
| container_audio, |
| contained_audio, |
| container_category, |
| contained_category, |
| min_margin_ms=self.min_margin_ms |
| ) |
| except ValueError as e: |
| self.logger.warning(f"Sample {sample_id}: Failed to build containment 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, container_category, contained_category, 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_container, filename_contained], |
| 'container_category': container_category, |
| 'contained_category': contained_category, |
| 'categories': categories, |
| 'question_type': question_type, |
| 'container_duration_ms': build_metadata['container_duration_ms'], |
| 'contained_duration_ms': build_metadata['contained_duration_ms'], |
| 'contained_start_ms': build_metadata['contained_start_ms'], |
| 'contained_end_ms': build_metadata['contained_end_ms'], |
| 'margin_before_ms': build_metadata['margin_before_ms'], |
| 'margin_after_ms': build_metadata['margin_after_ms'], |
| '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 during_contains sample {sample_id}: " |
| f"{container_category} contains {contained_category}, " |
| f"contained at {build_metadata['contained_start_ms']}-" |
| f"{build_metadata['contained_end_ms']}ms, type={question_type}" |
| ) |
| |
| return metadata |
| |
| def _generate_question( |
| self, |
| question_type: str, |
| container_category: str, |
| contained_category: 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 = [container_category, contained_category] |
| |
| if question_type == 'during': |
| |
| anchor = container_category |
| correct = contained_category |
| |
| 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 == 'contains': |
| |
| target = contained_category |
| correct = container_category |
| |
| formatted_mcq = mcq_template.format(target_sound=target) |
| formatted_open = open_template.format(target_sound=target) |
| |
| 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 == 'yes_no_during': |
| |
| |
| formatted_mcq = mcq_template.format( |
| target_sound=contained_category, |
| anchor_sound=container_category |
| ) |
| formatted_open = open_template.format( |
| target_sound=contained_category, |
| anchor_sound=container_category |
| ) |
| |
| 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 |
| ) |
| |
| else: |
| return None, None |
| |
| return mcq_data, open_data |
| |
| def generate_dataset(self) -> tuple: |
| """Generate the complete during_contains 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} during_contains 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 / 'during_contains_mcq.csv' |
| self._save_mcq_csv(all_metadata, mcq_csv_path) |
| |
| open_text_csv_path = self.output_base / 'during_contains_open_text.csv' |
| self._save_open_text_csv(all_metadata, open_text_csv_path) |
| |
| metadata_csv_path = self.output_base / 'during_contains_metadata.csv' |
| self._save_metadata_csv(all_metadata, metadata_csv_path) |
| |
| self.logger.info(f"During/Contains 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', |
| 'container_duration_ms', 'contained_duration_ms' |
| ]) |
| 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['container_duration_ms'], |
| meta['contained_duration_ms'] |
| ]) |
| |
| 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', |
| 'container_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['container_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', |
| 'container_category', 'contained_category', |
| 'container_duration_ms', 'contained_duration_ms', |
| 'contained_start_ms', 'contained_end_ms', |
| 'margin_before_ms', 'margin_after_ms', |
| '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['container_category'], |
| meta['contained_category'], |
| meta['container_duration_ms'], |
| meta['contained_duration_ms'], |
| meta['contained_start_ms'], |
| meta['contained_end_ms'], |
| meta['margin_before_ms'], |
| meta['margin_after_ms'], |
| meta['total_duration_ms'], |
| meta['actual_duration_s'], |
| meta['question_type'] |
| ]) |
|
|
|
|
| def main(config_path: str = None): |
| """Main entry point for during_contains 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( |
| 'during_contains_task', |
| log_file=str(Path(config['output']['base_path']) / config['logging']['log_file']), |
| level=config['logging']['level'], |
| console_output=config['logging']['console_output'] |
| ) |
| |
| generator = DuringContainsTaskGenerator(config, logger) |
| generator.generate_dataset() |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|