TREA_2.0_codebase / tasks /multihop_base.py
malay-36's picture
Upload updated pipeline codebase
7e6c03a verified
Raw
History Blame Contribute Delete
18 kB
"""
Shared base class for multihop temporal reasoning task generators.
All multihop tasks build "audio scenes" — sequential arrangements of sound
events with controlled order, durations, volume levels, and silences. The
base class provides the common scene-building infrastructure; subclasses
implement task-specific question generation.
"""
import csv
import random
from collections import Counter
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,
get_max_clip_num_to_be_joined,
build_clip_sequence_with_silences,
generate_controlled_gap_durations,
get_lufs_loudness,
normalize_to_lufs,
create_dataset,
)
class MultihopBaseGenerator:
"""
Base class for multihop temporal reasoning task generators.
Provides shared utilities for building multi-event audio scenes and
standard CSV output helpers. Subclasses must implement:
- generate_sample(sample_id, target_duration_seconds)
- TASK_NAME (class attribute)
"""
TASK_NAME: str = "multihop_base" # override in subclasses
# ------------------------------------------------------------------ init
def __init__(self, config: dict, logger=None):
self.config = config
self.logger = logger or setup_logger(self.TASK_NAME)
self.task_config = config["tasks"][self.TASK_NAME]
# Audio parameters
audio_cfg = config["audio"]
self.min_clip_duration = audio_cfg["min_clip_duration"]
self.max_clip_duration = audio_cfg["max_clip_duration"]
self.source_clip_duration = audio_cfg.get("source_clip_duration", 5.0)
self.min_silence_ms = audio_cfg.get("min_silence_duration", 100)
self.max_extra_silence_per_gap_ms = audio_cfg.get(
"max_extra_silence_per_gap", 500
)
self.crossfade_ms = audio_cfg.get("crossfade_duration", 0)
self.task_duration_hours = self.task_config["task_duration_size"]
# Dataset — subclasses may override with duration-aware adapter
self.dataset = create_dataset(config)
# Audio processor (no automatic silence — we control gaps manually)
self.audio_processor = AudioProcessor(
crossfade_duration=audio_cfg.get("crossfade_duration", 500),
silence_duration=audio_cfg.get("silence_duration", 1000),
with_silence=False,
normalize=audio_cfg.get("normalize", False),
normalize_target_dBFS=audio_cfg.get("normalize_target_dBFS", -20.0),
)
# Question generator
mcq_cfg = config.get("mcq", {})
self.question_generator = QuestionGenerator(
num_options=mcq_cfg.get("num_options", 4),
option_labels=mcq_cfg.get("option_labels", ["A", "B", "C", "D"]),
distractor_strategy=mcq_cfg.get("distractor_strategy", "balanced"),
)
# Output paths
self.output_base = Path(config["output"]["base_path"]) / self.TASK_NAME
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)
# --------------------------------------------------------- scene builder
def build_sequential_scene(
self,
n_events: int,
target_duration_s: float,
volume_levels: Optional[List[float]] = None,
gap_min_ms: int = 300,
gap_max_ms: int = 2000,
gap_multiplier: float = 2.0,
) -> Tuple[
"AudioSegment", # final audio
List[str], # categories in order
List[str], # source filenames
List[Dict], # per-event metadata
Dict, # build metadata
]:
"""
Build an audio scene with *n_events* sequential sound events.
Each event uses a unique ESC-50 category, extended to roughly fill
its time slot. Controlled silences separate events.
Args:
n_events: Number of sound events (unique categories).
target_duration_s: Target total audio length.
volume_levels: Per-event dB adjustments (optional).
gap_min_ms / gap_max_ms: Range for silence gaps.
gap_multiplier: Ensure longest gap >= shortest gap * multiplier.
Returns:
(final_audio, categories, source_filenames, events_meta, build_meta)
"""
from pydub import AudioSegment as PydubSegment
# 1. Sample categories
n_events = min(n_events, len(self.dataset.CATEGORIES))
categories = self.dataset.sample_categories(n_events)
random.shuffle(categories)
# 2. Load source audio for each event
source_files = []
audio_segments = []
num_gaps = n_events - 1
for i, cat in enumerate(categories):
fname, fpath = self.dataset.sample_file_from_category(cat)
audio = self.audio_processor.load_audio(fpath)
# Apply volume adjustment if specified
if volume_levels and i < len(volume_levels):
audio = audio.apply_gain(volume_levels[i])
audio_segments.append(audio)
source_files.append(fname)
# 3. Generate controlled gap durations
if num_gaps > 0:
gap_durations = generate_controlled_gap_durations(
num_gaps,
min_gap_ms=gap_min_ms,
max_gap_ms=gap_max_ms,
gap_multiplier=gap_multiplier,
)
# Distribute extra silence to meet target duration
from utils import distribute_remainder_as_silences
total_audio_ms = sum(len(seg) for seg in audio_segments)
total_gap_ms = sum(gap_durations)
target_ms = int(target_duration_s * 1000)
available_extra_ms = target_ms - total_audio_ms - total_gap_ms
if available_extra_ms > 0:
extra_silences = distribute_remainder_as_silences(
available_extra_ms, num_gaps, max_per_gap_ms=2000
)
gap_durations = [g + e for g, e in zip(gap_durations, extra_silences)]
else:
gap_durations = []
# 4. Build final audio
result = audio_segments[0]
current_ms = 0
events_meta = []
events_meta.append(
{
"index": 0,
"category": categories[0],
"start_ms": 0,
"end_ms": len(audio_segments[0]),
"duration_ms": len(audio_segments[0]),
"volume_db": volume_levels[0] if volume_levels else 0,
"source_file": source_files[0],
}
)
current_ms = len(audio_segments[0])
for i in range(1, n_events):
gap_ms = gap_durations[i - 1]
silence = PydubSegment.silent(duration=gap_ms)
result = result + silence
gap_start = current_ms
current_ms += gap_ms
event_start = current_ms
result = result + audio_segments[i]
event_end = current_ms + len(audio_segments[i])
current_ms = event_end
events_meta.append(
{
"index": i,
"category": categories[i],
"start_ms": event_start,
"end_ms": event_end,
"duration_ms": len(audio_segments[i]),
"volume_db": volume_levels[i] if volume_levels else 0,
"source_file": source_files[i],
"gap_before_ms": gap_ms,
}
)
# 5. Build metadata
build_meta = {
"n_events": n_events,
"gap_durations_ms": gap_durations,
"total_duration_ms": len(result),
}
if gap_durations:
longest_gap_idx = gap_durations.index(max(gap_durations))
shortest_gap_idx = gap_durations.index(min(gap_durations))
build_meta["longest_gap_idx"] = longest_gap_idx
build_meta["shortest_gap_idx"] = shortest_gap_idx
build_meta["longest_gap_ms"] = gap_durations[longest_gap_idx]
build_meta["shortest_gap_ms"] = gap_durations[shortest_gap_idx]
return result, categories, source_files, events_meta, build_meta
# -------------------------------------------------- dataset generation
def generate_dataset(self) -> tuple:
"""Generate the complete dataset for this task."""
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} {self.TASK_NAME} samples "
f"(target: {self.task_duration_hours}h)..."
)
# Balanced question type distribution
question_types = list(self.task_config.get("question_types", []))
if question_types:
balanced_qtypes = []
per_type = num_samples // len(question_types)
remainder = num_samples % len(question_types)
for qt in question_types:
count = per_type + (1 if remainder > 0 else 0)
balanced_qtypes.extend([qt] * count)
remainder = max(0, remainder - 1)
random.shuffle(balanced_qtypes)
else:
balanced_qtypes = [None] * num_samples
all_metadata = []
for i, duration in enumerate(sample_durations):
qtype = balanced_qtypes[i] if i < len(balanced_qtypes) else None
metadata = self.generate_sample(
i, target_duration_seconds=duration, question_type=qtype
)
if metadata is not None:
all_metadata.append(metadata)
# ── Compute reasoning traces ──
# Load trace configurations and templates
import json
import generate_reasoning_traces
trace_templates_path = Path(__file__).parent.parent / "trace_templates.json"
trace_templates = {}
if trace_templates_path.exists():
with open(trace_templates_path) as f:
trace_templates = json.load(f)
config_templates = generate_reasoning_traces.load_config_templates(
str(Path(__file__).parent.parent / "config.yaml")
)
# Load LLM if enabled
tokenizer, model, device = None, None, None
llm_enabled = self.config.get("llm", {}).get("enabled", False)
if llm_enabled:
# Lazy load model, maybe we only load once globally
# But here we just load if enabled
if not hasattr(MultihopBaseGenerator, "_shared_model"):
self.logger.info("Loading LLM for verbal traces...")
MultihopBaseGenerator._shared_tokenizer = generate_reasoning_traces.AutoTokenizer.from_pretrained(
"meta-llama/Llama-3.1-8B-Instruct", use_fast=False
)
MultihopBaseGenerator._shared_model = generate_reasoning_traces.AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3.1-8B-Instruct",
torch_dtype="auto",
device_map="auto",
)
MultihopBaseGenerator._shared_model.eval()
MultihopBaseGenerator._shared_device = next(MultihopBaseGenerator._shared_model.parameters()).device
tokenizer = MultihopBaseGenerator._shared_tokenizer
model = MultihopBaseGenerator._shared_model
device = MultihopBaseGenerator._shared_device
# Process each generated sample to append traces
for meta in all_metadata:
question = str(meta.get("open_text_question", ""))
answer = str(meta.get("open_text_answer", ""))
qtype = str(meta.get("question_type", ""))
categories = meta.get("categories", [])
sym_trace = generate_reasoning_traces.compute_symbolic_trace(
self.TASK_NAME, qtype, question, answer,
categories, trace_templates, config_templates
)
meta["symbolic_trace"] = json.dumps(sym_trace)
if llm_enabled and model:
clean_question = question.replace("_", " ")
clean_answer = answer.replace("_", " ")
clean_categories = [str(c).replace("_", " ") for c in categories]
verbal = generate_reasoning_traces.verbalize_trace(
tokenizer, model, device,
clean_question, clean_answer, sym_trace, clean_categories
)
meta["verbal_trace"] = verbal
else:
meta["verbal_trace"] = "[DRY RUN] " + " ".join(sym_trace)
self.logger.info(
f"Generated {len(all_metadata)}/{num_samples} samples successfully with reasoning traces"
)
# Save CSVs
mcq_path = self.output_base / f"{self.TASK_NAME}_mcq.csv"
self._save_mcq_csv(all_metadata, mcq_path)
open_path = self.output_base / f"{self.TASK_NAME}_open_text.csv"
self._save_open_text_csv(all_metadata, open_path)
meta_path = self.output_base / f"{self.TASK_NAME}_metadata.csv"
self._save_metadata_csv(all_metadata, meta_path)
self.logger.info(f"{self.TASK_NAME} task complete!")
self.logger.info(f" - MCQ CSV: {mcq_path}")
self.logger.info(f" - Open-text CSV: {open_path}")
self.logger.info(f" - Metadata CSV: {meta_path}")
return mcq_path, open_path
# Subclasses must implement this:
def generate_sample(
self,
sample_id: int,
target_duration_seconds: float = None,
question_type: str = None,
) -> Optional[Dict]:
raise NotImplementedError
# --------------------------------------------------------- CSV helpers
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",
"symbolic_trace", "verbal_trace"
]
)
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.get("source_files", [])),
str(meta.get("categories", [])),
meta.get("symbolic_trace", ""),
meta.get("verbal_trace", "")
]
)
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",
"symbolic_trace", "verbal_trace"
]
)
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.get("source_files", [])),
str(meta.get("categories", [])),
meta.get("symbolic_trace", ""),
meta.get("verbal_trace", "")
]
)
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", "n_events",
"categories", "source_files",
"question_type",
"target_duration_s", "actual_duration_s",
"symbolic_trace", "verbal_trace"
]
)
for meta in metadata_list:
writer.writerow(
[
meta["id"],
meta["audio_path"],
meta.get("n_events", ""),
str(meta.get("categories", [])),
str(meta.get("source_files", [])),
meta["question_type"],
meta.get("target_duration_s", ""),
meta.get("actual_duration_s", ""),
meta.get("symbolic_trace", ""),
meta.get("verbal_trace", "")
]
)