1
Fork 0

feat: adapted simulate to trained config

This commit is contained in:
Jona Reynaert 2026-04-22 18:47:35 +02:00
parent c4447976ab
commit 395b04d9a8
4 changed files with 334 additions and 30 deletions

View file

@ -4,13 +4,12 @@
# Path to the trained model (optional) # Path to the trained model (optional)
model_path: null model_path: null
# Type of model to use if no path is provided (e.g., random)
model_type: "random"
# Execution backend (MJX or BRAX)
backend: "MJC"
# Script behavior # Script behavior
headless: false headless: false
# In headless mode this is required; in viewer mode null means "infinite". # In headless mode this is required; in viewer mode null means "infinite".
max_steps: null max_steps: null
# Optional: path to the Hydra config.yaml used during training.
# When set, scripts/simulate.py will use it to default morphology/arena/environment/architecture
# to match training (unless you explicitly override those keys via CLI).
trained_config_path: null

View file

@ -3,6 +3,10 @@
Uses Hydra to load the same BrittleStarConfig that was used during training. Uses Hydra to load the same BrittleStarConfig that was used during training.
Override settings via CLI, e.g.: Override settings via CLI, e.g.:
python scripts/simulate.py morphology=3_arms python scripts/simulate.py morphology=3_arms
To replay a run using the *exact* Hydra config used during training, pass:
python scripts/simulate.py simulation.trained_config_path=runs/.../.hydra/config.yaml \
simulation.model_path=runs/.../final_model.flax
""" """
from __future__ import annotations from __future__ import annotations
@ -17,11 +21,16 @@ import hydra
import jax import jax
import jax.numpy as jnp import jax.numpy as jnp
import numpy as np import numpy as np
from omegaconf import DictConfig, OmegaConf import yaml
from omegaconf import DictConfig, OmegaConf, open_dict
from brittle_star_project import 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,
pad_observation,
)
_ALLOWED_OBS_KEYS = { _ALLOWED_OBS_KEYS = {
"joint_position", "joint_position",
@ -36,6 +45,74 @@ _ALLOWED_OBS_KEYS = {
"xy_distance_to_target", "xy_distance_to_target",
} }
def _dense_layer_sizes_from_params(params: Any) -> list[int]:
"""Infer GenericDenseLayersWithActivation.layer_sizes from a Flax params tree."""
try:
dense_params = params["params"]
except Exception as exc:
raise ValueError("Unexpected sensor params structure (missing 'params')") from exc
layer_sizes: list[int] = []
idx = 0
while True:
key = f"Dense_{idx}"
if key not in dense_params:
break
kernel = dense_params[key]["kernel"]
layer_sizes.append(int(np.asarray(kernel).shape[1]))
idx += 1
if not layer_sizes:
raise ValueError("Could not infer Dense_* layers from sensor params")
return layer_sizes
def _infer_action_dim_from_actor_params(params: Any) -> int | None:
"""Best-effort infer action_dim from a Flax Actor params tree."""
try:
dense0 = params["params"]["Dense_0"]
bias = dense0.get("bias")
kernel = dense0.get("kernel")
except Exception:
return None
if bias is not None:
try:
return int(np.asarray(bias).shape[0])
except Exception:
return None
if kernel is not None:
try:
return int(np.asarray(kernel).shape[1])
except Exception:
return None
return None
def _has_cli_override(overrides: list[str], key: str) -> bool:
prefixes = (f"{key}=", f"{key}.", f"+{key}=", f"+{key}.")
return any(str(o).startswith(prefixes) for o in overrides)
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 _transform_obs_dict(obs_dict: dict[str, Any]) -> jnp.ndarray: def _transform_obs_dict(obs_dict: dict[str, Any]) -> jnp.ndarray:
"""Flatten the env's observation dict into a 1D vector. """Flatten the env's observation dict into a 1D vector.
@ -69,9 +146,8 @@ class CleanRLPPOPolicy:
) -> None: ) -> None:
from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation
hidden_dim = int(sensor_params["params"]["Dense_0"]["kernel"].shape[1]) layer_sizes = _dense_layer_sizes_from_params(sensor_params)
self._sensor = GenericDenseLayersWithActivation(layer_sizes=layer_sizes)
self._sensor = GenericDenseLayersWithActivation(layer_sizes=[hidden_dim, hidden_dim])
self._actor = Actor(action_dim=action_dim) self._actor = Actor(action_dim=action_dim)
self._sensor_apply = jax.jit(self._sensor.apply) self._sensor_apply = jax.jit(self._sensor.apply)
self._actor_apply = jax.jit(self._actor.apply) self._actor_apply = jax.jit(self._actor.apply)
@ -156,6 +232,28 @@ class CleanRLPPOPolicy:
feature_extractor_params, feature_extractor_params,
) )
# Accept a plain dict-shaped Flax params mapping commonly produced
# by saving `agent_state.params` directly. Typical keys are
# 'sensor_params' and 'actor_params', or sometimes nested under 'params'.
if isinstance(restored_obj, dict):
# Top-level params dict
params_sub = restored_obj.get("params", {})
sensor_params = restored_obj.get("sensor_params") or params_sub.get("sensor_params")
actor_params = restored_obj.get("actor_params") or params_sub.get("actor_params")
critic_params = restored_obj.get("critic_params") or params_sub.get("critic_params")
feature_extractor_params = restored_obj.get(
"feature_extractor_params"
) or params_sub.get("feature_extractor_params")
# Some checkpoints only save actor+sensor as top-level
if sensor_params is not None and actor_params is not None:
return (
cfg_part,
sensor_params,
actor_params,
critic_params,
feature_extractor_params,
)
raise ValueError( raise ValueError(
f"Unexpected checkpoint structure in {path}. " f"Unexpected checkpoint structure in {path}. "
"Expected [config_dict, [sensor_params, actor_params, critic_params, " "Expected [config_dict, [sensor_params, actor_params, critic_params, "
@ -168,6 +266,16 @@ class CleanRLPPOPolicy:
_parse_checkpoint(restored) _parse_checkpoint(restored)
) )
ckpt_action_dim = _infer_action_dim_from_actor_params(actor_params)
if ckpt_action_dim is not None and ckpt_action_dim != action_dim:
raise ValueError(
"Checkpoint/env mismatch: "
f"checkpoint expects action_dim={ckpt_action_dim}, "
f"env provides action_dim={action_dim}. "
"Use the same Hydra config (morphology/arena/environment) "
"that was used during training."
)
return CleanRLPPOPolicy( return CleanRLPPOPolicy(
sensor_params=sensor_params, sensor_params=sensor_params,
actor_params=actor_params, actor_params=actor_params,
@ -193,7 +301,7 @@ def _get_xy_distance_to_target(observations: dict[str, Any]) -> float | None:
def _target_reached(*, state: Any) -> bool: def _target_reached(*, state: Any) -> bool:
return bool(getattr(state, "terminated", False)) return bool(getattr(state, "terminated", False) or getattr(state, "truncated", False))
def _rollout_one_episode_headless( def _rollout_one_episode_headless(
@ -202,6 +310,9 @@ def _rollout_one_episode_headless(
policy: CleanRLPPOPolicy, policy: CleanRLPPOPolicy,
seed: int, seed: int,
max_steps: int, max_steps: int,
action_low: np.ndarray | None,
action_high: np.ndarray | None,
padding_masks: dict[str, Any] | None,
) -> tuple[float, int, bool, float | None]: ) -> tuple[float, int, bool, float | None]:
"""Run one rollout up to max_steps. """Run one rollout up to max_steps.
@ -220,7 +331,12 @@ def _rollout_one_episode_headless(
steps = 0 steps = 0
for _ in range(int(max_steps)): for _ in range(int(max_steps)):
action = policy.act(observations=observations) obs_dict = observations or {}
if padding_masks is not None:
obs_dict = pad_observation(obs_dict, padding_masks)
action = policy.act(observations=obs_dict)
action = _maybe_clip_action(action, action_low, action_high)
nu = int(state.mj_model.nu) nu = int(state.mj_model.nu)
if nu > 0 and action.shape != (nu,): if nu > 0 and action.shape != (nu,):
@ -251,6 +367,9 @@ def _run_one_episode_viewer(
state: Any, state: Any,
control_dt: float, control_dt: float,
max_steps: int | None, max_steps: int | None,
action_low: np.ndarray | None,
action_high: np.ndarray | None,
padding_masks: dict[str, Any] | None,
) -> None: ) -> None:
import mujoco.viewer import mujoco.viewer
@ -272,7 +391,12 @@ def _run_one_episode_viewer(
break break
step_start = time.time() step_start = time.time()
action = policy.act(observations=observations or {}) obs_dict = observations or {}
if padding_masks is not None:
obs_dict = pad_observation(obs_dict, padding_masks)
action = policy.act(observations=obs_dict)
action = _maybe_clip_action(action, action_low, action_high)
if model.nu > 0 and action.shape != (int(model.nu),): if model.nu > 0 and action.shape != (int(model.nu),):
raise ValueError( raise ValueError(
f"Policy returned action shape {action.shape}, expected ({int(model.nu)},)" f"Policy returned action shape {action.shape}, expected ({int(model.nu)},)"
@ -321,16 +445,126 @@ def _infer_checkpoint_obs_dim(policy: CleanRLPPOPolicy) -> int | None:
return None return None
def _load_trained_config(path: Path) -> DictConfig:
"""Load a trained config YAML.
Supports both:
- Hydra's run config (e.g. runs/.../.hydra/config.yaml)
- This project's logger metadata YAMLs, which may contain
``!!python/object/apply:...`` tags for Enums.
For safety, we *do not* execute Python constructors from YAML; we only
treat these tags as data and extract their scalar arguments.
"""
if not path.exists():
raise FileNotFoundError(f"trained_config_path does not exist: '{path}'.")
if not path.is_file():
raise ValueError(f"trained_config_path must be a file, got: '{path}'.")
try:
return OmegaConf.load(path)
except Exception as exc:
python_apply_prefix = "tag:yaml.org,2002:python/object/apply:"
class _SafeLoaderWithPythonApply(yaml.SafeLoader):
pass
def _construct_python_apply(
loader: yaml.SafeLoader,
_tag_suffix: str,
node: yaml.Node,
) -> Any:
if isinstance(node, yaml.SequenceNode):
seq = loader.construct_sequence(node)
if len(seq) == 1:
return seq[0]
return seq
if isinstance(node, yaml.MappingNode):
return loader.construct_mapping(node)
return loader.construct_scalar(node)
_SafeLoaderWithPythonApply.add_multi_constructor(
python_apply_prefix, _construct_python_apply
)
try:
data = yaml.load(path.read_text(encoding="utf-8"), Loader=_SafeLoaderWithPythonApply)
except Exception as yaml_exc:
raise ValueError(
"Failed to load trained_config_path as YAML. "
"If this is a Hydra run, pass the run's '.hydra/config.yaml' file. "
f"Got: '{path}'."
) from yaml_exc
if not isinstance(data, dict):
raise ValueError(
"trained_config_path must contain a YAML mapping (dict-like) at the root. "
f"Got type={type(data).__name__} from '{path}'."
) from exc
# Normalize known enum-like strings to their Enum *names* so OmegaConf's
# structured config merge behaves like the normal Hydra config.
from brittle_star_project.environment.env_types import Task
env_cfg = data.get("environment")
if isinstance(env_cfg, dict) and isinstance(env_cfg.get("task"), str):
task_str = str(env_cfg["task"])
try:
env_cfg["task"] = Task[task_str].name
except Exception:
try:
env_cfg["task"] = Task(task_str).name
except Exception:
pass
return OmegaConf.create(data)
@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")
def main(dict_cfg: DictConfig) -> None: def main(dict_cfg: DictConfig) -> None:
# Convert DictConfig to structured dataclass, ensuring the root schema is applied. # Compose against the structured schema first, so missing keys are validated.
config: BrittleStarConfig = OmegaConf.to_object( cfg = OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg)
OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg)
)
backend = config.simulation.backend # Optional: override env-defining sections (morphology/arena/environment/architecture)
# using the exact Hydra config that was used for training.
trained_cfg_path = cfg.simulation.trained_config_path
if trained_cfg_path:
overrides_raw = OmegaConf.select(cfg, "hydra.overrides.task") or []
overrides = [str(o) for o in overrides_raw]
trained_cfg_path_abs = Path(hydra.utils.to_absolute_path(trained_cfg_path))
trained_cfg = _load_trained_config(trained_cfg_path_abs)
if "hydra" in trained_cfg:
with open_dict(trained_cfg):
del trained_cfg["hydra"]
with open_dict(cfg):
for key in ("morphology", "arena", "environment", "architecture"):
if key in trained_cfg and not _has_cli_override(overrides, key):
base_node = OmegaConf.select(cfg, key)
override_node = OmegaConf.select(trained_cfg, key)
try:
cfg[key] = OmegaConf.merge(base_node, override_node)
except Exception as exc:
raise ValueError(
"Failed to merge trained config into the active Hydra config. "
f"Key={key!r}, trained_config_path='{trained_cfg_path_abs}'."
) from exc
# Convert DictConfig to structured dataclass.
config: BrittleStarConfig = OmegaConf.to_object(cfg)
backend = Backend.MJC
seed = int(config.experiment.seed) seed = int(config.experiment.seed)
if getattr(config.architecture, "name", None) != "centralized":
raise ValueError(
"simulate.py currently only supports architecture=centralized. "
f"Got architecture.name={getattr(config.architecture, 'name', None)!r}. "
"(Training supports decentralized, but simulation wiring for it isn't implemented.)"
)
model_path_str = config.simulation.model_path model_path_str = config.simulation.model_path
if model_path_str is None: if model_path_str is None:
raise ValueError( raise ValueError(
@ -359,15 +593,28 @@ def main(dict_cfg: DictConfig) -> None:
state0 = env.reset(seed=seed) state0 = env.reset(seed=seed)
# Match training's padded observation layout for amputated morphologies.
padding_masks = compute_padding_masks(config.morphology.segments_per_arm)
# Match training's action clipping behavior.
action_space = getattr(raw_env, "action_space", None)
action_low = (
None if action_space is None else np.asarray(action_space.low, dtype=np.float32).ravel()
)
action_high = (
None if action_space is None else np.asarray(action_space.high, dtype=np.float32).ravel()
)
# ======= MODEL SETUP ======= # ======= MODEL SETUP =======
nu = int(state0.mj_model.nu) nu = int(state0.mj_model.nu)
policy = CleanRLPPOPolicy.load(model_path, action_dim=nu) policy = CleanRLPPOPolicy.load(model_path, action_dim=nu)
# Helpful early failure when configs don't match the checkpoint. # Helpful early failure when configs don't match the checkpoint.
observations0 = _get_observations(state0) observations0 = _get_observations(state0)
env_obs_dim = int(_transform_obs_dict(observations0 or {}).shape[0]) obs0_dict = pad_observation(observations0 or {}, padding_masks)
env_obs_dim = int(_transform_obs_dict(obs0_dict).shape[0])
ckpt_obs_dim = _infer_checkpoint_obs_dim(policy) ckpt_obs_dim = _infer_checkpoint_obs_dim(policy)
if ckpt_obs_dim is not None and ckpt_obs_dim != env_obs_dim: if ckpt_obs_dim is not None and ckpt_obs_dim != env_obs_dim:
raise ValueError( raise ValueError(
"Checkpoint/env mismatch: " "Checkpoint/env mismatch: "
@ -375,7 +622,7 @@ def main(dict_cfg: DictConfig) -> None:
"Use the same Hydra config (morphology/arena/environment) " "Use the same Hydra config (morphology/arena/environment) "
"that was used during training." "that was used during training."
) )
# ======= SIMULATION ======= # ======= SIMULATION =======
headless = bool(config.simulation.headless) headless = bool(config.simulation.headless)
max_steps = config.simulation.max_steps max_steps = config.simulation.max_steps
@ -392,6 +639,9 @@ def main(dict_cfg: DictConfig) -> None:
policy=policy, policy=policy,
seed=seed, seed=seed,
max_steps=max_steps_i, max_steps=max_steps_i,
action_low=action_low,
action_high=action_high,
padding_masks=padding_masks,
) )
final_dist_str = "n/a" if final_dist is None else f"{final_dist:.3f}" final_dist_str = "n/a" if final_dist is None else f"{final_dist:.3f}"
print( print(
@ -418,6 +668,9 @@ def main(dict_cfg: DictConfig) -> None:
state=state0, state=state0,
control_dt=control_dt, control_dt=control_dt,
max_steps=max_steps_val, max_steps=max_steps_val,
action_low=action_low,
action_high=action_high,
padding_masks=padding_masks,
) )
env.close() env.close()

View file

@ -1,18 +1,19 @@
from dataclasses import dataclass from dataclasses import dataclass
from typing import Optional from typing import Optional
from brittle_star_project.environment.env_types import Backend
@dataclass @dataclass
class SimulationSettings: class SimulationSettings:
"""Settings for the simulation script.""" """Settings for the simulation script."""
model_path: Optional[str] = None model_path: Optional[str] = None
model_type: str = "random"
backend: Backend = Backend.MJC
# Script behavior # Script behavior
headless: bool = False headless: bool = False
# If None, viewer mode runs until window closed or target reached. # If None, viewer mode runs until window closed or target reached.
max_steps: Optional[int] = None max_steps: Optional[int] = None
# Optional: point to a Hydra config.yaml from a training run (e.g. runs/.../.hydra/config.yaml).
# When set, the simulation script can override
# morphology/arena/environment/architecture to match.
trained_config_path: Optional[str] = None

View file

@ -6,6 +6,7 @@ This logger ensures all experimental data is preserved by writing to:
3. stdout (for real-time monitoring) 3. stdout (for real-time monitoring)
""" """
from enum import Enum
import logging import logging
import yaml import yaml
import sys import sys
@ -25,6 +26,33 @@ _active_logger: Optional[Any] = None
_proxy_instance: Optional["LoggerProxy"] = None _proxy_instance: Optional["LoggerProxy"] = None
def _sanitize_for_yaml(obj: Any) -> Any:
"""Convert non-primitive values into YAML-safe structures.
In particular, avoids PyYAML serializing Enums as
``!!python/object/apply:...`` which OmegaConf will not load.
"""
if isinstance(obj, Enum):
return obj.name
if isinstance(obj, Path):
return str(obj)
if isinstance(obj, (np.generic, jnp.ndarray)):
try:
return obj.item()
except Exception:
pass
if isinstance(obj, np.ndarray):
return obj.tolist()
if isinstance(obj, dict):
return {str(k): _sanitize_for_yaml(v) for k, v in obj.items()}
if isinstance(obj, list):
return [_sanitize_for_yaml(v) for v in obj]
if isinstance(obj, tuple):
return [_sanitize_for_yaml(v) for v in obj]
return obj
def get_logger() -> "LoggerProxy": def get_logger() -> "LoggerProxy":
"""Retrieve the global LoggerProxy. """Retrieve the global LoggerProxy.
@ -216,7 +244,13 @@ class UnifiedLogger:
"""Save configuration to disk.""" """Save configuration to disk."""
try: try:
with open(self.config_file, "w") as f: with open(self.config_file, "w") as f:
yaml.dump(self.full_config, f, default_flow_style=False, indent=2, sort_keys=False) yaml.safe_dump(
_sanitize_for_yaml(self.full_config),
f,
default_flow_style=False,
indent=2,
sort_keys=False,
)
self.info(f"Config saved to {self.config_file}") self.info(f"Config saved to {self.config_file}")
except Exception as e: except Exception as e:
self.error(f"Error saving config: {e}") self.error(f"Error saving config: {e}")
@ -296,7 +330,12 @@ class UnifiedLogger:
else: else:
serializable_metric[k] = v serializable_metric[k] = v
f.write("---\n") f.write("---\n")
yaml.dump(serializable_metric, f, default_flow_style=False) yaml.safe_dump(
_sanitize_for_yaml(serializable_metric),
f,
default_flow_style=False,
sort_keys=False,
)
self.metrics_buffer.clear() self.metrics_buffer.clear()
except Exception as e: except Exception as e:
self.error(f"Error flushing metrics: {e}") self.error(f"Error flushing metrics: {e}")
@ -321,7 +360,13 @@ class UnifiedLogger:
if metadata: if metadata:
metadata_path = self.checkpoints_dir / f"{prefix}_step_{step}_metadata.yaml" metadata_path = self.checkpoints_dir / f"{prefix}_step_{step}_metadata.yaml"
with open(metadata_path, "w") as f: with open(metadata_path, "w") as f:
yaml.dump(metadata, f, default_flow_style=False) yaml.safe_dump(
_sanitize_for_yaml(metadata),
f,
default_flow_style=False,
indent=2,
sort_keys=False,
)
self.info(f"Checkpoint saved: {checkpoint_path}") self.info(f"Checkpoint saved: {checkpoint_path}")
@ -357,7 +402,13 @@ class UnifiedLogger:
if metadata: if metadata:
metadata_path = self.run_dir / "final_model_metadata.yaml" metadata_path = self.run_dir / "final_model_metadata.yaml"
with open(metadata_path, "w") as f: with open(metadata_path, "w") as f:
yaml.dump(metadata, f, default_flow_style=False) yaml.safe_dump(
_sanitize_for_yaml(metadata),
f,
default_flow_style=False,
indent=2,
sort_keys=False,
)
self.info(f"Final model saved: {final_model_path}") self.info(f"Final model saved: {final_model_path}")