fix(PPOTrainer): fixed duplicate definition of conversion function, tracing error related to obs_mode, added adjacent matrix argument
This commit is contained in:
parent
a005c0ccad
commit
db5d4c3dc2
3 changed files with 66 additions and 20 deletions
|
|
@ -1,4 +1,4 @@
|
||||||
from .env_config import ArenaConfig, EnvConfig, MorphologyConfig
|
from .env_config import ArenaConfig, EnvConfig, MorphologyConfig, MorphMode
|
||||||
from .env_types import Backend, Task
|
from .env_types import Backend, Task
|
||||||
from .env_wrapper import BrittleStarEnv, StepResult
|
from .env_wrapper import BrittleStarEnv, StepResult
|
||||||
from .factory import BrittleStarEnvFactory
|
from .factory import BrittleStarEnvFactory
|
||||||
|
|
@ -12,4 +12,5 @@ __all__ = [
|
||||||
"BrittleStarEnv",
|
"BrittleStarEnv",
|
||||||
"StepResult",
|
"StepResult",
|
||||||
"BrittleStarEnvFactory",
|
"BrittleStarEnvFactory",
|
||||||
|
"MorphMode",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -1,10 +1,18 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
from .env_types import Task
|
from .env_types import Task
|
||||||
|
|
||||||
|
|
||||||
|
class MorphMode(Enum):
|
||||||
|
CENTRALIZED = 0
|
||||||
|
FULLY_CONNECTED = 1
|
||||||
|
RING = 2
|
||||||
|
SEGMENT = 3
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class MorphologyConfig:
|
class MorphologyConfig:
|
||||||
"""Brittle star morphology configuration.
|
"""Brittle star morphology configuration.
|
||||||
|
|
@ -20,6 +28,7 @@ class MorphologyConfig:
|
||||||
segments_per_arm: list[int] = field(default_factory=lambda: [4, 4, 4, 4, 4])
|
segments_per_arm: list[int] = field(default_factory=lambda: [4, 4, 4, 4, 4])
|
||||||
use_p_control: bool = True
|
use_p_control: bool = True
|
||||||
use_torque_control: bool = False
|
use_torque_control: bool = False
|
||||||
|
morph_mode: MorphMode = MorphMode.CENTRALIZED
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def num_arms(self) -> int:
|
def num_arms(self) -> int:
|
||||||
|
|
@ -60,3 +69,5 @@ class EnvConfig:
|
||||||
# Light escape
|
# Light escape
|
||||||
# Per docs in upstream env config: integer factors of 200.
|
# Per docs in upstream env config: integer factors of 200.
|
||||||
light_perlin_noise_scale: int = 0
|
light_perlin_noise_scale: int = 0
|
||||||
|
|
||||||
|
obs_mode: int = 0 # TODO: ????
|
||||||
|
|
|
||||||
|
|
@ -24,15 +24,7 @@ from brittle_star_project.MLPs.mlps import (
|
||||||
Storage,
|
Storage,
|
||||||
)
|
)
|
||||||
from brittle_star_project.ppo import PPO
|
from brittle_star_project.ppo import PPO
|
||||||
from enum import Enum
|
from brittle_star_project.environment import MorphMode
|
||||||
|
|
||||||
|
|
||||||
class MorphMode(Enum):
|
|
||||||
CENTRALIZED = 0
|
|
||||||
FULLY_CONNECTED = 1
|
|
||||||
RING = 2
|
|
||||||
SEGMENT = 3
|
|
||||||
|
|
||||||
|
|
||||||
# TODO: move to config
|
# TODO: move to config
|
||||||
_ALLOWED_OBS_KEYS = {
|
_ALLOWED_OBS_KEYS = {
|
||||||
|
|
@ -138,7 +130,7 @@ def _normalize_obs(obs, mean, var, eps=1e-8):
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def _convert_obs_dict_to_array(obs_dict, morph_mode, segments_per_arm):
|
def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm):
|
||||||
|
|
||||||
num_segments = sum(segments_per_arm)
|
num_segments = sum(segments_per_arm)
|
||||||
num_arms = len(segments_per_arm)
|
num_arms = len(segments_per_arm)
|
||||||
|
|
@ -221,7 +213,7 @@ _SEGMENT_SCALED_KEYS = frozenset(
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit(static_argnums=(1,))
|
||||||
def _convert_obs_dict_to_array(obs_dict, obs_mode, segments_per_arm):
|
def _convert_obs_dict_to_array(obs_dict, obs_mode, segments_per_arm):
|
||||||
|
|
||||||
num_segments = sum(segments_per_arm)
|
num_segments = sum(segments_per_arm)
|
||||||
|
|
@ -279,6 +271,7 @@ def _convert_obs_dict_to_array(obs_dict, obs_mode, segments_per_arm):
|
||||||
return jax.vmap(_filter_and_flatten)(obs_dict)
|
return jax.vmap(_filter_and_flatten)(obs_dict)
|
||||||
|
|
||||||
|
|
||||||
|
# TODO: update to work with extra dimension + message passing
|
||||||
def _get_action_and_value_noise(
|
def _get_action_and_value_noise(
|
||||||
sensor: GenericDenseLayersWithActivation,
|
sensor: GenericDenseLayersWithActivation,
|
||||||
feature_extractor: GenericDenseLayersWithActivation,
|
feature_extractor: GenericDenseLayersWithActivation,
|
||||||
|
|
@ -289,6 +282,7 @@ def _get_action_and_value_noise(
|
||||||
key: jax.random.PRNGKey,
|
key: jax.random.PRNGKey,
|
||||||
action_low,
|
action_low,
|
||||||
action_high,
|
action_high,
|
||||||
|
adj_matrix: jnp.ndarray,
|
||||||
):
|
):
|
||||||
hidden = sensor.apply(agent_state.params["sensor_params"], next_obs)
|
hidden = sensor.apply(agent_state.params["sensor_params"], next_obs)
|
||||||
hidden_critic = feature_extractor.apply(
|
hidden_critic = feature_extractor.apply(
|
||||||
|
|
@ -308,10 +302,12 @@ def _get_action_and_value_noise(
|
||||||
return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key
|
return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key
|
||||||
|
|
||||||
|
|
||||||
|
# TODO: update to work with extra dimension + message passing
|
||||||
def _step_once(
|
def _step_once(
|
||||||
carry,
|
carry,
|
||||||
_,
|
_,
|
||||||
env_step_fn,
|
env_step_fn,
|
||||||
|
adj_matrix,
|
||||||
sensor: GenericDenseLayersWithActivation,
|
sensor: GenericDenseLayersWithActivation,
|
||||||
feature_extractor: GenericDenseLayersWithActivation,
|
feature_extractor: GenericDenseLayersWithActivation,
|
||||||
actor: Actor,
|
actor: Actor,
|
||||||
|
|
@ -321,7 +317,16 @@ def _step_once(
|
||||||
):
|
):
|
||||||
agent_state, episode_stats, obs, done, key, env_state = carry
|
agent_state, episode_stats, obs, done, key, env_state = carry
|
||||||
clipped_action, raw_action, logprob, value, mean, std, key = _get_action_and_value_noise(
|
clipped_action, raw_action, logprob, value, mean, std, key = _get_action_and_value_noise(
|
||||||
sensor, feature_extractor, actor, critic, agent_state, obs, key, action_low, action_high
|
sensor,
|
||||||
|
feature_extractor,
|
||||||
|
actor,
|
||||||
|
critic,
|
||||||
|
agent_state,
|
||||||
|
obs,
|
||||||
|
key,
|
||||||
|
action_low,
|
||||||
|
action_high,
|
||||||
|
adj_matrix,
|
||||||
)
|
)
|
||||||
|
|
||||||
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
|
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
|
||||||
|
|
@ -361,7 +366,8 @@ 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):
|
# TODO: update to work with extra dimension + message passing
|
||||||
|
def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, obs_mode, segments_per_arm):
|
||||||
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)
|
||||||
|
|
@ -385,16 +391,22 @@ 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),
|
(
|
||||||
|
_convert_obs_dict_to_array(next_env_state.observations, obs_mode, segments_per_arm),
|
||||||
|
reward,
|
||||||
|
done,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# TODO: update to work with extra dimension + message passing
|
||||||
def _rollout_jit(
|
def _rollout_jit(
|
||||||
agent_state,
|
agent_state,
|
||||||
episode_stats,
|
episode_stats,
|
||||||
env_state,
|
env_state,
|
||||||
next_obs,
|
next_obs,
|
||||||
next_done,
|
next_done,
|
||||||
|
adj_matrix,
|
||||||
key,
|
key,
|
||||||
max_steps,
|
max_steps,
|
||||||
step_env_fn,
|
step_env_fn,
|
||||||
|
|
@ -415,6 +427,7 @@ def _rollout_jit(
|
||||||
env_step_fn=step_env_fn,
|
env_step_fn=step_env_fn,
|
||||||
action_low=action_low,
|
action_low=action_low,
|
||||||
action_high=action_high,
|
action_high=action_high,
|
||||||
|
adj_matrix=adj_matrix,
|
||||||
),
|
),
|
||||||
(agent_state, episode_stats, next_obs, next_done, key, env_state),
|
(agent_state, episode_stats, next_obs, next_done, key, env_state),
|
||||||
(),
|
(),
|
||||||
|
|
@ -442,6 +455,7 @@ def _compute_gae_jit(
|
||||||
num_envs,
|
num_envs,
|
||||||
feature_extractor,
|
feature_extractor,
|
||||||
critic,
|
critic,
|
||||||
|
adj_matrix: jnp.ndarray,
|
||||||
):
|
):
|
||||||
next_value = critic.apply(
|
next_value = critic.apply(
|
||||||
agent_state.params["critic_params"],
|
agent_state.params["critic_params"],
|
||||||
|
|
@ -484,7 +498,6 @@ class PPOTrainer:
|
||||||
env: BrittleStarJaxEnvWrapper,
|
env: BrittleStarJaxEnvWrapper,
|
||||||
run_dir: str,
|
run_dir: str,
|
||||||
run_name: str,
|
run_name: str,
|
||||||
morph: MorphMode = MorphMode.CENTRALIZED,
|
|
||||||
):
|
):
|
||||||
self.cfg = cfg
|
self.cfg = cfg
|
||||||
self.ppo = cfg.ppo
|
self.ppo = cfg.ppo
|
||||||
|
|
@ -501,6 +514,9 @@ class PPOTrainer:
|
||||||
|
|
||||||
self.key = jax.random.PRNGKey(self.experiment.seed)
|
self.key = jax.random.PRNGKey(self.experiment.seed)
|
||||||
|
|
||||||
|
morph = cfg.morphology.morph_mode
|
||||||
|
|
||||||
|
self.logger.info(f"[INIT]: Used morphology mode {morph}")
|
||||||
self.adj = build_adjacency(cfg.morphology.segments_per_arm, morph)
|
self.adj = build_adjacency(cfg.morphology.segments_per_arm, morph)
|
||||||
|
|
||||||
self.sensor, self.feature_extractor, self.actor, self.critic = self._init_agent()
|
self.sensor, self.feature_extractor, self.actor, self.critic = self._init_agent()
|
||||||
|
|
@ -509,6 +525,9 @@ class PPOTrainer:
|
||||||
self.actor.apply = jax.jit(self.actor.apply)
|
self.actor.apply = jax.jit(self.actor.apply)
|
||||||
self.critic.apply = jax.jit(self.critic.apply)
|
self.critic.apply = jax.jit(self.critic.apply)
|
||||||
|
|
||||||
|
self.obs_mode = self.cfg.environment.obs_mode
|
||||||
|
self.segments_per_arm = jnp.asarray(self.cfg.morphology.segments_per_arm, dtype=jnp.int32)
|
||||||
|
|
||||||
action_low = jnp.asarray(self.env.single_action_space.low, dtype=jnp.float32)
|
action_low = jnp.asarray(self.env.single_action_space.low, dtype=jnp.float32)
|
||||||
action_high = jnp.asarray(self.env.single_action_space.high, dtype=jnp.float32)
|
action_high = jnp.asarray(self.env.single_action_space.high, dtype=jnp.float32)
|
||||||
|
|
||||||
|
|
@ -516,13 +535,19 @@ 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_mode=self.obs_mode,
|
||||||
|
segments_per_arm=self.segments_per_arm,
|
||||||
|
),
|
||||||
sensor=self.sensor,
|
sensor=self.sensor,
|
||||||
feature_extractor=self.feature_extractor,
|
feature_extractor=self.feature_extractor,
|
||||||
actor=self.actor,
|
actor=self.actor,
|
||||||
critic=self.critic,
|
critic=self.critic,
|
||||||
action_low=action_low,
|
action_low=action_low,
|
||||||
action_high=action_high,
|
action_high=action_high,
|
||||||
|
adj_matrix=self.adj,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
self._compute_gae_jit = jax.jit(
|
self._compute_gae_jit = jax.jit(
|
||||||
|
|
@ -533,6 +558,7 @@ class PPOTrainer:
|
||||||
gae_lambda=self.ppo.gae_lambda,
|
gae_lambda=self.ppo.gae_lambda,
|
||||||
feature_extractor=self.feature_extractor,
|
feature_extractor=self.feature_extractor,
|
||||||
critic=self.critic,
|
critic=self.critic,
|
||||||
|
adj_matrix=self.adj,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -568,8 +594,12 @@ class PPOTrainer:
|
||||||
|
|
||||||
dummy_reset = self.env.reset(seed=0)
|
dummy_reset = self.env.reset(seed=0)
|
||||||
for k, v in dummy_reset.observations.items():
|
for k, v in dummy_reset.observations.items():
|
||||||
print(k, v.shape)
|
self.logger.debug(k, v.shape)
|
||||||
sample_obs = _convert_obs_dict_to_array(dummy_reset.observations)[0] # take first env
|
sample_obs = _convert_obs_dict_to_array(
|
||||||
|
dummy_reset.observations,
|
||||||
|
self.obs_mode,
|
||||||
|
self.segments_per_arm,
|
||||||
|
)[0] # take first env
|
||||||
self.obs_mean = jnp.zeros((len(sample_obs),))
|
self.obs_mean = jnp.zeros((len(sample_obs),))
|
||||||
self.obs_var = jnp.ones((len(sample_obs),))
|
self.obs_var = jnp.ones((len(sample_obs),))
|
||||||
self.obs_count = 1e-4
|
self.obs_count = 1e-4
|
||||||
|
|
@ -796,7 +826,11 @@ 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 = _convert_obs_dict_to_array(
|
||||||
|
env_state.observations,
|
||||||
|
self.obs_mode,
|
||||||
|
self.segments_per_arm,
|
||||||
|
)
|
||||||
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()}")
|
||||||
|
|
|
||||||
Reference in a new issue