fix(ppo training): fixed shape error while clipping action
This commit is contained in:
parent
bb14507fc7
commit
a43f3c0172
3 changed files with 58 additions and 50 deletions
|
|
@ -42,7 +42,7 @@ def main(dict_cfg: DictConfig):
|
||||||
base_dir=os.path.dirname(run_dir),
|
base_dir=os.path.dirname(run_dir),
|
||||||
)
|
)
|
||||||
logger = get_logger()
|
logger = get_logger()
|
||||||
logger.set_level(logging.INFO)
|
logger.set_level(logging.DEBUG)
|
||||||
logger.info(f"Hydra-initialized run: {run_name}")
|
logger.info(f"Hydra-initialized run: {run_name}")
|
||||||
logger.info(f"Output directory: {run_dir}")
|
logger.info(f"Output directory: {run_dir}")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -44,6 +44,7 @@ class MessagePasser(nn.Module):
|
||||||
@nn.compact
|
@nn.compact
|
||||||
def __call__(self, x: jnp.ndarray, adj_matrix: jnp.ndarray):
|
def __call__(self, x: jnp.ndarray, adj_matrix: jnp.ndarray):
|
||||||
for _ in range(self.num_propagation_steps):
|
for _ in range(self.num_propagation_steps):
|
||||||
|
# (n_nodes, feat)
|
||||||
messages = nn.Dense(self.hidden_dim)(x)
|
messages = nn.Dense(self.hidden_dim)(x)
|
||||||
messages = nn.tanh(messages)
|
messages = nn.tanh(messages)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -220,18 +220,16 @@ def _get_action_and_value_noise(
|
||||||
action_high,
|
action_high,
|
||||||
adj_matrix: jnp.ndarray,
|
adj_matrix: jnp.ndarray,
|
||||||
):
|
):
|
||||||
def apply_message_passing(p, h):
|
# (B, n_nodes, feat)
|
||||||
assert message_passer is not None, (
|
|
||||||
"message passing shouldn't be performed if morphology mode is Centralized"
|
|
||||||
)
|
|
||||||
|
|
||||||
return jax.vmap(lambda x: message_passer.apply(p, x, adj_matrix))(h)
|
|
||||||
|
|
||||||
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:
|
||||||
hidden = jax.vmap(apply_message_passing)(
|
params = agent_state.params["message_passer_params"]
|
||||||
agent_state.params["message_passer_params"], hidden
|
# (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)
|
||||||
|
|
||||||
|
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
|
||||||
|
|
@ -244,10 +242,11 @@ 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
|
||||||
flat_action = raw_action.reshape(
|
clipped_action = _clip_action(raw_action, action_low, action_high)
|
||||||
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)
|
||||||
|
|
@ -632,73 +631,81 @@ class PPOTrainer:
|
||||||
self.morph_mode,
|
self.morph_mode,
|
||||||
self.segments_per_arm,
|
self.segments_per_arm,
|
||||||
)[0] # take first env
|
)[0] # take first env
|
||||||
self.logger.info(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((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
|
||||||
self.logger.info(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.info(f"[_init_agent_state] obs_var: {self.obs_var.shape}")
|
self.logger.debug(f"[_init_agent_state] obs_var: {self.obs_var.shape}")
|
||||||
|
|
||||||
|
self.logger.debug(f"[_init_agent_state]: Needed copies: {self.needed_copies}")
|
||||||
sensor_keys = jax.random.split(sensor_key, self.needed_copies)
|
sensor_keys = jax.random.split(sensor_key, self.needed_copies)
|
||||||
actor_keys = jax.random.split(actor_key, self.needed_copies)
|
actor_keys = jax.random.split(actor_key, self.needed_copies)
|
||||||
|
|
||||||
# note: assumed only 1 message passer needed for now
|
# note: assumed only 1 message passer needed for now
|
||||||
# message_passer_keys = jax.random.split(message_passer_key, self.needed_copies)
|
# message_passer_keys = jax.random.split(message_passer_key, self.needed_copies)
|
||||||
message_passer_keys = jnp.asarray([message_passer_key], dtype=jnp.uint32)
|
|
||||||
|
|
||||||
# (needed_copies, 175)
|
# (needed_copies, 175)
|
||||||
sensor_params = jax.vmap(lambda k: self.sensor.init(k, sample_obs))(sensor_keys)
|
sensor_params = jax.vmap(lambda k: self.sensor.init(k, sample_obs))(sensor_keys)
|
||||||
self.logger.info(
|
# self.logger.debug(
|
||||||
f"[_init_agent_state] sensor_params: {jax.tree.map(lambda x: x.shape, sensor_params)}"
|
# f"[_init_agent_state] sensor_params: {jax.tree.map(lambda x: x.shape, sensor_params)}"
|
||||||
)
|
# )
|
||||||
|
|
||||||
single_sensor_param = jax.tree.map(lambda x: x[0], sensor_params)
|
single_sensor_param = jax.tree.map(lambda x: x[0], sensor_params)
|
||||||
self.logger.info(
|
# self.logger.debug(
|
||||||
f"[_init_agent_state] single_sensor_param: {
|
# f"[_init_agent_state] single_sensor_param: {
|
||||||
jax.tree.map(lambda x: x.shape, single_sensor_param)
|
# jax.tree.map(lambda x: x.shape, single_sensor_param)
|
||||||
}"
|
# }"
|
||||||
)
|
# )
|
||||||
|
|
||||||
sensor_params_sample = self.sensor.apply(single_sensor_param, sample_obs)
|
sensor_params_sample = self.sensor.apply(single_sensor_param, sample_obs)
|
||||||
self.logger.info(f"[_init_agent_state] sensor_params_sample: {sensor_params_sample.shape}")
|
self.logger.debug(
|
||||||
|
f"[_init_agent_state] sensor_params_sample shape: {sensor_params_sample.shape}"
|
||||||
|
)
|
||||||
|
|
||||||
actor_params = jax.vmap(lambda k: self.actor.init(k, sensor_params_sample))(actor_keys)
|
actor_params = jax.vmap(lambda k: self.actor.init(k, sensor_params_sample))(actor_keys)
|
||||||
self.logger.info(
|
# self.logger.debug(
|
||||||
f"[_init_agent_state] actor_params: {jax.tree.map(lambda x: x.shape, actor_params)}"
|
# f"[_init_agent_state] actor_params: {jax.tree.map(lambda x: x.shape, actor_params)}"
|
||||||
)
|
# )
|
||||||
|
|
||||||
message_passer_params = {}
|
message_passer_params = {}
|
||||||
if self.morph_mode != MorphMode.CENTRALIZED:
|
if self.morph_mode != MorphMode.CENTRALIZED:
|
||||||
assert self.message_passer is not None, "MessagePasser is None"
|
assert self.message_passer is not None, "MessagePasser is None"
|
||||||
|
|
||||||
message_passer_params = jax.vmap(
|
# message_passer_params = jax.vmap(
|
||||||
lambda k: self.message_passer.init(
|
# lambda k: self.message_passer.init(
|
||||||
k, self.sensor.apply(single_sensor_param, sample_obs), self.adj
|
# k, self.sensor.apply(single_sensor_param, sample_obs), self.adj
|
||||||
)
|
# )
|
||||||
)(message_passer_keys)
|
# )(message_passer_keys)
|
||||||
self.logger.info(
|
message_passer_params = self.message_passer.init(
|
||||||
f"[_init_agent_state] message_passer_params: {
|
message_passer_key,
|
||||||
jax.tree.map(lambda x: x.shape, message_passer_params)
|
self.sensor.apply(single_sensor_param, sample_obs),
|
||||||
}"
|
self.adj,
|
||||||
)
|
)
|
||||||
|
# self.logger.debug(
|
||||||
|
# f"[_init_agent_state] message_passer_params: {
|
||||||
|
# jax.tree.map(lambda x: x.shape, message_passer_params)
|
||||||
|
# }"
|
||||||
|
# )
|
||||||
|
|
||||||
flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic
|
flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic
|
||||||
self.logger.info(f"[_init_agent_state] flat_obs: {flat_obs.shape}")
|
self.logger.debug(f"[_init_agent_state] flat_obs: {flat_obs.shape}")
|
||||||
|
|
||||||
feature_extractor_params = self.feature_extractor.init(feature_extractor_key, flat_obs)
|
feature_extractor_params = self.feature_extractor.init(feature_extractor_key, flat_obs)
|
||||||
self.logger.info(
|
# self.logger.debug(
|
||||||
f"[_init_agent_state] feature_extractor_params: {
|
# f"[_init_agent_state] feature_extractor_params: {
|
||||||
jax.tree.map(lambda x: x.shape, feature_extractor_params)
|
# jax.tree.map(lambda x: x.shape, feature_extractor_params)
|
||||||
}"
|
# }"
|
||||||
)
|
# )
|
||||||
|
|
||||||
critic_input = self.feature_extractor.apply(feature_extractor_params, flat_obs)
|
critic_input = self.feature_extractor.apply(feature_extractor_params, flat_obs)
|
||||||
self.logger.info(f"[_init_agent_state] critic_input: {critic_input.shape}")
|
self.logger.debug(f"[_init_agent_state] critic_input: {critic_input.shape}")
|
||||||
|
|
||||||
critic_params = self.critic.init(critic_key, critic_input)
|
critic_params = self.critic.init(critic_key, critic_input)
|
||||||
self.logger.info(
|
# self.logger.debug(
|
||||||
f"[_init_agent_state] critic_params: {jax.tree.map(lambda x: x.shape, critic_params)}"
|
# f"[_init_agent_state] critic_params: {jax.tree.map(lambda x: x.shape, critic_params)}"
|
||||||
)
|
# )
|
||||||
|
|
||||||
return TrainState.create(
|
return TrainState.create(
|
||||||
apply_fn=None,
|
apply_fn=None,
|
||||||
|
|
|
||||||
Reference in a new issue