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:
parent
70facf51cf
commit
0ed2dfe4ef
2 changed files with 19 additions and 11 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
Reference in a new issue