1
Fork 0

chore: observation pipeline in training

This commit is contained in:
Tibo De Peuter 2026-04-28 11:15:16 +02:00
parent 0032496073
commit 036edf8aab
Signed by: tdpeuter
SSH key fingerprint: SHA256:u/h/LVoqKF1Iz02uOyxe6hcjmoZASCGV2HM0TG9ZMoU
2 changed files with 27 additions and 74 deletions

View file

@ -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()

View file

@ -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(