diff --git a/scripts/simulate.py b/scripts/simulate.py index 5564aa8..1a4aad2 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -34,6 +34,7 @@ from brittle_star_project.evaluation.video import ( create_evaluation_dir, save_evaluation_metadata, ) +from brittle_star_project.MLPs.adjancency_builder import build_adjacency @hydra.main(config_path="../configs", config_name="main_config", version_base="1.3") @@ -134,8 +135,23 @@ def main(dict_cfg: DictConfig) -> None: trained_action_dim = raw_env.action_space.shape[0] // needed_copies # 7. Load policy + message_passing_steps = ( + metadata.get("architecture", {}) or {} + ).get("message_passing_steps") + if message_passing_steps is None: + message_passing_steps = 4 + message_passing_steps = int(message_passing_steps) + + adj_matrix = None + if env_morphology.morph_mode != MorphMode.CENTRALIZED: + adj_matrix = build_adjacency(env_morphology.segments_per_arm, env_morphology.morph_mode) + 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, + message_passing_steps=message_passing_steps, + adj_matrix=adj_matrix, ) # Convert the JAX boolean mask to a numpy array for easy indexing diff --git a/src/brittle_star_project/evaluation/checkpoint.py b/src/brittle_star_project/evaluation/checkpoint.py index d10046b..9b868d5 100644 --- a/src/brittle_star_project/evaluation/checkpoint.py +++ b/src/brittle_star_project/evaluation/checkpoint.py @@ -4,6 +4,8 @@ import yaml from dataclasses import dataclass from pathlib import Path +from collections.abc import Mapping + import flax from omegaconf import OmegaConf @@ -32,15 +34,19 @@ def load_params(path: Path) -> dict: sensor_params = None actor_params = None + message_passer_params = None # Extract params from restored checkpoint - if isinstance(restored, dict): + if isinstance(restored, Mapping): 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") + message_passer_params = restored.get("message_passer_params") or params_sub.get( + "message_passer_params" + ) elif isinstance(restored, (list, tuple)) and len(restored) >= 2: params_part = restored[1] - if isinstance(params_part, dict): + if isinstance(params_part, Mapping): 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: @@ -53,6 +59,7 @@ def load_params(path: Path) -> dict: return { "sensor_params": sensor_params, "actor_params": actor_params, + "message_passer_params": message_passer_params, } diff --git a/src/brittle_star_project/evaluation/policy.py b/src/brittle_star_project/evaluation/policy.py index 6161d90..5b12d07 100644 --- a/src/brittle_star_project/evaluation/policy.py +++ b/src/brittle_star_project/evaluation/policy.py @@ -24,10 +24,17 @@ class PolicyAgent: *, sensor_params: Any, actor_params: Any, + message_passer_params: Any | None = None, + message_passing_steps: int | None = None, + adj_matrix: Any | None = None, action_dim: int, obs_processor: Any, ) -> None: - from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation + from brittle_star_project.MLPs.mlps import ( + Actor, + GenericDenseLayersWithActivation, + MessagePasser, + ) # Infer layer sizes from params try: @@ -54,11 +61,31 @@ class PolicyAgent: self._sensor = GenericDenseLayersWithActivation(layer_sizes=layer_sizes) self._actor = Actor(action_dim=action_dim) + + self._message_passer = None + if message_passer_params is not None and not ( + isinstance(message_passer_params, dict) and len(message_passer_params) == 0 + ): + if message_passing_steps is None or adj_matrix is None: + raise ValueError( + "Checkpoint contains message_passer_params but PolicyAgent was not given " + "message_passing_steps and adj_matrix. Pass these when constructing the agent " + "so decentralized evaluation matches training." + ) + + hidden_dim = int(layer_sizes[-1]) + self._message_passer = MessagePasser( + hidden_dim=hidden_dim, + num_propagation_steps=int(message_passing_steps), + adj_matrix=jnp.asarray(adj_matrix), + ) + self._message_passer.apply = jax.jit(self._message_passer.apply) 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, + "message_passer_params": message_passer_params, } self._obs_processor = obs_processor @@ -68,6 +95,9 @@ class PolicyAgent: *, sensor_params: Any, actor_params: Any, + message_passer_params: Any | None = None, + message_passing_steps: int | None = None, + adj_matrix: Any | None = None, action_dim: int, obs_processor: Any, ) -> "PolicyAgent": @@ -75,14 +105,24 @@ class PolicyAgent: return cls( sensor_params=sensor_params, actor_params=actor_params, + message_passer_params=message_passer_params, + message_passing_steps=message_passing_steps, + adj_matrix=adj_matrix, action_dim=action_dim, obs_processor=obs_processor, ) - def set_params(self, *, sensor_params: Any, actor_params: Any) -> None: + def set_params( + self, + *, + sensor_params: Any, + actor_params: Any, + message_passer_params: Any | None = None, + ) -> None: """Update parameters for evaluation without rebuilding the model.""" self._params["sensor_params"] = sensor_params self._params["actor_params"] = actor_params + self._params["message_passer_params"] = message_passer_params @classmethod def from_checkpoint( @@ -91,6 +131,8 @@ class PolicyAgent: *, action_dim: int, obs_processor: Any, + message_passing_steps: int | None = None, + adj_matrix: Any | None = None, ) -> "PolicyAgent": """Load params from .flax and construct the agent.""" params = load_params(model_path) @@ -98,6 +140,9 @@ class PolicyAgent: return cls( sensor_params=params["sensor_params"], actor_params=params["actor_params"], + message_passer_params=params.get("message_passer_params"), + message_passing_steps=message_passing_steps, + adj_matrix=adj_matrix, action_dim=action_dim, obs_processor=obs_processor, ) @@ -117,10 +162,16 @@ class PolicyAgent: batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations) obs = self._obs_processor(batched_obs) - # TODO: message passing hidden = self._apply_per_node(self._sensor, self._params["sensor_params"], obs) - # hidden = jax.vmap(...) + if self._message_passer is not None: + mp_params = self._params.get("message_passer_params") + if mp_params is None or (isinstance(mp_params, dict) and len(mp_params) == 0): + raise ValueError( + "PolicyAgent has a message passer but message_passer_params are missing/empty." + ) + hidden = jax.vmap(lambda x: self._message_passer.apply(mp_params, x))(hidden) + mean, _log_std = self._apply_per_node(self._actor, self._params["actor_params"], hidden) return np.asarray(mean, dtype=np.float32).ravel()