diff --git a/configs/simulation/default.yaml b/configs/simulation/default.yaml index 8645f00..84694e0 100644 --- a/configs/simulation/default.yaml +++ b/configs/simulation/default.yaml @@ -4,13 +4,12 @@ # Path to the trained model (optional) 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 headless: false # In headless mode this is required; in viewer mode null means "infinite". 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 diff --git a/scripts/simulate.py b/scripts/simulate.py index 28660b9..a9a51f4 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -3,6 +3,10 @@ Uses Hydra to load the same BrittleStarConfig that was used during training. Override settings via CLI, e.g.: 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 @@ -17,11 +21,16 @@ import hydra import jax import jax.numpy as jnp 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.register_configs import register_configs +from brittle_star_project.environment.padded_obs_wrapper import ( + compute_padding_masks, + pad_observation, +) _ALLOWED_OBS_KEYS = { "joint_position", @@ -36,6 +45,74 @@ _ALLOWED_OBS_KEYS = { "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: """Flatten the env's observation dict into a 1D vector. @@ -69,9 +146,8 @@ class CleanRLPPOPolicy: ) -> None: from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation - hidden_dim = int(sensor_params["params"]["Dense_0"]["kernel"].shape[1]) - - self._sensor = GenericDenseLayersWithActivation(layer_sizes=[hidden_dim, hidden_dim]) + layer_sizes = _dense_layer_sizes_from_params(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) @@ -156,6 +232,28 @@ class CleanRLPPOPolicy: 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, " @@ -168,6 +266,16 @@ class CleanRLPPOPolicy: _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( sensor_params=sensor_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: - return bool(getattr(state, "terminated", False)) + return bool(getattr(state, "terminated", False) or getattr(state, "truncated", False)) def _rollout_one_episode_headless( @@ -202,6 +310,9 @@ def _rollout_one_episode_headless( policy: CleanRLPPOPolicy, seed: 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]: """Run one rollout up to max_steps. @@ -220,7 +331,12 @@ def _rollout_one_episode_headless( steps = 0 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) if nu > 0 and action.shape != (nu,): @@ -251,6 +367,9 @@ def _run_one_episode_viewer( state: Any, control_dt: float, max_steps: int | None, + action_low: np.ndarray | None, + action_high: np.ndarray | None, + padding_masks: dict[str, Any] | None, ) -> None: import mujoco.viewer @@ -272,7 +391,12 @@ def _run_one_episode_viewer( break 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),): raise ValueError( 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 +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") def main(dict_cfg: DictConfig) -> None: - # Convert DictConfig to structured dataclass, ensuring the root schema is applied. - config: BrittleStarConfig = OmegaConf.to_object( - OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg) - ) + # Compose against the structured schema first, so missing keys are validated. + 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) + 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: raise ValueError( @@ -359,15 +593,28 @@ def main(dict_cfg: DictConfig) -> None: 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 ======= nu = int(state0.mj_model.nu) policy = CleanRLPPOPolicy.load(model_path, action_dim=nu) # Helpful early failure when configs don't match the checkpoint. 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) - + if ckpt_obs_dim is not None and ckpt_obs_dim != env_obs_dim: raise ValueError( "Checkpoint/env mismatch: " @@ -375,7 +622,7 @@ def main(dict_cfg: DictConfig) -> None: "Use the same Hydra config (morphology/arena/environment) " "that was used during training." ) - + # ======= SIMULATION ======= headless = bool(config.simulation.headless) max_steps = config.simulation.max_steps @@ -392,6 +639,9 @@ def main(dict_cfg: DictConfig) -> None: policy=policy, seed=seed, 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}" print( @@ -418,6 +668,9 @@ def main(dict_cfg: DictConfig) -> None: state=state0, control_dt=control_dt, max_steps=max_steps_val, + action_low=action_low, + action_high=action_high, + padding_masks=padding_masks, ) env.close() diff --git a/src/brittle_star_project/configs/config_simulation.py b/src/brittle_star_project/configs/config_simulation.py index 00488fc..872ca20 100644 --- a/src/brittle_star_project/configs/config_simulation.py +++ b/src/brittle_star_project/configs/config_simulation.py @@ -1,18 +1,19 @@ from dataclasses import dataclass from typing import Optional -from brittle_star_project.environment.env_types import Backend - @dataclass class SimulationSettings: """Settings for the simulation script.""" model_path: Optional[str] = None - model_type: str = "random" - backend: Backend = Backend.MJC # Script behavior headless: bool = False # If None, viewer mode runs until window closed or target reached. 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 diff --git a/src/experiment_logger/unified_logger.py b/src/experiment_logger/unified_logger.py index 634638d..f37136b 100644 --- a/src/experiment_logger/unified_logger.py +++ b/src/experiment_logger/unified_logger.py @@ -6,6 +6,7 @@ This logger ensures all experimental data is preserved by writing to: 3. stdout (for real-time monitoring) """ +from enum import Enum import logging import yaml import sys @@ -25,6 +26,33 @@ _active_logger: Optional[Any] = 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": """Retrieve the global LoggerProxy. @@ -216,7 +244,13 @@ class UnifiedLogger: """Save configuration to disk.""" try: 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}") except Exception as e: self.error(f"Error saving config: {e}") @@ -296,7 +330,12 @@ class UnifiedLogger: else: serializable_metric[k] = v 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() except Exception as e: self.error(f"Error flushing metrics: {e}") @@ -321,7 +360,13 @@ class UnifiedLogger: if metadata: metadata_path = self.checkpoints_dir / f"{prefix}_step_{step}_metadata.yaml" 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}") @@ -357,7 +402,13 @@ class UnifiedLogger: if metadata: metadata_path = self.run_dir / "final_model_metadata.yaml" 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}")