TREA_2.0_codebase / tasks /task_conditional_duration.py
malay-36's picture
Upload updated pipeline codebase
7e6c03a verified
Raw
History Blame Contribute Delete
10.4 kB
"""
Conditional Duration task generator for temporal reasoning dataset.
Multi-hop task: First apply a temporal condition (before/after anchor),
then compare durations of events satisfying that condition.
Uses PreprocessedESC50Dataset for accurate effective durations.
"""
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,
concatenate_to_target_duration,
generate_controlled_gap_durations,
create_preprocessed_dataset,
)
from tasks.multihop_base import MultihopBaseGenerator
class ConditionalDurationTaskGenerator(MultihopBaseGenerator):
"""Generates conditional_duration task dataset samples."""
TASK_NAME = "conditional_duration"
def __init__(self, config: dict, logger=None):
super().__init__(config, logger)
# Override dataset with preprocessed version for effective durations
preprocessed_path = self.task_config.get(
"preprocessed_data_path",
config["tasks"].get("duration", {}).get("preprocessed_data_path", ""),
)
self.dataset = create_preprocessed_dataset(config, preprocessed_path=preprocessed_path)
def generate_sample(
self,
sample_id: int,
target_duration_seconds: float = None,
question_type: str = None,
) -> Optional[Dict]:
"""
Generate a single conditional_duration sample.
Pipeline:
1. Build sequential scene with 5-8 events, each with known effective duration
2. Pick an anchor sound
3. Among events before/after anchor, find longest/shortest
4. Generate question
"""
n_events = random.randint(
self.task_config.get("min_events", 5),
self.task_config.get("max_events", 8),
)
# Sample unique categories
n_unique = min(n_events, len(self.dataset.CATEGORIES))
categories = self.dataset.sample_categories(n_unique)
# Pad if needed
while len(categories) < n_events:
categories.append(random.choice(categories[:n_unique]))
random.shuffle(categories)
# Load audio with known durations
from pydub import AudioSegment as PydubSegment
source_files = []
audio_segments = []
effective_durations = []
num_gaps = n_events - 1
for cat in categories:
fname, fpath, eff_dur = self.dataset.sample_file_from_category_with_duration(cat)
audio = self.audio_processor.load_audio(fpath)
# Use effective duration to extend clip
per_event_s = max(self.source_clip_duration, eff_dur)
audio_segments.append(audio)
source_files.append(fname)
effective_durations.append(eff_dur)
# Build gaps
if num_gaps > 0:
gap_durations = generate_controlled_gap_durations(
num_gaps, min_gap_ms=300, max_gap_ms=1500, gap_multiplier=2.0
)
else:
gap_durations = []
# Assemble audio
result = audio_segments[0]
events_meta = [
{"index": 0, "category": categories[0],
"start_ms": 0, "end_ms": len(audio_segments[0]),
"effective_duration_s": effective_durations[0]}
]
current_ms = len(audio_segments[0])
for i in range(1, n_events):
gap_ms = gap_durations[i - 1]
result = result + PydubSegment.silent(duration=gap_ms)
current_ms += gap_ms
event_start = current_ms
result = result + audio_segments[i]
current_ms += len(audio_segments[i])
events_meta.append(
{"index": i, "category": categories[i],
"start_ms": event_start, "end_ms": current_ms,
"effective_duration_s": effective_durations[i],
"gap_before_ms": gap_ms}
)
# Save audio
output_path = self.audio_output / f"{sample_id}.wav"
result.export(str(output_path), format="wav")
# Select question type
if question_type is None:
question_type = random.choice(self.task_config["question_types"])
# Pick anchor (not first or last for before/after questions)
anchor_idx = random.randint(1, n_events - 2) if n_events > 2 else 0
anchor_sound = categories[anchor_idx]
# Compute answer based on question type
mcq_data, open_data, q_meta = self._generate_question(
question_type, categories, effective_durations, anchor_idx, anchor_sound
)
if mcq_data is None:
return None
metadata = {
"id": sample_id,
"audio_path": str(output_path.relative_to(self.output_base.parent)),
"n_events": n_events,
"categories": categories,
"source_files": source_files,
"effective_durations": effective_durations,
"question_type": question_type,
"anchor_sound": anchor_sound,
"anchor_index": anchor_idx,
"target_duration_s": target_duration_seconds,
"actual_duration_s": len(result) / 1000.0,
"mcq_question": mcq_data["question"],
"mcq_options": mcq_data["options"],
"mcq_correct_answer": mcq_data["correct_answer"],
"open_text_question": open_data["question"],
"open_text_answer": open_data["correct_answer"],
**q_meta,
}
self.logger.info(
f"Generated conditional_duration sample {sample_id}: "
f"{n_events} events, type={question_type}, anchor={anchor_sound}"
)
return metadata
def _generate_question(
self,
question_type: str,
categories: List[str],
effective_durations: List[float],
anchor_idx: int,
anchor_sound: str,
):
"""Generate conditional duration question."""
n = len(categories)
if question_type in ("longest_after", "shortest_after"):
subset_indices = list(range(anchor_idx + 1, n))
elif question_type in ("longest_before", "shortest_before"):
subset_indices = list(range(0, anchor_idx))
elif question_type == "repeated_compare":
# Find a category that appears both before and after anchor
before_cats = set(categories[:anchor_idx])
after_cats = set(categories[anchor_idx + 1:])
common = before_cats & after_cats
if not common:
return None, None, {}
target_sound = random.choice(list(common))
# Find indices
before_idx = [i for i in range(anchor_idx) if categories[i] == target_sound]
after_idx = [i for i in range(anchor_idx + 1, n) if categories[i] == target_sound]
before_dur = effective_durations[before_idx[0]]
after_dur = effective_durations[after_idx[0]]
if before_dur >= after_dur:
correct = f"before {anchor_sound}"
else:
correct = f"after {anchor_sound}"
mcq_text = self.task_config["mcq_questions"]["repeated_compare"].format(target_sound=target_sound, anchor_sound=anchor_sound)
open_text = self.task_config["open_text_questions"]["repeated_compare"].format(target_sound=target_sound, anchor_sound=anchor_sound)
options = [f"before {anchor_sound}", f"after {anchor_sound}"]
# Add distractors
other = [c for c in self.dataset.CATEGORIES if c != target_sound and c != anchor_sound]
random.shuffle(other)
options.extend(other[:2])
random.shuffle(options)
option_labels = ["A", "B", "C", "D"]
correct_label = option_labels[options.index(correct)]
option_map = {l: v for l, v in zip(option_labels, options)}
mcq_data = {"question": mcq_text, "options": option_map,
"correct_answer": correct_label, "correct_value": correct}
open_data = {"question": open_text, "correct_answer": correct}
q_meta = {"target_sound": target_sound, "correct_value": correct}
return mcq_data, open_data, q_meta
else:
return None, None, {}
if not subset_indices:
return None, None, {}
# Find longest or shortest in subset
subset_durations = [(i, effective_durations[i]) for i in subset_indices]
if "longest" in question_type:
best_idx, best_dur = max(subset_durations, key=lambda x: x[1])
else:
best_idx, best_dur = min(subset_durations, key=lambda x: x[1])
correct_category = categories[best_idx]
present_cats = [categories[i] for i in subset_indices]
mcq_text = self.task_config["mcq_questions"][question_type].format(anchor_sound=anchor_sound)
open_text = self.task_config["open_text_questions"][question_type].format(anchor_sound=anchor_sound)
mcq_data = self.question_generator.generate_category_mcq(
mcq_text, correct_category, present_cats, self.dataset.CATEGORIES
)
open_data = self.question_generator.generate_category_open_text(
open_text, correct_category
)
q_meta = {
"correct_category": correct_category,
"correct_duration_s": best_dur,
"subset_size": len(subset_indices),
}
return mcq_data, open_data, q_meta
def main(config_path: str = None):
"""Main entry point for conditional_duration 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(
"conditional_duration_task",
log_file=str(Path(config["output"]["base_path"]) / config["logging"]["log_file"]),
level=config["logging"]["level"],
console_output=config["logging"]["console_output"],
)
generator = ConditionalDurationTaskGenerator(config, logger)
generator.generate_dataset()
if __name__ == "__main__":
main()