1
Fork 0

feat: implemented simulator

This commit is contained in:
Jona Reynaert 2026-04-02 10:26:52 +02:00
parent f594cafead
commit 6db1129ca7

View file

@ -1,47 +1,346 @@
from __future__ import annotations from __future__ import annotations
import argparse import argparse
import time
from pathlib import Path from pathlib import Path
from typing import Any
import flax
import jax
import jax.numpy as jnp
import numpy as np
from brittle_star_project import ( from brittle_star_project import (
Backend, Backend,
BrittleStarEnv,
BrittleStarEnvFactory,
SimulationConfig,
simulate_policy,
) )
from brittle_star_project.environment import from_json from brittle_star_project.environment import from_json
from brittle_star_project.rl import RLModel # imports concrete models via rl.__init__
from brittle_star_project.rl.base import get_rl_model_registry
MODEL_BY_NAME = get_rl_model_registry() def _flatten_obs_dict(obs_dict: dict[str, Any]) -> jnp.ndarray:
MODEL_OPTIONS = sorted(MODEL_BY_NAME) """Flatten the env's observation dict into a 1D vector.
concatenates values in the dict's iteration order and skips empty arrays.
"""
parts: list[jnp.ndarray] = []
for v in obs_dict.values():
arr = jnp.asarray(v)
if arr.size == 0:
continue
parts.append(arr.reshape((-1,)))
if not parts:
return jnp.zeros((0,), dtype=jnp.float32)
return jnp.concatenate(parts, axis=0)
# A minimal policy class to load a CleanRL/Flax checkpoint and run inference.
class CleanRLPPOPolicy:
def __init__(
self,
*,
network_params: Any,
actor_params: Any,
action_dim: int,
deterministic: bool = True,
seed: int = 0,
) -> None:
from brittle_star_project.rl import Actor, Network
self._network = Network()
self._actor = Actor(action_dim=action_dim)
self._network_apply = jax.jit(self._network.apply)
self._actor_apply = jax.jit(self._actor.apply)
self._params = {
"network_params": network_params,
"actor_params": actor_params,
}
self._deterministic = deterministic
self._rng = jax.random.PRNGKey(int(seed))
@staticmethod
def load(
path: Path,
*,
action_dim: int,
deterministic: bool,
seed: int,
) -> "CleanRLPPOPolicy":
def _get_index(container: Any, idx: int) -> Any:
if isinstance(container, (list, tuple)):
return container[idx]
if isinstance(container, dict):
return container.get(idx, container.get(str(idx)))
raise KeyError(idx)
def _looks_like_indexed_dict(container: Any) -> bool:
return (
isinstance(container, dict)
and container
and all(str(k).isdigit() for k in container.keys())
)
def _parse_checkpoint(restored_obj: Any) -> tuple[Any, Any, Any, Any]:
"""Extract (args_dict, network_params, actor_params, critic_params).
`src/train.py` saves:
flax.serialization.to_bytes([vars(args), [net, actor, critic]])
`msgpack_restore()` occasionally restores lists as dicts keyed by
string indices ("0", "1", ...), so we accept both shapes.
"""
args_part: Any | None = None
params_part: Any = restored_obj
if isinstance(restored_obj, (list, tuple)) and len(restored_obj) >= 2:
args_part = restored_obj[0]
params_part = restored_obj[1]
elif _looks_like_indexed_dict(restored_obj) and (
"0" in restored_obj or "1" in restored_obj
):
args_part = restored_obj.get("0", restored_obj.get(0))
params_part = restored_obj.get("1", restored_obj.get(1))
if _looks_like_indexed_dict(params_part):
network_params = _get_index(params_part, 0)
actor_params = _get_index(params_part, 1)
critic_params = _get_index(params_part, 2)
if network_params is None or actor_params is None:
raise ValueError("Missing required params in checkpoint")
return args_part, network_params, actor_params, critic_params
if isinstance(params_part, (list, tuple)) and len(params_part) >= 2:
network_params = params_part[0]
actor_params = params_part[1]
critic_params = params_part[2] if len(params_part) >= 3 else None
return args_part, network_params, actor_params, critic_params
raise ValueError(
f"Unexpected .cleanrl_model structure in {path}. "
"Expected [args_dict, [network_params, actor_params, critic_params]] "
"or an equivalent dict-indexed variant."
)
payload = path.read_bytes()
restored = flax.serialization.msgpack_restore(payload)
_args_dict, network_params, actor_params, _critic_params = _parse_checkpoint(restored)
return CleanRLPPOPolicy(
network_params=network_params,
actor_params=actor_params,
action_dim=action_dim,
deterministic=deterministic,
seed=seed,
)
def reset(self, seed: int) -> None:
self._rng = jax.random.PRNGKey(int(seed))
def act(self, *, observations: dict[str, Any]) -> np.ndarray:
obs = _flatten_obs_dict(observations)
hidden = self._network_apply(self._params["network_params"], obs)
mean, log_std = self._actor_apply(self._params["actor_params"], hidden)
if self._deterministic:
action = mean
else:
self._rng, sub = jax.random.split(self._rng)
noise = jax.random.normal(sub, shape=mean.shape)
action = mean + noise * jnp.exp(log_std)
return np.asarray(action, dtype=np.float32).ravel()
def _get_observations(state: Any) -> dict[str, Any]:
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))
def _rollout_one_episode_headless(
*,
env: Any,
policy: CleanRLPPOPolicy,
seed: int,
max_steps: int,
) -> tuple[float, int, bool, float | None]:
"""Run one rollout up to `max_steps`.
Returns (return, length, reached_target, final_xy_dist).
"""
policy.reset(seed)
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)
# NOTE: In the MJC backend, `state.reward` is always 0.0.
# To get a meaningful return, we compute a simple progress reward:
# r_t = d_{t-1} - d_t
# where d is `xy_distance_to_target`.
steps = 0
for _ in range(int(max_steps)):
action = policy.act(observations=observations)
nu = int(state.mj_model.nu)
if nu > 0 and action.shape != (nu,):
raise ValueError(f"Policy returned action shape {action.shape}, expected ({nu},)")
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 _run_one_episode_viewer(
*,
env: Any,
policy: CleanRLPPOPolicy,
seed: int,
state: Any,
control_dt: float,
max_steps: int,
) -> None:
import mujoco.viewer
model = state.mj_model
data = state.mj_data
seed = int(seed)
episode_return = 0.0
observations = _get_observations(state)
prev_dist = _get_xy_distance_to_target(observations)
reached_target = _target_reached(state=state)
viewer = mujoco.viewer.launch_passive(model, data)
try:
steps = 0
for _step_idx in range(int(max_steps)):
if not viewer.is_running():
break
step_start = time.time()
# One control step. We do the env step under the viewer lock.
action = policy.act(observations=observations)
if model.nu > 0 and action.shape != (int(model.nu),):
raise ValueError(
f"Policy returned action shape {action.shape}, expected ({int(model.nu)},)"
)
# 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
# Real-time pacing so the viewer doesn't run as fast as possible.
remaining = control_dt - (time.time() - step_start)
if remaining > 0:
time.sleep(remaining)
# Done: target reached, fixed horizon reached, or window closed.
if viewer.is_running():
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}"
)
viewer.close()
finally:
# Ensure the GUI thread stops before the env/model/data are torn down.
try:
viewer.close()
except Exception:
pass
for _ in range(200):
if not viewer.is_running():
break
time.sleep(0.01)
def parse_args() -> argparse.Namespace: def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Simulate a trained policy in the MuJoCo viewer.") from brittle_star_project.environment import Task
p = argparse.ArgumentParser(
description="Run a trained policy for exactly one episode (viewer or headless)."
)
p.add_argument( p.add_argument(
"--model", "--model",
type=str, type=str,
default=None, required=True,
help="Path to a saved model artifact. If omitted, a model is created from --model-type.", help=("Path to a CleanRL/Flax '.cleanrl_model' checkpoint (saved by src/train.py)."),
) )
p.add_argument( p.add_argument(
"--model-type", "--deterministic",
choices=MODEL_OPTIONS, action=argparse.BooleanOptionalAction,
default="random", default=True,
help="Which model class to instantiate when --model is omitted.", help="Use mean action (deterministic) or sample actions (stochastic).",
)
p.add_argument(
"--headless",
action="store_true",
help="Run without the MuJoCo viewer (still exactly one episode).",
)
p.add_argument(
"--max-steps",
type=int,
required=True,
help=(
"Number of control steps to run (fixed horizon). "
"This script stops when this many steps are reached, or earlier if "
"the target is reached (directed locomotion)."
),
) )
p.add_argument( p.add_argument(
"--backend", "--backend",
choices=[b for b in Backend], choices=[b for b in Backend],
default=Backend.MJX, default=Backend.MJC,
) )
p.add_argument("--seed", type=int, default=None) p.add_argument("--seed", type=int, default=0)
return p.parse_args() return p.parse_args()
def main() -> None: def main() -> None:
from brittle_star_project.environment import (
BrittleStarEnv,
BrittleStarEnvFactory,
)
args = parse_args() args = parse_args()
morphology_cfg, arena_cfg, env_cfg = from_json("../configs/test.json") morphology_cfg, arena_cfg, env_cfg = from_json("../configs/test.json")
@ -63,30 +362,57 @@ def main() -> None:
# policy/model. # policy/model.
nu = int(state.mj_model.nu) nu = int(state.mj_model.nu)
if args.model is not None: model_path = Path(args.model)
model_path = Path(args.model) if model_path.suffix != ".cleanrl_model":
policy = RLModel.load(model_path) raise ValueError(f"Expected a '.cleanrl_model' checkpoint, got '{model_path.name}'.")
if hasattr(policy, "nu"):
policy.nu = nu
else:
model_cls = MODEL_BY_NAME[str(args.model_type)]
policy = model_cls(seed=seed_for_env)
if hasattr(policy, "nu"):
policy.nu = nu
# If the policy/model has a `seed` attribute, use the provided seed (or default) to reset it. policy = CleanRLPPOPolicy.load(
default_seed = int(getattr(policy, "seed", seed_for_env)) model_path,
if args.seed is not None and hasattr(policy, "reset"): action_dim=nu,
deterministic=bool(args.deterministic),
seed=seed_for_env,
)
# Reset policy RNG if a seed was provided.
default_seed = seed_for_env
if args.seed is not None:
policy.reset(int(args.seed)) policy.reset(int(args.seed))
# ======= SIMULATION ======= # ======= SIMULATION =======
rollout_cfg = SimulationConfig( if args.headless:
realtime=True, max_steps = int(args.max_steps)
seed=int(args.seed) if args.seed is not None else default_seed, if max_steps <= 0:
) raise ValueError("--max-steps must be > 0")
simulate_policy(policy, rollout_cfg, state) ep_seed = int(args.seed) if args.seed is not None else default_seed
ep_return, ep_len, reached_target, final_dist = _rollout_one_episode_headless(
env=env,
policy=policy,
seed=ep_seed,
max_steps=max_steps,
)
final_dist_str = "n/a" if final_dist is None else f"{final_dist:.3f}"
print(
"episode done: "
f"return={ep_return:.6f}, len={ep_len}, "
f"target_reached={reached_target}, final_xy_dist={final_dist_str}"
)
else:
max_steps = int(args.max_steps)
if max_steps <= 0:
raise ValueError("--max-steps must be > 0")
model_dt = float(state.mj_model.opt.timestep)
control_dt = model_dt * float(env_cfg.num_physics_steps_per_control_step)
_run_one_episode_viewer(
env=env,
policy=policy,
seed=int(args.seed) if args.seed is not None else default_seed,
state=state,
control_dt=control_dt,
max_steps=max_steps,
)
env.close() env.close()