1
Fork 0

fix(PPOTrainer.py): fixed usage of sensor output in PPOTrainer, correct way uses the output of the feature extractor instead

This commit is contained in:
Robin Meersman 2026-04-09 19:55:39 +02:00
parent 70facf51cf
commit 0ed2dfe4ef
2 changed files with 19 additions and 11 deletions

View file

@ -49,14 +49,14 @@ class AgentParams:
@jax.tree_util.register_dataclass @jax.tree_util.register_dataclass
@dataclass @dataclass
class Storage: class Storage:
obs: jnp.array obs: jax.Array
actions: jnp.array actions: jax.Array
logprobs: jnp.array logprobs: jax.Array
dones: jnp.array dones: jax.Array
values: jnp.array values: jax.Array
advantages: jnp.array advantages: jax.Array
returns: jnp.array returns: jax.Array
rewards: jnp.array rewards: jax.Array
def replace(self, **kwargs) -> "Storage": def replace(self, **kwargs) -> "Storage":
fs = fields(self) fs = fields(self)

View file

@ -169,11 +169,19 @@ def _compute_gae_once(carry, inp, gamma, gae_lambda):
# jit applied on partial-wrapped wrapper method self._compute_gae_jit # jit applied on partial-wrapped wrapper method self._compute_gae_jit
def _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( next_value = critic.apply(
agent_state.params["critic_params"], 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) ).squeeze(-1)
advantages = jnp.zeros((num_envs,)) advantages = jnp.zeros((num_envs,))
@ -232,7 +240,7 @@ class PPOTrainer:
num_envs=self.args.num_envs, num_envs=self.args.num_envs,
gamma=self.args.gamma, gamma=self.args.gamma,
gae_lambda=self.args.gae_lambda, gae_lambda=self.args.gae_lambda,
sensor=self.sensor, feature_extractor=self.feature_extractor,
critic=self.critic, critic=self.critic,
) )
) )