Source code for bigym.loco.eval.runner

"""Reference evaluation runner — the executable form of protocol v1.

The protocol module defines the evaluation block (100 episodes, contiguous
seeds from 620000); this runner builds the task's official environment and
emits a leaderboard record only when ``protocol_violations`` finds no
result-affecting departure from the official configuration.

Library use::

    from bigym.loco.eval import evaluate

    result = evaluate(
        my_policy,                      # callable: ExtendedTimeStep -> action
        task_name="move_plate",
        method="my-cool-method",
        # The env is bigym.loco.make(task_name, config, **overrides): the
        # official configuration unless told otherwise. Overrides of
        # result-affecting fields suppress the record.
        overrides={"frame_stack": 2},
    )
    print(result["record"])             # leaderboard entry, fingerprint inside
    print(result["protocol_violations"])  # [] for an official env

CLI use (the policy is loaded from an import spec ``module:factory`` where
``factory()`` returns the policy callable)::

    python -m bigym.loco.eval.runner \
        --task move_plate --method my-cool-method \
        --policy my_package.eval:make_policy \
        [--overrides '{"camera_keys": ["head"]}'] \
        --out record.json

Policy contract: a callable ``policy(timestep) -> np.ndarray`` (action for
``env.action_spec()``); if it has a ``reset()`` method it is called between
episodes (stateful policies: action chunking, RNNs).

Protocol parameters are deliberately NOT configurable. Running fewer
episodes (``episodes=N``) is supported for smoke-testing the plumbing, but
such runs return a summary WITHOUT a leaderboard record — a number produced
off-protocol should never look like a benchmark number. The same holds for
an env built with result-affecting overrides
(``bigym.loco.eval.protocol_violations``): the summary is returned, the
departing fields are listed, and no record is produced.

Per-episode records (seed, success, length, reward, termination reason,
fell) are returned under ``episode_records`` and written next to
``--out`` as ``<out>.episodes.csv``.
"""

from __future__ import annotations

import importlib
import json
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable, Mapping

import numpy as np
import tyro

from bigym.loco import make
from bigym.loco.config import EnvConfig
from bigym.loco.eval.protocol import (
    EVAL_EPISODES,
    eval_seeds,
    is_success,
    leaderboard_record,
    protocol_violations,
)


[docs] def evaluate( policy: Callable[..., np.ndarray], *, task_name: str, method: str, checkpoint: str = "n/a", episodes: int = EVAL_EPISODES, config: EnvConfig | None = None, overrides: Mapping[str, Any] | None = None, env: Any = None, verbose: bool = True, ) -> dict[str, Any]: """Run one evaluation block; return summary (+ leaderboard record). Pass either ``config`` / ``overrides`` (forwarded to :func:`bigym.loco.make` with ``task_name``; the env is created and closed here) or a ready ``env`` (kept open). Returns ``{"task", "method", "episodes", "success_rate", "episode_rewards", "substrate", "env_config", "protocol_violations", ["record"]}``; ``record`` (the protocol-v1 leaderboard entry) only for a full block on an official env. """ episodes = int(episodes) if episodes < 1: raise ValueError("episodes must be >= 1") owns_env = env is None if owns_env: env = make(task_name, config, **dict(overrides or {})) try: seeds = eval_seeds(episodes) episode_rewards: list[float] = [] episode_records: list[dict[str, Any]] = [] successes = 0 for index, seed in enumerate(seeds): if hasattr(policy, "reset"): policy.reset() # ty: ignore[call-non-callable] timestep = env.reset(seed=seed) total_reward = 0.0 length = 0 while not timestep.last(): action = policy(timestep) timestep = env.step(action) total_reward += float(timestep.reward or 0.0) length += 1 fell = bool(env.episode_fell()) success = is_success(env) successes += int(success) episode_rewards.append(total_reward) termination = env.episode_termination() episode_records.append( { "episode": index, "seed": int(seed), "success": int(success), "length": int(length), "reward": float(total_reward), "termination": termination, "fell": int(fell), } ) if verbose: print( f"[eval] ep {index + 1}/{episodes} seed={seed} " f"reward={total_reward:.3f} success={int(success)} " f"end={termination} running_sr={successes / (index + 1):.2%}" ) success_rate = successes / episodes fingerprint = env.substrate_fingerprint() violations = protocol_violations(env) env_config = getattr(env, "config", None) result: dict[str, Any] = { "task": str(task_name), "method": str(method), "episodes": episodes, "success_rate": success_rate, "episode_rewards": episode_rewards, "episode_records": episode_records, "substrate": dict(fingerprint), "env_config": env_config.to_dict() if env_config is not None else None, "protocol_violations": violations, } if violations and verbose: print( "[eval] env departs from the official configuration, NO " "leaderboard record: " + "; ".join(violations) ) if episodes == EVAL_EPISODES and not violations: result["record"] = leaderboard_record( task=task_name, method=method, backend=str(fingerprint.get("lowerbody_backend") or "none"), success_rate=success_rate, episodes=episodes, checkpoint=checkpoint, substrate_fingerprint=fingerprint, seeds=seeds, ) elif verbose: print( f"[eval] {episodes} episodes < protocol {EVAL_EPISODES}: " "smoke run, NO leaderboard record produced." ) return result finally: if owns_env: env.close()
def _load_policy(spec: str) -> Callable[..., np.ndarray]: module_name, sep, attr = spec.partition(":") if not sep: raise SystemExit( f"--policy must be 'module:factory' (a zero-arg factory returning " f"the policy callable), got {spec!r}" ) factory = getattr(importlib.import_module(module_name), attr) return factory() @dataclass class EvalConfig: """Settings of the protocol evaluation run.""" task: str """Task name (see TASKS).""" method: str """Method name for the record.""" policy: str """Import spec 'module:factory'; factory() returns the policy callable.""" overrides: str = "{}" """JSON dict of EnvConfig fields applied last, e.g. '{"frame_stack": 2}'.""" checkpoint: str = "n/a" """Checkpoint name for the record.""" episodes: int = EVAL_EPISODES """Episodes to run; fewer than the protocol's 100 is a smoke run with no record.""" out: Path | None = None """Write result JSON here.""" def main() -> None: """Command-line entry point: run the protocol eval and report results.""" args = tyro.cli(EvalConfig, description=__doc__) result = evaluate( _load_policy(args.policy), task_name=args.task, method=args.method, checkpoint=args.checkpoint, episodes=args.episodes, overrides=json.loads(args.overrides), ) print( f"\n{args.method} on {args.task}: success_rate=" f"{result['success_rate']:.2%} over {result['episodes']} episodes" ) if args.out is not None: args.out.write_text(json.dumps(result, indent=2)) print(f"wrote {args.out}") episodes_csv = args.out.with_suffix(args.out.suffix + ".episodes.csv") with episodes_csv.open("w") as f: f.write("episode,seed,success,length,reward,termination,fell\n") for rec in result["episode_records"]: f.write( f"{rec['episode']},{rec['seed']},{rec['success']},{rec['length']}," f"{rec['reward']},{rec['termination']},{rec['fell']}\n" ) print(f"wrote {episodes_csv}") if __name__ == "__main__": main()