fix: adapted checkpoint eval to handle decentralized models
This commit is contained in:
parent
2594c53d49
commit
d0e40db648
2 changed files with 15 additions and 4 deletions
|
|
@ -52,6 +52,7 @@ def build_eval_rollout_fn(
|
||||||
obs_processor: Callable,
|
obs_processor: Callable,
|
||||||
sensor_apply: Callable,
|
sensor_apply: Callable,
|
||||||
actor_apply: Callable,
|
actor_apply: Callable,
|
||||||
|
message_passer_apply: Callable | None = None,
|
||||||
action_low: jnp.ndarray,
|
action_low: jnp.ndarray,
|
||||||
action_high: jnp.ndarray,
|
action_high: jnp.ndarray,
|
||||||
reward_fn: Callable,
|
reward_fn: Callable,
|
||||||
|
|
@ -67,6 +68,9 @@ def build_eval_rollout_fn(
|
||||||
returned by ``create_obs_processor``.
|
returned by ``create_obs_processor``.
|
||||||
sensor_apply: The sensor network's ``apply`` method (JIT-compiled).
|
sensor_apply: The sensor network's ``apply`` method (JIT-compiled).
|
||||||
actor_apply: The actor network's ``apply`` method (JIT-compiled).
|
actor_apply: The actor network's ``apply`` method (JIT-compiled).
|
||||||
|
message_passer_apply: Optional message-passing module apply method.
|
||||||
|
When provided, it is applied between the sensor and actor, using
|
||||||
|
``params["message_passer_params"]``.
|
||||||
action_low: Per-joint action lower bound (JAX array, shape ``(action_dim,)``).
|
action_low: Per-joint action lower bound (JAX array, shape ``(action_dim,)``).
|
||||||
action_high: Per-joint action upper bound (JAX array, shape ``(action_dim,)``).
|
action_high: Per-joint action upper bound (JAX array, shape ``(action_dim,)``).
|
||||||
reward_fn: Shaped reward function with signature
|
reward_fn: Shaped reward function with signature
|
||||||
|
|
@ -101,10 +105,14 @@ def build_eval_rollout_fn(
|
||||||
|
|
||||||
obs = obs_processor(state.observations)
|
obs = obs_processor(state.observations)
|
||||||
hidden = sensor_apply(params["sensor_params"], obs)
|
hidden = sensor_apply(params["sensor_params"], obs)
|
||||||
|
if message_passer_apply is not None:
|
||||||
|
mp_params = params["message_passer_params"]
|
||||||
|
hidden = jax.vmap(lambda x: message_passer_apply(mp_params, x))(hidden)
|
||||||
mean, _log_std = actor_apply(params["actor_params"], hidden)
|
mean, _log_std = actor_apply(params["actor_params"], hidden)
|
||||||
|
|
||||||
# Deterministic action: use the actor mean, no exploration noise.
|
# Deterministic action: use the actor mean, no exploration noise.
|
||||||
action = jnp.clip(mean, action_low, action_high)
|
flat_mean = mean.reshape(mean.shape[0], -1)
|
||||||
|
action = jnp.clip(flat_mean, action_low, action_high)
|
||||||
next_state = step_1(state=state, action=action)
|
next_state = step_1(state=state, action=action)
|
||||||
|
|
||||||
shaped_reward = reward_fn(state, next_state)
|
shaped_reward = reward_fn(state, next_state)
|
||||||
|
|
|
||||||
|
|
@ -30,8 +30,8 @@ from brittle_star_project.MLPs.mlps import (
|
||||||
MessagePasser,
|
MessagePasser,
|
||||||
OneDenseLayerMLP,
|
OneDenseLayerMLP,
|
||||||
Storage,
|
Storage,
|
||||||
build_adjacency,
|
|
||||||
)
|
)
|
||||||
|
from brittle_star_project.MLPs.adjancency_builder import build_adjacency
|
||||||
from brittle_star_project.ppo import PPO
|
from brittle_star_project.ppo import PPO
|
||||||
from brittle_star_project.environment import MorphMode
|
from brittle_star_project.environment import MorphMode
|
||||||
from brittle_star_project.utils import logged_jit
|
from brittle_star_project.utils import logged_jit
|
||||||
|
|
@ -901,8 +901,11 @@ class PPOTrainer:
|
||||||
self._eval_fn = build_eval_rollout_fn(
|
self._eval_fn = build_eval_rollout_fn(
|
||||||
env=self.env,
|
env=self.env,
|
||||||
obs_processor=self.obs_processor,
|
obs_processor=self.obs_processor,
|
||||||
sensor_apply=self.sensor.apply,
|
sensor_apply=lambda p, x: apply_per_node(self.sensor, p, x),
|
||||||
actor_apply=self.actor.apply,
|
actor_apply=lambda p, x: apply_per_node(self.actor, p, x),
|
||||||
|
message_passer_apply=(
|
||||||
|
None if self.message_passer is None else self.message_passer.apply
|
||||||
|
),
|
||||||
action_low=self._action_low,
|
action_low=self._action_low,
|
||||||
action_high=self._action_high,
|
action_high=self._action_high,
|
||||||
reward_fn=reward_fn,
|
reward_fn=reward_fn,
|
||||||
|
|
|
||||||
Reference in a new issue