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
@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)

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
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,
)
)