diff --git a/src/brittle_star_project/MLPs/mlps.py b/src/brittle_star_project/MLPs/mlps.py index 6abb540..83d7389 100644 --- a/src/brittle_star_project/MLPs/mlps.py +++ b/src/brittle_star_project/MLPs/mlps.py @@ -49,14 +49,14 @@ class AgentParams: @jax.tree_util.register_dataclass @dataclass class Storage: - obs: jnp.array - actions: jnp.array - logprobs: jnp.array - dones: jnp.array - values: jnp.array - advantages: jnp.array - returns: jnp.array - rewards: jnp.array + obs: jax.Array + actions: jax.Array + logprobs: jax.Array + dones: jax.Array + values: jax.Array + advantages: jax.Array + returns: jax.Array + rewards: jax.Array def replace(self, **kwargs) -> "Storage": fs = fields(self) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 8c238f1..197036b 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -169,11 +169,19 @@ def _compute_gae_once(carry, inp, gamma, gae_lambda): # jit applied on partial-wrapped wrapper method self._compute_gae_jit def _compute_gae_jit( - agent_state, storage, next_obs, next_done, gamma, gae_lambda, num_envs, sensor, critic + agent_state, + storage, + next_obs, + next_done, + gamma, + gae_lambda, + num_envs, + feature_extractor, + critic, ): next_value = critic.apply( agent_state.params["critic_params"], - sensor.apply(agent_state.params["sensor_params"], next_obs), + feature_extractor.apply(agent_state.params["sensor_params"], next_obs), ).squeeze(-1) advantages = jnp.zeros((num_envs,)) @@ -232,7 +240,7 @@ class PPOTrainer: num_envs=self.args.num_envs, gamma=self.args.gamma, gae_lambda=self.args.gae_lambda, - sensor=self.sensor, + feature_extractor=self.feature_extractor, critic=self.critic, ) )