fix(PPOTrainer.py): bug fixes regarding jax.jit
This commit is contained in:
parent
a9302cfd4e
commit
62e11ce243
2 changed files with 44 additions and 18 deletions
|
|
@ -40,7 +40,7 @@ def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def get_action_and_value_noise(
|
def _get_action_and_value_noise(
|
||||||
sensor: GenericDenseLayersWithActivation,
|
sensor: GenericDenseLayersWithActivation,
|
||||||
feature_extractor: GenericDenseLayersWithActivation,
|
feature_extractor: GenericDenseLayersWithActivation,
|
||||||
actor: Actor,
|
actor: Actor,
|
||||||
|
|
@ -76,7 +76,7 @@ def _step_once(
|
||||||
critic: OneDenseLayerMLP,
|
critic: OneDenseLayerMLP,
|
||||||
):
|
):
|
||||||
agent_state, episode_stats, obs, done, key, env_state = carry
|
agent_state, episode_stats, obs, done, key, env_state = carry
|
||||||
action, logprob, value, key = get_action_and_value_noise(
|
action, logprob, value, key = _get_action_and_value_noise(
|
||||||
sensor, feature_extractor, actor, critic, agent_state, obs, key
|
sensor, feature_extractor, actor, critic, agent_state, obs, key
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -127,7 +127,10 @@ def _step_env_wrapped(env_step_fn, env_state, action, episode_stats):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@partial(
|
||||||
|
jax.jit,
|
||||||
|
static_argnames=("env_step_fn", "sensor", "feature_extractor", "actor", "critic"),
|
||||||
|
)
|
||||||
def _rollout_jit(
|
def _rollout_jit(
|
||||||
agent_state,
|
agent_state,
|
||||||
episode_stats,
|
episode_stats,
|
||||||
|
|
@ -135,7 +138,7 @@ def _rollout_jit(
|
||||||
next_obs,
|
next_obs,
|
||||||
next_done,
|
next_done,
|
||||||
key,
|
key,
|
||||||
args,
|
max_steps,
|
||||||
env_step_fn,
|
env_step_fn,
|
||||||
sensor: GenericDenseLayersWithActivation,
|
sensor: GenericDenseLayersWithActivation,
|
||||||
feature_extractor: GenericDenseLayersWithActivation,
|
feature_extractor: GenericDenseLayersWithActivation,
|
||||||
|
|
@ -153,13 +156,13 @@ def _rollout_jit(
|
||||||
),
|
),
|
||||||
(agent_state, episode_stats, next_obs, next_done, key, env_state),
|
(agent_state, episode_stats, next_obs, next_done, key, env_state),
|
||||||
(),
|
(),
|
||||||
args.max_steps,
|
max_steps,
|
||||||
)
|
)
|
||||||
return agent_state, episode_stats, next_obs, next_done, storage, key, env_state
|
return agent_state, episode_stats, next_obs, next_done, storage, key, env_state
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def compute_gae_once(carry, inp, gamma, gae_lambda):
|
def _compute_gae_once(carry, inp, gamma, gae_lambda):
|
||||||
advantages = carry
|
advantages = carry
|
||||||
nextdone, nextvalues, curvalues, reward = inp
|
nextdone, nextvalues, curvalues, reward = inp
|
||||||
nextnonterminal = 1.0 - nextdone
|
nextnonterminal = 1.0 - nextdone
|
||||||
|
|
@ -168,18 +171,23 @@ def compute_gae_once(carry, inp, gamma, gae_lambda):
|
||||||
return advantages, advantages
|
return advantages, advantages
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@partial(
|
||||||
def compute_gae_jit(agent_state, storage, next_obs, next_done, sensor, critic, args):
|
jax.jit,
|
||||||
|
static_argnames=("sensor", "critic"),
|
||||||
|
)
|
||||||
|
def _compute_gae_jit(
|
||||||
|
agent_state, storage, next_obs, next_done, sensor, critic, gamma, gae_lambda, num_envs
|
||||||
|
):
|
||||||
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),
|
sensor.apply(agent_state.params["sensor_params"], next_obs),
|
||||||
).squeeze(-1)
|
).squeeze(-1)
|
||||||
|
|
||||||
advantages = jnp.zeros((args.num_envs,))
|
advantages = jnp.zeros((num_envs,))
|
||||||
dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0)
|
dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0)
|
||||||
values = jnp.concatenate([storage.values, next_value[None, :]], axis=0)
|
values = jnp.concatenate([storage.values, next_value[None, :]], axis=0)
|
||||||
_, advantages = jax.lax.scan(
|
_, advantages = jax.lax.scan(
|
||||||
partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda),
|
partial(_compute_gae_once, gamma=gamma, gae_lambda=gae_lambda),
|
||||||
advantages,
|
advantages,
|
||||||
(dones[1:], values[1:], values[:-1], storage.rewards),
|
(dones[1:], values[1:], values[:-1], storage.rewards),
|
||||||
reverse=True,
|
reverse=True,
|
||||||
|
|
@ -213,6 +221,19 @@ class PPOTrainer:
|
||||||
self.actor.apply = jax.jit(self.actor.apply)
|
self.actor.apply = jax.jit(self.actor.apply)
|
||||||
self.critic.apply = jax.jit(self.critic.apply)
|
self.critic.apply = jax.jit(self.critic.apply)
|
||||||
|
|
||||||
|
self._rollout_jit = partial(
|
||||||
|
_rollout_jit,
|
||||||
|
sensor=self.sensor,
|
||||||
|
feature_extractor=self.feature_extractor,
|
||||||
|
actor=self.actor,
|
||||||
|
critic=self.critic,
|
||||||
|
)
|
||||||
|
self._compute_gae_jit = partial(
|
||||||
|
_compute_gae_jit,
|
||||||
|
sensor=self.sensor,
|
||||||
|
critic=self.critic,
|
||||||
|
)
|
||||||
|
|
||||||
self._ppo = PPO(self.args, self.sensor, self.actor, self.critic, self.feature_extractor)
|
self._ppo = PPO(self.args, self.sensor, self.actor, self.critic, self.feature_extractor)
|
||||||
|
|
||||||
self.agent_state = self._init_agent_state()
|
self.agent_state = self._init_agent_state()
|
||||||
|
|
@ -287,24 +308,26 @@ class PPOTrainer:
|
||||||
)
|
)
|
||||||
|
|
||||||
def _rollout(self, env_state, next_obs, next_done) -> tuple[Storage, ...]:
|
def _rollout(self, env_state, next_obs, next_done) -> tuple[Storage, ...]:
|
||||||
return _rollout_jit(
|
return self._rollout_jit(
|
||||||
self.agent_state,
|
self.agent_state,
|
||||||
self.episode_stats,
|
self.episode_stats,
|
||||||
env_state,
|
env_state,
|
||||||
next_obs,
|
next_obs,
|
||||||
next_done,
|
next_done,
|
||||||
self.key,
|
self.key,
|
||||||
self.args,
|
self.args.num_steps,
|
||||||
self.env.step,
|
self.env.step,
|
||||||
self.sensor,
|
|
||||||
self.feature_extractor,
|
|
||||||
self.actor,
|
|
||||||
self.critic,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _compute_gae(self, storage, next_obs, next_done) -> Storage:
|
def _compute_gae(self, storage, next_obs, next_done) -> Storage:
|
||||||
return compute_gae_jit(
|
return self._compute_gae_jit(
|
||||||
self.agent_state, storage, next_obs, next_done, self.sensor, self.critic, self.args
|
self.agent_state,
|
||||||
|
storage,
|
||||||
|
next_obs,
|
||||||
|
next_done,
|
||||||
|
self.args.gamma,
|
||||||
|
self.args.gae_lambda,
|
||||||
|
self.args.num_envs,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _log(
|
def _log(
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,9 @@
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
import jax
|
||||||
|
|
||||||
|
|
||||||
|
@jax.tree_util.register_dataclass
|
||||||
@dataclass
|
@dataclass
|
||||||
class PPOArgs:
|
class PPOArgs:
|
||||||
"""
|
"""
|
||||||
|
|
|
||||||
Reference in a new issue