transformer / inference.py
Mohith202's picture
Deploy BrainRL OpenEnv server
32d14f4 verified
Raw
History Blame Contribute Delete
5.37 kB
#!/usr/bin/env python
"""Prompt-based LLM/static-prior runner for BrainRL.
Runs one (or more) full episodes through the OpenEnv environment with the
selected condition / subject context. With ``--use-llm`` it asks a chat model
for one JSON action per OpenEnv step; without, it uses the deterministic
prompt fallback so the demo always works.
"""
from __future__ import annotations
import argparse
from pathlib import Path
from data_split import DEFAULT_PARTICIPANT_INFO, build_condition_split
from prompts import build_action_prompt, llm_prompt_policy_action, static_prompt_policy_action
from server.brain_environment import BrainRegionSelectionEnvironment
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description="Run prompt-based inference on BrainRL")
parser.add_argument("--use-llm", action="store_true", help="Ask an LLM for each action")
parser.add_argument("--show-prompts", action="store_true", help="Print prompt every step")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--episodes", type=int, default=1, help="How many episodes to roll out")
parser.add_argument(
"--condition",
choices=("single_m", "single_f", "mixed_m", "mixed_f"),
default=None,
)
parser.add_argument(
"--participant-info",
type=str,
default=str(DEFAULT_PARTICIPANT_INFO),
)
parser.add_argument("--train-subjects", type=str, default=None)
parser.add_argument("--test-subjects", type=str, default=None)
parser.add_argument(
"--exclude-subjects",
type=str,
default=None,
help="Comma list / range of subjects to drop (e.g. corrupted recordings).",
)
parser.add_argument(
"--split",
choices=("train", "test", "all"),
default="all",
)
parser.add_argument("--subject-id", type=str, default=None, help="Override single subject id")
parser.add_argument("--run-id", type=str, default=None, help="Override single run id (run1..run4)")
return parser
def _episode_pairs(args: argparse.Namespace) -> list[dict[str, str | None]]:
if args.subject_id:
return [
{
"subject_id": args.subject_id,
"run_id": args.run_id,
"condition": args.condition,
}
]
if args.condition is None:
return [{"subject_id": None, "run_id": None, "condition": None}]
split = build_condition_split(
condition=args.condition,
participant_info_path=args.participant_info,
train_subjects_spec=args.train_subjects,
test_subjects_spec=args.test_subjects,
exclude_subjects=args.exclude_subjects,
)
pairs = split.pairs_for(args.split)
if not pairs:
raise SystemExit(f"No pairs for {args.condition}/{args.split}.")
return [pair.as_dict() for pair in pairs]
def run_one_episode(env: BrainRegionSelectionEnvironment, args: argparse.Namespace) -> None:
stim = env._stimulus_features
stim_label = (
f" stimulus=window={stim['window_index'] + 1}/{stim['n_windows']} "
f"dominant_pos={stim.get('dominant_pos')} n_words={stim.get('n_words')} "
f"density={stim.get('speech_density'):.2f}"
if stim
else " stimulus=unavailable"
)
print(
f"[INIT] subject={env._subject_id} run={env._run_id} condition={env._condition} "
f"candidates={env._subset.n_regions} budget={env._subset.selection_budget} "
f"prompt_top_k={env._subset.prompt_top_k}" + stim_label
)
rewards: list[float] = []
done = False
while not done:
state = env._build_selection_state()
if args.show_prompts:
print("[PROMPT]")
print(build_action_prompt(state))
if args.use_llm:
region_id, raw_text = llm_prompt_policy_action(state)
else:
region_id = static_prompt_policy_action(state)
raw_text = f'{{"region_id": "{region_id}"}}'
result = env._process_action(region_id)
rewards.append(float(result["reward"]))
print(
f"[STEP {env._timestep}/{env._subset.selection_budget}] region={region_id} "
f"reward={float(result['reward']):.4f} r2={float(result['current_r2']):.4f} "
f"done={result['done']}"
)
if args.use_llm:
print(f"[MODEL] {raw_text}")
done = bool(result["done"])
print(
f"[END] final_r2={env._current_r2:.4f} total_reward={sum(rewards):.4f} "
f"selected={len(env._selected_region_ids)} "
f"order={' -> '.join(env._selected_region_ids[:10])}"
+ (" ..." if len(env._selected_region_ids) > 10 else "")
)
def main() -> None:
args = build_parser().parse_args()
pairs = _episode_pairs(args)
env = BrainRegionSelectionEnvironment()
for episode_idx in range(int(args.episodes)):
pair = pairs[episode_idx % len(pairs)]
env.reset(
seed=int(args.seed) + episode_idx,
subject_id=pair.get("subject_id"),
run_id=pair.get("run_id"),
condition=pair.get("condition"),
)
if int(args.episodes) > 1:
print(f"\n=== Episode {episode_idx + 1}/{args.episodes} ===")
run_one_episode(env, args)
if __name__ == "__main__":
main()