From 036edf8aab4e85b42d0b38c650d430ad94e6b79d Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Tue, 28 Apr 2026 11:15:16 +0200 Subject: [PATCH] chore: observation pipeline in training --- .../environment/BrittleStarJaxEnvWrapper.py | 22 +++--- .../trainers/PPOTrainer.py | 79 ++++--------------- 2 files changed, 27 insertions(+), 74 deletions(-) diff --git a/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py b/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py index a5175a7..b514c11 100644 --- a/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py +++ b/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py @@ -5,7 +5,7 @@ from experiment_logger import get_logger from .env_config import EnvConfig, MorphologyConfig, ArenaConfig from .env_types import Backend from .factory import BrittleStarEnvFactory -from .padded_obs_wrapper import compute_padding_masks, pad_observations_batched +from .padded_obs_wrapper import compute_padding_masks class BrittleStarJaxEnvWrapper: @@ -48,6 +48,15 @@ class BrittleStarJaxEnvWrapper: def raw(self): return self._env + @property + def padding_masks(self) -> dict: + """Pre-computed boolean masks for amputated limb padding. + + Pass to create_obs_processor so the processor handles padding + after normalization in the correct pipeline order. + """ + return self._padding_masks + @property def single_action_space(self): return self._env.action_space @@ -61,10 +70,6 @@ class BrittleStarJaxEnvWrapper: self._action_rng, env_rng = jax.random.split(jax.random.PRNGKey(seed), 2) env_rngs = jnp.array(jax.random.split(env_rng, self._num_envs)) state = self._vectorized_reset(rng=env_rngs) - - state = state.replace( - observations=pad_observations_batched(state.observations, self._padding_masks) - ) return state def sample_actions(self): @@ -75,12 +80,7 @@ class BrittleStarJaxEnvWrapper: return self._vectorized_action_sample(rng=jnp.array(sub_rngs)) def step(self, state, action): - next_state = self._vectorized_step(state=state, action=action) - - next_state = next_state.replace( - observations=pad_observations_batched(next_state.observations, self._padding_masks) - ) - return next_state + return self._vectorized_step(state=state, action=action) def close(self): self._env.close() diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 7cdc1ff..a3e4eac 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -16,6 +16,7 @@ from experiment_logger import get_logger from brittle_star_project.configs.main_config import BrittleStarConfig from brittle_star_project.dataclasses import EpisodeStatistics from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper +from brittle_star_project.environment.obs_processing import create_obs_processor from brittle_star_project.MLPs.mlps import ( Actor, AgentParams, @@ -25,19 +26,6 @@ from brittle_star_project.MLPs.mlps import ( ) from brittle_star_project.ppo import PPO -# 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? @@ -65,27 +53,6 @@ def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, lear return learning_rate * frac -@jax.jit -def _normalize_obs(obs, mean, var, eps=1e-8): - return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0) - - -@jax.jit -def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: - """Convert the raw observation dict → flat array, filtering unwanted keys.""" - - def _filter_and_flatten(o: dict) -> jnp.ndarray: - values = [] - for key in sorted(o.keys()): - if key in _ALLOWED_OBS_KEYS: # TODO: NORMALIZATION or .. of observations?? - v = o[key] - if v.size > 0: - values.append(jnp.asarray(v).flatten()) - return jnp.concatenate(values) - - return jax.vmap(_filter_and_flatten)(obs_dict) - - def _get_action_and_value_noise( sensor: GenericDenseLayersWithActivation, feature_extractor: GenericDenseLayersWithActivation, @@ -168,7 +135,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): +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) @@ -192,7 +159,7 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn): return ( episode_stats, next_env_state, - (_convert_obs_dict_to_array(next_env_state.observations), reward, done), + (obs_processor(next_env_state.observations), reward, done), ) @@ -303,6 +270,12 @@ class PPOTrainer: self.key = jax.random.PRNGKey(self.experiment.seed) + # Build the centralized observation processor: derive -> normalize -> pad -> flatten. + self.obs_processor = create_obs_processor( + bounds_dict=self.cfg.obs_bounds.to_bounds_dict(), + padding_masks=self.env.padding_masks, + ) + self.sensor, self.feature_extractor, self.actor, self.critic = self._init_agent() self.sensor.apply = jax.jit(self.sensor.apply) self.feature_extractor.apply = jax.jit(self.feature_extractor.apply) @@ -316,7 +289,11 @@ class PPOTrainer: partial( _rollout_jit, max_steps=self.ppo.num_steps, - step_env_fn=partial(_step_env_wrapped, env_step_fn=self.env.step), + step_env_fn=partial( + _step_env_wrapped, + env_step_fn=self.env.step, + obs_processor=self.obs_processor, + ), sensor=self.sensor, feature_extractor=self.feature_extractor, actor=self.actor, @@ -367,10 +344,7 @@ class PPOTrainer: ) dummy_reset = self.env.reset(seed=0) - sample_obs = _convert_obs_dict_to_array(dummy_reset.observations)[0] # take first env - self.obs_mean = jnp.zeros((len(sample_obs),)) - self.obs_var = jnp.ones((len(sample_obs),)) - self.obs_count = 1e-4 + sample_obs = self.obs_processor(dummy_reset.observations)[0] # take first env sensor_params = self.sensor.init(sensor_key, sample_obs) feature_extractor_params = self.feature_extractor.init(feature_extractor_key, sample_obs) actor_params = self.actor.init(actor_key, self.sensor.apply(sensor_params, sample_obs)) @@ -410,25 +384,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) - batch_var = jnp.var(obs, axis=0) - batch_count = obs.shape[0] - - 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, @@ -594,7 +549,7 @@ 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(env_state.observations) + next_obs = self.obs_processor(env_state.observations) next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_) self.logger.log_non_interactive(f"Initial reset completed: {time.ctime()}") @@ -609,8 +564,6 @@ class PPOTrainer: env_state, next_obs, next_done, training_measurements, storage = self._step( env_state, next_obs, next_done, iteration=iteration ) - 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(