diff --git a/configs/ppo/smoke_test.yaml b/configs/ppo/smoke_test.yaml index 40e4c07..4b7bb84 100644 --- a/configs/ppo/smoke_test.yaml +++ b/configs/ppo/smoke_test.yaml @@ -2,9 +2,9 @@ # Lower timestep count for quick iterations/testing. learning_rate: 0.0005 -total_timesteps: 65536 -num_envs: 512 -num_steps: 128 +total_timesteps: 1024 +num_envs: 32 +num_steps: 32 anneal_lr: true gamma: 0.99 gae_lambda: 0.95 diff --git a/scripts/train.py b/scripts/train.py index c367000..8e61bd0 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -8,6 +8,7 @@ from brittle_star_project.configs.register_configs import register_configs from brittle_star_project.trainers.PPOTrainer import PPOTrainer from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from experiment_logger import init_logger, get_logger +import logging def make_env(cfg: BrittleStarConfig) -> BrittleStarJaxEnvWrapper: @@ -41,6 +42,7 @@ def main(dict_cfg: DictConfig): base_dir=os.path.dirname(run_dir), ) logger = get_logger() + logger.set_level(logging.DEBUG) logger.info(f"Hydra-initialized run: {run_name}") logger.info(f"Output directory: {run_dir}") diff --git a/src/brittle_star_project/dataclasses/EpisodeStatistics.py b/src/brittle_star_project/dataclasses/EpisodeStatistics.py index 0b2832f..ff2982e 100644 --- a/src/brittle_star_project/dataclasses/EpisodeStatistics.py +++ b/src/brittle_star_project/dataclasses/EpisodeStatistics.py @@ -4,7 +4,7 @@ import jax.numpy as jnp @flax.struct.dataclass class EpisodeStatistics: - episode_returns: jnp.array - episode_lengths: jnp.array - returned_episode_returns: jnp.array - returned_episode_lengths: jnp.array + episode_returns: jnp.ndarray + episode_lengths: jnp.ndarray + returned_episode_returns: jnp.ndarray + returned_episode_lengths: jnp.ndarray diff --git a/src/brittle_star_project/ppo.py b/src/brittle_star_project/ppo.py index 6eb5221..87e1dfe 100644 --- a/src/brittle_star_project/ppo.py +++ b/src/brittle_star_project/ppo.py @@ -1,15 +1,27 @@ from functools import partial -import flax import jax import jax.numpy as jnp +from jax import debug +from flax.core import FrozenDict from experiment_logger import get_logger +from brittle_star_project.utils import logged_jit logger = get_logger() + + # Chose to use a class as it seemed the easiest way to integrate the CleanRL code style # with our need to seperate concerns class PPO: - def __init__(self, args, sensor_apply, actor_apply, critic_apply, feature_extractor_apply, message_passer=None): + def __init__( + self, + args, + sensor_apply, + actor_apply, + critic_apply, + feature_extractor_apply, + message_passer=None, + ): self.args = args if not message_passer: @@ -30,13 +42,13 @@ class PPO: # This PPO class should be initialized only once, # or this function will need to recompile - @partial(jax.jit, static_argnums=0) + @partial(logged_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}") + debug.callback(logger.debug, f"[PPO] storage.obs shape: {storage.obs.shape}") + debug.callback(logger.debug, f"[PPO] storage.actions shape: {storage.actions.shape}") + debug.callback(logger.debug, f"[PPO] storage.logprobs shape: {storage.logprobs.shape}") + debug.callback(logger.debug, f"[PPO] storage.advantages shape: {storage.advantages.shape}") + debug.callback(logger.debug, f"[PPO] storage.returns shape: {storage.returns.shape}") args = self.args ppo_loss_grad_fn = self.ppo_loss_grad_fn @@ -56,11 +68,16 @@ 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}") + debug.callback(logger.debug, f"[PPO] minibatch.obs: {minibatch.obs.shape}") + debug.callback(logger.debug, f"[PPO] minibatch.actions: {minibatch.actions.shape}") + debug.callback( + logger.debug, f"[PPO] minibatch.logprobs: {minibatch.logprobs.shape}" + ) + debug.callback( + logger.debug, f"[PPO] minibatch.advantages: {minibatch.advantages.shape}" + ) + debug.callback(logger.debug, f"[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, @@ -70,13 +87,7 @@ class PPO: minibatch.returns, ) agent_state = agent_state.apply_gradients(grads=grads) - return agent_state, ( - loss, - pg_loss, - v_loss, - entropy_loss, - approx_kl - ) + return agent_state, (loss, pg_loss, v_loss, entropy_loss, approx_kl) agent_state, metrics = jax.lax.scan(update_minibatch, agent_state, shuffled_storage) return (agent_state, key), metrics @@ -95,14 +106,14 @@ that are now not in the same scope """ -@partial(jax.jit, static_argnums=(0, 1, 2, 3, 4)) +@partial(logged_jit, static_argnums=(0, 1, 2, 3, 4)) def get_action_and_value( sensor_apply, actor_apply, message_passer, critic_apply, feature_extractor_apply, - params: flax.core.FrozenDict, + params: FrozenDict, x: jnp.ndarray, action: jnp.ndarray, ): @@ -110,25 +121,28 @@ 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}") + debug.callback(logger.debug, f"[SHAPE] hidden_sensor: {hidden_sensor.shape}") + debug.callback(logger.debug, f"[SHAPE] 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}") + debug.callback(logger.debug, f"[SHAPE] mean: {mean.shape}") + debug.callback(logger.debug, f"[SHAPE] log_std: {log_std.shape}") + debug.callback(logger.debug, f"[SHAPE] 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}") + debug.callback(logger.debug, f"[SHAPE] logprob pre-sum: {logprob.shape}") + logprob = logprob.sum(axis=(-2, -1)) - logger.info(f"[SHAPE] logprob final: {logprob.shape}") + debug.callback(logger.debug, f"[SHAPE] 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}") + debug.callback(logger.debug, f"[SHAPE] 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..fe0eb02 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -26,6 +26,7 @@ from brittle_star_project.MLPs.mlps import ( ) from brittle_star_project.ppo import PPO from brittle_star_project.environment import MorphMode +from brittle_star_project.utils import logged_jit # TODO: move to config _ALLOWED_OBS_KEYS = { @@ -108,7 +109,7 @@ def build_adjacency(segments_per_arm, mode: MorphMode): return adj -@jax.jit +@logged_jit def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray: return jnp.clip(action, low, high) @@ -119,13 +120,13 @@ def _compute_explained_variance(values: jnp.ndarray, returns: jnp.ndarray) -> fl return float(explained_var) -@jax.jit +@logged_jit def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate): frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations return learning_rate * frac -@jax.jit +@logged_jit def _normalize_obs(obs, mean, var, eps=1e-8): return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0) @@ -134,7 +135,7 @@ def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm num_segments = segments_per_arm.sum() num_arms = jnp.where(segments_per_arm > 0, 1, 0).sum() - @jax.jit + @logged_jit def _filter_and_flatten(o) -> jnp.ndarray: # vmap feeds one env at a time — v has NO batch dim here # shapes are e.g. (n_features,) or (n_nodes, feat) @@ -230,11 +231,12 @@ def _get_action_and_value_noise( raw_action = mean + noise * std clipped_action = _clip_action(raw_action, action_low, action_high) - logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) value = apply_shared(critic, agent_state.params["critic_params"], hidden_critic) - - raw_action = raw_action.reshape(raw_action.shape[0], -1) # concat the per agent, keep the envs dim + + raw_action = raw_action.reshape( + raw_action.shape[0], -1 + ) # concat the per agent, keep the envs dim clipped_action = _clip_action(raw_action, action_low, action_high) return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key @@ -338,6 +340,7 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, morph_mode, ), ) + def apply_per_node(net, params, x): # params: (nodes, ...) # x: (batch, nodes, feat) @@ -348,6 +351,7 @@ def apply_per_node(net, params, x): return jax.vmap(apply_single_node, in_axes=(0, 1), out_axes=1)(params, x) + def apply_shared(net, params, x): # x: (batch, nodes, feat) # If the critic expects a single vector per environment: @@ -355,6 +359,7 @@ def apply_shared(net, params, x): x_flattened = x.reshape(batch_size, -1) return jax.vmap(lambda xi: net.apply(params, xi))(x_flattened) + # TODO: update to work with extra dimension + message passing def _rollout_jit( agent_state, @@ -415,9 +420,10 @@ def _compute_gae_jit( critic, adj_matrix: jnp.ndarray, ): - next_value = apply_shared(critic, + next_value = apply_shared( + critic, agent_state.params["critic_params"], - apply_shared(feature_extractor,agent_state.params["feature_extractor_params"], next_obs), + apply_shared(feature_extractor, agent_state.params["feature_extractor_params"], next_obs), ).squeeze(-1) advantages = jnp.zeros((num_envs,)) @@ -487,15 +493,15 @@ class PPOTrainer: self.needed_copies, ) = self._init_agent() - self.sensor.apply = jax.jit(self.sensor.apply) - self.feature_extractor.apply = jax.jit(self.feature_extractor.apply) - self.actor.apply = jax.jit(self.actor.apply) - self.critic.apply = jax.jit(self.critic.apply) + self.sensor.apply = logged_jit(self.sensor.apply) + self.feature_extractor.apply = logged_jit(self.feature_extractor.apply) + self.actor.apply = logged_jit(self.actor.apply) + self.critic.apply = logged_jit(self.critic.apply) 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) - self._rollout_jit = jax.jit( + self._rollout_jit = logged_jit( partial( _rollout_jit, max_steps=self.ppo.num_steps, @@ -515,7 +521,7 @@ class PPOTrainer: adj_matrix=self.adj, ) ) - self._compute_gae_jit = jax.jit( + self._compute_gae_jit = logged_jit( partial( _compute_gae_jit, num_envs=self.ppo.num_envs, @@ -527,10 +533,17 @@ class PPOTrainer: ) ) - apply_sensor = lambda p, x: apply_per_node(self.sensor, p, x) - apply_actor = lambda p, x: apply_per_node(self.actor, p, x) - apply_critic = lambda p, x: apply_shared(self.critic, p, x) - apply_feature = lambda p, x: apply_shared(self.feature_extractor, p, x) + def apply_sensor(p, x): + return apply_per_node(self.sensor, p, x) + + def apply_actor(p, x): + return apply_per_node(self.actor, p, x) + + def apply_critic(p, x): + return apply_shared(self.critic, p, x) + + def apply_feature(p, x): + return apply_shared(self.feature_extractor, p, x) self._ppo = PPO(self.ppo, apply_sensor, apply_actor, apply_critic, apply_feature) @@ -553,11 +566,11 @@ class PPOTrainer: case MorphMode.CENTRALIZED: needed_copies = 1 case MorphMode.FULLY_CONNECTED | MorphMode.RING: - needed_copies = jnp.where(self.segments_per_arm > 0, 1, 0).sum() + needed_copies = jnp.where(self.segments_per_arm > 0, 1, 0).sum().item() case MorphMode.SEGMENT: needed_copies = ( self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum() - ) + ).item() actor = Actor(action_dim=self.env.single_action_space.shape[0]) sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) @@ -604,12 +617,11 @@ class PPOTrainer: ) )(message_passer_keys) - flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic + flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic feature_extractor_params = self.feature_extractor.init(feature_extractor_key, flat_obs) critic_params = self.critic.init( - critic_key, - self.feature_extractor.apply(feature_extractor_params, flat_obs) + critic_key, self.feature_extractor.apply(feature_extractor_params, flat_obs) ) return TrainState.create( diff --git a/src/brittle_star_project/utils/__init__.py b/src/brittle_star_project/utils/__init__.py new file mode 100644 index 0000000..ed72e5a --- /dev/null +++ b/src/brittle_star_project/utils/__init__.py @@ -0,0 +1,3 @@ +from .logged_jit import logged_jit + +__all__ = ["logged_jit"] diff --git a/src/brittle_star_project/utils/logged_jit.py b/src/brittle_star_project/utils/logged_jit.py new file mode 100644 index 0000000..3d29ba6 --- /dev/null +++ b/src/brittle_star_project/utils/logged_jit.py @@ -0,0 +1,17 @@ +import jax +from experiment_logger import get_logger + + +def logged_jit(fn, **jit_kwargs): + logger = get_logger() + name = getattr(fn, "__name__", getattr(fn, "__qualname__", repr(fn))) + + def decorator(func): + def traced_func(*args, **kwargs): + logger.debug(f"[JIT] Compiling {name}...") + return func(*args, **kwargs) + + jitted = jax.jit(traced_func, **jit_kwargs) + return jitted + + return decorator(fn)