fix(ppo training): more shape-related bug fixes + moved computation of #arms and segments outside of jit compiled function
This commit is contained in:
parent
a43f3c0172
commit
dd2612e2ec
1 changed files with 25 additions and 21 deletions
|
|
@ -133,10 +133,7 @@ def _normalize_obs(obs, mean, var, eps=1e-8):
|
||||||
return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0)
|
return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0)
|
||||||
|
|
||||||
|
|
||||||
def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm: jnp.ndarray):
|
def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, num_segments: int, num_arms: int):
|
||||||
num_segments = int(segments_per_arm.sum())
|
|
||||||
num_arms = int(jnp.where(segments_per_arm > 0, 1, 0).sum())
|
|
||||||
|
|
||||||
@logged_jit
|
@logged_jit
|
||||||
def _filter_and_flatten(o) -> jnp.ndarray:
|
def _filter_and_flatten(o) -> jnp.ndarray:
|
||||||
# vmap feeds one env at a time — v has NO batch dim here
|
# vmap feeds one env at a time — v has NO batch dim here
|
||||||
|
|
@ -222,15 +219,12 @@ def _get_action_and_value_noise(
|
||||||
):
|
):
|
||||||
# (B, n_nodes, feat)
|
# (B, n_nodes, feat)
|
||||||
hidden = apply_per_node(sensor, agent_state.params["sensor_params"], next_obs)
|
hidden = apply_per_node(sensor, agent_state.params["sensor_params"], next_obs)
|
||||||
get_logger().debug(f"[_get_action_and_value_noise] hidden (before): {hidden.shape}")
|
|
||||||
|
|
||||||
if message_passer is not None:
|
if message_passer is not None:
|
||||||
params = agent_state.params["message_passer_params"]
|
params = agent_state.params["message_passer_params"]
|
||||||
# (n_nodes, feat) --> let each node talk with its neighbours ==> vmap over B dimension
|
# (n_nodes, feat) --> let each node talk with its neighbours ==> vmap over B dimension
|
||||||
hidden = jax.vmap(lambda x: message_passer.apply(params, x, adj_matrix))(hidden)
|
hidden = jax.vmap(lambda x: message_passer.apply(params, x, adj_matrix))(hidden)
|
||||||
|
|
||||||
get_logger().debug(f"[_get_action_and_value_noise] hidden (after): {hidden.shape}")
|
|
||||||
|
|
||||||
hidden_critic = apply_shared(
|
hidden_critic = apply_shared(
|
||||||
feature_extractor, agent_state.params["feature_extractor_params"], next_obs
|
feature_extractor, agent_state.params["feature_extractor_params"], next_obs
|
||||||
)
|
)
|
||||||
|
|
@ -242,11 +236,10 @@ def _get_action_and_value_noise(
|
||||||
std = jnp.exp(log_std)
|
std = jnp.exp(log_std)
|
||||||
|
|
||||||
raw_action = mean + noise * std
|
raw_action = mean + noise * std
|
||||||
clipped_action = _clip_action(raw_action, action_low, action_high)
|
flat_action = raw_action.reshape(
|
||||||
|
raw_action.shape[0], -1
|
||||||
flat_clipped_action = clipped_action.reshape(
|
|
||||||
clipped_action.shape[0], -1
|
|
||||||
) # concat the per agent, keep the envs dim (batch, agent * action)
|
) # concat the per agent, keep the envs dim (batch, agent * action)
|
||||||
|
flat_clipped_action = _clip_action(flat_action, action_low, action_high)
|
||||||
|
|
||||||
logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(
|
logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(
|
||||||
axis=(-2, -1)
|
axis=(-2, -1)
|
||||||
|
|
@ -335,7 +328,9 @@ 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, morph_mode, segments_per_arm):
|
def _step_env_wrapped(
|
||||||
|
episode_stats, env_state, action, env_step_fn, morph_mode, num_segments: int, num_arms: int
|
||||||
|
):
|
||||||
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)
|
||||||
|
|
@ -361,7 +356,7 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, morph_mode,
|
||||||
next_env_state,
|
next_env_state,
|
||||||
(
|
(
|
||||||
_convert_obs_dict_to_array_morphology(
|
_convert_obs_dict_to_array_morphology(
|
||||||
next_env_state.observations, morph_mode, segments_per_arm
|
next_env_state.observations, morph_mode, num_segments, num_arms
|
||||||
),
|
),
|
||||||
reward,
|
reward,
|
||||||
done,
|
done,
|
||||||
|
|
@ -508,6 +503,9 @@ class PPOTrainer:
|
||||||
self.morph_mode = self.cfg.morphology.morph_mode
|
self.morph_mode = self.cfg.morphology.morph_mode
|
||||||
self.segments_per_arm = jnp.asarray(self.cfg.morphology.segments_per_arm, dtype=jnp.int32)
|
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()
|
||||||
|
|
||||||
self.logger.info(f"[INIT]: Used morphology mode {self.morph_mode}")
|
self.logger.info(f"[INIT]: Used morphology mode {self.morph_mode}")
|
||||||
self.adj = build_adjacency(cfg.morphology.segments_per_arm, self.morph_mode)
|
self.adj = build_adjacency(cfg.morphology.segments_per_arm, self.morph_mode)
|
||||||
|
|
||||||
|
|
@ -536,7 +534,8 @@ class PPOTrainer:
|
||||||
_step_env_wrapped,
|
_step_env_wrapped,
|
||||||
env_step_fn=self.env.step,
|
env_step_fn=self.env.step,
|
||||||
morph_mode=self.morph_mode,
|
morph_mode=self.morph_mode,
|
||||||
segments_per_arm=self.segments_per_arm,
|
num_segments=self.num_segments,
|
||||||
|
num_arms=self.num_arms,
|
||||||
),
|
),
|
||||||
sensor=self.sensor,
|
sensor=self.sensor,
|
||||||
feature_extractor=self.feature_extractor,
|
feature_extractor=self.feature_extractor,
|
||||||
|
|
@ -601,7 +600,7 @@ class PPOTrainer:
|
||||||
self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum()
|
self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum()
|
||||||
).item()
|
).item()
|
||||||
|
|
||||||
actor = Actor(action_dim=self.env.single_action_space.shape[0])
|
actor = Actor(action_dim=self.env.single_action_space.shape[0] // needed_copies)
|
||||||
sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
|
sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
|
||||||
message_passer: Optional[nn.Module] = (
|
message_passer: Optional[nn.Module] = (
|
||||||
MessagePasser(
|
MessagePasser(
|
||||||
|
|
@ -629,12 +628,13 @@ class PPOTrainer:
|
||||||
sample_obs = _convert_obs_dict_to_array_morphology(
|
sample_obs = _convert_obs_dict_to_array_morphology(
|
||||||
dummy_reset.observations,
|
dummy_reset.observations,
|
||||||
self.morph_mode,
|
self.morph_mode,
|
||||||
self.segments_per_arm,
|
self.num_segments,
|
||||||
|
self.num_arms,
|
||||||
)[0] # take first env
|
)[0] # take first env
|
||||||
|
|
||||||
self.logger.debug(f"[_init_agent_state] sample_obs: {sample_obs.shape}")
|
self.logger.debug(f"[_init_agent_state] sample_obs: {sample_obs.shape}")
|
||||||
self.obs_mean = jnp.zeros((len(sample_obs),))
|
self.obs_mean = jnp.zeros((sample_obs.shape[-1],))
|
||||||
self.obs_var = jnp.ones((len(sample_obs),))
|
self.obs_var = jnp.ones((sample_obs.shape[-1],))
|
||||||
self.obs_count = 1e-4
|
self.obs_count = 1e-4
|
||||||
self.logger.debug(f"[_init_agent_state] obs_mean: {self.obs_mean.shape}")
|
self.logger.debug(f"[_init_agent_state] obs_mean: {self.obs_mean.shape}")
|
||||||
self.logger.debug(f"[_init_agent_state] obs_var: {self.obs_var.shape}")
|
self.logger.debug(f"[_init_agent_state] obs_var: {self.obs_var.shape}")
|
||||||
|
|
@ -746,10 +746,13 @@ class PPOTrainer:
|
||||||
)
|
)
|
||||||
|
|
||||||
def _update_obs_stats(self, obs: jnp.ndarray):
|
def _update_obs_stats(self, obs: jnp.ndarray):
|
||||||
batch_mean = jnp.mean(obs, axis=0)
|
batch_mean = jnp.mean(obs, axis=(0, 1))
|
||||||
batch_var = jnp.var(obs, axis=0)
|
batch_var = jnp.var(obs, axis=(0, 1))
|
||||||
batch_count = obs.shape[0]
|
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
|
delta = batch_mean - self.obs_mean
|
||||||
total_count = self.obs_count + batch_count
|
total_count = self.obs_count + batch_count
|
||||||
|
|
||||||
|
|
@ -932,7 +935,8 @@ class PPOTrainer:
|
||||||
next_obs = _convert_obs_dict_to_array_morphology(
|
next_obs = _convert_obs_dict_to_array_morphology(
|
||||||
env_state.observations,
|
env_state.observations,
|
||||||
self.morph_mode,
|
self.morph_mode,
|
||||||
self.segments_per_arm,
|
self.num_segments,
|
||||||
|
self.num_arms,
|
||||||
)
|
)
|
||||||
self.logger.info(f"[train] next_obs: {next_obs.shape}")
|
self.logger.info(f"[train] next_obs: {next_obs.shape}")
|
||||||
next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_)
|
next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_)
|
||||||
|
|
|
||||||
Reference in a new issue