1
Fork 0

fix: uniform naming of models accross code

This commit is contained in:
cedric 2026-04-02 18:37:37 +00:00
parent 79221372df
commit 1c4e40b5fb
3 changed files with 39 additions and 34 deletions

View file

@ -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

View file

@ -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,

View file

@ -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"],
], ],
] ]
) )