From 6114e36e277a39e23328f38122646e48a48a3ff4 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Sat, 2 May 2026 21:50:42 +0200 Subject: [PATCH] feat(PPOTrainer): decentralized message passing babyyyy --- configs/architecture/decentralized.yaml | 2 +- configs/environment/dir_loc_further.yaml | 2 + configs/experiment/long_2arm.yaml | 6 + configs/ppo/chickendinnerwinner.yaml | 16 +++ src/brittle_star_project/MLPs/mlps.py | 23 +++- .../trainers/PPOTrainer.py | 107 ++++++++++++------ 6 files changed, 122 insertions(+), 34 deletions(-) create mode 100644 configs/environment/dir_loc_further.yaml create mode 100644 configs/experiment/long_2arm.yaml create mode 100644 configs/ppo/chickendinnerwinner.yaml diff --git a/configs/architecture/decentralized.yaml b/configs/architecture/decentralized.yaml index 4b9fffb..9e3ae53 100644 --- a/configs/architecture/decentralized.yaml +++ b/configs/architecture/decentralized.yaml @@ -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" diff --git a/configs/environment/dir_loc_further.yaml b/configs/environment/dir_loc_further.yaml new file mode 100644 index 0000000..236f0c5 --- /dev/null +++ b/configs/environment/dir_loc_further.yaml @@ -0,0 +1,2 @@ +simulation_time: 50000.0 +target_distance: 3.0 \ No newline at end of file diff --git a/configs/experiment/long_2arm.yaml b/configs/experiment/long_2arm.yaml new file mode 100644 index 0000000..ae4f4b0 --- /dev/null +++ b/configs/experiment/long_2arm.yaml @@ -0,0 +1,6 @@ +# Testing chicken dinner 4 but further distance. + +exp_name: "long2arm" +seed: 123 +torch_deterministic: true +cuda: true \ No newline at end of file diff --git a/configs/ppo/chickendinnerwinner.yaml b/configs/ppo/chickendinnerwinner.yaml new file mode 100644 index 0000000..cf5e3d6 --- /dev/null +++ b/configs/ppo/chickendinnerwinner.yaml @@ -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 \ No newline at end of file diff --git a/src/brittle_star_project/MLPs/mlps.py b/src/brittle_star_project/MLPs/mlps.py index 45a9359..cf488c1 100644 --- a/src/brittle_star_project/MLPs/mlps.py +++ b/src/brittle_star_project/MLPs/mlps.py @@ -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: diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 146c184..fef6f1c 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -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,