| """Annotation harmonization for A1 baseline event tables.""" |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import asdict, dataclass |
| from pathlib import Path |
| from typing import Any |
|
|
| import numpy as np |
| import pandas as pd |
|
|
| from .constants import ( |
| DEFAULT_MIXED_RUN_WORD_TABLE_FALLBACK_MAP, |
| DEFAULT_RUN_WORD_TABLE_MAP, |
| ) |
|
|
|
|
| @dataclass(frozen=True) |
| class AnnotationSourceResolution: |
| """Resolved source table for one run.""" |
|
|
| run: int |
| source_word_table: str |
| annotation_used_mixed_fallback: bool |
| provenance: str |
|
|
|
|
| REQUIRED_WORD_TABLE_COLUMNS: tuple[str, ...] = ("word", "onset", "offset") |
|
|
|
|
| def _validate_required_columns(df: pd.DataFrame, table_name: str) -> None: |
| missing = [column for column in REQUIRED_WORD_TABLE_COLUMNS if column not in df.columns] |
| if missing: |
| raise ValueError(f"Annotation table {table_name} missing columns: {missing}") |
|
|
|
|
| def _load_and_normalize_word_table( |
| annotation_dir: Path, |
| table_name: str, |
| time_scale_seconds: float, |
| ) -> pd.DataFrame: |
| table_path = annotation_dir / table_name |
| if not table_path.exists(): |
| raise FileNotFoundError(f"Annotation table not found: {table_path}") |
|
|
| df = pd.read_csv(table_path) |
| _validate_required_columns(df, table_name=table_name) |
|
|
| out = df.copy() |
| out["word"] = out["word"].astype(str).fillna("").str.strip() |
| out = out[out["word"] != ""].reset_index(drop=True) |
|
|
| out["onset"] = pd.to_numeric(out["onset"], errors="coerce") |
| out["offset"] = pd.to_numeric(out["offset"], errors="coerce") |
|
|
| if out["onset"].isna().any() or out["offset"].isna().any(): |
| raise ValueError(f"Annotation table {table_name} contains non-numeric onset/offset values") |
|
|
| out["onset_s"] = out["onset"].astype(float) * float(time_scale_seconds) |
| out["offset_s"] = out["offset"].astype(float) * float(time_scale_seconds) |
| out["duration_s"] = out["offset_s"] - out["onset_s"] |
| out["word_index"] = np.arange(len(out), dtype=np.int64) |
| out["source_word_table"] = table_name |
|
|
| return out[["word_index", "word", "onset_s", "offset_s", "duration_s", "source_word_table"]] |
|
|
|
|
| def _validate_timing_table(df: pd.DataFrame, label: str) -> list[str]: |
| errors: list[str] = [] |
|
|
| if (df["onset_s"] < 0).any(): |
| errors.append(f"{label}: negative onset_s found") |
|
|
| if (df["offset_s"] < 0).any(): |
| errors.append(f"{label}: negative offset_s found") |
|
|
| if (df["duration_s"] < 0).any(): |
| errors.append(f"{label}: negative duration_s found") |
|
|
| if (df["offset_s"] < df["onset_s"]).any(): |
| errors.append(f"{label}: offset_s earlier than onset_s") |
|
|
| onset_diff = df["onset_s"].diff().dropna() |
| if (onset_diff < 0).any(): |
| errors.append(f"{label}: onset_s is not monotonic non-decreasing") |
|
|
| offset_diff = df["offset_s"].diff().dropna() |
| if (offset_diff < 0).any(): |
| errors.append(f"{label}: offset_s is not monotonic non-decreasing") |
|
|
| return errors |
|
|
|
|
| def _resolve_run_source( |
| run: int, |
| enable_mixed_fallback: bool, |
| run_word_table_map: dict[int, str | None], |
| mixed_fallback_map: dict[int, str], |
| ) -> AnnotationSourceResolution: |
| source_table = run_word_table_map.get(run) |
| if source_table is not None: |
| return AnnotationSourceResolution( |
| run=run, |
| source_word_table=source_table, |
| annotation_used_mixed_fallback=False, |
| provenance=f"fixed_run_source:{source_table}", |
| ) |
|
|
| if enable_mixed_fallback and run in mixed_fallback_map: |
| fallback_table = mixed_fallback_map[run] |
| return AnnotationSourceResolution( |
| run=run, |
| source_word_table=fallback_table, |
| annotation_used_mixed_fallback=True, |
| provenance=f"mixed_run_fallback_source:{fallback_table}", |
| ) |
|
|
| raise ValueError( |
| "No annotation source available for run " |
| f"{run}. Provide a direct run mapping or enable mixed fallback." |
| ) |
|
|
|
|
| def build_harmonized_annotation_tables( |
| manifest_df: pd.DataFrame, |
| annotation_dir: Path, |
| time_scale_seconds: float, |
| enable_mixed_fallback: bool, |
| run_word_table_map: dict[int, str | None] | None = None, |
| mixed_fallback_map: dict[int, str] | None = None, |
| ) -> tuple[pd.DataFrame, pd.DataFrame, dict[str, Any]]: |
| """Build run-level templates and subject-run unified annotation tables.""" |
| if manifest_df.empty: |
| raise ValueError("Manifest is empty; cannot harmonize annotations") |
|
|
| required_columns = { |
| "subject", |
| "run", |
| "condition_fixed", |
| "condition_effective", |
| "speaker_stream", |
| "used_mixed_fallback", |
| } |
| missing = required_columns.difference(manifest_df.columns) |
| if missing: |
| raise ValueError(f"Manifest missing required columns for annotation harmonization: {sorted(missing)}") |
|
|
| annotation_dir = annotation_dir.resolve() |
| run_word_table_map = run_word_table_map or DEFAULT_RUN_WORD_TABLE_MAP |
| mixed_fallback_map = mixed_fallback_map or DEFAULT_MIXED_RUN_WORD_TABLE_FALLBACK_MAP |
|
|
| source_cache: dict[str, pd.DataFrame] = {} |
| run_template_rows: list[pd.DataFrame] = [] |
| source_resolutions: list[AnnotationSourceResolution] = [] |
|
|
| runs = sorted({int(value) for value in manifest_df["run"].tolist()}) |
|
|
| for run in runs: |
| resolution = _resolve_run_source( |
| run=run, |
| enable_mixed_fallback=enable_mixed_fallback, |
| run_word_table_map=run_word_table_map, |
| mixed_fallback_map=mixed_fallback_map, |
| ) |
| source_resolutions.append(resolution) |
|
|
| if resolution.source_word_table not in source_cache: |
| source_cache[resolution.source_word_table] = _load_and_normalize_word_table( |
| annotation_dir=annotation_dir, |
| table_name=resolution.source_word_table, |
| time_scale_seconds=time_scale_seconds, |
| ) |
|
|
| template_df = source_cache[resolution.source_word_table].copy() |
| template_df["run"] = int(run) |
| template_df["annotation_used_mixed_fallback"] = bool(resolution.annotation_used_mixed_fallback) |
| template_df["provenance"] = resolution.provenance |
|
|
| run_template_rows.append(template_df) |
|
|
| run_events_df = pd.concat(run_template_rows, ignore_index=True) |
|
|
| template_errors: list[str] = [] |
| for run, run_df in run_events_df.groupby("run"): |
| template_errors.extend(_validate_timing_table(run_df, label=f"run_template_run{run}")) |
|
|
| if template_errors: |
| raise ValueError("Annotation template validation failed: " + "; ".join(template_errors)) |
|
|
| merged_rows: list[pd.DataFrame] = [] |
| for row in manifest_df.itertuples(index=False): |
| subject = str(getattr(row, "subject")) |
| run = int(getattr(row, "run")) |
|
|
| run_template = run_events_df[run_events_df["run"] == run].copy() |
| run_template["subject"] = subject |
| run_template["condition_fixed"] = str(getattr(row, "condition_fixed")) |
| run_template["condition_effective"] = str(getattr(row, "condition_effective")) |
| run_template["speaker_stream"] = str(getattr(row, "speaker_stream")) |
| run_template["used_mixed_fallback"] = bool(getattr(row, "used_mixed_fallback")) |
|
|
| merged_rows.append(run_template) |
|
|
| unified_df = pd.concat(merged_rows, ignore_index=True) |
| unified_df = unified_df[ |
| [ |
| "subject", |
| "run", |
| "word_index", |
| "word", |
| "onset_s", |
| "offset_s", |
| "duration_s", |
| "condition_fixed", |
| "condition_effective", |
| "speaker_stream", |
| "used_mixed_fallback", |
| "annotation_used_mixed_fallback", |
| "source_word_table", |
| "provenance", |
| ] |
| ] |
|
|
| unified_errors: list[str] = [] |
| for (subject, run), group_df in unified_df.groupby(["subject", "run"]): |
| unified_errors.extend(_validate_timing_table(group_df, label=f"unified_{subject}_run{run}")) |
|
|
| if unified_errors: |
| raise ValueError("Unified annotation validation failed: " + "; ".join(unified_errors)) |
|
|
| source_resolution_df = pd.DataFrame([asdict(value) for value in source_resolutions]) |
| source_resolution_df = source_resolution_df.sort_values(["run"]).reset_index(drop=True) |
|
|
| run_word_counts = ( |
| run_events_df.groupby("run")["word_index"].max().add(1).astype(int).to_dict() |
| if not run_events_df.empty |
| else {} |
| ) |
|
|
| annotation_qc: dict[str, Any] = { |
| "annotation_dir": str(annotation_dir), |
| "time_scale_seconds": float(time_scale_seconds), |
| "enable_mixed_fallback": bool(enable_mixed_fallback), |
| "n_manifest_rows": int(len(manifest_df)), |
| "n_run_templates": int(run_events_df["run"].nunique()) if not run_events_df.empty else 0, |
| "n_unified_rows": int(len(unified_df)), |
| "run_word_counts": {str(run): int(count) for run, count in sorted(run_word_counts.items())}, |
| "source_resolution": [asdict(value) for value in source_resolutions], |
| } |
|
|
| if not unified_df.empty: |
| annotation_qc["duration_s_min"] = float(unified_df["duration_s"].min()) |
| annotation_qc["duration_s_max"] = float(unified_df["duration_s"].max()) |
|
|
| run_events_df = run_events_df.sort_values(["run", "word_index"]).reset_index(drop=True) |
| unified_df = unified_df.sort_values(["subject", "run", "word_index"]).reset_index(drop=True) |
|
|
| return run_events_df, unified_df, source_resolution_df, annotation_qc |
|
|