TREA_2.0_codebase / tasks /task_during_contains.py
malay-36's picture
Upload updated pipeline codebase
7e6c03a verified
Raw
History Blame Contribute Delete
17 kB
"""
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 parameters
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)
# Task-specific parameters
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)
# Dataset
self.dataset = create_dataset(config)
# Audio processor
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)
)
# Question generator
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')
)
# Output paths
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
"""
# Step 1: Sample categories
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]
# Step 2: Load audio
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)
# Step 3: Extend container to target duration
container_duration_s = random.uniform(
self.min_container_duration_s,
self.max_container_duration_s
)
# Step 4: Contained stays at source clip duration (5s typically)
# Make sure contained is shorter than container minus margins
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]
# Step 5: Build containment audio
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
# Save audio
output_audio_path = self.audio_output / f"{sample_id}.wav"
final_audio.export(str(output_audio_path), format="wav")
# Step 6: Generate questions
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':
# "Which sound occurs during {anchor_sound}?"
anchor = container_category # The container is the anchor
correct = contained_category # The contained occurs "during" the container
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':
# "Which longer sound contains {target_sound}?"
target = contained_category # The shorter contained sound
correct = container_category # The container "contains" the target
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':
# "Does {target_sound} happen entirely during {anchor_sound}?"
# Always True by construction
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")
# Save CSVs
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()