""" 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", "") ] )