diff --git a/src/MLPs/mlps.py b/src/MLPs/mlps.py index 1fdc6a2..6abb540 100644 --- a/src/MLPs/mlps.py +++ b/src/MLPs/mlps.py @@ -40,10 +40,10 @@ class Actor(nn.Module): @jax.tree_util.register_dataclass @dataclass class AgentParams: - network_params: flax.core.FrozenDict + sensor_params: flax.core.FrozenDict actor_params: flax.core.FrozenDict critic_params: flax.core.FrozenDict - critic_network_params: flax.core.FrozenDict + feature_extractor_params: flax.core.FrozenDict @jax.tree_util.register_dataclass diff --git a/src/ppo.py b/src/ppo.py index 23c1e45..f2dfb82 100644 --- a/src/ppo.py +++ b/src/ppo.py @@ -8,9 +8,7 @@ import jax.numpy as jnp # 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, input_network, action_network, critic, critic_network, message_passer=None - ): + def __init__(self, args, sensor, actor, critic, feature_extractor, message_passer=None): self.args = args if not message_passer: @@ -20,10 +18,10 @@ class PPO: partial( ppo_loss, args=args, - input_network_apply=input_network.apply, - action_network_apply=action_network.apply, + sensor_apply=sensor.apply, + actor_apply=actor.apply, critic_apply=critic.apply, - critic_network_apply=critic_network.apply, + feature_extractor_apply=feature_extractor.apply, message_passer=message_passer, ), has_aux=True, @@ -92,15 +90,15 @@ def get_action_and_value2( action_apply, message_passer, critic_apply, - critic_network_apply, + feature_extractor_apply, params: flax.core.FrozenDict, x: jnp.ndarray, action: jnp.ndarray, ): - hidden_network = input_apply(params["network_params"], x) - hidden_critic = critic_network_apply(params["critic_network_params"], x) - hidden_network = message_passer(hidden_network) - mean, log_std = action_apply(params["actor_params"], hidden_network) + hidden_sensor = input_apply(params["sensor_params"], x) + hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x) + hidden_sensor = message_passer(hidden_sensor) + mean, log_std = action_apply(params["actor_params"], hidden_sensor) std = jnp.exp(log_std) logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) @@ -118,18 +116,18 @@ def ppo_loss( mb_advantages, mb_returns, args, - input_network_apply, - action_network_apply, + sensor_apply, + actor_apply, message_passer, critic_apply, - critic_network_apply, + feature_extractor_apply, ): newlogprob, entropy, newvalue = get_action_and_value2( - input_network_apply, - action_network_apply, + sensor_apply, + actor_apply, message_passer, critic_apply, - critic_network_apply, + feature_extractor_apply, params, x, a, diff --git a/src/train.py b/src/train.py index 10f14bf..1d3bb37 100644 --- a/src/train.py +++ b/src/train.py @@ -72,7 +72,7 @@ def train(args: PPOArgs): random.seed(args.seed) np.random.seed(args.seed) key = jax.random.PRNGKey(args.seed) - key, network_key, actor_key, critic_key, critic_network_key = jax.random.split(key, 5) + key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5) torch.backends.cudnn.deterministic = args.torch_deterministic device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu") @@ -122,8 +122,8 @@ def train(args: PPOArgs): return args.learning_rate * frac print("Initializing the models...") - network = GenericDenseLayersWithActivation() - critic_network = GenericDenseLayersWithActivation() + sensor = GenericDenseLayersWithActivation() + feature_extractor = GenericDenseLayersWithActivation() actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX critic = OneDenseLayerMLP() # messager = OneDenseLayerMLP() @@ -135,15 +135,17 @@ def train(args: PPOArgs): if v.size > 0 ] ) - network_params = network.init(network_key, sample_obs) - critic_network_params = critic_network.init(critic_network_key, sample_obs) - actor_params = actor.init(actor_key, network.apply(network_params, sample_obs)) - critic_params = critic.init(critic_key, critic_network.apply(critic_network_params, sample_obs)) + sensor_params = sensor.init(sensor_key, sample_obs) + feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs) + actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs)) + critic_params = critic.init( + critic_key, feature_extractor.apply(feature_extractor_params, sample_obs) + ) agent_state = TrainState.create( apply_fn=None, params=asdict( - AgentParams(network_params, actor_params, critic_params, critic_network_params) + AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params) ), tx=optax.chain( optax.clip_by_global_norm(args.max_grad_norm), @@ -153,11 +155,11 @@ def train(args: PPOArgs): ), ) - network.apply = jax.jit(network.apply) - critic_network.apply = jax.jit(critic_network.apply) + sensor.apply = jax.jit(sensor.apply) + feature_extractor.apply = jax.jit(feature_extractor.apply) actor.apply = jax.jit(actor.apply) critic.apply = jax.jit(critic.apply) - ppo_instance = PPO(args, network, actor, critic, critic_network) + ppo_instance = PPO(args, sensor, actor, critic, feature_extractor) @jax.jit def get_action_and_value_noise( @@ -165,7 +167,11 @@ def train(args: PPOArgs): next_obs: jnp.ndarray, key: jax.random.PRNGKey, ): - hidden = network.apply(agent_state.params["network_params"], next_obs) + hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) + hidden_critic = feature_extractor.apply( + agent_state.params["feature_extractor_params"], next_obs + ) + # Continuous actions: sample from a Gaussian parameterized by the actor mean, log_std = actor.apply(agent_state.params["actor_params"], hidden) key, subkey = jax.random.split(key) @@ -173,7 +179,7 @@ def train(args: PPOArgs): std = jnp.exp(log_std) action = mean + noise * std logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) - value = critic.apply(agent_state.params["critic_params"], hidden) + value = critic.apply(agent_state.params["critic_params"], hidden_critic) return action, logprob, value.squeeze(-1), key @jax.jit @@ -189,7 +195,7 @@ def train(args: PPOArgs): def compute_gae(agent_state, next_obs, next_done, storage): next_value = critic.apply( agent_state.params["critic_params"], - network.apply(agent_state.params["network_params"], next_obs), + sensor.apply(agent_state.params["sensor_params"], next_obs), ).squeeze(-1) advantages = jnp.zeros((args.num_envs,)) @@ -307,9 +313,10 @@ def train(args: PPOArgs): [ vars(args), [ - agent_state.params["network_params"], + agent_state.params["sensor_params"], agent_state.params["actor_params"], agent_state.params["critic_params"], + agent_state.params["feature_extractor_params"], ], ] )