fix:logging shapes
This commit is contained in:
parent
bceb7fe4cf
commit
f77766ba08
2 changed files with 58 additions and 26 deletions
|
|
@ -32,11 +32,11 @@ class PPO:
|
||||||
# or this function will need to recompile
|
# or this function will need to recompile
|
||||||
@partial(jax.jit, static_argnums=0)
|
@partial(jax.jit, static_argnums=0)
|
||||||
def update_ppo(self, agent_state, storage, key):
|
def update_ppo(self, agent_state, storage, key):
|
||||||
logger.info(f"[PPO] storage.obs shape: {getattr(storage, 'obs', None).shape}")
|
logger.info(f"[update_ppo] storage.obs shape: {getattr(storage, 'obs', None).shape}")
|
||||||
logger.info(f"[PPO] storage.actions shape: {storage.actions.shape}")
|
logger.info(f"[update_ppo] storage.actions shape: {storage.actions.shape}")
|
||||||
logger.info(f"[PPO] storage.logprobs shape: {storage.logprobs.shape}")
|
logger.info(f"[update_ppo] storage.logprobs shape: {storage.logprobs.shape}")
|
||||||
logger.info(f"[PPO] storage.advantages shape: {storage.advantages.shape}")
|
logger.info(f"[update_ppo] storage.advantages shape: {storage.advantages.shape}")
|
||||||
logger.info(f"[PPO] storage.returns shape: {storage.returns.shape}")
|
logger.info(f"[update_ppo] storage.returns shape: {storage.returns.shape}")
|
||||||
|
|
||||||
args = self.args
|
args = self.args
|
||||||
ppo_loss_grad_fn = self.ppo_loss_grad_fn
|
ppo_loss_grad_fn = self.ppo_loss_grad_fn
|
||||||
|
|
@ -56,11 +56,11 @@ class PPO:
|
||||||
shuffled_storage = jax.tree.map(convert_data, flatten_storage)
|
shuffled_storage = jax.tree.map(convert_data, flatten_storage)
|
||||||
|
|
||||||
def update_minibatch(agent_state, minibatch):
|
def update_minibatch(agent_state, minibatch):
|
||||||
logger.info(f"[PPO] minibatch.obs: {minibatch.obs.shape}")
|
logger.info(f"[update_ppo] minibatch.obs: {minibatch.obs.shape}")
|
||||||
logger.info(f"[PPO] minibatch.actions: {minibatch.actions.shape}")
|
logger.info(f"[update_ppo] minibatch.actions: {minibatch.actions.shape}")
|
||||||
logger.info(f"[PPO] minibatch.logprobs: {minibatch.logprobs.shape}")
|
logger.info(f"[update_ppo] minibatch.logprobs: {minibatch.logprobs.shape}")
|
||||||
logger.info(f"[PPO] minibatch.advantages: {minibatch.advantages.shape}")
|
logger.info(f"[update_ppo] minibatch.advantages: {minibatch.advantages.shape}")
|
||||||
logger.info(f"[PPO] minibatch.returns: {minibatch.returns.shape}")
|
logger.info(f"[update_ppo] minibatch.returns: {minibatch.returns.shape}")
|
||||||
(loss, (pg_loss, v_loss, entropy_loss, approx_kl)), grads = ppo_loss_grad_fn(
|
(loss, (pg_loss, v_loss, entropy_loss, approx_kl)), grads = ppo_loss_grad_fn(
|
||||||
agent_state.params,
|
agent_state.params,
|
||||||
minibatch.obs,
|
minibatch.obs,
|
||||||
|
|
@ -110,25 +110,25 @@ def get_action_and_value(
|
||||||
hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x)
|
hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x)
|
||||||
hidden_sensor = message_passer(hidden_sensor)
|
hidden_sensor = message_passer(hidden_sensor)
|
||||||
|
|
||||||
logger.info(f"[SHAPE] hidden_sensor: {hidden_sensor.shape}")
|
logger.info(f"[get_action_and_value] hidden_sensor: {hidden_sensor.shape}")
|
||||||
logger.info(f"[SHAPE] hidden_critic: {hidden_critic.shape}")
|
logger.info(f"[get_action_and_value] hidden_critic: {hidden_critic.shape}")
|
||||||
|
|
||||||
mean, log_std = actor_apply(params["actor_params"], hidden_sensor)
|
mean, log_std = actor_apply(params["actor_params"], hidden_sensor)
|
||||||
|
|
||||||
logger.info(f"[SHAPE] mean: {mean.shape}")
|
logger.info(f"[get_action_and_value] mean: {mean.shape}")
|
||||||
logger.info(f"[SHAPE] log_std: {log_std.shape}")
|
logger.info(f"[get_action_and_value] log_std: {log_std.shape}")
|
||||||
logger.info(f"[SHAPE] action: {action.shape}")
|
logger.info(f"[get_action_and_value] action: {action.shape}")
|
||||||
|
|
||||||
log_std = jnp.clip(log_std, -5, 2)
|
log_std = jnp.clip(log_std, -5, 2)
|
||||||
std = jnp.exp(log_std)
|
std = jnp.exp(log_std)
|
||||||
|
|
||||||
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi))
|
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi))
|
||||||
logger.info(f"[SHAPE] logprob pre-sum: {logprob.shape}")
|
logger.info(f"[get_action_and_value] logprob pre-sum: {logprob.shape}")
|
||||||
logprob = logprob.sum(axis=(-2, -1))
|
logprob = logprob.sum(axis=(-2, -1))
|
||||||
logger.info(f"[SHAPE] logprob final: {logprob.shape}")
|
logger.info(f"[get_action_and_value] logprob final: {logprob.shape}")
|
||||||
entropy = (0.5 + 0.5 * jnp.log(2 * jnp.pi) + log_std).sum(axis=(-2, -1))
|
entropy = (0.5 + 0.5 * jnp.log(2 * jnp.pi) + log_std).sum(axis=(-2, -1))
|
||||||
value = critic_apply(params["critic_params"], hidden_critic).squeeze(-1)
|
value = critic_apply(params["critic_params"], hidden_critic).squeeze(-1)
|
||||||
logger.info(f"[SHAPE] value: {value.shape}")
|
logger.info(f"[get_action_and_value] value: {value.shape}")
|
||||||
return logprob, entropy, value
|
return logprob, entropy, value
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -27,6 +27,7 @@ from brittle_star_project.MLPs.mlps import (
|
||||||
from brittle_star_project.ppo import PPO
|
from brittle_star_project.ppo import PPO
|
||||||
from brittle_star_project.environment import MorphMode
|
from brittle_star_project.environment import MorphMode
|
||||||
|
|
||||||
|
logger11 = get_logger()
|
||||||
# TODO: move to config
|
# TODO: move to config
|
||||||
_ALLOWED_OBS_KEYS = {
|
_ALLOWED_OBS_KEYS = {
|
||||||
"joint_position",
|
"joint_position",
|
||||||
|
|
@ -266,7 +267,23 @@ def _step_once(
|
||||||
action_high,
|
action_high,
|
||||||
adj_matrix,
|
adj_matrix,
|
||||||
)
|
)
|
||||||
|
logger11.info(f"[_step_once] raw_action: {raw_action.shape}")
|
||||||
|
logger11.info(f"[_step_once] clipped_action: {clipped_action.shape}")
|
||||||
|
|
||||||
|
# Supporting signals (often where mismatch originates)
|
||||||
|
logger11.info(f"[_step_once] logprob: {logprob.shape}")
|
||||||
|
logger11.info(f"[_step_once] value: {value.shape}")
|
||||||
|
logger11.info(f"[_step_once] mean: {mean.shape}")
|
||||||
|
logger11.info(f"[_step_once] std: {std.shape}")
|
||||||
|
|
||||||
|
# ---- ENV STEP ----
|
||||||
|
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
|
||||||
|
episode_stats, env_state, clipped_action
|
||||||
|
)
|
||||||
|
|
||||||
|
logger11.info(f"[_step_once] next_obs: {next_obs.shape}")
|
||||||
|
logger11.info(f"[_step_once] reward: {reward.shape}")
|
||||||
|
logger11.info(f"[_step_once] next_done: {next_done.shape}")
|
||||||
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
|
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
|
||||||
episode_stats, env_state, clipped_action
|
episode_stats, env_state, clipped_action
|
||||||
)
|
)
|
||||||
|
|
@ -582,10 +599,12 @@ 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.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.info(f"[_init_agent_state] obs_var: {self.obs_var.shape}")
|
||||||
|
|
||||||
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)
|
||||||
|
|
@ -593,24 +612,35 @@ class PPOTrainer:
|
||||||
|
|
||||||
# (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(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(f"[_init_agent_state] 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}")
|
||||||
|
|
||||||
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(f"[_init_agent_state] actor_params: {jax.tree.map(lambda x: x.shape, actor_params)}")
|
||||||
|
|
||||||
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)
|
k, self.sensor.apply(single_sensor_param, sample_obs)
|
||||||
)
|
)
|
||||||
)(message_passer_keys)
|
)(message_passer_keys)
|
||||||
|
self.logger.info(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
|
||||||
|
self.logger.info(f"[_init_agent_state] flat_obs: {flat_obs.shape}")
|
||||||
|
|
||||||
flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic
|
|
||||||
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(f"[_init_agent_state] feature_extractor_params: {jax.tree.map(lambda x: x.shape, feature_extractor_params)}")
|
||||||
|
|
||||||
critic_params = self.critic.init(
|
critic_input = self.feature_extractor.apply(feature_extractor_params, flat_obs)
|
||||||
critic_key,
|
self.logger.info(f"[_init_agent_state] critic_input: {critic_input.shape}")
|
||||||
self.feature_extractor.apply(feature_extractor_params, flat_obs)
|
|
||||||
)
|
critic_params = self.critic.init(critic_key, critic_input)
|
||||||
|
self.logger.info(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,
|
||||||
|
|
@ -743,7 +773,7 @@ class PPOTrainer:
|
||||||
def _step(self, env_state, next_obs, next_done, iteration: int) -> tuple:
|
def _step(self, env_state, next_obs, next_done, iteration: int) -> tuple:
|
||||||
if iteration == 1:
|
if iteration == 1:
|
||||||
self.logger.log_non_interactive(f"Starting first rollout (JIT): {time.ctime()}")
|
self.logger.log_non_interactive(f"Starting first rollout (JIT): {time.ctime()}")
|
||||||
|
self.logger.info(f"[_step] next_obs (in): {next_obs.shape}")
|
||||||
(
|
(
|
||||||
self.agent_state,
|
self.agent_state,
|
||||||
self.episode_stats,
|
self.episode_stats,
|
||||||
|
|
@ -753,7 +783,7 @@ class PPOTrainer:
|
||||||
self.key,
|
self.key,
|
||||||
next_env_state,
|
next_env_state,
|
||||||
) = self._rollout(env_state, next_obs, next_done)
|
) = self._rollout(env_state, next_obs, next_done)
|
||||||
|
self.logger.info(f"[_step] next_obs (post-rollout): {next_obs.shape}")
|
||||||
if iteration == 1:
|
if iteration == 1:
|
||||||
self.logger.log_non_interactive(f"First rollout completed: {time.ctime()}")
|
self.logger.log_non_interactive(f"First rollout completed: {time.ctime()}")
|
||||||
|
|
||||||
|
|
@ -839,6 +869,7 @@ class PPOTrainer:
|
||||||
self.morph_mode,
|
self.morph_mode,
|
||||||
self.segments_per_arm,
|
self.segments_per_arm,
|
||||||
)
|
)
|
||||||
|
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_)
|
||||||
|
|
||||||
self.logger.log_non_interactive(f"Initial reset completed: {time.ctime()}")
|
self.logger.log_non_interactive(f"Initial reset completed: {time.ctime()}")
|
||||||
|
|
@ -853,6 +884,7 @@ 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.logger.info(f"[train] next_obs (post-step): {next_obs.shape}")
|
||||||
self._update_obs_stats(next_obs)
|
self._update_obs_stats(next_obs)
|
||||||
next_obs = _normalize_obs(next_obs, self.obs_mean, self.obs_var)
|
next_obs = _normalize_obs(next_obs, self.obs_mean, self.obs_var)
|
||||||
|
|
||||||
|
|
|
||||||
Reference in a new issue