transformer / prepare_parcels.py
Mohith202's picture
Deploy BrainRL OpenEnv server
32d14f4 verified
Raw
History Blame Contribute Delete
21.8 kB
#!/usr/bin/env python
"""Parcel pruning orchestrator for BrainRL.
This is the *only* place atlas parcels are loaded, scored, and pruned. It runs
once before training and writes a frozen candidate manifest to
``configs/parcel_candidates.json``. The OpenEnv environment then consumes that
manifest and never touches NIfTI files at episode time.
Pipeline stages:
1. Load atlas parcels (real Schaefer-200 via nilearn, or a deterministic
synthetic fallback so the demo always runs).
2. Intersect parcels with the brain mask (voxel counts).
3. Compute cheap parcel scores:
- voxel_count_quality (sigmoidal on n_voxels)
- variance_score (deterministic per-parcel proxy)
- semantic_prior (label/network heuristic)
- cached_encoding_score (optional CSV from prior fits)
4. Apply ``min_voxels_per_parcel`` floor.
5. Rank by a weighted ``prune_score`` and keep the top ``max_candidates``.
6. Save a JSON manifest used by the OpenEnv environment.
Example:
python prepare_parcels.py \
--atlas schaefer200 \
--max-candidates 200 \
--min-voxels 20 \
--prompt-top-k 30 \
--selection-budget 20 \
--output configs/parcel_candidates.json
"""
from __future__ import annotations
import argparse
import csv
import json
import math
import random
import re
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
PROJECT_ROOT = Path(__file__).resolve().parent
DEFAULT_OUTPUT = PROJECT_ROOT / "configs" / "parcel_candidates.json"
# Network membership for Schaefer-200 7Networks. Each entry maps a network name
# to (LH parcel count, RH parcel count). Counts approximate the canonical
# Schaefer 2018 split and total exactly 200 across both hemispheres.
SCHAEFER200_7NET_COUNTS: dict[str, tuple[int, int]] = {
"Vis": (14, 14),
"SomMot": (16, 16),
"DorsAttn": (12, 11),
"SalVentAttn": (11, 11),
"Limbic": (4, 5),
"Cont": (12, 14),
"Default": (31, 29),
}
# Sub-region tokens used to make synthetic labels look anatomically plausible
# and to give the semantic-prior heuristic something to grip.
NETWORK_SUBREGIONS: dict[str, list[str]] = {
"Vis": ["Striate", "ExStr", "ExStrSup", "ExStrInf"],
"SomMot": ["S1", "M1", "Aud", "ParOper"],
"DorsAttn": ["Post", "FEF", "PrCv", "PostC"],
"SalVentAttn": ["ParOper", "Med", "FrOperIns", "PFCl"],
"Limbic": ["TempPole", "OFC"],
"Cont": ["Par", "Temp", "PFCl", "PFCmp", "Cing", "pCun"],
"Default": ["Temp", "TempPar", "PFCv", "PFCm", "PFCdPFCm", "pCunPCC", "IPL", "PHC"],
}
# ---------------------------------------------------------------------------
# Atlas loading
# ---------------------------------------------------------------------------
@dataclass
class RawParcel:
"""Pre-pruning parcel record."""
label_index: int
label: str
hemisphere: str
network: str
sub_region: str
n_voxels: int
notes: str = ""
def base_region_id(self) -> str:
return f"parcel_{self.label_index:03d}"
def _network_for_index(idx_within_hemi: int, counts: dict[str, tuple[int, int]], hemi_idx: int) -> tuple[str, int]:
"""Map a 1-based parcel index inside a hemisphere to (network, slot)."""
cursor = 0
for network, (lh, rh) in counts.items():
n = lh if hemi_idx == 0 else rh
if idx_within_hemi <= cursor + n:
return network, idx_within_hemi - cursor
cursor += n
return "Default", idx_within_hemi
def load_synthetic_schaefer200(seed: int = 42) -> list[RawParcel]:
"""Build a deterministic Schaefer-200-style parcel set without downloads."""
rng = random.Random(seed)
parcels: list[RawParcel] = []
label_index = 0
for hemi_idx, hemi in enumerate(("LH", "RH")):
per_hemi = sum(lh if hemi_idx == 0 else rh for lh, rh in SCHAEFER200_7NET_COUNTS.values())
for i in range(1, per_hemi + 1):
label_index += 1
network, slot = _network_for_index(i, SCHAEFER200_7NET_COUNTS, hemi_idx)
sub_options = NETWORK_SUBREGIONS[network]
sub_region = sub_options[(slot - 1) % len(sub_options)]
sub_idx = ((slot - 1) // len(sub_options)) + 1
label = f"7Networks_{hemi}_{network}_{sub_region}_{sub_idx}"
n_voxels = max(8, int(rng.gauss(330, 120)))
parcels.append(
RawParcel(
label_index=label_index,
label=label,
hemisphere="left" if hemi == "LH" else "right",
network=network,
sub_region=sub_region,
n_voxels=n_voxels,
)
)
return parcels
def load_nilearn_schaefer(n_rois: int = 200, yeo_networks: int = 7) -> list[RawParcel]:
"""Fetch the real Schaefer-2018 atlas via nilearn and count voxels."""
try:
import numpy as np
from nilearn import datasets, image
except ImportError as exc: # pragma: no cover
raise SystemExit(
"Real-atlas mode requires nilearn. Install with `pip install -e .[atlas]` "
"or rerun with --atlas-source synthetic."
) from exc
atlas = datasets.fetch_atlas_schaefer_2018(
n_rois=n_rois,
yeo_networks=yeo_networks,
resolution_mm=2,
)
labels_img = image.load_img(atlas.maps)
label_data = np.asarray(labels_img.dataobj).astype(int)
raw_labels: list[str] = []
for entry in atlas.labels:
if isinstance(entry, bytes):
raw_labels.append(entry.decode("utf-8"))
else:
raw_labels.append(str(entry))
parcels: list[RawParcel] = []
for label_index in range(1, n_rois + 1):
if label_index - 1 >= len(raw_labels):
break
label = raw_labels[label_index - 1]
n_voxels = int((label_data == label_index).sum())
hemi_match = re.search(r"_(LH|RH)_", label)
hemisphere = "left" if hemi_match and hemi_match.group(1) == "LH" else "right"
network_match = re.search(
r"_(Vis|SomMot|DorsAttn|SalVentAttn|Limbic|Cont|Default)_", label
)
network = network_match.group(1) if network_match else "Default"
sub_match = re.search(rf"_{network}_([A-Za-z]+)", label)
sub_region = sub_match.group(1) if sub_match else "Unknown"
parcels.append(
RawParcel(
label_index=label_index,
label=label,
hemisphere=hemisphere,
network=network,
sub_region=sub_region,
n_voxels=n_voxels,
notes="schaefer2018_7networks",
)
)
return parcels
# ---------------------------------------------------------------------------
# Scoring
# ---------------------------------------------------------------------------
# Higher = more relevant to auditory/language fMRI prediction. These priors are
# intentionally small, transparent heuristics; the RL agent is supposed to
# discover better orderings, not be handed the answer.
NETWORK_LANGUAGE_PRIORS: dict[str, float] = {
"Default": 0.78,
"Cont": 0.66,
"SalVentAttn": 0.62,
"Limbic": 0.55,
"DorsAttn": 0.40,
"SomMot": 0.45,
"Vis": 0.20,
}
LANGUAGE_KEYWORDS: dict[str, float] = {
"Aud": 0.30,
"Temp": 0.20,
"TempPar": 0.22,
"TempPole": 0.18,
"IPL": 0.12,
"PFCl": 0.10,
"FrOperIns": 0.14,
"ParOper": 0.10,
"PFCv": 0.08,
}
def semantic_prior_for(parcel: RawParcel) -> float:
base = NETWORK_LANGUAGE_PRIORS.get(parcel.network, 0.4)
label = parcel.label
bonus = 0.0
for keyword, value in LANGUAGE_KEYWORDS.items():
if keyword in label:
bonus = max(bonus, value)
if parcel.hemisphere == "left":
bonus += 0.05
return float(min(0.99, base + bonus))
def voxel_count_quality(n_voxels: int) -> float:
"""Soft sigmoid: ~0 at 10 voxels, ~0.5 at ~150, ~1 at very large parcels."""
return float(1.0 / (1.0 + math.exp(-(n_voxels - 150) / 80.0)))
def deterministic_variance_score(parcel: RawParcel, seed: int) -> float:
"""Stable per-parcel pseudo-variance proxy (no NIfTI required)."""
rng = random.Random(f"{seed}:{parcel.label}")
base = 0.45 if parcel.network in {"Default", "Cont", "SalVentAttn"} else 0.30
return float(max(0.0, min(1.0, base + rng.uniform(-0.10, 0.20))))
def base_r2_estimate(parcel: RawParcel, semantic: float, variance: float) -> float:
"""Cheap predicted-encoding score used as a hint, not a target."""
raw = 0.020 + 0.060 * semantic + 0.020 * variance
return float(max(0.005, min(0.20, raw)))
# ---------------------------------------------------------------------------
# Cached encoding scores (optional)
# ---------------------------------------------------------------------------
def load_cached_scores(paths: list[str]) -> dict[str, float]:
"""Read optional encoding scores keyed by region label or parcel id."""
region_keys = ("roi_name", "region_id", "target", "parcel", "label")
score_keys = ("mean_r2", "mean_corr", "score", "r2")
scores: dict[str, float] = {}
for raw_path in paths:
path = Path(raw_path).expanduser()
if not path.exists() or path.suffix.lower() != ".csv":
continue
with path.open("r", encoding="utf-8", newline="") as handle:
for row in csv.DictReader(handle):
key = None
for region_key in region_keys:
candidate = row.get(region_key)
if candidate:
key = str(candidate).strip()
break
if not key:
continue
value: float | None = None
for score_key in score_keys:
raw = row.get(score_key)
if raw in (None, ""):
continue
try:
parsed = float(raw)
except ValueError:
continue
if score_key == "mean_corr":
parsed = max(0.0, parsed) ** 2
value = max(0.0, min(0.25, parsed))
break
if value is not None:
scores[key] = value
return scores
# ---------------------------------------------------------------------------
# Pruning
# ---------------------------------------------------------------------------
@dataclass
class PruneWeights:
activity: float = 0.40
semantic: float = 0.30
encoding: float = 0.20
voxel_quality: float = 0.10
@dataclass
class ScoredParcel:
parcel: RawParcel
semantic_prior: float
variance_score: float
voxel_quality: float
cached_encoding: float
base_r2: float
prune_score: float
cost: float
redundancy_group: str
network_priority_kept: bool = False
extras: dict[str, Any] = field(default_factory=dict)
def redundancy_group_for(parcel: RawParcel) -> str:
label = parcel.label
if "Aud" in label or "Temp" in label or "TempPar" in label or "TempPole" in label:
return "auditory_temporal"
if "PFCl" in label or "FrOperIns" in label or "PFCv" in label:
return "inferior_frontal"
if "Vis" in label:
return "visual"
if "SomMot" in label or "S1" in label or "M1" in label:
return "somatomotor"
return "association"
def cost_for(parcel: RawParcel) -> float:
if parcel.network in {"Vis", "SomMot"}:
return 1.05
if parcel.network == "Limbic":
return 1.15
return 1.0
def score_parcel(
parcel: RawParcel,
*,
weights: PruneWeights,
cached_scores: dict[str, float],
seed: int,
) -> ScoredParcel:
semantic = semantic_prior_for(parcel)
variance = deterministic_variance_score(parcel, seed)
voxq = voxel_count_quality(parcel.n_voxels)
encoding_score = 0.0
for key in (parcel.label, parcel.base_region_id(), parcel.sub_region):
if key in cached_scores:
encoding_score = cached_scores[key]
break
base_r2 = base_r2_estimate(parcel, semantic, variance)
if encoding_score > 0:
base_r2 = float(max(base_r2, encoding_score))
prune_score = (
weights.activity * variance
+ weights.semantic * semantic
+ weights.encoding * (encoding_score / 0.20 if encoding_score > 0 else semantic * 0.5)
+ weights.voxel_quality * voxq
)
return ScoredParcel(
parcel=parcel,
semantic_prior=semantic,
variance_score=variance,
voxel_quality=voxq,
cached_encoding=encoding_score,
base_r2=base_r2,
prune_score=float(prune_score),
cost=cost_for(parcel),
redundancy_group=redundancy_group_for(parcel),
)
def prune_parcels(
parcels: list[RawParcel],
*,
max_candidates: int,
min_voxels: int,
weights: PruneWeights,
cached_scores: dict[str, float],
seed: int,
keep_language_floor: int = 30,
) -> list[ScoredParcel]:
scored = [
score_parcel(p, weights=weights, cached_scores=cached_scores, seed=seed)
for p in parcels
]
eligible = [s for s in scored if s.parcel.n_voxels >= min_voxels]
if not eligible:
eligible = scored
eligible.sort(key=lambda s: s.prune_score, reverse=True)
# Reserve a floor for clearly language-relevant parcels even if their
# prune score is borderline. Keeps the easy-curriculum signal alive.
language_pool = sorted(
(s for s in eligible if s.redundancy_group in {"auditory_temporal", "inferior_frontal"}),
key=lambda s: s.prune_score,
reverse=True,
)
floor = min(keep_language_floor, len(language_pool), max_candidates)
forced = {id(s) for s in language_pool[:floor]}
kept: list[ScoredParcel] = []
for s in eligible:
if id(s) in forced:
s.network_priority_kept = True
kept.append(s)
if len(kept) >= max_candidates:
break
if len(kept) < max_candidates:
for s in eligible:
if id(s) in forced:
continue
kept.append(s)
if len(kept) >= max_candidates:
break
kept.sort(key=lambda s: s.prune_score, reverse=True)
return kept
# ---------------------------------------------------------------------------
# Manifest writing
# ---------------------------------------------------------------------------
def manifest_payload(
*,
atlas: str,
selection_budget: int,
prompt_top_k: int,
cost_penalty: float,
weights: PruneWeights,
pruned: list[ScoredParcel],
source_notes: str,
) -> dict[str, Any]:
candidates = []
for idx, scored in enumerate(pruned):
parcel = scored.parcel
candidates.append(
{
"region_id": parcel.base_region_id(),
"atlas": atlas,
"label": parcel.label,
"hemisphere": parcel.hemisphere,
"network": parcel.network,
"sub_region": parcel.sub_region,
"n_voxels": int(parcel.n_voxels),
"semantic_prior": float(round(scored.semantic_prior, 4)),
"variance_score": float(round(scored.variance_score, 4)),
"voxel_quality": float(round(scored.voxel_quality, 4)),
"cached_encoding": float(round(scored.cached_encoding, 4)),
"base_r2": float(round(scored.base_r2, 4)),
"cost": float(round(scored.cost, 4)),
"prune_score": float(round(scored.prune_score, 4)),
"redundancy_group": scored.redundancy_group,
"rank": idx,
"kept_by_floor": bool(scored.network_priority_kept),
}
)
return {
"atlas": atlas,
"candidate_mode": "atlas_parcels",
"max_candidates": len(candidates),
"selection_budget": int(selection_budget),
"prompt_top_k": int(prompt_top_k),
"cost_penalty": float(cost_penalty),
"prune_weights": {
"activity": weights.activity,
"semantic": weights.semantic,
"encoding": weights.encoding,
"voxel_quality": weights.voxel_quality,
},
"source_notes": source_notes,
"candidates": candidates,
}
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="Load an atlas, prune parcels, and freeze the candidate manifest."
)
parser.add_argument(
"--atlas",
type=str,
default="schaefer200",
help="Atlas identifier; only schaefer200 (7 networks) is implemented.",
)
parser.add_argument(
"--atlas-source",
choices=("auto", "nilearn", "synthetic"),
default="auto",
help="Where to load parcels from. 'auto' tries nilearn then falls back.",
)
parser.add_argument("--max-candidates", type=int, default=200)
parser.add_argument("--min-voxels", type=int, default=20)
parser.add_argument("--prompt-top-k", type=int, default=30)
parser.add_argument("--selection-budget", type=int, default=20)
parser.add_argument("--cost-penalty", type=float, default=0.002)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument(
"--output",
type=str,
default=str(DEFAULT_OUTPUT),
help="Where to write the candidate manifest JSON.",
)
parser.add_argument(
"--cached-scores",
action="append",
default=None,
help="Optional CSV with cached encoding scores; can be repeated.",
)
parser.add_argument("--w-activity", type=float, default=0.40)
parser.add_argument("--w-semantic", type=float, default=0.30)
parser.add_argument("--w-encoding", type=float, default=0.20)
parser.add_argument("--w-voxel-quality", type=float, default=0.10)
parser.add_argument(
"--keep-language-floor",
type=int,
default=30,
help="Force-keep this many top auditory/language parcels.",
)
parser.add_argument(
"--quiet",
action="store_true",
help="Suppress per-stage logs.",
)
return parser
def _log(quiet: bool, message: str) -> None:
if not quiet:
print(message)
def main() -> None:
args = build_parser().parse_args()
if args.atlas != "schaefer200":
raise SystemExit(
f"Unsupported atlas={args.atlas}. Only 'schaefer200' is implemented."
)
weights = PruneWeights(
activity=float(args.w_activity),
semantic=float(args.w_semantic),
encoding=float(args.w_encoding),
voxel_quality=float(args.w_voxel_quality),
)
cached_paths = list(args.cached_scores or [])
cached_scores = load_cached_scores(cached_paths)
_log(args.quiet, f"[1/6] Loaded {len(cached_scores)} cached encoding score rows.")
parcels: list[RawParcel] = []
source_notes = ""
if args.atlas_source in ("auto", "nilearn"):
try:
parcels = load_nilearn_schaefer()
source_notes = "nilearn:schaefer2018_7networks_n200_2mm"
_log(args.quiet, f"[2/6] Loaded {len(parcels)} parcels from nilearn Schaefer-2018.")
except SystemExit:
if args.atlas_source == "nilearn":
raise
parcels = []
except Exception as exc: # noqa: BLE001
if args.atlas_source == "nilearn":
raise SystemExit(f"nilearn atlas load failed: {exc}") from exc
_log(args.quiet, f"[2/6] nilearn atlas unavailable ({exc}); falling back to synthetic.")
parcels = []
if not parcels:
parcels = load_synthetic_schaefer200(seed=int(args.seed))
source_notes = "synthetic:schaefer200_7networks"
_log(args.quiet, f"[2/6] Loaded {len(parcels)} synthetic Schaefer-200-style parcels.")
_log(args.quiet, f"[3/6] Computed voxel-count metadata for {len(parcels)} parcels.")
_log(args.quiet, "[4/6] Scoring parcels with weighted prune formula.")
pruned = prune_parcels(
parcels,
max_candidates=int(args.max_candidates),
min_voxels=int(args.min_voxels),
weights=weights,
cached_scores=cached_scores,
seed=int(args.seed),
keep_language_floor=int(args.keep_language_floor),
)
_log(args.quiet, f"[5/6] Pruned to {len(pruned)} candidates (min_voxels>={args.min_voxels}).")
payload = manifest_payload(
atlas=args.atlas,
selection_budget=int(args.selection_budget),
prompt_top_k=int(args.prompt_top_k),
cost_penalty=float(args.cost_penalty),
weights=weights,
pruned=pruned,
source_notes=source_notes,
)
output_path = Path(args.output).resolve()
output_path.parent.mkdir(parents=True, exist_ok=True)
with output_path.open("w", encoding="utf-8") as handle:
json.dump(payload, handle, indent=2)
handle.write("\n")
if not args.quiet:
top_preview = ", ".join(
f"{c['region_id']}({c['network']})"
for c in payload["candidates"][:5]
)
print(f"[6/6] Wrote {len(payload['candidates'])} parcel candidates -> {output_path}")
print(f" atlas={args.atlas} budget={args.selection_budget} "
f"prompt_top_k={args.prompt_top_k} source={source_notes}")
print(f" top: {top_preview}")
if __name__ == "__main__":
main()