1
Fork 0

feat(PPOTrainer): decentralized message passing babyyyy

This commit is contained in:
Robin Meersman 2026-05-02 21:50:42 +02:00
parent 9398592e34
commit 6114e36e27
6 changed files with 122 additions and 34 deletions

View file

@ -26,7 +26,7 @@ critic:
activation: "tanh" activation: "tanh"
# Synchronous message-passing rounds per control step # Synchronous message-passing rounds per control step
message_passing_steps: 1 message_passing_steps: 4
# Connectivity topology (e.g., ring, fully_connected) # Connectivity topology (e.g., ring, fully_connected)
topology_type: "ring" topology_type: "ring"

View file

@ -0,0 +1,2 @@
simulation_time: 50000.0
target_distance: 3.0

View file

@ -0,0 +1,6 @@
# Testing chicken dinner 4 but further distance.
exp_name: "long2arm"
seed: 123
torch_deterministic: true
cuda: true

View file

@ -0,0 +1,16 @@
anneal_lr: true
clip_coef: 0.2
clip_vloss: true
ent_coef: 0.001
gae_lambda: 0.95
gamma: 0.99
learning_rate: 0.0001
max_grad_norm: 0.5
norm_adv: true
num_envs: 32
num_minibatches: 32
num_steps: 64
target_kl: 0.02
total_timesteps: 12288000
update_epochs: 4
vf_coef: 1.0

View file

@ -1,6 +1,5 @@
from dataclasses import dataclass, fields, field from dataclasses import dataclass, fields, field
import flax
import flax.linen as nn import flax.linen as nn
import jax.numpy as jnp import jax.numpy as jnp
import jax.tree_util import jax.tree_util
@ -38,6 +37,28 @@ class Actor(nn.Module):
return mean, log_std return mean, log_std
class MessagePasser(nn.Module):
hidden_dim: int
num_propagation_steps: int
@nn.compact
def __call__(self, x: jnp.ndarray, adj_matrix: jnp.ndarray):
for _ in range(self.num_propagation_steps):
messages = nn.Dense(self.hidden_dim)(x)
messages = nn.tanh(messages)
agg = adj_matrix # if mean is wanted: adj_matrix / (adj.sum(axis=-1, keepdims=True) + 1e-8)
aggregated = agg @ messages
x_concat = jnp.concatenate([x, aggregated], axis=-1)
gate = nn.sigmoid(nn.Dense(self.hidden_dim)(x_concat))
candidate = nn.tanh(nn.Dense(self.hidden_dim)(x_concat))
x = gate * x + (1 - gate) * candidate
return x
@jax.tree_util.register_dataclass @jax.tree_util.register_dataclass
@dataclass @dataclass
class AgentParams: class AgentParams:

View file

@ -3,7 +3,7 @@ import random
import time import time
from dataclasses import asdict, dataclass from dataclasses import asdict, dataclass
from functools import partial from functools import partial
from typing import Any from typing import Any, Optional
import jax import jax
import jax.numpy as jnp import jax.numpy as jnp
@ -21,6 +21,7 @@ from brittle_star_project.MLPs.mlps import (
Actor, Actor,
AgentParams, AgentParams,
GenericDenseLayersWithActivation, GenericDenseLayersWithActivation,
MessagePasser,
OneDenseLayerMLP, OneDenseLayerMLP,
Storage, Storage,
) )
@ -211,14 +212,20 @@ def _get_action_and_value_noise(
feature_extractor: nn.Module, feature_extractor: nn.Module,
actor: nn.Module, actor: nn.Module,
critic: nn.Module, critic: nn.Module,
message_passer: Optional[nn.Module],
agent_state: TrainState, agent_state: TrainState,
next_obs: jnp.ndarray, next_obs: jnp.ndarray,
key: jax.random.PRNGKey, key,
action_low, action_low,
action_high, action_high,
adj_matrix: jnp.ndarray, adj_matrix: jnp.ndarray,
): ):
def apply_message_passing(h):
return message_passer.apply(agent_state.params["message_passer_params"], h, adj_matrix)
hidden = apply_per_node(sensor, agent_state.params["sensor_params"], next_obs) hidden = apply_per_node(sensor, agent_state.params["sensor_params"], next_obs)
hidden = jax.vmap(apply_message_passing)(hidden)
hidden_critic = apply_shared( hidden_critic = apply_shared(
feature_extractor, agent_state.params["feature_extractor_params"], next_obs feature_extractor, agent_state.params["feature_extractor_params"], next_obs
) )
@ -230,10 +237,14 @@ def _get_action_and_value_noise(
std = jnp.exp(log_std) std = jnp.exp(log_std)
raw_action = mean + noise * std raw_action = mean + noise * std
flat_action = raw_action.reshape(raw_action.shape[0], -1) # concat the per agent, keep the envs dim flat_action = raw_action.reshape(
raw_action.shape[0], -1
) # concat the per agent, keep the envs dim (batch, agent * action)
flat_clipped_action = _clip_action(flat_action, action_low, action_high) flat_clipped_action = _clip_action(flat_action, action_low, action_high)
logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(axis=(-2, -1)) logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(
axis=(-2, -1)
)
value = apply_shared(critic, agent_state.params["critic_params"], hidden_critic) value = apply_shared(critic, agent_state.params["critic_params"], hidden_critic)
return flat_clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key return flat_clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key
@ -249,7 +260,7 @@ def _step_once(
feature_extractor: nn.Module, feature_extractor: nn.Module,
actor: nn.Module, actor: nn.Module,
critic: nn.Module, critic: nn.Module,
message_passer: nn.Module, message_passer: Optional[nn.Module],
action_low, action_low,
action_high, action_high,
): ):
@ -259,6 +270,7 @@ def _step_once(
feature_extractor, feature_extractor,
actor, actor,
critic, critic,
message_passer,
agent_state, agent_state,
obs, obs,
key, key,
@ -266,23 +278,23 @@ def _step_once(
action_high, action_high,
adj_matrix, adj_matrix,
) )
logger11.info(f"[_step_once] raw_action: {raw_action.shape}") logger11.debug(f"[_step_once] raw_action: {raw_action.shape}")
logger11.info(f"[_step_once] clipped_action: {flat_clipped_action.shape}") logger11.debug(f"[_step_once] clipped_action: {flat_clipped_action.shape}")
# Supporting signals (often where mismatch originates) # Supporting signals (often where mismatch originates)
logger11.info(f"[_step_once] logprob: {logprob.shape}") logger11.debug(f"[_step_once] logprob: {logprob.shape}")
logger11.info(f"[_step_once] value: {value.shape}") logger11.debug(f"[_step_once] value: {value.shape}")
logger11.info(f"[_step_once] mean: {mean.shape}") logger11.debug(f"[_step_once] mean: {mean.shape}")
logger11.info(f"[_step_once] std: {std.shape}") logger11.debug(f"[_step_once] std: {std.shape}")
# ---- ENV STEP ---- # ---- ENV STEP ----
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn( episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
episode_stats, env_state, flat_clipped_action episode_stats, env_state, flat_clipped_action
) )
logger11.info(f"[_step_once] next_obs: {next_obs.shape}") logger11.debug(f"[_step_once] next_obs: {next_obs.shape}")
logger11.info(f"[_step_once] reward: {reward.shape}") logger11.debug(f"[_step_once] reward: {reward.shape}")
logger11.info(f"[_step_once] next_done: {next_done.shape}") logger11.debug(f"[_step_once] next_done: {next_done.shape}")
storage = Storage( storage = Storage(
obs=obs, obs=obs,
@ -317,7 +329,7 @@ def _reward_fn(env_state, next_env_state):
return jnp.where(next_env_state.terminated, 50.0, clipped_env_reward - penalty) return jnp.where(next_env_state.terminated, 50.0, clipped_env_reward - penalty)
# TODO: update to work with extra dimension + message passing # TODO: message passing
def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, morph_mode, segments_per_arm): def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, morph_mode, segments_per_arm):
next_env_state = env_step_fn(env_state, action) next_env_state = env_step_fn(env_state, action)
@ -385,7 +397,7 @@ def _rollout_jit(
feature_extractor: nn.Module, feature_extractor: nn.Module,
actor: nn.Module, actor: nn.Module,
critic: nn.Module, critic: nn.Module,
message_passer: nn.Module, message_passer: Optional[nn.Module],
action_low, action_low,
action_high, action_high,
adj_matrix, adj_matrix,
@ -429,7 +441,6 @@ def _compute_gae_jit(
num_envs, num_envs,
feature_extractor, feature_extractor,
critic, critic,
adj_matrix: jnp.ndarray,
): ):
next_value = apply_shared( next_value = apply_shared(
critic, critic,
@ -540,7 +551,6 @@ class PPOTrainer:
gae_lambda=self.ppo.gae_lambda, gae_lambda=self.ppo.gae_lambda,
feature_extractor=self.feature_extractor, feature_extractor=self.feature_extractor,
critic=self.critic, critic=self.critic,
adj_matrix=self.adj,
) )
) )
@ -556,6 +566,9 @@ class PPOTrainer:
def apply_feature(p, x): def apply_feature(p, x):
return apply_shared(self.feature_extractor, p, x) return apply_shared(self.feature_extractor, p, x)
def apply_message_passer(p, x):
return apply_shared(self.message_passer, p, x)
self._ppo = PPO(self.ppo, apply_sensor, apply_actor, apply_critic, apply_feature) self._ppo = PPO(self.ppo, apply_sensor, apply_actor, apply_critic, apply_feature)
self.agent_state = self._init_agent_state() self.agent_state = self._init_agent_state()
@ -585,7 +598,14 @@ class PPOTrainer:
actor = Actor(action_dim=self.env.single_action_space.shape[0]) actor = Actor(action_dim=self.env.single_action_space.shape[0])
sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
message_passer = OneDenseLayerMLP() message_passer: Optional[nn.Module] = (
MessagePasser(
hidden_dim=300,
num_propagation_steps=self.cfg.architecture.message_passing_steps or 4,
)
if self.morph_mode != MorphMode.CENTRALIZED
else None
)
feature_extractor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) feature_extractor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
critic = OneDenseLayerMLP() critic = OneDenseLayerMLP()
@ -615,39 +635,62 @@ class PPOTrainer:
sensor_keys = jax.random.split(sensor_key, self.needed_copies) sensor_keys = jax.random.split(sensor_key, self.needed_copies)
actor_keys = jax.random.split(actor_key, self.needed_copies) actor_keys = jax.random.split(actor_key, self.needed_copies)
message_passer_keys = jax.random.split(message_passer_key, self.needed_copies)
# note: assumed only 1 message passer needed for now
# message_passer_keys = jax.random.split(message_passer_key, self.needed_copies)
message_passer_keys = jnp.asarray([message_passer_key], dtype=jnp.uint32)
# (needed_copies, 175) # (needed_copies, 175)
sensor_params = jax.vmap(lambda k: self.sensor.init(k, sample_obs))(sensor_keys) sensor_params = jax.vmap(lambda k: self.sensor.init(k, sample_obs))(sensor_keys)
self.logger.info(f"[_init_agent_state] sensor_params: {jax.tree.map(lambda x: x.shape, sensor_params)}") self.logger.info(
f"[_init_agent_state] sensor_params: {jax.tree.map(lambda x: x.shape, sensor_params)}"
)
single_sensor_param = jax.tree.map(lambda x: x[0], sensor_params) single_sensor_param = jax.tree.map(lambda x: x[0], sensor_params)
self.logger.info(f"[_init_agent_state] single_sensor_param: {jax.tree.map(lambda x: x.shape, single_sensor_param)}") self.logger.info(
f"[_init_agent_state] single_sensor_param: {
jax.tree.map(lambda x: x.shape, single_sensor_param)
}"
)
sensor_params_sample = self.sensor.apply(single_sensor_param, sample_obs) sensor_params_sample = self.sensor.apply(single_sensor_param, sample_obs)
self.logger.info(f"[_init_agent_state] sensor_params_sample: {sensor_params_sample.shape}") self.logger.info(f"[_init_agent_state] sensor_params_sample: {sensor_params_sample.shape}")
actor_params = jax.vmap(lambda k: self.actor.init(k, sensor_params_sample))(actor_keys) actor_params = jax.vmap(lambda k: self.actor.init(k, sensor_params_sample))(actor_keys)
self.logger.info(f"[_init_agent_state] actor_params: {jax.tree.map(lambda x: x.shape, actor_params)}") self.logger.info(
f"[_init_agent_state] actor_params: {jax.tree.map(lambda x: x.shape, actor_params)}"
)
message_passer_params = jax.vmap( message_passer_params = {}
lambda k: self.message_passer.init( if self.morph_mode != MorphMode.CENTRALIZED:
k, self.sensor.apply(single_sensor_param, sample_obs) message_passer_params = jax.vmap(
) lambda k: self.message_passer.init(
)(message_passer_keys) k, self.sensor.apply(single_sensor_param, sample_obs)
self.logger.info(f"[_init_agent_state] message_passer_params: {jax.tree.map(lambda x: x.shape, message_passer_params)}") )
)(message_passer_keys)
self.logger.info(
f"[_init_agent_state] message_passer_params: {
jax.tree.map(lambda x: x.shape, message_passer_params)
}"
)
flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic
self.logger.info(f"[_init_agent_state] flat_obs: {flat_obs.shape}") self.logger.info(f"[_init_agent_state] flat_obs: {flat_obs.shape}")
feature_extractor_params = self.feature_extractor.init(feature_extractor_key, flat_obs) feature_extractor_params = self.feature_extractor.init(feature_extractor_key, flat_obs)
self.logger.info(f"[_init_agent_state] feature_extractor_params: {jax.tree.map(lambda x: x.shape, feature_extractor_params)}") self.logger.info(
f"[_init_agent_state] feature_extractor_params: {
jax.tree.map(lambda x: x.shape, feature_extractor_params)
}"
)
critic_input = self.feature_extractor.apply(feature_extractor_params, flat_obs) critic_input = self.feature_extractor.apply(feature_extractor_params, flat_obs)
self.logger.info(f"[_init_agent_state] critic_input: {critic_input.shape}") self.logger.info(f"[_init_agent_state] critic_input: {critic_input.shape}")
critic_params = self.critic.init(critic_key, critic_input) critic_params = self.critic.init(critic_key, critic_input)
self.logger.info(f"[_init_agent_state] critic_params: {jax.tree.map(lambda x: x.shape, critic_params)}") self.logger.info(
f"[_init_agent_state] critic_params: {jax.tree.map(lambda x: x.shape, critic_params)}"
)
return TrainState.create( return TrainState.create(
apply_fn=None, apply_fn=None,