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"
# Synchronous message-passing rounds per control step
message_passing_steps: 1
message_passing_steps: 4
# Connectivity topology (e.g., ring, fully_connected)
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
import flax
import flax.linen as nn
import jax.numpy as jnp
import jax.tree_util
@ -38,6 +37,28 @@ class Actor(nn.Module):
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
@dataclass
class AgentParams:

View file

@ -3,7 +3,7 @@ import random
import time
from dataclasses import asdict, dataclass
from functools import partial
from typing import Any
from typing import Any, Optional
import jax
import jax.numpy as jnp
@ -21,6 +21,7 @@ from brittle_star_project.MLPs.mlps import (
Actor,
AgentParams,
GenericDenseLayersWithActivation,
MessagePasser,
OneDenseLayerMLP,
Storage,
)
@ -211,14 +212,20 @@ def _get_action_and_value_noise(
feature_extractor: nn.Module,
actor: nn.Module,
critic: nn.Module,
message_passer: Optional[nn.Module],
agent_state: TrainState,
next_obs: jnp.ndarray,
key: jax.random.PRNGKey,
key,
action_low,
action_high,
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 = jax.vmap(apply_message_passing)(hidden)
hidden_critic = apply_shared(
feature_extractor, agent_state.params["feature_extractor_params"], next_obs
)
@ -230,12 +237,16 @@ def _get_action_and_value_noise(
std = jnp.exp(log_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)
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)
return flat_clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key
@ -249,7 +260,7 @@ def _step_once(
feature_extractor: nn.Module,
actor: nn.Module,
critic: nn.Module,
message_passer: nn.Module,
message_passer: Optional[nn.Module],
action_low,
action_high,
):
@ -259,6 +270,7 @@ def _step_once(
feature_extractor,
actor,
critic,
message_passer,
agent_state,
obs,
key,
@ -266,23 +278,23 @@ def _step_once(
action_high,
adj_matrix,
)
logger11.info(f"[_step_once] raw_action: {raw_action.shape}")
logger11.info(f"[_step_once] clipped_action: {flat_clipped_action.shape}")
logger11.debug(f"[_step_once] raw_action: {raw_action.shape}")
logger11.debug(f"[_step_once] clipped_action: {flat_clipped_action.shape}")
# Supporting signals (often where mismatch originates)
logger11.info(f"[_step_once] logprob: {logprob.shape}")
logger11.info(f"[_step_once] value: {value.shape}")
logger11.info(f"[_step_once] mean: {mean.shape}")
logger11.info(f"[_step_once] std: {std.shape}")
logger11.debug(f"[_step_once] logprob: {logprob.shape}")
logger11.debug(f"[_step_once] value: {value.shape}")
logger11.debug(f"[_step_once] mean: {mean.shape}")
logger11.debug(f"[_step_once] std: {std.shape}")
# ---- ENV STEP ----
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
episode_stats, env_state, flat_clipped_action
)
logger11.info(f"[_step_once] next_obs: {next_obs.shape}")
logger11.info(f"[_step_once] reward: {reward.shape}")
logger11.info(f"[_step_once] next_done: {next_done.shape}")
logger11.debug(f"[_step_once] next_obs: {next_obs.shape}")
logger11.debug(f"[_step_once] reward: {reward.shape}")
logger11.debug(f"[_step_once] next_done: {next_done.shape}")
storage = Storage(
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)
# 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):
next_env_state = env_step_fn(env_state, action)
@ -385,7 +397,7 @@ def _rollout_jit(
feature_extractor: nn.Module,
actor: nn.Module,
critic: nn.Module,
message_passer: nn.Module,
message_passer: Optional[nn.Module],
action_low,
action_high,
adj_matrix,
@ -429,7 +441,6 @@ def _compute_gae_jit(
num_envs,
feature_extractor,
critic,
adj_matrix: jnp.ndarray,
):
next_value = apply_shared(
critic,
@ -540,7 +551,6 @@ class PPOTrainer:
gae_lambda=self.ppo.gae_lambda,
feature_extractor=self.feature_extractor,
critic=self.critic,
adj_matrix=self.adj,
)
)
@ -556,6 +566,9 @@ class PPOTrainer:
def apply_feature(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.agent_state = self._init_agent_state()
@ -585,7 +598,14 @@ class PPOTrainer:
actor = Actor(action_dim=self.env.single_action_space.shape[0])
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])
critic = OneDenseLayerMLP()
@ -615,39 +635,62 @@ class PPOTrainer:
sensor_keys = jax.random.split(sensor_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)
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)
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)
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)
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(
lambda k: self.message_passer.init(
k, self.sensor.apply(single_sensor_param, sample_obs)
)
)(message_passer_keys)
self.logger.info(f"[_init_agent_state] message_passer_params: {jax.tree.map(lambda x: x.shape, message_passer_params)}")
message_passer_params = {}
if self.morph_mode != MorphMode.CENTRALIZED:
message_passer_params = jax.vmap(
lambda k: self.message_passer.init(
k, self.sensor.apply(single_sensor_param, sample_obs)
)
)(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
self.logger.info(f"[_init_agent_state] flat_obs: {flat_obs.shape}")
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)
self.logger.info(f"[_init_agent_state] critic_input: {critic_input.shape}")
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(
apply_fn=None,