csai / code /a1_pipeline /annotations.py
Mohith202's picture
Add core ROI masks, visualization script, and model profiles
8c5a642
Raw
History Blame Contribute Delete
9.51 kB
"""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