refactor: evaluation subpackage
This commit is contained in:
parent
8f2a5d25ed
commit
9d1e2c9bff
8 changed files with 409 additions and 382 deletions
|
|
@ -4,280 +4,29 @@ Automatically extracts the training configuration (morphology, environment, etc.
|
||||||
from the sidecar metadata YAML file to ensure simulation perfectly matches training.
|
from the sidecar metadata YAML file to ensure simulation perfectly matches training.
|
||||||
Override simulation settings via CLI, e.g.:
|
Override simulation settings via CLI, e.g.:
|
||||||
uv run scripts/simulate.py \
|
uv run scripts/simulate.py \
|
||||||
simulation.morphology_override=config/morphology/3_arms.yaml \
|
simulation.morphology_override=configs/morphology/3_arms.yaml \
|
||||||
simulation.model_path=runs/.../final_model.flax
|
simulation.model_path=runs/.../final_model.flax
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import itertools
|
|
||||||
import time
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
import flax
|
|
||||||
import hydra
|
import hydra
|
||||||
import jax
|
|
||||||
import jax.numpy as jnp
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import yaml
|
|
||||||
from omegaconf import DictConfig, OmegaConf
|
from omegaconf import DictConfig, OmegaConf
|
||||||
|
import yaml
|
||||||
|
|
||||||
from brittle_star_project import Backend, BrittleStarEnv, BrittleStarEnvFactory
|
from brittle_star_project import Backend, BrittleStarEnv, BrittleStarEnvFactory
|
||||||
from brittle_star_project.configs.main_config import BrittleStarConfig
|
from brittle_star_project.configs.main_config import BrittleStarConfig
|
||||||
from brittle_star_project.configs.register_configs import register_configs
|
from brittle_star_project.configs.register_configs import register_configs
|
||||||
from brittle_star_project.environment.padded_obs_wrapper import compute_padding_masks
|
from brittle_star_project.environment.padded_obs_wrapper import compute_padding_masks
|
||||||
from brittle_star_project.environment.obs_processing import create_obs_processor
|
from brittle_star_project.environment.obs_processing import create_obs_processor
|
||||||
from brittle_star_project.environment.env_config import (
|
from brittle_star_project.environment.env_config import MorphologyConfig
|
||||||
MorphologyConfig,
|
|
||||||
ArenaConfig,
|
|
||||||
EnvConfig,
|
|
||||||
ObservationBoundsConfig,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs
|
||||||
class PolicyAgent:
|
from brittle_star_project.evaluation.policy import PolicyAgent
|
||||||
"""Wraps a trained Flax actor for deterministic inference."""
|
from brittle_star_project.evaluation.rollout import rollout_headless, rollout_viewer
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
sensor_params: Any,
|
|
||||||
actor_params: Any,
|
|
||||||
action_dim: int,
|
|
||||||
obs_processor: Any,
|
|
||||||
) -> None:
|
|
||||||
from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation
|
|
||||||
|
|
||||||
# Infer layer sizes from params
|
|
||||||
try:
|
|
||||||
dense_params = (
|
|
||||||
sensor_params.get("params", {})
|
|
||||||
if isinstance(sensor_params, dict)
|
|
||||||
else sensor_params["params"]
|
|
||||||
)
|
|
||||||
except Exception:
|
|
||||||
dense_params = sensor_params
|
|
||||||
|
|
||||||
layer_sizes = []
|
|
||||||
idx = 0
|
|
||||||
while True:
|
|
||||||
key = f"Dense_{idx}"
|
|
||||||
if key not in dense_params:
|
|
||||||
break
|
|
||||||
layer_sizes.append(int(np.asarray(dense_params[key]["kernel"]).shape[1]))
|
|
||||||
idx += 1
|
|
||||||
|
|
||||||
if not layer_sizes:
|
|
||||||
raise ValueError("Could not infer Dense_* layers from sensor params")
|
|
||||||
|
|
||||||
self._sensor = GenericDenseLayersWithActivation(layer_sizes=layer_sizes)
|
|
||||||
self._actor = Actor(action_dim=action_dim)
|
|
||||||
self._sensor_apply = jax.jit(self._sensor.apply)
|
|
||||||
self._actor_apply = jax.jit(self._actor.apply)
|
|
||||||
self._params = {
|
|
||||||
"sensor_params": sensor_params,
|
|
||||||
"actor_params": actor_params,
|
|
||||||
}
|
|
||||||
self._obs_processor = obs_processor
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def load(
|
|
||||||
path: Path,
|
|
||||||
*,
|
|
||||||
action_dim: int,
|
|
||||||
obs_processor: Any,
|
|
||||||
) -> "PolicyAgent":
|
|
||||||
payload = path.read_bytes()
|
|
||||||
restored = flax.serialization.msgpack_restore(payload)
|
|
||||||
|
|
||||||
sensor_params = None
|
|
||||||
actor_params = None
|
|
||||||
|
|
||||||
# Extract params from restored checkpoint
|
|
||||||
if isinstance(restored, dict):
|
|
||||||
params_sub = restored.get("params", {})
|
|
||||||
sensor_params = restored.get("sensor_params") or params_sub.get("sensor_params")
|
|
||||||
actor_params = restored.get("actor_params") or params_sub.get("actor_params")
|
|
||||||
elif isinstance(restored, (list, tuple)) and len(restored) >= 2:
|
|
||||||
params_part = restored[1]
|
|
||||||
if isinstance(params_part, dict):
|
|
||||||
sensor_params = params_part.get("0", params_part.get(0))
|
|
||||||
actor_params = params_part.get("1", params_part.get(1))
|
|
||||||
elif isinstance(params_part, (list, tuple)) and len(params_part) >= 2:
|
|
||||||
sensor_params = params_part[0]
|
|
||||||
actor_params = params_part[1]
|
|
||||||
|
|
||||||
if sensor_params is None or actor_params is None:
|
|
||||||
raise ValueError(f"Could not extract sensor and actor params from checkpoint: {path}")
|
|
||||||
|
|
||||||
return PolicyAgent(
|
|
||||||
sensor_params=sensor_params,
|
|
||||||
actor_params=actor_params,
|
|
||||||
action_dim=action_dim,
|
|
||||||
obs_processor=obs_processor,
|
|
||||||
)
|
|
||||||
|
|
||||||
def act(self, *, observations: dict[str, Any]) -> np.ndarray:
|
|
||||||
batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations)
|
|
||||||
obs = self._obs_processor(batched_obs)[0]
|
|
||||||
hidden = self._sensor_apply(self._params["sensor_params"], obs)
|
|
||||||
mean, _log_std = self._actor_apply(self._params["actor_params"], hidden)
|
|
||||||
|
|
||||||
# Always evaluate with the actor mean.
|
|
||||||
# (Sampling adds exploration noise, which is useful for training but not for evaluation.)
|
|
||||||
return np.asarray(mean, dtype=np.float32).ravel()
|
|
||||||
|
|
||||||
|
|
||||||
def _get_observations(state: Any) -> dict[str, Any] | None:
|
|
||||||
return getattr(state, "observations", None)
|
|
||||||
|
|
||||||
|
|
||||||
def _get_xy_distance_to_target(observations: dict[str, Any]) -> float | None:
|
|
||||||
return float(np.asarray(observations["xy_distance_to_target"]).reshape(-1)[0])
|
|
||||||
|
|
||||||
|
|
||||||
def _target_reached(*, state: Any) -> bool:
|
|
||||||
return bool(getattr(state, "terminated", False) or getattr(state, "truncated", False))
|
|
||||||
|
|
||||||
|
|
||||||
def _maybe_clip_action(
|
|
||||||
action: np.ndarray,
|
|
||||||
low: np.ndarray | None,
|
|
||||||
high: np.ndarray | None,
|
|
||||||
) -> np.ndarray:
|
|
||||||
if low is None or high is None:
|
|
||||||
return action
|
|
||||||
low = np.asarray(low, dtype=np.float32).ravel()
|
|
||||||
high = np.asarray(high, dtype=np.float32).ravel()
|
|
||||||
if low.shape != action.shape or high.shape != action.shape:
|
|
||||||
return action
|
|
||||||
return np.clip(action, low, high)
|
|
||||||
|
|
||||||
|
|
||||||
def _rollout_headless(
|
|
||||||
*,
|
|
||||||
env: BrittleStarEnv,
|
|
||||||
policy: PolicyAgent,
|
|
||||||
seed: int,
|
|
||||||
max_steps: int,
|
|
||||||
action_low: np.ndarray | None,
|
|
||||||
action_high: np.ndarray | None,
|
|
||||||
action_mask: np.ndarray | None = None,
|
|
||||||
) -> tuple[float, int, bool, float | None]:
|
|
||||||
state = env.reset(seed=seed)
|
|
||||||
|
|
||||||
ep_return = 0.0
|
|
||||||
observations = _get_observations(state)
|
|
||||||
prev_dist = _get_xy_distance_to_target(observations)
|
|
||||||
reached_target = _target_reached(state=state)
|
|
||||||
|
|
||||||
steps = 0
|
|
||||||
for _ in range(int(max_steps)):
|
|
||||||
obs_dict = observations or {}
|
|
||||||
|
|
||||||
action = policy.act(observations=obs_dict)
|
|
||||||
if action_mask is not None:
|
|
||||||
action = action[action_mask]
|
|
||||||
action = _maybe_clip_action(action, action_low, action_high)
|
|
||||||
|
|
||||||
state = env.step(state=state, action=action)
|
|
||||||
steps += 1
|
|
||||||
|
|
||||||
observations = _get_observations(state)
|
|
||||||
cur_dist = _get_xy_distance_to_target(observations)
|
|
||||||
if prev_dist is not None and cur_dist is not None:
|
|
||||||
ep_return += prev_dist - cur_dist
|
|
||||||
prev_dist = cur_dist
|
|
||||||
|
|
||||||
reached_target = _target_reached(state=state)
|
|
||||||
if reached_target:
|
|
||||||
break
|
|
||||||
|
|
||||||
final_dist = _get_xy_distance_to_target(observations)
|
|
||||||
return ep_return, steps, reached_target, final_dist
|
|
||||||
|
|
||||||
|
|
||||||
def _rollout_viewer(
|
|
||||||
*,
|
|
||||||
env: BrittleStarEnv,
|
|
||||||
policy: PolicyAgent,
|
|
||||||
seed: int,
|
|
||||||
state: Any,
|
|
||||||
control_dt: float,
|
|
||||||
max_steps: int | None,
|
|
||||||
action_low: np.ndarray | None,
|
|
||||||
action_high: np.ndarray | None,
|
|
||||||
action_mask: np.ndarray | None = None,
|
|
||||||
) -> None:
|
|
||||||
import mujoco.viewer
|
|
||||||
|
|
||||||
model = state.mj_model
|
|
||||||
data = state.mj_data
|
|
||||||
|
|
||||||
episode_return = 0.0
|
|
||||||
observations = _get_observations(state)
|
|
||||||
prev_dist = _get_xy_distance_to_target(observations)
|
|
||||||
reached_target = _target_reached(state=state)
|
|
||||||
|
|
||||||
steps = 0
|
|
||||||
# Use the viewer as a context manager to avoid GLX teardown races.
|
|
||||||
with mujoco.viewer.launch_passive(model, data) as viewer:
|
|
||||||
step_iter = range(int(max_steps)) if max_steps is not None else itertools.count()
|
|
||||||
for _step_idx in step_iter:
|
|
||||||
if not viewer.is_running():
|
|
||||||
break
|
|
||||||
step_start = time.time()
|
|
||||||
|
|
||||||
obs_dict = observations or {}
|
|
||||||
|
|
||||||
action = policy.act(observations=obs_dict)
|
|
||||||
if action_mask is not None:
|
|
||||||
action = action[action_mask]
|
|
||||||
action = _maybe_clip_action(action, action_low, action_high)
|
|
||||||
|
|
||||||
# The passive viewer runs a GUI thread; protect MuJoCo state mutation.
|
|
||||||
with viewer.lock():
|
|
||||||
state = env.step(state=state, action=action)
|
|
||||||
|
|
||||||
if not viewer.is_running():
|
|
||||||
break
|
|
||||||
viewer.sync()
|
|
||||||
|
|
||||||
steps += 1
|
|
||||||
|
|
||||||
observations = _get_observations(state)
|
|
||||||
cur_dist = _get_xy_distance_to_target(observations)
|
|
||||||
if prev_dist is not None and cur_dist is not None:
|
|
||||||
episode_return += prev_dist - cur_dist
|
|
||||||
prev_dist = cur_dist
|
|
||||||
|
|
||||||
reached_target = _target_reached(state=state)
|
|
||||||
if reached_target:
|
|
||||||
break
|
|
||||||
|
|
||||||
remaining = control_dt - (time.time() - step_start)
|
|
||||||
if remaining > 0:
|
|
||||||
time.sleep(remaining)
|
|
||||||
|
|
||||||
dist = _get_xy_distance_to_target(observations)
|
|
||||||
dist_str = "n/a" if dist is None else f"{dist:.3f}"
|
|
||||||
print(
|
|
||||||
"episode done: "
|
|
||||||
f"return={episode_return:.6f}, len={steps}, "
|
|
||||||
f"target_reached={reached_target}, final_xy_dist={dist_str}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _load_metadata_yaml(model_path: Path) -> dict:
|
|
||||||
"""Discover and load the sidecar metadata YAML file."""
|
|
||||||
metadata_path = model_path.with_name(model_path.stem + "_metadata.yaml")
|
|
||||||
if not metadata_path.exists():
|
|
||||||
raise FileNotFoundError(
|
|
||||||
f"Could not find metadata YAML for {model_path.name}. Expected it at {metadata_path}"
|
|
||||||
)
|
|
||||||
with open(metadata_path, "r") as f:
|
|
||||||
return yaml.safe_load(f)
|
|
||||||
|
|
||||||
|
|
||||||
@hydra.main(config_path="../configs", config_name="main_config", version_base="1.3")
|
@hydra.main(config_path="../configs", config_name="main_config", version_base="1.3")
|
||||||
|
|
@ -297,36 +46,10 @@ def main(dict_cfg: DictConfig) -> None:
|
||||||
raise ValueError(f"Expected a '.flax' checkpoint, got '{model_path.name}'.")
|
raise ValueError(f"Expected a '.flax' checkpoint, got '{model_path.name}'.")
|
||||||
|
|
||||||
# 2. Discover + load sidecar metadata YAML
|
# 2. Discover + load sidecar metadata YAML
|
||||||
metadata = _load_metadata_yaml(model_path)
|
metadata = load_metadata(model_path)
|
||||||
|
|
||||||
# 3. Reconstruct typed configs from metadata
|
# 3. Reconstruct typed configs from metadata
|
||||||
trained_morphology = OmegaConf.to_object(
|
training = metadata_to_configs(metadata)
|
||||||
OmegaConf.merge(OmegaConf.structured(MorphologyConfig), metadata.get("morphology", {}))
|
|
||||||
)
|
|
||||||
trained_arena = OmegaConf.to_object(
|
|
||||||
OmegaConf.merge(OmegaConf.structured(ArenaConfig), metadata.get("arena", {}))
|
|
||||||
)
|
|
||||||
|
|
||||||
env_dict = metadata.get("environment", {})
|
|
||||||
if isinstance(env_dict.get("task"), str):
|
|
||||||
from brittle_star_project.environment.env_types import Task
|
|
||||||
|
|
||||||
try:
|
|
||||||
env_dict["task"] = Task[env_dict["task"]].name
|
|
||||||
except Exception:
|
|
||||||
try:
|
|
||||||
env_dict["task"] = Task(env_dict["task"]).name
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
trained_environment = OmegaConf.to_object(
|
|
||||||
OmegaConf.merge(OmegaConf.structured(EnvConfig), env_dict)
|
|
||||||
)
|
|
||||||
trained_obs_bounds = OmegaConf.to_object(
|
|
||||||
OmegaConf.merge(
|
|
||||||
OmegaConf.structured(ObservationBoundsConfig), metadata.get("obs_bounds", {})
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
# 4. Determine environment morphology
|
# 4. Determine environment morphology
|
||||||
if sim_cfg.morphology_override is not None:
|
if sim_cfg.morphology_override is not None:
|
||||||
|
|
@ -339,14 +62,15 @@ def main(dict_cfg: DictConfig) -> None:
|
||||||
OmegaConf.merge(OmegaConf.structured(MorphologyConfig), override_dict)
|
OmegaConf.merge(OmegaConf.structured(MorphologyConfig), override_dict)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
env_morphology = trained_morphology
|
env_morphology = training.morphology
|
||||||
|
|
||||||
# 5. Build obs_processor with TRAINING morphology padding masks always
|
# 5. Build obs_processor with TRAINING morphology padding masks always
|
||||||
padding_masks = compute_padding_masks(
|
padding_masks = compute_padding_masks(
|
||||||
segments_per_arm=env_morphology.segments_per_arm,
|
segments_per_arm=env_morphology.segments_per_arm,
|
||||||
|
reference_segments_per_arm=training.morphology.segments_per_arm,
|
||||||
)
|
)
|
||||||
obs_processor = create_obs_processor(
|
obs_processor = create_obs_processor(
|
||||||
bounds_dict=trained_obs_bounds.to_bounds_dict(),
|
bounds_dict=training.obs_bounds.to_bounds_dict(),
|
||||||
padding_masks=padding_masks,
|
padding_masks=padding_masks,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -358,23 +82,23 @@ def main(dict_cfg: DictConfig) -> None:
|
||||||
raw_env = factory.create_environment(
|
raw_env = factory.create_environment(
|
||||||
backend,
|
backend,
|
||||||
env_morphology,
|
env_morphology,
|
||||||
trained_arena,
|
training.arena,
|
||||||
trained_environment,
|
training.environment,
|
||||||
)
|
)
|
||||||
env = BrittleStarEnv(
|
env = BrittleStarEnv(
|
||||||
raw_env,
|
raw_env,
|
||||||
backend=backend,
|
backend=backend,
|
||||||
config=trained_environment,
|
config=training.environment,
|
||||||
morphology_config=env_morphology,
|
morphology_config=env_morphology,
|
||||||
)
|
)
|
||||||
|
|
||||||
state0 = env.reset(seed=seed)
|
state0 = env.reset(seed=seed)
|
||||||
|
|
||||||
# Calculate the action dimension the model was trained with
|
# Calculate the action dimension the model was trained with
|
||||||
trained_action_dim = sum(trained_morphology.segments_per_arm) * 2
|
trained_action_dim = sum(training.morphology.segments_per_arm) * 2
|
||||||
|
|
||||||
# 7. Load policy
|
# 7. Load policy
|
||||||
policy = PolicyAgent.load(
|
policy = PolicyAgent.from_checkpoint(
|
||||||
model_path, action_dim=trained_action_dim, obs_processor=obs_processor
|
model_path, action_dim=trained_action_dim, obs_processor=obs_processor
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -402,7 +126,7 @@ def main(dict_cfg: DictConfig) -> None:
|
||||||
if max_steps_i <= 0:
|
if max_steps_i <= 0:
|
||||||
raise ValueError("simulation.max_steps must be > 0")
|
raise ValueError("simulation.max_steps must be > 0")
|
||||||
|
|
||||||
ep_return, ep_len, reached_target, final_dist = _rollout_headless(
|
result = rollout_headless(
|
||||||
env=env,
|
env=env,
|
||||||
policy=policy,
|
policy=policy,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
|
|
@ -411,11 +135,11 @@ def main(dict_cfg: DictConfig) -> None:
|
||||||
action_high=action_high,
|
action_high=action_high,
|
||||||
action_mask=action_mask,
|
action_mask=action_mask,
|
||||||
)
|
)
|
||||||
final_dist_str = "n/a" if final_dist is None else f"{final_dist:.3f}"
|
final_dist_str = "n/a" if result.final_xy_dist is None else f"{result.final_xy_dist:.3f}"
|
||||||
print(
|
print(
|
||||||
"episode done: "
|
"episode done: "
|
||||||
f"return={ep_return:.6f}, len={ep_len}, "
|
f"return={result.return_:.6f}, len={result.length}, "
|
||||||
f"target_reached={reached_target}, final_xy_dist={final_dist_str}"
|
f"target_reached={result.reached_target}, final_xy_dist={final_dist_str}"
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
max_steps_val = None
|
max_steps_val = None
|
||||||
|
|
@ -426,9 +150,9 @@ def main(dict_cfg: DictConfig) -> None:
|
||||||
max_steps_val = max_steps_i
|
max_steps_val = max_steps_i
|
||||||
|
|
||||||
model_dt = float(state0.mj_model.opt.timestep)
|
model_dt = float(state0.mj_model.opt.timestep)
|
||||||
control_dt = model_dt * float(trained_environment.num_physics_steps_per_control_step)
|
control_dt = model_dt * float(training.environment.num_physics_steps_per_control_step)
|
||||||
|
|
||||||
_rollout_viewer(
|
rollout_viewer(
|
||||||
env=env,
|
env=env,
|
||||||
policy=policy,
|
policy=policy,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,14 @@ from .environment.env_types import Backend, Task
|
||||||
from .environment.env_config import ArenaConfig, EnvConfig, MorphologyConfig
|
from .environment.env_config import ArenaConfig, EnvConfig, MorphologyConfig
|
||||||
from .environment.factory import BrittleStarEnvFactory
|
from .environment.factory import BrittleStarEnvFactory
|
||||||
from .environment.env_wrapper import BrittleStarEnv
|
from .environment.env_wrapper import BrittleStarEnv
|
||||||
from .render import simulate_policy, SimulationConfig, ControlPolicy
|
from .evaluation import (
|
||||||
|
PolicyAgent,
|
||||||
|
ControlPolicy,
|
||||||
|
load_metadata,
|
||||||
|
rollout_headless,
|
||||||
|
rollout_viewer,
|
||||||
|
EpisodeResult,
|
||||||
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"ArenaConfig",
|
"ArenaConfig",
|
||||||
|
|
@ -12,7 +19,10 @@ __all__ = [
|
||||||
"EnvConfig",
|
"EnvConfig",
|
||||||
"MorphologyConfig",
|
"MorphologyConfig",
|
||||||
"Task",
|
"Task",
|
||||||
"simulate_policy",
|
"PolicyAgent",
|
||||||
"SimulationConfig",
|
|
||||||
"ControlPolicy",
|
"ControlPolicy",
|
||||||
|
"load_metadata",
|
||||||
|
"rollout_headless",
|
||||||
|
"rollout_viewer",
|
||||||
|
"EpisodeResult",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
17
src/brittle_star_project/evaluation/__init__.py
Normal file
17
src/brittle_star_project/evaluation/__init__.py
Normal file
|
|
@ -0,0 +1,17 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from .checkpoint import load_metadata, load_params, metadata_to_configs, TrainingConfig
|
||||||
|
from .policy import PolicyAgent, ControlPolicy
|
||||||
|
from .rollout import rollout_headless, rollout_viewer, EpisodeResult
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"load_metadata",
|
||||||
|
"load_params",
|
||||||
|
"metadata_to_configs",
|
||||||
|
"TrainingConfig",
|
||||||
|
"PolicyAgent",
|
||||||
|
"ControlPolicy",
|
||||||
|
"rollout_headless",
|
||||||
|
"rollout_viewer",
|
||||||
|
"EpisodeResult",
|
||||||
|
]
|
||||||
105
src/brittle_star_project/evaluation/checkpoint.py
Normal file
105
src/brittle_star_project/evaluation/checkpoint.py
Normal file
|
|
@ -0,0 +1,105 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import flax
|
||||||
|
from omegaconf import OmegaConf
|
||||||
|
|
||||||
|
from brittle_star_project.environment.env_config import (
|
||||||
|
MorphologyConfig,
|
||||||
|
ArenaConfig,
|
||||||
|
EnvConfig,
|
||||||
|
ObservationBoundsConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TrainingConfig:
|
||||||
|
"""Holds typed configurations extracted from a training run's metadata."""
|
||||||
|
|
||||||
|
morphology: MorphologyConfig
|
||||||
|
arena: ArenaConfig
|
||||||
|
environment: EnvConfig
|
||||||
|
obs_bounds: ObservationBoundsConfig
|
||||||
|
|
||||||
|
|
||||||
|
def load_params(path: Path) -> dict:
|
||||||
|
"""Load model parameters from a .flax checkpoint file."""
|
||||||
|
payload = path.read_bytes()
|
||||||
|
restored = flax.serialization.msgpack_restore(payload)
|
||||||
|
|
||||||
|
sensor_params = None
|
||||||
|
actor_params = None
|
||||||
|
|
||||||
|
# Extract params from restored checkpoint
|
||||||
|
if isinstance(restored, dict):
|
||||||
|
params_sub = restored.get("params", {})
|
||||||
|
sensor_params = restored.get("sensor_params") or params_sub.get("sensor_params")
|
||||||
|
actor_params = restored.get("actor_params") or params_sub.get("actor_params")
|
||||||
|
elif isinstance(restored, (list, tuple)) and len(restored) >= 2:
|
||||||
|
params_part = restored[1]
|
||||||
|
if isinstance(params_part, dict):
|
||||||
|
sensor_params = params_part.get("0", params_part.get(0))
|
||||||
|
actor_params = params_part.get("1", params_part.get(1))
|
||||||
|
elif isinstance(params_part, (list, tuple)) and len(params_part) >= 2:
|
||||||
|
sensor_params = params_part[0]
|
||||||
|
actor_params = params_part[1]
|
||||||
|
|
||||||
|
if sensor_params is None or actor_params is None:
|
||||||
|
raise ValueError(f"Could not extract sensor and actor params from checkpoint: {path}")
|
||||||
|
|
||||||
|
return {
|
||||||
|
"sensor_params": sensor_params,
|
||||||
|
"actor_params": actor_params,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def load_metadata(model_path: Path) -> dict:
|
||||||
|
"""Discover and load the sidecar metadata YAML file."""
|
||||||
|
metadata_path = model_path.with_name(model_path.stem + "_metadata.yaml")
|
||||||
|
if not metadata_path.exists():
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Could not find metadata YAML for {model_path.name}. Expected it at {metadata_path}"
|
||||||
|
)
|
||||||
|
with open(metadata_path, "r") as f:
|
||||||
|
return yaml.safe_load(f)
|
||||||
|
|
||||||
|
|
||||||
|
def metadata_to_configs(metadata: dict) -> TrainingConfig:
|
||||||
|
"""Reconstruct typed configuration objects from a metadata dictionary."""
|
||||||
|
trained_morphology = OmegaConf.to_object(
|
||||||
|
OmegaConf.merge(OmegaConf.structured(MorphologyConfig), metadata.get("morphology", {}))
|
||||||
|
)
|
||||||
|
trained_arena = OmegaConf.to_object(
|
||||||
|
OmegaConf.merge(OmegaConf.structured(ArenaConfig), metadata.get("arena", {}))
|
||||||
|
)
|
||||||
|
|
||||||
|
env_dict = metadata.get("environment", {})
|
||||||
|
if isinstance(env_dict.get("task"), str):
|
||||||
|
from brittle_star_project.environment.env_types import Task
|
||||||
|
|
||||||
|
try:
|
||||||
|
env_dict["task"] = Task[env_dict["task"]].name
|
||||||
|
except Exception:
|
||||||
|
try:
|
||||||
|
env_dict["task"] = Task(env_dict["task"]).name
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
trained_environment = OmegaConf.to_object(
|
||||||
|
OmegaConf.merge(OmegaConf.structured(EnvConfig), env_dict)
|
||||||
|
)
|
||||||
|
trained_obs_bounds = OmegaConf.to_object(
|
||||||
|
OmegaConf.merge(
|
||||||
|
OmegaConf.structured(ObservationBoundsConfig), metadata.get("obs_bounds", {})
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return TrainingConfig(
|
||||||
|
morphology=trained_morphology,
|
||||||
|
arena=trained_arena,
|
||||||
|
environment=trained_environment,
|
||||||
|
obs_bounds=trained_obs_bounds,
|
||||||
|
)
|
||||||
89
src/brittle_star_project/evaluation/policy.py
Normal file
89
src/brittle_star_project/evaluation/policy.py
Normal file
|
|
@ -0,0 +1,89 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Protocol
|
||||||
|
|
||||||
|
import jax
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from brittle_star_project.evaluation.checkpoint import load_params
|
||||||
|
|
||||||
|
|
||||||
|
class ControlPolicy(Protocol):
|
||||||
|
"""Protocol for any policy that can produce actions from observations."""
|
||||||
|
|
||||||
|
def act(self, *, observations: dict[str, Any]) -> np.ndarray: ...
|
||||||
|
|
||||||
|
|
||||||
|
class PolicyAgent:
|
||||||
|
"""Wraps a trained Flax actor for deterministic inference."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
sensor_params: Any,
|
||||||
|
actor_params: Any,
|
||||||
|
action_dim: int,
|
||||||
|
obs_processor: Any,
|
||||||
|
) -> None:
|
||||||
|
from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation
|
||||||
|
|
||||||
|
# Infer layer sizes from params
|
||||||
|
try:
|
||||||
|
dense_params = (
|
||||||
|
sensor_params.get("params", {})
|
||||||
|
if isinstance(sensor_params, dict)
|
||||||
|
else sensor_params["params"]
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
dense_params = sensor_params
|
||||||
|
|
||||||
|
layer_sizes = []
|
||||||
|
idx = 0
|
||||||
|
while True:
|
||||||
|
key = f"Dense_{idx}"
|
||||||
|
if key not in dense_params:
|
||||||
|
break
|
||||||
|
layer_sizes.append(int(np.asarray(dense_params[key]["kernel"]).shape[1]))
|
||||||
|
idx += 1
|
||||||
|
|
||||||
|
if not layer_sizes:
|
||||||
|
raise ValueError("Could not infer Dense_* layers from sensor params")
|
||||||
|
|
||||||
|
self._sensor = GenericDenseLayersWithActivation(layer_sizes=layer_sizes)
|
||||||
|
self._actor = Actor(action_dim=action_dim)
|
||||||
|
self._sensor_apply = jax.jit(self._sensor.apply)
|
||||||
|
self._actor_apply = jax.jit(self._actor.apply)
|
||||||
|
self._params = {
|
||||||
|
"sensor_params": sensor_params,
|
||||||
|
"actor_params": actor_params,
|
||||||
|
}
|
||||||
|
self._obs_processor = obs_processor
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_checkpoint(
|
||||||
|
cls,
|
||||||
|
model_path: Path,
|
||||||
|
*,
|
||||||
|
action_dim: int,
|
||||||
|
obs_processor: Any,
|
||||||
|
) -> "PolicyAgent":
|
||||||
|
"""Load params from .flax and construct the agent."""
|
||||||
|
params = load_params(model_path)
|
||||||
|
|
||||||
|
return cls(
|
||||||
|
sensor_params=params["sensor_params"],
|
||||||
|
actor_params=params["actor_params"],
|
||||||
|
action_dim=action_dim,
|
||||||
|
obs_processor=obs_processor,
|
||||||
|
)
|
||||||
|
|
||||||
|
def act(self, *, observations: dict[str, Any]) -> np.ndarray:
|
||||||
|
"""Return deterministic action (actor mean, no exploration noise)."""
|
||||||
|
batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations)
|
||||||
|
obs = self._obs_processor(batched_obs)[0]
|
||||||
|
hidden = self._sensor_apply(self._params["sensor_params"], obs)
|
||||||
|
mean, _log_std = self._actor_apply(self._params["actor_params"], hidden)
|
||||||
|
|
||||||
|
return np.asarray(mean, dtype=np.float32).ravel()
|
||||||
163
src/brittle_star_project/evaluation/rollout.py
Normal file
163
src/brittle_star_project/evaluation/rollout.py
Normal file
|
|
@ -0,0 +1,163 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import itertools
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from brittle_star_project import BrittleStarEnv
|
||||||
|
from brittle_star_project.evaluation.policy import ControlPolicy
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class EpisodeResult:
|
||||||
|
return_: float
|
||||||
|
length: int
|
||||||
|
reached_target: bool
|
||||||
|
final_xy_dist: float | None
|
||||||
|
|
||||||
|
|
||||||
|
def _get_observations(state: Any) -> dict[str, Any] | None:
|
||||||
|
return getattr(state, "observations", None)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_xy_distance_to_target(observations: dict[str, Any]) -> float | None:
|
||||||
|
return float(np.asarray(observations["xy_distance_to_target"]).reshape(-1)[0])
|
||||||
|
|
||||||
|
|
||||||
|
def _target_reached(*, state: Any) -> bool:
|
||||||
|
return bool(getattr(state, "terminated", False) or getattr(state, "truncated", False))
|
||||||
|
|
||||||
|
|
||||||
|
def _maybe_clip_action(
|
||||||
|
action: np.ndarray,
|
||||||
|
low: np.ndarray | None,
|
||||||
|
high: np.ndarray | None,
|
||||||
|
) -> np.ndarray:
|
||||||
|
if low is None or high is None:
|
||||||
|
return action
|
||||||
|
low = np.asarray(low, dtype=np.float32).ravel()
|
||||||
|
high = np.asarray(high, dtype=np.float32).ravel()
|
||||||
|
if low.shape != action.shape or high.shape != action.shape:
|
||||||
|
return action
|
||||||
|
return np.clip(action, low, high)
|
||||||
|
|
||||||
|
|
||||||
|
def rollout_headless(
|
||||||
|
*,
|
||||||
|
env: BrittleStarEnv,
|
||||||
|
policy: ControlPolicy,
|
||||||
|
seed: int,
|
||||||
|
max_steps: int,
|
||||||
|
action_low: np.ndarray | None,
|
||||||
|
action_high: np.ndarray | None,
|
||||||
|
action_mask: np.ndarray | None = None,
|
||||||
|
) -> EpisodeResult:
|
||||||
|
"""Run an episode headlessly and return the result."""
|
||||||
|
state = env.reset(seed=seed)
|
||||||
|
|
||||||
|
ep_return = 0.0
|
||||||
|
observations = _get_observations(state)
|
||||||
|
prev_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||||
|
reached_target = _target_reached(state=state)
|
||||||
|
|
||||||
|
steps = 0
|
||||||
|
for _ in range(int(max_steps)):
|
||||||
|
obs_dict = observations or {}
|
||||||
|
|
||||||
|
action = policy.act(observations=obs_dict)
|
||||||
|
if action_mask is not None:
|
||||||
|
action = action[action_mask]
|
||||||
|
action = _maybe_clip_action(action, action_low, action_high)
|
||||||
|
|
||||||
|
state = env.step(state=state, action=action)
|
||||||
|
steps += 1
|
||||||
|
|
||||||
|
observations = _get_observations(state)
|
||||||
|
cur_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||||
|
if prev_dist is not None and cur_dist is not None:
|
||||||
|
ep_return += prev_dist - cur_dist
|
||||||
|
prev_dist = cur_dist
|
||||||
|
|
||||||
|
reached_target = _target_reached(state=state)
|
||||||
|
if reached_target:
|
||||||
|
break
|
||||||
|
|
||||||
|
final_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||||
|
return EpisodeResult(
|
||||||
|
return_=ep_return,
|
||||||
|
length=steps,
|
||||||
|
reached_target=reached_target,
|
||||||
|
final_xy_dist=final_dist,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def rollout_viewer(
|
||||||
|
*,
|
||||||
|
env: BrittleStarEnv,
|
||||||
|
policy: ControlPolicy,
|
||||||
|
seed: int,
|
||||||
|
state: Any,
|
||||||
|
control_dt: float,
|
||||||
|
max_steps: int | None,
|
||||||
|
action_low: np.ndarray | None,
|
||||||
|
action_high: np.ndarray | None,
|
||||||
|
action_mask: np.ndarray | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Run an episode using the interactive MuJoCo viewer."""
|
||||||
|
import mujoco.viewer
|
||||||
|
|
||||||
|
model = state.mj_model
|
||||||
|
data = state.mj_data
|
||||||
|
|
||||||
|
episode_return = 0.0
|
||||||
|
observations = _get_observations(state)
|
||||||
|
prev_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||||
|
reached_target = _target_reached(state=state)
|
||||||
|
|
||||||
|
steps = 0
|
||||||
|
with mujoco.viewer.launch_passive(model, data) as viewer:
|
||||||
|
step_iter = range(int(max_steps)) if max_steps is not None else itertools.count()
|
||||||
|
for _step_idx in step_iter:
|
||||||
|
if not viewer.is_running():
|
||||||
|
break
|
||||||
|
step_start = time.time()
|
||||||
|
|
||||||
|
obs_dict = observations or {}
|
||||||
|
action = policy.act(observations=obs_dict)
|
||||||
|
if action_mask is not None:
|
||||||
|
action = action[action_mask]
|
||||||
|
action = _maybe_clip_action(action, action_low, action_high)
|
||||||
|
|
||||||
|
with viewer.lock():
|
||||||
|
state = env.step(state=state, action=action)
|
||||||
|
|
||||||
|
if not viewer.is_running():
|
||||||
|
break
|
||||||
|
viewer.sync()
|
||||||
|
|
||||||
|
steps += 1
|
||||||
|
|
||||||
|
observations = _get_observations(state)
|
||||||
|
cur_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||||
|
if prev_dist is not None and cur_dist is not None:
|
||||||
|
episode_return += prev_dist - cur_dist
|
||||||
|
prev_dist = cur_dist
|
||||||
|
|
||||||
|
reached_target = _target_reached(state=state)
|
||||||
|
if reached_target:
|
||||||
|
break
|
||||||
|
|
||||||
|
remaining = control_dt - (time.time() - step_start)
|
||||||
|
if remaining > 0:
|
||||||
|
time.sleep(remaining)
|
||||||
|
|
||||||
|
dist = _get_xy_distance_to_target(observations) if observations else None
|
||||||
|
dist_str = "n/a" if dist is None else f"{dist:.3f}"
|
||||||
|
print(
|
||||||
|
"episode done: "
|
||||||
|
f"return={episode_return:.6f}, len={steps}, "
|
||||||
|
f"target_reached={reached_target}, final_xy_dist={dist_str}"
|
||||||
|
)
|
||||||
|
|
@ -1,3 +0,0 @@
|
||||||
from .renderer import simulate_policy, SimulationConfig, ControlPolicy
|
|
||||||
|
|
||||||
__all__ = ["simulate_policy", "SimulationConfig", "ControlPolicy"]
|
|
||||||
|
|
@ -1,78 +0,0 @@
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import time
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from typing import Any, Protocol
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class SimulationConfig:
|
|
||||||
realtime: bool = True
|
|
||||||
seed: int = 0
|
|
||||||
|
|
||||||
|
|
||||||
class ControlPolicy(Protocol):
|
|
||||||
def act(self, *, obs: np.ndarray | None = None, t: float = 0.0) -> np.ndarray: ...
|
|
||||||
|
|
||||||
|
|
||||||
def _default_observations(data: Any) -> np.ndarray:
|
|
||||||
qpos = np.asarray(data.qpos, dtype=np.float32).ravel()
|
|
||||||
qvel = np.asarray(data.qvel, dtype=np.float32).ravel()
|
|
||||||
return np.concatenate([qpos, qvel], axis=0)
|
|
||||||
|
|
||||||
|
|
||||||
def simulate_policy(
|
|
||||||
policy: ControlPolicy,
|
|
||||||
config: SimulationConfig,
|
|
||||||
state: Any | None = None,
|
|
||||||
) -> None:
|
|
||||||
"""Open MuJoCo's native viewer and step using actions from a policy.
|
|
||||||
|
|
||||||
This path drives MuJoCo physics directly (mj_step) and uses the policy output
|
|
||||||
as `data.ctrl`.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import mujoco.viewer
|
|
||||||
|
|
||||||
if state is None:
|
|
||||||
raise ValueError("A valid environment state must be provided.")
|
|
||||||
|
|
||||||
model = state.mj_model
|
|
||||||
data = state.mj_data
|
|
||||||
|
|
||||||
start = time.time()
|
|
||||||
with mujoco.viewer.launch_passive(model, data) as viewer:
|
|
||||||
while viewer.is_running():
|
|
||||||
step_start = time.time()
|
|
||||||
|
|
||||||
t = time.time() - start
|
|
||||||
|
|
||||||
# Input vector for the policy
|
|
||||||
# TODO: custom input
|
|
||||||
obs = _default_observations(data)
|
|
||||||
|
|
||||||
# Policy action
|
|
||||||
ctrl = policy.act(obs=obs, t=t)
|
|
||||||
|
|
||||||
# Check if the policy output vector give an input for each actuator (nu)
|
|
||||||
# TODO: what if model trained on full morphology but we want to test on a damaged one?
|
|
||||||
# (nu mismatch)
|
|
||||||
if model.nu > 0:
|
|
||||||
ctrl = np.asarray(ctrl, dtype=np.float32).ravel()
|
|
||||||
if ctrl.shape != (model.nu,):
|
|
||||||
raise ValueError(
|
|
||||||
f"Policy returned ctrl shape {ctrl.shape}, expected ({model.nu},)"
|
|
||||||
)
|
|
||||||
data.ctrl[:] = ctrl
|
|
||||||
|
|
||||||
# Step the simulation and update the viewer
|
|
||||||
mujoco.mj_step(model, data)
|
|
||||||
viewer.sync()
|
|
||||||
|
|
||||||
# If we're running in realtime mode, sleep to maintain real-time pacing.
|
|
||||||
if config.realtime:
|
|
||||||
remaining = model.opt.timestep - (time.time() - step_start)
|
|
||||||
if remaining > 0:
|
|
||||||
time.sleep(remaining)
|
|
||||||
Reference in a new issue