chore: observation pipeline in training
This commit is contained in:
parent
0032496073
commit
036edf8aab
2 changed files with 27 additions and 74 deletions
|
|
@ -5,7 +5,7 @@ from experiment_logger import get_logger
|
||||||
from .env_config import EnvConfig, MorphologyConfig, ArenaConfig
|
from .env_config import EnvConfig, MorphologyConfig, ArenaConfig
|
||||||
from .env_types import Backend
|
from .env_types import Backend
|
||||||
from .factory import BrittleStarEnvFactory
|
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:
|
class BrittleStarJaxEnvWrapper:
|
||||||
|
|
@ -48,6 +48,15 @@ class BrittleStarJaxEnvWrapper:
|
||||||
def raw(self):
|
def raw(self):
|
||||||
return self._env
|
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
|
@property
|
||||||
def single_action_space(self):
|
def single_action_space(self):
|
||||||
return self._env.action_space
|
return self._env.action_space
|
||||||
|
|
@ -61,10 +70,6 @@ class BrittleStarJaxEnvWrapper:
|
||||||
self._action_rng, env_rng = jax.random.split(jax.random.PRNGKey(seed), 2)
|
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))
|
env_rngs = jnp.array(jax.random.split(env_rng, self._num_envs))
|
||||||
state = self._vectorized_reset(rng=env_rngs)
|
state = self._vectorized_reset(rng=env_rngs)
|
||||||
|
|
||||||
state = state.replace(
|
|
||||||
observations=pad_observations_batched(state.observations, self._padding_masks)
|
|
||||||
)
|
|
||||||
return state
|
return state
|
||||||
|
|
||||||
def sample_actions(self):
|
def sample_actions(self):
|
||||||
|
|
@ -75,12 +80,7 @@ class BrittleStarJaxEnvWrapper:
|
||||||
return self._vectorized_action_sample(rng=jnp.array(sub_rngs))
|
return self._vectorized_action_sample(rng=jnp.array(sub_rngs))
|
||||||
|
|
||||||
def step(self, state, action):
|
def step(self, state, action):
|
||||||
next_state = self._vectorized_step(state=state, action=action)
|
return 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
|
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
self._env.close()
|
self._env.close()
|
||||||
|
|
|
||||||
|
|
@ -16,6 +16,7 @@ from experiment_logger import get_logger
|
||||||
from brittle_star_project.configs.main_config import BrittleStarConfig
|
from brittle_star_project.configs.main_config import BrittleStarConfig
|
||||||
from brittle_star_project.dataclasses import EpisodeStatistics
|
from brittle_star_project.dataclasses import EpisodeStatistics
|
||||||
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
|
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 (
|
from brittle_star_project.MLPs.mlps import (
|
||||||
Actor,
|
Actor,
|
||||||
AgentParams,
|
AgentParams,
|
||||||
|
|
@ -25,19 +26,6 @@ from brittle_star_project.MLPs.mlps import (
|
||||||
)
|
)
|
||||||
from brittle_star_project.ppo import PPO
|
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?
|
# TODO: clip scaled reward?
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -65,27 +53,6 @@ def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, lear
|
||||||
return learning_rate * frac
|
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(
|
def _get_action_and_value_noise(
|
||||||
sensor: GenericDenseLayersWithActivation,
|
sensor: GenericDenseLayersWithActivation,
|
||||||
feature_extractor: 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)
|
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)
|
next_env_state = env_step_fn(env_state, action)
|
||||||
|
|
||||||
reward = _reward_fn(env_state, next_env_state)
|
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 (
|
return (
|
||||||
episode_stats,
|
episode_stats,
|
||||||
next_env_state,
|
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)
|
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, self.feature_extractor, self.actor, self.critic = self._init_agent()
|
||||||
self.sensor.apply = jax.jit(self.sensor.apply)
|
self.sensor.apply = jax.jit(self.sensor.apply)
|
||||||
self.feature_extractor.apply = jax.jit(self.feature_extractor.apply)
|
self.feature_extractor.apply = jax.jit(self.feature_extractor.apply)
|
||||||
|
|
@ -316,7 +289,11 @@ class PPOTrainer:
|
||||||
partial(
|
partial(
|
||||||
_rollout_jit,
|
_rollout_jit,
|
||||||
max_steps=self.ppo.num_steps,
|
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,
|
sensor=self.sensor,
|
||||||
feature_extractor=self.feature_extractor,
|
feature_extractor=self.feature_extractor,
|
||||||
actor=self.actor,
|
actor=self.actor,
|
||||||
|
|
@ -367,10 +344,7 @@ class PPOTrainer:
|
||||||
)
|
)
|
||||||
|
|
||||||
dummy_reset = self.env.reset(seed=0)
|
dummy_reset = self.env.reset(seed=0)
|
||||||
sample_obs = _convert_obs_dict_to_array(dummy_reset.observations)[0] # take first env
|
sample_obs = self.obs_processor(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
|
|
||||||
sensor_params = self.sensor.init(sensor_key, sample_obs)
|
sensor_params = self.sensor.init(sensor_key, sample_obs)
|
||||||
feature_extractor_params = self.feature_extractor.init(feature_extractor_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))
|
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),
|
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, ...]:
|
def _rollout(self, env_state, next_obs, next_done) -> tuple[Any, ...]:
|
||||||
return self._rollout_jit(
|
return self._rollout_jit(
|
||||||
self.agent_state,
|
self.agent_state,
|
||||||
|
|
@ -594,7 +549,7 @@ class PPOTrainer:
|
||||||
self.logger.log_non_interactive(f"Initial reset started: {time.ctime()}")
|
self.logger.log_non_interactive(f"Initial reset started: {time.ctime()}")
|
||||||
|
|
||||||
env_state = self.env.reset(seed=self.experiment.seed)
|
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_)
|
next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_)
|
||||||
|
|
||||||
self.logger.log_non_interactive(f"Initial reset completed: {time.ctime()}")
|
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, training_measurements, storage = self._step(
|
||||||
env_state, next_obs, next_done, iteration=iteration
|
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
|
global_step += self.ppo.num_steps * self.ppo.num_envs
|
||||||
self._log(
|
self._log(
|
||||||
|
|
|
||||||
Reference in a new issue