1
Fork 0

feat(simulate): metadata config loading

This commit is contained in:
Tibo De Peuter 2026-04-28 13:17:24 +02:00
parent a3a1f3643f
commit 8f2a5d25ed
Signed by: tdpeuter
SSH key fingerprint: SHA256:u/h/LVoqKF1Iz02uOyxe6hcjmoZASCGV2HM0TG9ZMoU
5 changed files with 190 additions and 476 deletions

View file

@ -9,7 +9,7 @@ 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. # Optional: override morphology for amputation experiments.
# When set, scripts/simulate.py will use it to default morphology/arena/environment/architecture # Points to a morphology config YAML file (e.g. configs/morphology/3_arms.yaml).
# to match training (unless you explicitly override those keys via CLI). # If null, the training morphology from the model's metadata is used.
trained_config_path: null morphology_override: null

View file

@ -1,11 +1,10 @@
"""Simulate a trained policy in the MuJoCo viewer. """Simulate a trained policy in the MuJoCo viewer.
Uses Hydra to load the same BrittleStarConfig that was used during training. Automatically extracts the training configuration (morphology, environment, etc.)
Override settings via CLI, e.g.: from the sidecar metadata YAML file to ensure simulation perfectly matches training.
python scripts/simulate.py morphology=3_arms Override simulation settings via CLI, e.g.:
uv run scripts/simulate.py \
To replay a run using the *exact* Hydra config used during training, pass: simulation.morphology_override=config/morphology/3_arms.yaml \
python scripts/simulate.py simulation.trained_config_path=runs/.../.hydra/config.yaml \
simulation.model_path=runs/.../final_model.flax simulation.model_path=runs/.../final_model.flax
""" """
@ -22,85 +21,24 @@ import jax
import jax.numpy as jnp import jax.numpy as jnp
import numpy as np import numpy as np
import yaml import yaml
from omegaconf import DictConfig, OmegaConf, open_dict from omegaconf import DictConfig, OmegaConf
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 ObservationBoundsConfig from brittle_star_project.environment.env_config import (
MorphologyConfig,
ArenaConfig,
EnvConfig,
ObservationBoundsConfig,
)
def _dense_layer_sizes_from_params(params: Any) -> list[int]: class PolicyAgent:
"""Infer GenericDenseLayersWithActivation.layer_sizes from a Flax params tree.""" """Wraps a trained Flax actor for deterministic inference."""
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)
# A minimal policy class to load a CleanRL/Flax checkpoint and run inference.
class CleanRLPPOPolicy:
def __init__( def __init__(
self, self,
*, *,
@ -111,7 +49,28 @@ class CleanRLPPOPolicy:
) -> None: ) -> None:
from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation
layer_sizes = _dense_layer_sizes_from_params(sensor_params) # 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._sensor = GenericDenseLayersWithActivation(layer_sizes=layer_sizes)
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)
@ -128,122 +87,31 @@ class CleanRLPPOPolicy:
*, *,
action_dim: int, action_dim: int,
obs_processor: Any, obs_processor: Any,
) -> "CleanRLPPOPolicy": ) -> "PolicyAgent":
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, Any]:
"""Extract checkpoint parts.
Returns (config_dict, sensor_params, actor_params, critic_params,
feature_extractor_params).
PPOTrainer saves:
flax.serialization.to_bytes([
config_dict,
[sensor_params, actor_params, critic_params, feature_extractor_params],
])
msgpack_restore() may restore lists as dicts keyed by string indices
("0", "1", ...), so we accept both shapes.
"""
cfg_part: Any | None = None
params_part: Any = restored_obj
if isinstance(restored_obj, (list, tuple)) and len(restored_obj) >= 2:
cfg_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
):
cfg_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):
sensor_params = _get_index(params_part, 0)
actor_params = _get_index(params_part, 1)
critic_params = _get_index(params_part, 2)
feature_extractor_params = _get_index(params_part, 3)
if sensor_params is None or actor_params is None:
raise ValueError("Missing required params in checkpoint")
return (
cfg_part,
sensor_params,
actor_params,
critic_params,
feature_extractor_params,
)
if isinstance(params_part, (list, tuple)) and len(params_part) >= 2:
sensor_params = params_part[0]
actor_params = params_part[1]
critic_params = params_part[2] if len(params_part) >= 3 else None
feature_extractor_params = params_part[3] if len(params_part) >= 4 else None
return (
cfg_part,
sensor_params,
actor_params,
critic_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(
f"Unexpected checkpoint structure in {path}. "
"Expected [config_dict, [sensor_params, actor_params, critic_params, "
"feature_extractor_params]] or an equivalent dict-indexed variant."
)
payload = path.read_bytes() payload = path.read_bytes()
restored = flax.serialization.msgpack_restore(payload) restored = flax.serialization.msgpack_restore(payload)
_cfg_dict, sensor_params, actor_params, _critic_params, _feature_extractor_params = (
_parse_checkpoint(restored)
)
ckpt_action_dim = _infer_action_dim_from_actor_params(actor_params) sensor_params = None
if ckpt_action_dim is not None and ckpt_action_dim != action_dim: actor_params = None
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( # 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, sensor_params=sensor_params,
actor_params=actor_params, actor_params=actor_params,
action_dim=action_dim, action_dim=action_dim,
@ -273,23 +141,30 @@ def _target_reached(*, state: Any) -> bool:
return bool(getattr(state, "terminated", False) or getattr(state, "truncated", False)) return bool(getattr(state, "terminated", False) or getattr(state, "truncated", False))
def _rollout_one_episode_headless( 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, env: BrittleStarEnv,
policy: CleanRLPPOPolicy, policy: PolicyAgent,
seed: int, seed: int,
max_steps: int, max_steps: int,
action_low: np.ndarray | None, action_low: np.ndarray | None,
action_high: np.ndarray | None, action_high: np.ndarray | None,
action_mask: np.ndarray | None = None,
) -> tuple[float, int, bool, float | None]: ) -> tuple[float, int, bool, float | None]:
"""Run one rollout up to max_steps.
Returns (return, length, reached_target, final_xy_dist).
Note: In the MJC backend, the raw env reward can be 0.0; we compute a simple
progress reward based on xy_distance_to_target.
"""
state = env.reset(seed=seed) state = env.reset(seed=seed)
ep_return = 0.0 ep_return = 0.0
@ -302,12 +177,10 @@ def _rollout_one_episode_headless(
obs_dict = observations or {} obs_dict = observations or {}
action = policy.act(observations=obs_dict) 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) action = _maybe_clip_action(action, action_low, action_high)
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) state = env.step(state=state, action=action)
steps += 1 steps += 1
@ -325,23 +198,23 @@ def _rollout_one_episode_headless(
return ep_return, steps, reached_target, final_dist return ep_return, steps, reached_target, final_dist
def _run_one_episode_viewer( def _rollout_viewer(
*, *,
env: BrittleStarEnv, env: BrittleStarEnv,
policy: CleanRLPPOPolicy, policy: PolicyAgent,
seed: int, seed: int,
state: Any, state: Any,
control_dt: float, control_dt: float,
max_steps: int | None, max_steps: int | None,
action_low: np.ndarray | None, action_low: np.ndarray | None,
action_high: np.ndarray | None, action_high: np.ndarray | None,
action_mask: np.ndarray | None = None,
) -> None: ) -> None:
import mujoco.viewer import mujoco.viewer
model = state.mj_model model = state.mj_model
data = state.mj_data data = state.mj_data
_ = int(seed)
episode_return = 0.0 episode_return = 0.0
observations = _get_observations(state) observations = _get_observations(state)
prev_dist = _get_xy_distance_to_target(observations) prev_dist = _get_xy_distance_to_target(observations)
@ -359,11 +232,9 @@ def _run_one_episode_viewer(
obs_dict = observations or {} obs_dict = observations or {}
action = policy.act(observations=obs_dict) 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) action = _maybe_clip_action(action, action_low, action_high)
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. # The passive viewer runs a GUI thread; protect MuJoCo state mutation.
with viewer.lock(): with viewer.lock():
@ -398,199 +269,117 @@ def _run_one_episode_viewer(
) )
def _infer_checkpoint_obs_dim(policy: CleanRLPPOPolicy) -> int | None: def _load_metadata_yaml(model_path: Path) -> dict:
"""Best-effort read of the first Dense kernel input dim (obs dim).""" """Discover and load the sidecar metadata YAML file."""
metadata_path = model_path.with_name(model_path.stem + "_metadata.yaml")
try: if not metadata_path.exists():
kernel = policy._params["sensor_params"]["params"]["Dense_0"]["kernel"] raise FileNotFoundError(
return int(getattr(kernel, "shape")[0]) f"Could not find metadata YAML for {model_path.name}. Expected it at {metadata_path}"
except Exception:
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
) )
with open(metadata_path, "r") as f:
try: return yaml.safe_load(f)
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:
# Compose against the structured schema first, so missing keys are validated. # 1. Hydra composes ONLY SimulationSettings
cfg = OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg) cfg = OmegaConf.to_object(OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg))
sim_cfg = cfg.simulation
# Optional: override env-defining sections (morphology/arena/environment/architecture) model_path_str = sim_cfg.model_path
# 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)
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
if model_path_str is None: if model_path_str is None:
raise ValueError( raise ValueError(
"simulation.model_path must be set to a .flax checkpoint (e.g. final_model.flax)" "simulation.model_path must be set to a .flax checkpoint (e.g. final_model.flax)"
) )
# Hydra chdir changes CWD; resolve relative paths relative to the invocation.
model_path = Path(hydra.utils.to_absolute_path(model_path_str)) model_path = Path(hydra.utils.to_absolute_path(model_path_str))
if model_path.suffix != ".flax": if model_path.suffix != ".flax":
raise ValueError(f"Expected a '.flax' checkpoint, got '{model_path.name}'.") raise ValueError(f"Expected a '.flax' checkpoint, got '{model_path.name}'.")
# ======= ENVIRONMENT SETUP ======= # 2. Discover + load sidecar metadata YAML
metadata = _load_metadata_yaml(model_path)
# 3. Reconstruct typed configs from metadata
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", {})
)
)
# 4. Determine environment morphology
if sim_cfg.morphology_override is not None:
override_path = Path(hydra.utils.to_absolute_path(sim_cfg.morphology_override))
if not override_path.exists():
raise FileNotFoundError(f"Could not find morphology override YAML at {override_path}")
with open(override_path, "r") as f:
override_dict = yaml.safe_load(f)
env_morphology = OmegaConf.to_object(
OmegaConf.merge(OmegaConf.structured(MorphologyConfig), override_dict)
)
else:
env_morphology = trained_morphology
# 5. Build obs_processor with TRAINING morphology padding masks always
padding_masks = compute_padding_masks(
segments_per_arm=env_morphology.segments_per_arm,
)
obs_processor = create_obs_processor(
bounds_dict=trained_obs_bounds.to_bounds_dict(),
padding_masks=padding_masks,
)
# 6. Build environment
backend = Backend.MJC
seed = int(cfg.experiment.seed)
factory = BrittleStarEnvFactory() factory = BrittleStarEnvFactory()
raw_env = factory.create_environment( raw_env = factory.create_environment(
backend, backend,
config.morphology, env_morphology,
config.arena, trained_arena,
config.environment, trained_environment,
) )
env = BrittleStarEnv( env = BrittleStarEnv(
raw_env, raw_env,
backend=backend, backend=backend,
config=config.environment, config=trained_environment,
morphology_config=config.morphology, morphology_config=env_morphology,
) )
state0 = env.reset(seed=seed) state0 = env.reset(seed=seed)
# Match training's padded observation layout for amputated morphologies. # Calculate the action dimension the model was trained with
padding_masks = compute_padding_masks(config.morphology.segments_per_arm) trained_action_dim = sum(trained_morphology.segments_per_arm) * 2
training_bounds = ObservationBoundsConfig().to_bounds_dict() # 7. Load policy
if trained_cfg_path and "obs_bounds" in trained_cfg: policy = PolicyAgent.load(
try: model_path, action_dim=trained_action_dim, obs_processor=obs_processor
training_bounds = OmegaConf.to_object(trained_cfg.obs_bounds).to_bounds_dict()
except Exception:
pass
else:
training_bounds = config.obs_bounds.to_bounds_dict()
obs_processor = create_obs_processor(
bounds_dict=training_bounds,
padding_masks=padding_masks,
) )
# ======= MODEL SETUP ======= # Convert the JAX boolean mask to a numpy array for easy indexing
nu = int(state0.mj_model.nu) action_mask = np.asarray(padding_masks["mask_2x"])
policy = CleanRLPPOPolicy.load(model_path, action_dim=nu, obs_processor=obs_processor)
# Helpful early failure when configs don't match the checkpoint.
observations0 = _get_observations(state0)
obs0_dict = observations0 or {}
batched_obs0 = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], obs0_dict)
env_obs_dim = int(obs_processor(batched_obs0).shape[1])
ckpt_obs_dim = _infer_checkpoint_obs_dim(policy)
if ckpt_obs_dim is not None and ckpt_obs_dim != env_obs_dim:
raise ValueError(
"Checkpoint/env mismatch: "
f"checkpoint expects obs_dim={ckpt_obs_dim}, env provides obs_dim={env_obs_dim}. "
"Use the same Hydra config (morphology/arena/environment) "
"that was used during training."
)
# Match training's action clipping behavior. # Match training's action clipping behavior.
action_space = getattr(raw_env, "action_space", None) action_space = getattr(raw_env, "action_space", None)
@ -601,24 +390,26 @@ def main(dict_cfg: DictConfig) -> None:
None if action_space is None else np.asarray(action_space.high, dtype=np.float32).ravel() None if action_space is None else np.asarray(action_space.high, dtype=np.float32).ravel()
) )
# ======= SIMULATION ======= # 8. Run simulation
headless = bool(config.simulation.headless) headless = bool(sim_cfg.headless)
max_steps = config.simulation.max_steps max_steps = sim_cfg.max_steps
if headless: if headless:
if max_steps is None: if max_steps is None:
raise ValueError("simulation.max_steps is required when simulation.headless=true") raise ValueError("simulation.max_steps is required when simulation.headless=true")
max_steps_i = int(max_steps) max_steps_i = int(max_steps)
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_one_episode_headless( ep_return, ep_len, reached_target, final_dist = _rollout_headless(
env=env, env=env,
policy=policy, policy=policy,
seed=seed, seed=seed,
max_steps=max_steps_i, max_steps=max_steps_i,
action_low=action_low, action_low=action_low,
action_high=action_high, action_high=action_high,
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 final_dist is None else f"{final_dist:.3f}"
print( print(
@ -627,18 +418,17 @@ def main(dict_cfg: DictConfig) -> None:
f"target_reached={reached_target}, final_xy_dist={final_dist_str}" f"target_reached={reached_target}, final_xy_dist={final_dist_str}"
) )
else: else:
max_steps_val = None
if max_steps is not None: if max_steps is not None:
max_steps_i = int(max_steps) max_steps_i = int(max_steps)
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")
max_steps_val: int | None = max_steps_i max_steps_val = max_steps_i
else:
max_steps_val = None
model_dt = float(state0.mj_model.opt.timestep) model_dt = float(state0.mj_model.opt.timestep)
control_dt = model_dt * float(config.environment.num_physics_steps_per_control_step) control_dt = model_dt * float(trained_environment.num_physics_steps_per_control_step)
_run_one_episode_viewer( _rollout_viewer(
env=env, env=env,
policy=policy, policy=policy,
seed=seed, seed=seed,
@ -647,6 +437,7 @@ def main(dict_cfg: DictConfig) -> None:
max_steps=max_steps_val, max_steps=max_steps_val,
action_low=action_low, action_low=action_low,
action_high=action_high, action_high=action_high,
action_mask=action_mask,
) )
env.close() env.close()

View file

@ -13,7 +13,9 @@ class SimulationSettings:
# 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). # Override morphology for amputation experiments.
# When set, the simulation script can override # When set, the environment uses this morphology instead of the trained one.
# morphology/arena/environment/architecture to match. # Points to a morphology config YAML file (e.g. configs/morphology/3_arms.yaml).
trained_config_path: Optional[str] = None # Observations are padded from the override morphology UP TO the training
# morphology's shape via compute_padding_masks(override, reference=training).
morphology_override: Optional[str] = None

View file

@ -1,7 +1,9 @@
from .env_config import ArenaConfig, EnvConfig, MorphologyConfig from .env_config import ArenaConfig, EnvConfig, MorphologyConfig
from .env_types import Backend, Task from .env_types import Backend, Task
from .env_wrapper import BrittleStarEnv, StepResult from .env_wrapper import BrittleStarEnv
from .factory import BrittleStarEnvFactory from .factory import BrittleStarEnvFactory
from .obs_processing import create_obs_processor
from .padded_obs_wrapper import compute_padding_masks
__all__ = [ __all__ = [
"ArenaConfig", "ArenaConfig",
@ -10,6 +12,7 @@ __all__ = [
"Backend", "Backend",
"Task", "Task",
"BrittleStarEnv", "BrittleStarEnv",
"StepResult",
"BrittleStarEnvFactory", "BrittleStarEnvFactory",
"create_obs_processor",
"compute_padding_masks",
] ]

View file

@ -1,35 +1,11 @@
"""Observation padding wrapper for amputated brittle star morphologies. """Observation padding masks for amputated brittle star morphologies."""
When using a centralized controller, the global observation vector must remain
a constant size regardless of how many segments are amputated. This wrapper pads
the observation dictionary values with zeros using spatial insertion so that the
flattened observation maintains the correct physical mapping to the neural network.
"""
from __future__ import annotations from __future__ import annotations
from typing import Any, Sequence from typing import Any, Sequence
import jax
import jax.numpy as jnp import jax.numpy as jnp
# Observation keys whose size scales with the number of joints (2 per segment).
_JOINT_SCALED_KEYS = frozenset(
{
"joint_position",
"joint_velocity",
"joint_actuator_force",
"actuator_force",
}
)
# Observation keys whose size scales with the number of segments (1 per segment).
_SEGMENT_SCALED_KEYS = frozenset(
{
"segment_contact",
}
)
def compute_padding_masks( def compute_padding_masks(
segments_per_arm: Sequence[int], segments_per_arm: Sequence[int],
@ -60,7 +36,6 @@ def compute_padding_masks(
f"actual segments ({actual}) must be between 0 and reference ({ref})." f"actual segments ({actual}) must be between 0 and reference ({ref})."
) )
# 1x scaling (e.g., contacts: 1 value per segment) # 1x scaling (e.g., contacts: 1 value per segment)
# 1x scaling (e.g., contacts: 1 value per segment)
mask_1x.extend([True] * actual + [False] * (ref - actual)) mask_1x.extend([True] * actual + [False] * (ref - actual))
# 2x scaling (e.g., joints: 2 values per segment) # 2x scaling (e.g., joints: 2 values per segment)
mask_2x.extend([True] * (actual * 2) + [False] * ((ref - actual) * 2)) mask_2x.extend([True] * (actual * 2) + [False] * ((ref - actual) * 2))
@ -71,60 +46,3 @@ def compute_padding_masks(
"target_size_1x": sum(reference_segments_per_arm), "target_size_1x": sum(reference_segments_per_arm),
"target_size_2x": sum(reference_segments_per_arm) * 2, "target_size_2x": sum(reference_segments_per_arm) * 2,
} }
def pad_observation(
obs: dict[str, Any],
masks: dict[str, Any],
) -> dict[str, Any]:
"""Pad an observation dict using spatial insertion."""
padded = {}
for key, value in obs.items():
padded_dtype = _padding_dtype(value)
if key in _JOINT_SCALED_KEYS:
out = jnp.zeros(masks["target_size_2x"], dtype=padded_dtype)
padded[key] = out.at[masks["mask_2x"]].set(value)
elif key in _SEGMENT_SCALED_KEYS:
out = jnp.zeros(masks["target_size_1x"], dtype=padded_dtype)
padded[key] = out.at[masks["mask_1x"]].set(value)
else:
padded[key] = value
return padded
def pad_observations_batched(
obs: dict[str, Any],
masks: dict[str, Any],
) -> dict[str, Any]:
"""Pad a batched observation dict (leading batch dimension) using spatial insertion."""
padded = {}
for key, value in obs.items():
batch_size = value.shape[0]
padded_dtype = _padding_dtype(value)
if key in _JOINT_SCALED_KEYS:
out = jnp.zeros((batch_size, masks["target_size_2x"]), dtype=padded_dtype)
padded[key] = out.at[:, masks["mask_2x"]].set(value)
elif key in _SEGMENT_SCALED_KEYS:
out = jnp.zeros((batch_size, masks["target_size_1x"]), dtype=padded_dtype)
padded[key] = out.at[:, masks["mask_1x"]].set(value)
else:
padded[key] = value
return padded
def _padding_dtype(value: Any) -> jnp.dtype:
"""Choose a JAX-safe dtype for padding arrays.
When JAX x64 is disabled, allocating float64 zeros emits a warning. We
preserve the original dtype whenever it is supported, and otherwise fall
back to float32 for padding buffers.
"""
dtype = getattr(value, "dtype", None)
if dtype is None:
dtype = jnp.asarray(value).dtype
else:
dtype = jnp.dtype(dtype)
if dtype == jnp.float64 and not jax.config.read("jax_enable_x64"):
return jnp.float32
return dtype