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
|
@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)
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
|
||||||
Reference in a new issue