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,
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

View file

@ -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,
}

View file

@ -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()