fix: uniform naming of models accross code
This commit is contained in:
parent
79221372df
commit
1c4e40b5fb
3 changed files with 39 additions and 34 deletions
|
|
@ -40,10 +40,10 @@ class Actor(nn.Module):
|
||||||
@jax.tree_util.register_dataclass
|
@jax.tree_util.register_dataclass
|
||||||
@dataclass
|
@dataclass
|
||||||
class AgentParams:
|
class AgentParams:
|
||||||
network_params: flax.core.FrozenDict
|
sensor_params: flax.core.FrozenDict
|
||||||
actor_params: flax.core.FrozenDict
|
actor_params: flax.core.FrozenDict
|
||||||
critic_params: flax.core.FrozenDict
|
critic_params: flax.core.FrozenDict
|
||||||
critic_network_params: flax.core.FrozenDict
|
feature_extractor_params: flax.core.FrozenDict
|
||||||
|
|
||||||
|
|
||||||
@jax.tree_util.register_dataclass
|
@jax.tree_util.register_dataclass
|
||||||
|
|
|
||||||
32
src/ppo.py
32
src/ppo.py
|
|
@ -8,9 +8,7 @@ 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__(
|
def __init__(self, args, sensor, actor, critic, feature_extractor, message_passer=None):
|
||||||
self, args, input_network, action_network, critic, critic_network, message_passer=None
|
|
||||||
):
|
|
||||||
self.args = args
|
self.args = args
|
||||||
|
|
||||||
if not message_passer:
|
if not message_passer:
|
||||||
|
|
@ -20,10 +18,10 @@ class PPO:
|
||||||
partial(
|
partial(
|
||||||
ppo_loss,
|
ppo_loss,
|
||||||
args=args,
|
args=args,
|
||||||
input_network_apply=input_network.apply,
|
sensor_apply=sensor.apply,
|
||||||
action_network_apply=action_network.apply,
|
actor_apply=actor.apply,
|
||||||
critic_apply=critic.apply,
|
critic_apply=critic.apply,
|
||||||
critic_network_apply=critic_network.apply,
|
feature_extractor_apply=feature_extractor.apply,
|
||||||
message_passer=message_passer,
|
message_passer=message_passer,
|
||||||
),
|
),
|
||||||
has_aux=True,
|
has_aux=True,
|
||||||
|
|
@ -92,15 +90,15 @@ def get_action_and_value2(
|
||||||
action_apply,
|
action_apply,
|
||||||
message_passer,
|
message_passer,
|
||||||
critic_apply,
|
critic_apply,
|
||||||
critic_network_apply,
|
feature_extractor_apply,
|
||||||
params: flax.core.FrozenDict,
|
params: flax.core.FrozenDict,
|
||||||
x: jnp.ndarray,
|
x: jnp.ndarray,
|
||||||
action: jnp.ndarray,
|
action: jnp.ndarray,
|
||||||
):
|
):
|
||||||
hidden_network = input_apply(params["network_params"], x)
|
hidden_sensor = input_apply(params["sensor_params"], x)
|
||||||
hidden_critic = critic_network_apply(params["critic_network_params"], x)
|
hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x)
|
||||||
hidden_network = message_passer(hidden_network)
|
hidden_sensor = message_passer(hidden_sensor)
|
||||||
mean, log_std = action_apply(params["actor_params"], hidden_network)
|
mean, log_std = action_apply(params["actor_params"], hidden_sensor)
|
||||||
std = jnp.exp(log_std)
|
std = jnp.exp(log_std)
|
||||||
|
|
||||||
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
|
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
|
||||||
|
|
@ -118,18 +116,18 @@ def ppo_loss(
|
||||||
mb_advantages,
|
mb_advantages,
|
||||||
mb_returns,
|
mb_returns,
|
||||||
args,
|
args,
|
||||||
input_network_apply,
|
sensor_apply,
|
||||||
action_network_apply,
|
actor_apply,
|
||||||
message_passer,
|
message_passer,
|
||||||
critic_apply,
|
critic_apply,
|
||||||
critic_network_apply,
|
feature_extractor_apply,
|
||||||
):
|
):
|
||||||
newlogprob, entropy, newvalue = get_action_and_value2(
|
newlogprob, entropy, newvalue = get_action_and_value2(
|
||||||
input_network_apply,
|
sensor_apply,
|
||||||
action_network_apply,
|
actor_apply,
|
||||||
message_passer,
|
message_passer,
|
||||||
critic_apply,
|
critic_apply,
|
||||||
critic_network_apply,
|
feature_extractor_apply,
|
||||||
params,
|
params,
|
||||||
x,
|
x,
|
||||||
a,
|
a,
|
||||||
|
|
|
||||||
37
src/train.py
37
src/train.py
|
|
@ -72,7 +72,7 @@ def train(args: PPOArgs):
|
||||||
random.seed(args.seed)
|
random.seed(args.seed)
|
||||||
np.random.seed(args.seed)
|
np.random.seed(args.seed)
|
||||||
key = jax.random.PRNGKey(args.seed)
|
key = jax.random.PRNGKey(args.seed)
|
||||||
key, network_key, actor_key, critic_key, critic_network_key = jax.random.split(key, 5)
|
key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5)
|
||||||
|
|
||||||
torch.backends.cudnn.deterministic = args.torch_deterministic
|
torch.backends.cudnn.deterministic = args.torch_deterministic
|
||||||
device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu")
|
device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu")
|
||||||
|
|
@ -122,8 +122,8 @@ def train(args: PPOArgs):
|
||||||
return args.learning_rate * frac
|
return args.learning_rate * frac
|
||||||
|
|
||||||
print("Initializing the models...")
|
print("Initializing the models...")
|
||||||
network = GenericDenseLayersWithActivation()
|
sensor = GenericDenseLayersWithActivation()
|
||||||
critic_network = GenericDenseLayersWithActivation()
|
feature_extractor = GenericDenseLayersWithActivation()
|
||||||
actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX
|
actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX
|
||||||
critic = OneDenseLayerMLP()
|
critic = OneDenseLayerMLP()
|
||||||
# messager = OneDenseLayerMLP()
|
# messager = OneDenseLayerMLP()
|
||||||
|
|
@ -135,15 +135,17 @@ def train(args: PPOArgs):
|
||||||
if v.size > 0
|
if v.size > 0
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
network_params = network.init(network_key, sample_obs)
|
sensor_params = sensor.init(sensor_key, sample_obs)
|
||||||
critic_network_params = critic_network.init(critic_network_key, sample_obs)
|
feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs)
|
||||||
actor_params = actor.init(actor_key, network.apply(network_params, sample_obs))
|
actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs))
|
||||||
critic_params = critic.init(critic_key, critic_network.apply(critic_network_params, sample_obs))
|
critic_params = critic.init(
|
||||||
|
critic_key, feature_extractor.apply(feature_extractor_params, sample_obs)
|
||||||
|
)
|
||||||
|
|
||||||
agent_state = TrainState.create(
|
agent_state = TrainState.create(
|
||||||
apply_fn=None,
|
apply_fn=None,
|
||||||
params=asdict(
|
params=asdict(
|
||||||
AgentParams(network_params, actor_params, critic_params, critic_network_params)
|
AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params)
|
||||||
),
|
),
|
||||||
tx=optax.chain(
|
tx=optax.chain(
|
||||||
optax.clip_by_global_norm(args.max_grad_norm),
|
optax.clip_by_global_norm(args.max_grad_norm),
|
||||||
|
|
@ -153,11 +155,11 @@ def train(args: PPOArgs):
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
network.apply = jax.jit(network.apply)
|
sensor.apply = jax.jit(sensor.apply)
|
||||||
critic_network.apply = jax.jit(critic_network.apply)
|
feature_extractor.apply = jax.jit(feature_extractor.apply)
|
||||||
actor.apply = jax.jit(actor.apply)
|
actor.apply = jax.jit(actor.apply)
|
||||||
critic.apply = jax.jit(critic.apply)
|
critic.apply = jax.jit(critic.apply)
|
||||||
ppo_instance = PPO(args, network, actor, critic, critic_network)
|
ppo_instance = PPO(args, sensor, actor, critic, feature_extractor)
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def get_action_and_value_noise(
|
def get_action_and_value_noise(
|
||||||
|
|
@ -165,7 +167,11 @@ def train(args: PPOArgs):
|
||||||
next_obs: jnp.ndarray,
|
next_obs: jnp.ndarray,
|
||||||
key: jax.random.PRNGKey,
|
key: jax.random.PRNGKey,
|
||||||
):
|
):
|
||||||
hidden = network.apply(agent_state.params["network_params"], next_obs)
|
hidden = sensor.apply(agent_state.params["sensor_params"], next_obs)
|
||||||
|
hidden_critic = feature_extractor.apply(
|
||||||
|
agent_state.params["feature_extractor_params"], next_obs
|
||||||
|
)
|
||||||
|
|
||||||
# Continuous actions: sample from a Gaussian parameterized by the actor
|
# Continuous actions: sample from a Gaussian parameterized by the actor
|
||||||
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
|
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
|
||||||
key, subkey = jax.random.split(key)
|
key, subkey = jax.random.split(key)
|
||||||
|
|
@ -173,7 +179,7 @@ def train(args: PPOArgs):
|
||||||
std = jnp.exp(log_std)
|
std = jnp.exp(log_std)
|
||||||
action = mean + noise * std
|
action = mean + noise * std
|
||||||
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
|
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
|
||||||
value = critic.apply(agent_state.params["critic_params"], hidden)
|
value = critic.apply(agent_state.params["critic_params"], hidden_critic)
|
||||||
return action, logprob, value.squeeze(-1), key
|
return action, logprob, value.squeeze(-1), key
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
|
|
@ -189,7 +195,7 @@ def train(args: PPOArgs):
|
||||||
def compute_gae(agent_state, next_obs, next_done, storage):
|
def compute_gae(agent_state, next_obs, next_done, storage):
|
||||||
next_value = critic.apply(
|
next_value = critic.apply(
|
||||||
agent_state.params["critic_params"],
|
agent_state.params["critic_params"],
|
||||||
network.apply(agent_state.params["network_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((args.num_envs,))
|
||||||
|
|
@ -307,9 +313,10 @@ def train(args: PPOArgs):
|
||||||
[
|
[
|
||||||
vars(args),
|
vars(args),
|
||||||
[
|
[
|
||||||
agent_state.params["network_params"],
|
agent_state.params["sensor_params"],
|
||||||
agent_state.params["actor_params"],
|
agent_state.params["actor_params"],
|
||||||
agent_state.params["critic_params"],
|
agent_state.params["critic_params"],
|
||||||
|
agent_state.params["feature_extractor_params"],
|
||||||
],
|
],
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
|
||||||
Reference in a new issue