feature: optional message passer function and clearer naming
This commit is contained in:
parent
ccbbfc18d7
commit
31c2d188b7
1 changed files with 58 additions and 12 deletions
70
src/ppo.py
70
src/ppo.py
|
|
@ -8,14 +8,31 @@ import jax.numpy as jnp
|
||||||
# Chose to use a class as it seemed the easiest way to integrate the CleanRL code style
|
# Chose to use a class as it seemed the easiest way to integrate the CleanRL code style
|
||||||
# with our need to seperate concerns
|
# with our need to seperate concerns
|
||||||
class PPO:
|
class PPO:
|
||||||
def __init__(self, args, network, actor, critic):
|
"""
|
||||||
|
args: params for training
|
||||||
|
input_network: network that takes input state and returns hidden state
|
||||||
|
action_network: network that takes hidden state and returns action distribution
|
||||||
|
critic: value network used for ppo
|
||||||
|
message_passer: function that executes message passing and state agregation X times
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, args, input_network, action_network, critic, message_passer=None
|
||||||
|
):
|
||||||
self.args = args
|
self.args = args
|
||||||
self.network = network
|
|
||||||
self.actor = actor
|
if not message_passer:
|
||||||
self.critic = critic
|
message_passer = identity
|
||||||
|
|
||||||
self.ppo_loss_grad_fn = jax.value_and_grad(
|
self.ppo_loss_grad_fn = jax.value_and_grad(
|
||||||
partial(ppo_loss, args=args, network=network, actor=actor, critic=critic),
|
partial(
|
||||||
|
ppo_loss,
|
||||||
|
args=args,
|
||||||
|
input_network_apply=input_network.apply,
|
||||||
|
action_network_apply=action_network.apply,
|
||||||
|
critic_apply=critic.apply,
|
||||||
|
message_passer=message_passer,
|
||||||
|
),
|
||||||
has_aux=True,
|
has_aux=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -85,17 +102,19 @@ class values would ever change? Better safe than sorry.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
@partial(jax.jit, static_argnums=(0, 1, 2))
|
@partial(jax.jit, static_argnums=(0, 1, 2, 3))
|
||||||
def get_action_and_value2(
|
def get_action_and_value2(
|
||||||
network_apply,
|
input_apply,
|
||||||
actor_apply,
|
action_apply,
|
||||||
|
message_passer,
|
||||||
critic_apply,
|
critic_apply,
|
||||||
params: flax.core.FrozenDict,
|
params: flax.core.FrozenDict,
|
||||||
x: jnp.ndarray,
|
x: jnp.ndarray,
|
||||||
action: jnp.ndarray,
|
action: jnp.ndarray,
|
||||||
):
|
):
|
||||||
hidden = network_apply(params["network_params"], x)
|
hidden = input_apply(params["network_params"], x)
|
||||||
mean, log_std = actor_apply(params["actor_params"], hidden)
|
hidden = message_passer(hidden)
|
||||||
|
mean, log_std = action_apply(params["actor_params"], hidden)
|
||||||
std = jnp.exp(log_std)
|
std = jnp.exp(log_std)
|
||||||
|
|
||||||
logprob = -0.5 * (
|
logprob = -0.5 * (
|
||||||
|
|
@ -108,10 +127,26 @@ def get_action_and_value2(
|
||||||
|
|
||||||
|
|
||||||
def ppo_loss(
|
def ppo_loss(
|
||||||
params, x, a, logp, mb_advantages, mb_returns, args, network, actor, critic
|
params,
|
||||||
|
x,
|
||||||
|
a,
|
||||||
|
logp,
|
||||||
|
mb_advantages,
|
||||||
|
mb_returns,
|
||||||
|
args,
|
||||||
|
input_network_apply,
|
||||||
|
action_network_apply,
|
||||||
|
message_passer,
|
||||||
|
critic_apply,
|
||||||
):
|
):
|
||||||
newlogprob, entropy, newvalue = get_action_and_value2(
|
newlogprob, entropy, newvalue = get_action_and_value2(
|
||||||
network.apply, actor.apply, critic.apply, params, x, a
|
input_network_apply,
|
||||||
|
action_network_apply,
|
||||||
|
message_passer,
|
||||||
|
critic_apply,
|
||||||
|
params,
|
||||||
|
x,
|
||||||
|
a,
|
||||||
)
|
)
|
||||||
logratio = newlogprob - logp
|
logratio = newlogprob - logp
|
||||||
ratio = jnp.exp(logratio)
|
ratio = jnp.exp(logratio)
|
||||||
|
|
@ -129,3 +164,14 @@ def ppo_loss(
|
||||||
entropy_loss = entropy.mean()
|
entropy_loss = entropy.mean()
|
||||||
loss = pg_loss - args.ent_coef * entropy_loss + v_loss * args.vf_coef
|
loss = pg_loss - args.ent_coef * entropy_loss + v_loss * args.vf_coef
|
||||||
return loss, (pg_loss, v_loss, entropy_loss, jax.lax.stop_gradient(approx_kl))
|
return loss, (pg_loss, v_loss, entropy_loss, jax.lax.stop_gradient(approx_kl))
|
||||||
|
|
||||||
|
|
||||||
|
def identity(hidden):
|
||||||
|
"""
|
||||||
|
Used for seamless jax integration,
|
||||||
|
avoids having branching inside jitted function,
|
||||||
|
used as message_passer in case it is not given,
|
||||||
|
(in case of centralized lvl)
|
||||||
|
"""
|
||||||
|
|
||||||
|
return hidden
|
||||||
|
|
|
||||||
Reference in a new issue