From 43de6a967002f5dd82d11425daa854e824bd3a64 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 4 May 2026 22:23:58 +0200 Subject: [PATCH] merge: merged dev into branch --- src/brittle_star_project/MLPs/mlps.py | 3 +- .../environment/obs_processing.py | 65 ++++++- .../trainers/PPOTrainer.py | 164 ++---------------- 3 files changed, 79 insertions(+), 153 deletions(-) diff --git a/src/brittle_star_project/MLPs/mlps.py b/src/brittle_star_project/MLPs/mlps.py index 27d81c9..cc8c226 100644 --- a/src/brittle_star_project/MLPs/mlps.py +++ b/src/brittle_star_project/MLPs/mlps.py @@ -48,7 +48,8 @@ class MessagePasser(nn.Module): 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) + # note: if mean is wanted: adj_matrix / (adj.sum(axis=-1, keepdims=True) + 1e-8) + agg = adj_matrix aggregated = agg @ messages x_concat = jnp.concatenate([x, aggregated], axis=-1) diff --git a/src/brittle_star_project/environment/obs_processing.py b/src/brittle_star_project/environment/obs_processing.py index b18ee8c..6b7cbe8 100644 --- a/src/brittle_star_project/environment/obs_processing.py +++ b/src/brittle_star_project/environment/obs_processing.py @@ -2,6 +2,9 @@ import jax import jax.numpy as jnp from typing import Dict, Tuple, Optional +from brittle_star_project.environment.env_config import MorphMode +from experiment_logger import get_logger + _JOINT_SCALED_KEYS = frozenset( { "joint_position", @@ -19,7 +22,11 @@ _SEGMENT_SCALED_KEYS = frozenset( def create_obs_processor( - bounds_dict: Dict[str, Tuple[float, float]], padding_masks: Optional[Dict] = None + bounds_dict: Dict[str, Tuple[float, float]], + num_segments: int, + num_arms: int, + padding_masks: Optional[Dict] = None, + morph_mode: MorphMode = MorphMode.CENTRALIZED, ): def _add_derived_features(obs: dict) -> dict: new_obs = dict(obs) @@ -52,6 +59,8 @@ def create_obs_processor( return normalized def _pad_features(obs: dict) -> dict: + assert padding_masks is not None + padded = {} for key, arr in obs.items(): if key in _JOINT_SCALED_KEYS: @@ -73,13 +82,55 @@ def create_obs_processor( "robot_direction_to_target", "segment_contact", ] + values = [] - for key in ordered_keys: - if key in obs: - arr = jnp.asarray(obs[key]).flatten() - if arr.size > 0: - values.append(arr) - return jnp.concatenate(values) + + for key in sorted(obs.keys()): + if key not in ordered_keys: + continue + v = obs[key] + + # skip empty arrays and scalars + if v.size == 0: + continue + + # reshape scalars + if v.ndim == 0: + v = v.reshape(1) + + # -------- CENTRALIZED -------- + if morph_mode == MorphMode.CENTRALIZED: + values.append(v.reshape(1, -1)) # (1, feat) + continue + + # -------- SPLIT TO SEGMENTS -------- + if key in _JOINT_SCALED_KEYS: + if morph_mode == MorphMode.SEGMENT: + center_size = num_arms * 3 * 2 + v_center = v[:center_size].reshape(num_arms, 3 * 2) # (arms, 6) + v_segs = v[center_size:].reshape(-1, 2) # (segs, 2) + values.append(jnp.concatenate([v_center, v_segs], axis=0)) # (arms+segs, ?) + continue + v = v.reshape(num_arms, -1) # (n_arms, 2) + + elif key in _SEGMENT_SCALED_KEYS: + v = v[:, None] # (segments, 1) + + else: + # global key, broadcast to all nodes + n_nodes = (num_segments + num_arms) if morph_mode == MorphMode.SEGMENT else num_arms + v = jnp.repeat(v[None, :], n_nodes, axis=0) # (n_nodes, feat) + + # -------- SEGMENT MODE -------- + if morph_mode == MorphMode.SEGMENT: + values.append(v) # (n_nodes, feat) + continue + + # -------- ARM MODE -------- + v = v.reshape(num_arms, -1) + values.append(v) # (n_arms, feat) + + return jnp.concatenate(values, axis=-1) def _process_single(obs_dict: dict) -> jnp.ndarray: processed = _add_derived_features(obs_dict) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 569f276..a0af8b7 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -32,18 +32,7 @@ from brittle_star_project.utils import logged_jit logger11 = get_logger() # TODO: move to config -_ALLOWED_OBS_KEYS = { - "joint_position", - "joint_velocity", - "joint_actuator_force", - "actuator_force", - "disk_position", - "disk_rotation", - "disk_linear_velocity", - "disk_angular_velocity", - "unit_xy_direction_to_target", - "xy_distance_to_target", -} + # TODO: clip scaled reward? @@ -129,81 +118,6 @@ def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, lear return learning_rate * frac -@logged_jit -def _normalize_obs(obs, mean, var, eps=1e-8): - return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0) - -# TODO: update to work with new obs_processor -def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, num_segments: int, num_arms: int): - @logged_jit - def _filter_and_flatten(o) -> jnp.ndarray: - # vmap feeds one env at a time — v has NO batch dim here - # shapes are e.g. (n_features,) or (n_nodes, feat) - values = [] - - for key in sorted(o.keys()): - if key not in _ALLOWED_OBS_KEYS: - continue - v = o[key] - if v.size == 0: - continue - - # -------- CENTRALIZED -------- - if morph_mode == MorphMode.CENTRALIZED: - values.append(v.reshape(1, -1)) # (1, feat) - continue - - # -------- SPLIT TO SEGMENTS -------- - if key in _JOINT_SCALED_KEYS: - if morph_mode == MorphMode.SEGMENT: - center_size = num_arms * 3 * 2 - v_center = v[:center_size].reshape(num_arms, 3 * 2) # (arms, 6) - v_segs = v[center_size:].reshape(-1, 2) # (segs, 2) - values.append(jnp.concatenate([v_center, v_segs], axis=0)) # (arms+segs, ?) - continue - v = v.reshape(num_arms, -1) # (n_arms, 2) - - elif key in _SEGMENT_SCALED_KEYS: - v = v[:, None] # (segments, 1) - - else: - # global key, broadcast to all nodes - n_nodes = (num_segments + num_arms) if morph_mode == MorphMode.SEGMENT else num_arms - v = jnp.repeat(v[None, :], n_nodes, axis=0) # (n_nodes, feat) - - # -------- SEGMENT MODE -------- - if morph_mode == MorphMode.SEGMENT: - values.append(v) # (n_nodes, feat) - continue - - # -------- ARM MODE -------- - v = v.reshape(num_arms, -1) - values.append(v) # (n_arms, feat) - - return jnp.concatenate(values, axis=-1) # (n_nodes, total_feat) - - return jax.vmap(_filter_and_flatten)(obs_dict) - # output: (batch, n_nodes, total_feat) - - -# Observation keys whose size scales with the number of joints (2 per segment). -_JOINT_SCALED_KEYS = frozenset( - { # TODO CODE SMELL - "joint_position", - "joint_velocity", - "joint_actuator_force", - "actuator_force", - } -) - -# Observation keys whose size scales with the number of segments (1 per segment). -_SEGMENT_SCALED_KEYS = frozenset( - { - "segment_contact", - } -) - - def _get_action_and_value_noise( sensor: nn.Module, feature_extractor: nn.Module, @@ -249,7 +163,6 @@ def _get_action_and_value_noise( return flat_clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key -# TODO: update to work vectorized (sensor, actor, message passer) + message passing def _step_once( carry, _, @@ -328,9 +241,7 @@ def _reward_fn(env_state, next_env_state): return jnp.where(next_env_state.terminated, 50.0, clipped_env_reward - penalty) -def _step_env_wrapped( - episode_stats, env_state, action, env_step_fn, morph_mode, num_segments: int, num_arms: int, # TODO: obs_processor -): +def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, obs_processor): next_env_state = env_step_fn(env_state, action) reward = _reward_fn(env_state, next_env_state) @@ -355,13 +266,10 @@ def _step_env_wrapped( episode_stats, next_env_state, ( - _convert_obs_dict_to_array_morphology( - next_env_state.observations, morph_mode, num_segments, num_arms - ), + obs_processor(next_env_state.observations), reward, done, ), - # TODO (obs_processor(next_env_state.observations), reward, done), ) @@ -502,7 +410,7 @@ class PPOTrainer: self.key = jax.random.PRNGKey(self.experiment.seed) self.morph_mode = self.cfg.morphology.morph_mode - + self.segments_per_arm = jnp.asarray(self.cfg.morphology.segments_per_arm, dtype=jnp.int32) self.num_segments = self.segments_per_arm.sum().item() self.num_arms = jnp.where(self.segments_per_arm > 0, 1, 0).sum().item() @@ -527,6 +435,9 @@ class PPOTrainer: # Build the centralized observation processor: derive -> normalize -> pad -> flatten. self.obs_processor = create_obs_processor( bounds_dict=self.cfg.obs_bounds.to_bounds_dict(), + num_segments=self.num_segments, + num_arms=self.num_arms, + morph_mode=self.morph_mode, padding_masks=self.env.padding_masks, ) @@ -540,10 +451,7 @@ class PPOTrainer: step_env_fn=partial( _step_env_wrapped, env_step_fn=self.env.step, - morph_mode=self.morph_mode, - num_segments=self.num_segments, - num_arms=self.num_arms, - # TODO: obs_processor=self.obs_processor, + obs_processor=self.obs_processor, ), sensor=self.sensor, feature_extractor=self.feature_extractor, @@ -605,6 +513,7 @@ class PPOTrainer: self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum() ).item() + # scale actor output with size of model --> more models ==> less actions needed per model actor = Actor(action_dim=self.env.single_action_space.shape[0] // needed_copies) sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) message_passer: Optional[nn.Module] = ( @@ -628,15 +537,11 @@ class PPOTrainer: ) dummy_reset = self.env.reset(seed=0) - + for k, v in dummy_reset.observations.items(): self.logger.debug(k, v.shape) - sample_obs = _convert_obs_dict_to_array_morphology( - dummy_reset.observations, - self.morph_mode, - self.num_segments, - self.num_arms, - )[0] # take first env + + sample_obs = self.obs_processor(dummy_reset.observations)[0] # take first env self.logger.debug(f"[_init_agent_state] sample_obs: {sample_obs.shape}") self.obs_mean = jnp.zeros((sample_obs.shape[-1],)) @@ -703,7 +608,6 @@ class PPOTrainer: critic_params = self.critic.init(critic_key, critic_input) self.logger.debug( f"[_init_agent_state] critic_params: {jax.tree.map(lambda x: x.shape, critic_params)}" - # TODO: sample_obs = self.obs_processor(dummy_reset.observations)[0] # take first env ) return TrainState.create( @@ -744,28 +648,6 @@ class PPOTrainer: returned_episode_lengths=jnp.zeros(self.ppo.num_envs, dtype=jnp.int32), ) - def _update_obs_stats(self, obs: jnp.ndarray): - batch_mean = jnp.mean(obs, axis=(0, 1)) - batch_var = jnp.var(obs, axis=(0, 1)) - batch_count = obs.shape[0] - - self.logger.debug(f"Batch mean shape: {batch_mean.shape}") - self.logger.debug(f"obs_mean shape: {self.obs_mean.shape}") - - delta = batch_mean - self.obs_mean - total_count = self.obs_count + batch_count - - new_mean = self.obs_mean + delta * batch_count / total_count - - m_a = self.obs_var * self.obs_count - m_b = batch_var * batch_count - M2 = m_a + m_b + delta**2 * self.obs_count * batch_count / total_count - new_var = M2 / total_count - - self.obs_mean = new_mean - self.obs_var = new_var - self.obs_count = total_count - def _rollout(self, env_state, next_obs, next_done) -> tuple[Any, ...]: return self._rollout_jit( self.agent_state, @@ -840,7 +722,7 @@ class PPOTrainer: def _step(self, env_state, next_obs, next_done, iteration: int) -> tuple: if iteration == 1: self.logger.log_non_interactive(f"Starting first rollout (JIT): {time.ctime()}") - self.logger.info(f"[_step] next_obs (in): {next_obs.shape}") + self.logger.debug(f"[_step] next_obs (in): {next_obs.shape}") ( self.agent_state, self.episode_stats, @@ -850,12 +732,12 @@ class PPOTrainer: self.key, next_env_state, ) = self._rollout(env_state, next_obs, next_done) - self.logger.info(f"[_step] next_obs (post-rollout): {next_obs.shape}") + self.logger.debug(f"[_step] next_obs (post-rollout): {next_obs.shape}") if iteration == 1: self.logger.log_non_interactive(f"First rollout completed: {time.ctime()}") storage = self._compute_gae(storage, next_obs, next_done) - self.logger.info(f"[_step] storage.obs (post-gae): {storage.obs.shape}") + self.logger.debug(f"[_step] storage.obs (post-gae): {storage.obs.shape}") if iteration == 1: self.logger.log_non_interactive(f"Starting first PPO update (JIT): {time.ctime()}") @@ -931,15 +813,10 @@ class PPOTrainer: self.logger.log_non_interactive(f"Initial reset started: {time.ctime()}") env_state = self.env.reset(seed=self.experiment.seed) - - next_obs = _convert_obs_dict_to_array_morphology( - env_state.observations, - self.morph_mode, - self.num_segments, - self.num_arms, - ) - self.logger.info(f"[train] next_obs: {next_obs.shape}") - # TODO: next_obs = self.obs_processor(env_state.observations) + + next_obs = self.obs_processor(env_state.observations) + self.logger.debug(f"[train] next_obs: {next_obs.shape}") + next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_) self.logger.log_non_interactive(f"Initial reset completed: {time.ctime()}") @@ -954,9 +831,6 @@ class PPOTrainer: env_state, next_obs, next_done, training_measurements, storage = self._step( env_state, next_obs, next_done, iteration=iteration ) - self.logger.debug(f"[train] next_obs (post-step): {next_obs.shape}") - self._update_obs_stats(next_obs) - next_obs = _normalize_obs(next_obs, self.obs_mean, self.obs_var) global_step += self.ppo.num_steps * self.ppo.num_envs self._log(