TREA_2.0_codebase / tasks /task_multi_hop.py
malay-36's picture
Upload updated pipeline codebase
7e6c03a verified
Raw
History Blame Contribute Delete
14.2 kB
"""
Multi-Hop task generator for temporal reasoning dataset.
Multi-hop inter-task: Explicit two-step or three-step temporal reasoning
over combinations of order, duration, volume, silence, and count.
Example chains:
- "What sound occurs after the longest sound?" (order + duration)
- "What sound occurs before the loudest sound?" (order + volume)
- "How many sounds occur after the longest silence?" (count + silence)
"""
import random
from collections import Counter
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,
get_lufs_loudness,
create_preprocessed_dataset,
)
from tasks.multihop_base import MultihopBaseGenerator
class MultiHopTaskGenerator(MultihopBaseGenerator):
"""Generates multi_hop task dataset samples."""
TASK_NAME = "multi_hop"
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 multi_hop sample.
Builds a rich scene with duration, volume, and silence variation,
then asks questions requiring chained reasoning.
"""
n_events = random.randint(
self.task_config.get("min_events", 5),
self.task_config.get("max_events", 8),
)
# Allow some repeats for count_before_loudest / count_after_longest
n_unique = random.randint(
max(3, n_events // 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)
# Volume variation
volume_range = self.task_config.get("volume_range_db", [-10, 6])
volume_levels = [
round(random.uniform(volume_range[0], volume_range[1]), 1)
for _ in range(n_events)
]
# Ensure min volume difference
min_diff = self.task_config.get("min_volume_diff_db", 4.0)
vmin, vmax = min(volume_levels), max(volume_levels)
if vmax - vmin < min_diff:
idx_min = volume_levels.index(vmin)
volume_levels[idx_min] = vmax - min_diff
from pydub import AudioSegment as PydubSegment
source_files = []
audio_segments = []
effective_durations_ms = []
measured_loudness = []
for i, cat in enumerate(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 = audio.apply_gain(volume_levels[i])
loudness = get_lufs_loudness(audio)
audio_segments.append(audio)
source_files.append(fname)
effective_durations_ms.append(int(eff_dur * 1000))
measured_loudness.append(loudness)
# Gaps
num_gaps = n_events - 1
gap_min = self.task_config.get("min_gap_ms", 500)
gap_max = self.task_config.get("max_gap_ms", 3000)
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.5,
)
else:
gap_durations = []
# Assemble
result = audio_segments[0]
for i in range(1, n_events):
result = result + PydubSegment.silent(duration=gap_durations[i - 1])
result = result + audio_segments[i]
output_path = self.audio_output / f"{sample_id}.wav"
result.export(str(output_path), format="wav")
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, effective_durations_ms,
measured_loudness, volume_levels, 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,
"volume_levels_db": volume_levels,
"effective_durations_ms": effective_durations_ms,
"measured_loudness_lufs": measured_loudness,
"gap_durations_ms": gap_durations,
"question_type": question_type,
"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 multi_hop sample {sample_id}: "
f"{n_events} events, type={question_type}"
)
return metadata
def _generate_question(
self, question_type, categories, durations_ms,
loudness_levels, volume_levels, gap_durations
):
"""Generate multi-hop chained reasoning question."""
n = len(categories)
# ---- Step 1: Identify the "pivot" event based on first property ----
if question_type in ("after_longest", "before_longest"):
pivot_idx = durations_ms.index(max(durations_ms))
pivot_label = "longest sound"
elif question_type == "after_shortest":
pivot_idx = durations_ms.index(min(durations_ms))
pivot_label = "shortest sound"
elif question_type == "before_loudest":
pivot_idx = loudness_levels.index(max(loudness_levels))
pivot_label = "loudest sound"
elif question_type == "after_longest_gap":
if not gap_durations:
return None, None, {}
longest_gap_idx = gap_durations.index(max(gap_durations))
# Events after the longest gap
after_idx = longest_gap_idx + 1
events_after = categories[after_idx:]
count_after = len(events_after)
# Pick template variant — count or identify
if random.random() < 0.5 and events_after:
# Identify first event after longest gap
correct = categories[after_idx]
mcq_text = self.task_config["mcq_questions"]["after_longest_gap"]
open_text = self.task_config["open_text_questions"]["after_longest_gap"]
# If template asks "how many" use count, else use category
if "how many" in mcq_text.lower():
mcq_data = self.question_generator.generate_count_mcq(
mcq_text, count_after, self.dataset.CATEGORIES
)
open_data = self.question_generator.generate_count_open_text(
open_text, count_after
)
else:
mcq_data = self.question_generator.generate_category_mcq(
mcq_text, correct, categories, self.dataset.CATEGORIES
)
open_data = self.question_generator.generate_category_open_text(
open_text, correct
)
q_meta = {"longest_gap_idx": longest_gap_idx, "correct_value": correct,
"count_after_longest_gap": count_after}
return mcq_data, open_data, q_meta
else:
mcq_text = f"How many sounds occur after the longest silence?"
open_text = f"How many sounds occur after the longest silence?"
mcq_data = self.question_generator.generate_count_mcq(
mcq_text, count_after, self.dataset.CATEGORIES
)
open_data = self.question_generator.generate_count_open_text(
open_text, count_after
)
q_meta = {"longest_gap_idx": longest_gap_idx,
"count_after_longest_gap": count_after}
return mcq_data, open_data, q_meta
elif question_type in ("count_before_loudest", "count_after_longest"):
if question_type == "count_before_loudest":
pivot_idx = loudness_levels.index(max(loudness_levels))
region = categories[:pivot_idx]
else:
pivot_idx = durations_ms.index(max(durations_ms))
region = categories[pivot_idx + 1:]
# Count occurrences of a target sound in region
unique_in_region = list(set(region))
if unique_in_region:
target_sound = random.choice(unique_in_region)
else:
target_sound = random.choice(list(set(categories)))
count = region.count(target_sound)
mcq_text = self.task_config["mcq_questions"][question_type].format(target_sound=target_sound)
open_text = self.task_config["open_text_questions"][question_type].format(target_sound=target_sound)
mcq_data = self.question_generator.generate_count_mcq(
mcq_text, count, self.dataset.CATEGORIES
)
open_data = self.question_generator.generate_count_open_text(
open_text, count
)
q_meta = {"target_sound": target_sound, "correct_count": count,
"pivot_index": pivot_idx}
return mcq_data, open_data, q_meta
elif question_type in ("overlap_after_anchor", "loudest_after_anchor",
"longest_before_anchor"):
# Pick anchor
anchor_idx = random.randint(1, n - 2) if n > 2 else 0
anchor_sound = categories[anchor_idx]
if question_type == "loudest_after_anchor":
subset = [(i, categories[i], loudness_levels[i])
for i in range(anchor_idx + 1, n)]
if not subset:
return None, None, {}
best = max(subset, key=lambda x: x[2])
correct = best[1]
elif question_type == "longest_before_anchor":
subset = [(i, categories[i], durations_ms[i])
for i in range(0, anchor_idx)]
if not subset:
return None, None, {}
best = max(subset, key=lambda x: x[2])
correct = best[1]
elif question_type == "overlap_after_anchor":
# In sequential audio, no true overlap — use event right after anchor
if anchor_idx + 1 < n:
correct = categories[anchor_idx + 1]
else:
return None, None, {}
else:
return None, None, {}
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)
present = categories[anchor_idx + 1:] if "after" in question_type else categories[:anchor_idx]
mcq_data = self.question_generator.generate_category_mcq(
mcq_text, correct, present or categories, self.dataset.CATEGORIES
)
open_data = self.question_generator.generate_category_open_text(
open_text, correct
)
q_meta = {"anchor_sound": anchor_sound, "correct_category": correct}
return mcq_data, open_data, q_meta
else:
# Default: not one of the special cases above
return None, None, {}
# ---- Step 2: Get the event after/before the pivot ----
if "after" in question_type:
target_idx = pivot_idx + 1
if target_idx >= n:
return None, None, {}
else: # "before"
target_idx = pivot_idx - 1
if target_idx < 0:
return None, None, {}
correct = categories[target_idx]
mcq_text = self.task_config["mcq_questions"][question_type]
open_text = self.task_config["open_text_questions"][question_type]
mcq_data = self.question_generator.generate_category_mcq(
mcq_text, correct, categories, self.dataset.CATEGORIES
)
open_data = self.question_generator.generate_category_open_text(
open_text, correct
)
q_meta = {
"pivot_index": pivot_idx,
"pivot_category": categories[pivot_idx],
"correct_category": correct,
}
return mcq_data, open_data, q_meta
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(
"multi_hop_task",
log_file=str(Path(config["output"]["base_path"]) / config["logging"]["log_file"]),
level=config["logging"]["level"],
console_output=config["logging"]["console_output"],
)
generator = MultiHopTaskGenerator(config, logger)
generator.generate_dataset()
if __name__ == "__main__":
main()