transformer / server /brain_environment.py
Mohith202's picture
Deploy BrainRL OpenEnv server
82614b0 verified
Raw
History Blame Contribute Delete
20.2 kB
"""BrainRL region selection environment implementation."""
from __future__ import annotations
import json
import math
import os
import random
from typing import Any
from uuid import uuid4
try:
from openenv.core.env_server.mcp_environment import MCPEnvironment
from openenv.core.env_server.types import Action, Observation, State
except ImportError: # pragma: no cover - keeps local smoke tests lightweight
class MCPEnvironment: # type: ignore[no-redef]
def __init__(self, mcp: Any | None = None):
self.mcp = mcp
def step(self, action: Any, timeout_s: float | None = None, **kwargs: Any) -> Any:
return self._step_impl(action, timeout_s=timeout_s, **kwargs)
async def step_async(self, action: Any, timeout_s: float | None = None, **kwargs: Any) -> Any:
return self._step_impl(action, timeout_s=timeout_s, **kwargs)
class Action: # type: ignore[no-redef]
pass
class Observation: # type: ignore[no-redef]
def __init__(self, done: bool = False, reward: float = 0.0, metadata: dict | None = None):
self.done = done
self.reward = reward
self.metadata = metadata or {}
class State: # type: ignore[no-redef]
def __init__(self, episode_id: str, step_count: int = 0):
self.episode_id = episode_id
self.step_count = step_count
try:
from fastmcp import FastMCP
except ImportError: # pragma: no cover
class FastMCP: # type: ignore[no-redef]
def __init__(self, name: str):
self.name = name
def tool(self, func):
return func
from data_loader import (
BrainSubset,
DEFAULT_CONFIG_PATH,
RegionCandidate,
_parse_simple_yaml,
candidate_table,
load_brain_subset,
)
from rewards import format_reward_breakdown, reward_columns, score_region_action
from stimulus_loader import (
compact_for_prompt as compact_stimulus_for_prompt,
n_windows as stimulus_n_windows,
stimulus_bias_for_group,
summarize_window as summarize_stimulus_window,
)
def _env_int(name: str, default: int) -> int:
raw = os.environ.get(name)
if not raw:
return default
try:
return int(raw)
except (TypeError, ValueError):
return default
def _yaml_stimulus_defaults() -> tuple[int, int]:
"""Read stimulus_window_size / stimulus_top_words from subset_config.yaml."""
try:
config = _parse_simple_yaml(DEFAULT_CONFIG_PATH)
except Exception: # pragma: no cover - defensive: never fail env construction
config = {}
window = int(config.get("stimulus_window_size", 30) or 30)
top = int(config.get("stimulus_top_words", 10) or 10)
return max(1, window), max(1, top)
_yaml_window, _yaml_top = _yaml_stimulus_defaults()
_DEFAULT_STIMULUS_WINDOW_SIZE = _env_int("BRAINRL_STIMULUS_WINDOW_SIZE", _yaml_window)
_DEFAULT_STIMULUS_TOP_WORDS = _env_int("BRAINRL_STIMULUS_TOP_WORDS", _yaml_top)
TASKS = {
"roi_selection": {
"description": (
"Sequentially select brain regions that improve prediction of auditory "
"stimulus responses from a compact Le Petit Prince fMRI summary."
),
"difficulty": "medium",
"reward": "independent_verifier_components",
"valid_actions": "Any unselected candidate region_id",
}
}
class BrainRegionSelectionEnvironment(MCPEnvironment):
"""OpenEnv-style environment for active brain region acquisition."""
def __init__(self):
mcp = FastMCP("brain_rl_region_selection")
@mcp.tool
def get_selection_state() -> dict:
"""Return selected regions, current score, budget, and candidates."""
return self._build_selection_state()
@mcp.tool
def get_task_info() -> dict:
"""Return task metadata and reward definition."""
return self._build_task_info()
@mcp.tool
def take_action(region_id: str) -> dict:
"""Select the next region and advance one acquisition step."""
return self._process_action(region_id)
super().__init__(mcp)
self._state = State(episode_id=str(uuid4()), step_count=0)
self._rng = random.Random(42)
self._subset: BrainSubset = load_brain_subset()
self._task = "roi_selection"
self._episode_id = str(uuid4())
self._timestep = 0
self._selected_region_ids: list[str] = []
self._current_r2 = 0.0
self._last_feedback = "Environment initialized"
self._subject_id: str | None = None
self._run_id: str | None = None
self._condition: str | None = None
self._stimulus_window_size: int = _DEFAULT_STIMULUS_WINDOW_SIZE
self._stimulus_top_words: int = _DEFAULT_STIMULUS_TOP_WORDS
self._stimulus_window_index: int | None = None
self._stimulus_features: dict | None = None
self._effective_base_r2: dict[str, float] = self._compute_effective_base_r2()
def reset(
self,
seed: int | None = None,
episode_id: str | None = None,
task: str = "roi_selection",
subject_id: str | None = None,
run_id: str | None = None,
condition: str | None = None,
stimulus_window: int | None = None,
stimulus_window_size: int | None = None,
stimulus_top_words: int | None = None,
**_: Any,
) -> Observation:
if task not in TASKS:
raise ValueError(f"Unknown task={task}. Valid tasks: {sorted(TASKS)}")
if seed is not None:
self._rng = random.Random(seed)
self._subset = load_brain_subset()
self._task = task
self._episode_id = episode_id or str(uuid4())
self._state = State(episode_id=self._episode_id, step_count=0)
self._timestep = 0
self._selected_region_ids = []
self._current_r2 = 0.0
self._subject_id = subject_id
self._run_id = run_id
self._condition = condition
if stimulus_window_size is not None:
self._stimulus_window_size = max(1, int(stimulus_window_size))
if stimulus_top_words is not None:
self._stimulus_top_words = max(1, int(stimulus_top_words))
self._stimulus_window_index = (
int(stimulus_window) if stimulus_window is not None else None
)
self._stimulus_features = self._build_stimulus_features(seed=seed)
# ``_effective_base_r2`` depends on the active stimulus window so it
# has to be recomputed every reset(), not just on subject/condition
# changes.
self._effective_base_r2 = self._compute_effective_base_r2()
ctx = self._context_label()
self._last_feedback = f"Select the first brain region. ({ctx})" if ctx else "Select the first brain region."
return Observation(done=False, reward=0.0, metadata=self._build_observation())
def _build_stimulus_features(self, *, seed: int | None) -> dict | None:
"""Resolve the stimulus window for this episode (if data is available)."""
if not self._condition:
return None
# Deterministic key: episode varies stimulus across resets even when
# subject/run repeat, while staying reproducible for a given seed.
deterministic_key = (
self._subject_id or "_",
self._run_id or "_",
int(seed) if seed is not None else 0,
self._episode_id,
)
return summarize_stimulus_window(
self._condition,
window_index=self._stimulus_window_index,
window_size=self._stimulus_window_size,
top_words=self._stimulus_top_words,
deterministic_key=deterministic_key,
)
def _step_impl(
self,
action: Action,
timeout_s: float | None = None,
**kwargs: Any,
) -> Observation:
region_id = getattr(action, "region_id", None)
if region_id is None and isinstance(action, dict):
region_id = action.get("region_id")
result = self._process_action(str(region_id))
return Observation(
done=bool(result["done"]),
reward=float(result["reward"]),
metadata=self._build_observation(extra=result),
)
def step(self, action: Action, timeout_s: float | None = None, **kwargs: Any) -> Observation:
self._state.step_count += 1
return super().step(action, timeout_s=timeout_s, **kwargs)
async def step_async(
self,
action: Action,
timeout_s: float | None = None,
**kwargs: Any,
) -> Observation:
self._state.step_count += 1
return await super().step_async(action, timeout_s=timeout_s, **kwargs)
@property
def state(self) -> State:
return self._state
def _candidate_by_id(self) -> dict[str, RegionCandidate]:
return {candidate.region_id: candidate for candidate in self._subset.candidates}
def _selected_candidates(self) -> list[RegionCandidate]:
by_id = self._candidate_by_id()
return [by_id[region_id] for region_id in self._selected_region_ids if region_id in by_id]
def _context_label(self) -> str:
parts = [
f"subject={self._subject_id}" if self._subject_id else "",
f"run={self._run_id}" if self._run_id else "",
f"condition={self._condition}" if self._condition else "",
]
return ", ".join(p for p in parts if p)
def _condition_boost(self, candidate: RegionCandidate) -> float:
"""Per-condition multiplier so train/test conditions reward differently.
``single_m`` (single male narrator) emphasizes auditory/language ROIs;
``single_f`` does the same with a slight twist; ``mixed_*`` rewards a
broader set of association regions.
"""
condition = (self._condition or "").lower()
group = candidate.redundancy_group
if condition == "single_m":
return 1.20 if group == "auditory_temporal" else (1.05 if group == "inferior_frontal" else 0.95)
if condition == "single_f":
return 1.18 if group == "auditory_temporal" else (1.04 if group == "inferior_frontal" else 0.96)
if condition == "mixed_m":
return 1.10 if group in {"auditory_temporal", "inferior_frontal"} else 1.02
if condition == "mixed_f":
return 1.08 if group in {"auditory_temporal", "association"} else 1.0
return 1.0
def _subject_perturbation(self, candidate: RegionCandidate) -> float:
"""Deterministic per-subject jitter in [0.7, 1.3].
Hash-based so the same (subject, parcel) always gives the same value
but different subjects experience different reward landscapes - which
is what makes train/test generalization meaningful.
"""
if not self._subject_id:
return 1.0
h = abs(hash((self._subject_id, candidate.region_id))) % 10_000
return 0.7 + (h / 10_000.0) * 0.6
def _stimulus_bias(self, candidate: RegionCandidate) -> float:
"""Per-window stimulus multiplier for this candidate's group.
Bounded to a small range so it nudges the policy toward parcels
that match the current stimulus content (e.g. nouns/density →
auditory_temporal, function/syntactic words → inferior_frontal)
without overwhelming the underlying base_r2 ranking.
"""
return stimulus_bias_for_group(candidate.redundancy_group, self._stimulus_features)
def _compute_effective_base_r2(self) -> dict[str, float]:
effective: dict[str, float] = {}
for candidate in self._subset.candidates:
value = (
candidate.base_r2
* self._condition_boost(candidate)
* self._subject_perturbation(candidate)
* self._stimulus_bias(candidate)
)
effective[candidate.region_id] = float(max(0.001, min(0.30, value)))
return effective
def _score_regions(self, region_ids: list[str]) -> float:
by_id = self._candidate_by_id()
selected = [by_id[region_id] for region_id in region_ids if region_id in by_id]
if not selected:
return 0.0
total = 0.0
group_counts: dict[str, int] = {}
for candidate in selected:
group_count = group_counts.get(candidate.redundancy_group, 0)
diminishing_return = 0.72 ** group_count
base = self._effective_base_r2.get(candidate.region_id, candidate.base_r2)
total += base * diminishing_return
group_counts[candidate.redundancy_group] = group_count + 1
# Bound cumulative explained variance to keep rewards stable.
return float(1.0 - math.exp(-total))
def _process_action(self, region_id: str) -> dict:
by_id = self._candidate_by_id()
previous_r2 = self._current_r2
if self._timestep >= self._subset.selection_budget:
self._last_feedback = "Budget exhausted; episode already complete."
return self._result(
done=True,
error="budget_exhausted",
previous_r2=previous_r2,
cost_penalty=0.0,
)
if region_id not in by_id:
self._last_feedback = f"Invalid region_id={region_id}."
return self._result(
done=False,
error="invalid_region",
previous_r2=previous_r2,
cost_penalty=0.0,
)
if region_id in self._selected_region_ids:
self._last_feedback = f"Region {region_id} was already selected."
self._timestep += 1
return self._result(
done=self._is_done(),
error="duplicate_region",
previous_r2=previous_r2,
cost_penalty=0.0,
)
candidate = by_id[region_id]
self._selected_region_ids.append(region_id)
self._timestep += 1
self._current_r2 = self._score_regions(self._selected_region_ids)
delta_r2 = self._current_r2 - previous_r2
cost_penalty = self._subset.cost_penalty * candidate.cost
done = self._is_done()
reward_breakdown = score_region_action(
previous_r2=previous_r2,
current_r2=self._current_r2,
cost_penalty=cost_penalty,
error=None,
done=done,
selected_count=len(self._selected_region_ids),
)
self._last_feedback = (
f"Selected {region_id}: delta_r2={delta_r2:.4f}, "
f"reward_components=[{format_reward_breakdown(reward_breakdown.as_dict())}]."
)
return self._result(
done=done,
error=None,
previous_r2=previous_r2,
cost_penalty=cost_penalty,
reward_components=reward_breakdown.as_dict(),
)
def _is_done(self) -> bool:
return self._timestep >= self._subset.selection_budget or (
len(self._selected_region_ids) >= self._subset.n_regions
)
def _result(
self,
done: bool,
error: str | None,
previous_r2: float,
cost_penalty: float,
reward_components: dict[str, float] | None = None,
) -> dict:
if reward_components is None:
reward_components = score_region_action(
previous_r2=previous_r2,
current_r2=self._current_r2,
cost_penalty=cost_penalty,
error=error,
done=done,
selected_count=len(self._selected_region_ids),
).as_dict()
return {
"episode_id": self._episode_id,
"reward": float(reward_components["total_reward"]),
"reward_components": reward_components,
"done": bool(done),
"error": error,
"previous_r2": float(previous_r2),
"current_r2": float(self._current_r2),
"score": float(self._current_r2),
"selection_state": self._build_selection_state(),
"feedback": self._last_feedback,
}
def _build_task_info(self) -> dict:
task = TASKS[self._task]
return {
"task_name": self._task,
"description": task["description"],
"difficulty": task["difficulty"],
"reward": task["reward"],
"reward_components": reward_columns(),
"valid_actions": task["valid_actions"],
"dataset_name": self._subset.dataset_name,
"data_source": self._subset.source,
"candidate_mode": self._subset.candidate_mode,
"atlas": self._subset.atlas,
"selection_budget": int(self._subset.selection_budget),
"prompt_top_k": int(self._subset.prompt_top_k),
"cost_penalty": float(self._subset.cost_penalty),
"n_candidate_regions": int(self._subset.n_regions),
"subject_id": self._subject_id,
"run_id": self._run_id,
"condition": self._condition,
"stimulus": compact_stimulus_for_prompt(self._stimulus_features),
"stimulus_window_size": int(self._stimulus_window_size),
"stimulus_n_windows": int(
stimulus_n_windows(self._condition or "", window_size=self._stimulus_window_size)
if self._condition
else 0
),
}
def _build_selection_state(self) -> dict:
selected_set = set(self._selected_region_ids)
candidates = []
for candidate in self._subset.candidates:
item = candidate.as_dict()
item["selected"] = candidate.region_id in selected_set
candidates.append(item)
payload: dict[str, Any] = {
"episode_id": self._episode_id,
"task_name": self._task,
"timestep": int(self._timestep),
"selection_budget": int(self._subset.selection_budget),
"remaining_budget": int(max(0, self._subset.selection_budget - self._timestep)),
"selected_regions": list(self._selected_region_ids),
"current_r2": float(self._current_r2),
"candidate_regions": candidates,
"candidate_count": int(self._subset.n_regions),
"candidate_mode": self._subset.candidate_mode,
"atlas": self._subset.atlas,
"prompt_top_k": int(self._subset.prompt_top_k),
"dataset_name": self._subset.dataset_name,
"data_source": self._subset.source,
"subject_id": self._subject_id,
"run_id": self._run_id,
"condition": self._condition,
"feedback": self._last_feedback,
}
if self._stimulus_features:
payload["stimulus"] = compact_stimulus_for_prompt(self._stimulus_features)
payload["stimulus_window"] = int(self._stimulus_features.get("window_index", 0))
payload["stimulus_n_windows"] = int(self._stimulus_features.get("n_windows", 1))
else:
payload["stimulus"] = None
return payload
def _build_observation(self, extra: dict | None = None) -> dict:
payload = {
"selection_state": json.dumps(self._build_selection_state()),
"task_name": self._task,
"timestep": int(self._timestep),
"max_timesteps": int(self._subset.selection_budget),
"feedback": self._last_feedback,
"score": float(self._current_r2),
}
if extra:
payload.update(extra)
return payload
def render_text(self) -> str:
selected = ", ".join(self._selected_region_ids) or "none"
return (
f"BrainRL step={self._timestep}/{self._subset.selection_budget} "
f"r2={self._current_r2:.4f} selected=[{selected}]"
)
def load_default_candidates() -> list[dict[str, Any]]:
"""Convenience helper for scripts that only need candidate metadata."""
return candidate_table(load_brain_subset())