diff --git a/src/brittle_star_project/ppo.py b/src/brittle_star_project/ppo.py index 6eb5221..38dc18b 100644 --- a/src/brittle_star_project/ppo.py +++ b/src/brittle_star_project/ppo.py @@ -32,11 +32,11 @@ class PPO: # or this function will need to recompile @partial(jax.jit, static_argnums=0) def update_ppo(self, agent_state, storage, key): - logger.info(f"[PPO] storage.obs shape: {getattr(storage, 'obs', None).shape}") - logger.info(f"[PPO] storage.actions shape: {storage.actions.shape}") - logger.info(f"[PPO] storage.logprobs shape: {storage.logprobs.shape}") - logger.info(f"[PPO] storage.advantages shape: {storage.advantages.shape}") - logger.info(f"[PPO] storage.returns shape: {storage.returns.shape}") + logger.info(f"[update_ppo] storage.obs shape: {getattr(storage, 'obs', None).shape}") + logger.info(f"[update_ppo] storage.actions shape: {storage.actions.shape}") + logger.info(f"[update_ppo] storage.logprobs shape: {storage.logprobs.shape}") + logger.info(f"[update_ppo] storage.advantages shape: {storage.advantages.shape}") + logger.info(f"[update_ppo] storage.returns shape: {storage.returns.shape}") args = self.args ppo_loss_grad_fn = self.ppo_loss_grad_fn @@ -56,11 +56,11 @@ class PPO: shuffled_storage = jax.tree.map(convert_data, flatten_storage) def update_minibatch(agent_state, minibatch): - logger.info(f"[PPO] minibatch.obs: {minibatch.obs.shape}") - logger.info(f"[PPO] minibatch.actions: {minibatch.actions.shape}") - logger.info(f"[PPO] minibatch.logprobs: {minibatch.logprobs.shape}") - logger.info(f"[PPO] minibatch.advantages: {minibatch.advantages.shape}") - logger.info(f"[PPO] minibatch.returns: {minibatch.returns.shape}") + logger.info(f"[update_ppo] minibatch.obs: {minibatch.obs.shape}") + logger.info(f"[update_ppo] minibatch.actions: {minibatch.actions.shape}") + logger.info(f"[update_ppo] minibatch.logprobs: {minibatch.logprobs.shape}") + logger.info(f"[update_ppo] minibatch.advantages: {minibatch.advantages.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( agent_state.params, minibatch.obs, @@ -110,25 +110,25 @@ def get_action_and_value( hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x) hidden_sensor = message_passer(hidden_sensor) - logger.info(f"[SHAPE] hidden_sensor: {hidden_sensor.shape}") - logger.info(f"[SHAPE] hidden_critic: {hidden_critic.shape}") + logger.info(f"[get_action_and_value] hidden_sensor: {hidden_sensor.shape}") + logger.info(f"[get_action_and_value] hidden_critic: {hidden_critic.shape}") mean, log_std = actor_apply(params["actor_params"], hidden_sensor) - logger.info(f"[SHAPE] mean: {mean.shape}") - logger.info(f"[SHAPE] log_std: {log_std.shape}") - logger.info(f"[SHAPE] action: {action.shape}") + logger.info(f"[get_action_and_value] mean: {mean.shape}") + logger.info(f"[get_action_and_value] log_std: {log_std.shape}") + logger.info(f"[get_action_and_value] action: {action.shape}") log_std = jnp.clip(log_std, -5, 2) std = jnp.exp(log_std) 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)) - 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)) 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 diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 4a71b69..4684c17 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -27,6 +27,7 @@ from brittle_star_project.MLPs.mlps import ( from brittle_star_project.ppo import PPO from brittle_star_project.environment import MorphMode +logger11 = get_logger() # TODO: move to config _ALLOWED_OBS_KEYS = { "joint_position", @@ -266,7 +267,23 @@ def _step_once( action_high, 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, clipped_action ) @@ -582,10 +599,12 @@ class PPOTrainer: self.morph_mode, self.segments_per_arm, )[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_var = jnp.ones((len(sample_obs),)) 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) actor_keys = jax.random.split(actor_key, self.needed_copies) @@ -593,24 +612,35 @@ class PPOTrainer: # (needed_copies, 175) 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) + 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) + 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) + self.logger.info(f"[_init_agent_state] actor_params: {jax.tree.map(lambda x: x.shape, actor_params)}") message_passer_params = jax.vmap( lambda k: self.message_passer.init( k, self.sensor.apply(single_sensor_param, sample_obs) ) )(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) + 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_key, - 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}") + + 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( apply_fn=None, @@ -743,7 +773,7 @@ class PPOTrainer: def _step(self, env_state, next_obs, next_done, iteration: int) -> tuple: if iteration == 1: 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.episode_stats, @@ -753,7 +783,7 @@ class PPOTrainer: self.key, next_env_state, ) = self._rollout(env_state, next_obs, next_done) - + self.logger.info(f"[_step] next_obs (post-rollout): {next_obs.shape}") if iteration == 1: self.logger.log_non_interactive(f"First rollout completed: {time.ctime()}") @@ -839,6 +869,7 @@ class PPOTrainer: self.morph_mode, 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_) 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, iteration=iteration ) + self.logger.info(f"[train] next_obs (post-step): {next_obs.shape}") self._update_obs_stats(next_obs) next_obs = _normalize_obs(next_obs, self.obs_mean, self.obs_var)