TREA_2.0_codebase / tasks /task_temporal_arithmetic.py
malay-36's picture
Upload updated pipeline codebase
7e6c03a verified
Raw
History Blame Contribute Delete
10.2 kB
"""
Temporal Arithmetic task generator for temporal reasoning dataset.
Multi-hop inter-task: Combines duration, count/repetition, and silence
reasoning by comparing total active time across labels or vs silence.
"""
import random
from pathlib import Path
from typing import Dict, List, Optional
from utils import (
setup_logger,
set_random_seed,
concatenate_to_target_duration,
generate_controlled_gap_durations,
create_preprocessed_dataset,
)
from tasks.multihop_base import MultihopBaseGenerator
class TemporalArithmeticTaskGenerator(MultihopBaseGenerator):
"""Generates temporal_arithmetic task dataset samples."""
TASK_NAME = "temporal_arithmetic"
def __init__(self, config: dict, logger=None):
super().__init__(config, logger)
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 temporal_arithmetic sample.
Pipeline:
1. Build scene with 4-8 events (some repeated categories)
2. Track total duration per label and total silence
3. Generate arithmetic comparison questions
"""
n_events = random.randint(
self.task_config.get("min_events", 4),
self.task_config.get("max_events", 8),
)
# Allow repeated categories for aggregation
n_unique = random.randint(2, min(n_events, len(self.dataset.CATEGORIES)))
categories_pool = self.dataset.sample_categories(n_unique)
categories = list(categories_pool)
while len(categories) < n_events:
categories.append(random.choice(categories_pool))
random.shuffle(categories)
from pydub import AudioSegment as PydubSegment
source_files = []
audio_segments = []
effective_durations_ms = []
for cat in categories:
fname, fpath, eff_dur = self.dataset.sample_file_from_category_with_duration(cat)
audio = self.audio_processor.load_audio(fpath)
audio = concatenate_to_target_duration(audio, max(self.source_clip_duration, eff_dur))
audio_segments.append(audio)
source_files.append(fname)
effective_durations_ms.append(int(eff_dur * 1000))
# Generate gaps
num_gaps = n_events - 1
gap_min = self.task_config.get("min_gap_ms", 300)
gap_max = self.task_config.get("max_gap_ms", 2000)
if num_gaps > 0:
gap_durations = generate_controlled_gap_durations(
num_gaps, min_gap_ms=gap_min, max_gap_ms=gap_max, gap_multiplier=2.0
)
else:
gap_durations = []
# Assemble
result = audio_segments[0]
for i in range(1, n_events):
gap_ms = gap_durations[i - 1]
result = result + PydubSegment.silent(duration=gap_ms)
result = result + audio_segments[i]
output_path = self.audio_output / f"{sample_id}.wav"
result.export(str(output_path), format="wav")
# Compute aggregates
# Total duration per label
label_totals = {}
for cat, dur in zip(categories, effective_durations_ms):
label_totals[cat] = label_totals.get(cat, 0) + dur
total_silence_ms = sum(gap_durations) if gap_durations else 0
longest_silence_ms = max(gap_durations) if gap_durations else 0
if question_type is None:
question_type = random.choice(self.task_config["question_types"])
mcq_data, open_data, q_meta = self._generate_question(
question_type, categories, label_totals, total_silence_ms,
longest_silence_ms, gap_durations
)
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,
"question_type": question_type,
"label_totals_ms": label_totals,
"total_silence_ms": total_silence_ms,
"gap_durations_ms": gap_durations,
"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 temporal_arithmetic sample {sample_id}: "
f"{n_events} events, type={question_type}"
)
return metadata
def _generate_question(
self, question_type, categories, label_totals,
total_silence_ms, longest_silence_ms, gap_durations
):
"""Generate temporal arithmetic question."""
unique_cats = list(label_totals.keys())
if question_type == "total_label_duration":
if len(unique_cats) < 2:
return None, None, {}
sound1, sound2 = random.sample(unique_cats, 2)
dur1 = label_totals[sound1]
dur2 = label_totals[sound2]
correct = sound1 if dur1 >= dur2 else sound2
mcq_text = self.task_config["mcq_questions"]["total_label_duration"].format(sound1=sound1, sound2=sound2)
open_text = self.task_config["open_text_questions"]["total_label_duration"].format(sound1=sound1, sound2=sound2)
mcq_data = self.question_generator.generate_pairwise_comparison_mcq(
mcq_text, correct, (sound1, sound2), self.dataset.CATEGORIES
)
open_data = self.question_generator.generate_category_open_text(
open_text, correct
)
q_meta = {"dur1_ms": dur1, "dur2_ms": dur2}
return mcq_data, open_data, q_meta
elif question_type == "label_vs_silence":
target_sound = random.choice(unique_cats)
label_dur = label_totals[target_sound]
if label_dur >= total_silence_ms:
correct = target_sound
else:
correct = "total silence"
mcq_text = self.task_config["mcq_questions"]["label_vs_silence"].format(target_sound=target_sound)
open_text = self.task_config["open_text_questions"]["label_vs_silence"].format(target_sound=target_sound)
options = [target_sound, "total silence"]
other = [c for c in self.dataset.CATEGORIES if c != target_sound][:2]
options.extend(other)
random.shuffle(options)
options = options[:4]
if correct not in options:
options[0] = correct
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}
open_data = {"question": open_text, "correct_answer": correct}
q_meta = {"label_duration_ms": label_dur, "total_silence_ms": total_silence_ms}
return mcq_data, open_data, q_meta
elif question_type == "combined_duration":
target_sound = random.choice(unique_cats)
combined_dur = label_totals[target_sound]
if combined_dur >= longest_silence_ms:
correct = f"all {target_sound} sounds combined"
else:
correct = "the longest silence"
mcq_text = self.task_config["mcq_questions"]["combined_duration"].format(target_sound=target_sound, sound1=target_sound,
sound2=random.choice([c for c in unique_cats if c != target_sound] or unique_cats))
open_text = self.task_config["open_text_questions"]["combined_duration"].format(target_sound=target_sound, sound1=target_sound,
sound2=random.choice([c for c in unique_cats if c != target_sound] or unique_cats))
options = [f"all {target_sound} sounds combined", "the longest silence"]
other = [c for c in self.dataset.CATEGORIES if c != target_sound][:2]
options.extend(other)
random.shuffle(options)
options = options[:4]
if correct not in options:
options[0] = correct
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}
open_data = {"question": open_text, "correct_answer": correct}
q_meta = {"combined_duration_ms": combined_dur,
"longest_silence_ms": longest_silence_ms}
return mcq_data, open_data, q_meta
return None, None, {}
def main(config_path: str = None):
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(
"temporal_arithmetic_task",
log_file=str(Path(config["output"]["base_path"]) / config["logging"]["log_file"]),
level=config["logging"]["level"],
console_output=config["logging"]["console_output"],
)
generator = TemporalArithmeticTaskGenerator(config, logger)
generator.generate_dataset()
if __name__ == "__main__":
main()