1
Fork 0

fix: resolved todo in checkpoint evaluation regarding message passing

This commit is contained in:
Jona Reynaert 2026-05-11 18:02:58 +02:00
parent 05138bb7dd
commit 8b18a9f550
3 changed files with 81 additions and 7 deletions

View file

@ -34,6 +34,7 @@ from brittle_star_project.evaluation.video import (
create_evaluation_dir, create_evaluation_dir,
save_evaluation_metadata, 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") @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 trained_action_dim = raw_env.action_space.shape[0] // needed_copies
# 7. Load policy # 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( 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 # Convert the JAX boolean mask to a numpy array for easy indexing

View file

@ -4,6 +4,8 @@ import yaml
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from collections.abc import Mapping
import flax import flax
from omegaconf import OmegaConf from omegaconf import OmegaConf
@ -32,15 +34,19 @@ def load_params(path: Path) -> dict:
sensor_params = None sensor_params = None
actor_params = None actor_params = None
message_passer_params = None
# Extract params from restored checkpoint # Extract params from restored checkpoint
if isinstance(restored, dict): if isinstance(restored, Mapping):
params_sub = restored.get("params", {}) params_sub = restored.get("params", {})
sensor_params = restored.get("sensor_params") or params_sub.get("sensor_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") 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: elif isinstance(restored, (list, tuple)) and len(restored) >= 2:
params_part = restored[1] params_part = restored[1]
if isinstance(params_part, dict): if isinstance(params_part, Mapping):
sensor_params = params_part.get("0", params_part.get(0)) sensor_params = params_part.get("0", params_part.get(0))
actor_params = params_part.get("1", params_part.get(1)) actor_params = params_part.get("1", params_part.get(1))
elif isinstance(params_part, (list, tuple)) and len(params_part) >= 2: elif isinstance(params_part, (list, tuple)) and len(params_part) >= 2:
@ -53,6 +59,7 @@ def load_params(path: Path) -> dict:
return { return {
"sensor_params": sensor_params, "sensor_params": sensor_params,
"actor_params": actor_params, "actor_params": actor_params,
"message_passer_params": message_passer_params,
} }

View file

@ -24,10 +24,17 @@ class PolicyAgent:
*, *,
sensor_params: Any, sensor_params: Any,
actor_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, action_dim: int,
obs_processor: Any, obs_processor: Any,
) -> None: ) -> 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 # Infer layer sizes from params
try: try:
@ -54,11 +61,31 @@ class PolicyAgent:
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._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._sensor.apply = jax.jit(self._sensor.apply)
self._actor.apply = jax.jit(self._actor.apply) self._actor.apply = jax.jit(self._actor.apply)
self._params = { self._params = {
"sensor_params": sensor_params, "sensor_params": sensor_params,
"actor_params": actor_params, "actor_params": actor_params,
"message_passer_params": message_passer_params,
} }
self._obs_processor = obs_processor self._obs_processor = obs_processor
@ -68,6 +95,9 @@ class PolicyAgent:
*, *,
sensor_params: Any, sensor_params: Any,
actor_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, action_dim: int,
obs_processor: Any, obs_processor: Any,
) -> "PolicyAgent": ) -> "PolicyAgent":
@ -75,14 +105,24 @@ class PolicyAgent:
return cls( return cls(
sensor_params=sensor_params, sensor_params=sensor_params,
actor_params=actor_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, action_dim=action_dim,
obs_processor=obs_processor, 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.""" """Update parameters for evaluation without rebuilding the model."""
self._params["sensor_params"] = sensor_params self._params["sensor_params"] = sensor_params
self._params["actor_params"] = actor_params self._params["actor_params"] = actor_params
self._params["message_passer_params"] = message_passer_params
@classmethod @classmethod
def from_checkpoint( def from_checkpoint(
@ -91,6 +131,8 @@ class PolicyAgent:
*, *,
action_dim: int, action_dim: int,
obs_processor: Any, obs_processor: Any,
message_passing_steps: int | None = None,
adj_matrix: Any | None = None,
) -> "PolicyAgent": ) -> "PolicyAgent":
"""Load params from .flax and construct the agent.""" """Load params from .flax and construct the agent."""
params = load_params(model_path) params = load_params(model_path)
@ -98,6 +140,9 @@ class PolicyAgent:
return cls( return cls(
sensor_params=params["sensor_params"], sensor_params=params["sensor_params"],
actor_params=params["actor_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, action_dim=action_dim,
obs_processor=obs_processor, obs_processor=obs_processor,
) )
@ -117,10 +162,16 @@ class PolicyAgent:
batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations) batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations)
obs = self._obs_processor(batched_obs) obs = self._obs_processor(batched_obs)
# TODO: message passing
hidden = self._apply_per_node(self._sensor, self._params["sensor_params"], obs) 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) mean, _log_std = self._apply_per_node(self._actor, self._params["actor_params"], hidden)
return np.asarray(mean, dtype=np.float32).ravel() return np.asarray(mean, dtype=np.float32).ravel()