1
Fork 0

feat(PPOTrainer): vectorized init agent state, partially added message passer param to rollout, ...

This commit is contained in:
Robin Meersman 2026-04-29 11:23:01 +02:00
parent bc3408b59a
commit e85fc09ed6
2 changed files with 75 additions and 44 deletions

View file

@ -6,6 +6,7 @@ import jax.numpy as jnp
import jax.tree_util
from typing import Sequence, Callable
from flax.linen.initializers import constant, orthogonal
from flax.core import FrozenDict
# semi generic so we can easily make a config for it in experiments
@ -40,27 +41,28 @@ class Actor(nn.Module):
@jax.tree_util.register_dataclass
@dataclass
class AgentParams:
sensor_params: flax.core.FrozenDict
actor_params: flax.core.FrozenDict
critic_params: flax.core.FrozenDict
feature_extractor_params: flax.core.FrozenDict
sensor_params: FrozenDict | dict
actor_params: FrozenDict | dict
critic_params: FrozenDict | dict
feature_extractor_params: FrozenDict | dict
message_passer_params: FrozenDict | dict
@jax.tree_util.register_dataclass
@dataclass
class Storage:
obs: jnp.array
actions: jnp.array
logprobs: jnp.array
dones: jnp.array
values: jnp.array
advantages: jnp.array
returns: jnp.array
rewards: jnp.array
obs: jnp.ndarray
actions: jnp.ndarray
logprobs: jnp.ndarray
dones: jnp.ndarray
values: jnp.ndarray
advantages: jnp.ndarray
returns: jnp.ndarray
rewards: jnp.ndarray
raw_actions: jnp.ndarray = None # before clipping
means: jnp.ndarray = None # policy mean
stds: jnp.ndarray = None # policy std
raw_actions: jnp.ndarray | None = None # before clipping
means: jnp.ndarray | None = None # policy mean
stds: jnp.ndarray | None = None # policy std
def replace(self, **kwargs) -> "Storage":
fs = fields(self)

View file

@ -9,6 +9,7 @@ import jax
import jax.numpy as jnp
import numpy as np
import optax
import flax.linen as nn
from flax.training.train_state import TrainState
from experiment_logger import get_logger
@ -217,10 +218,10 @@ _SEGMENT_SCALED_KEYS = frozenset(
# TODO: update to work with extra dimension + message passing
def _get_action_and_value_noise(
sensor: GenericDenseLayersWithActivation,
feature_extractor: GenericDenseLayersWithActivation,
actor: Actor,
critic: OneDenseLayerMLP,
sensor: nn.Module,
feature_extractor: nn.Module,
actor: nn.Module,
critic: nn.Module,
agent_state: TrainState,
next_obs: jnp.ndarray,
key: jax.random.PRNGKey,
@ -252,10 +253,11 @@ def _step_once(
_,
env_step_fn,
adj_matrix,
sensor: GenericDenseLayersWithActivation,
feature_extractor: GenericDenseLayersWithActivation,
actor: Actor,
critic: OneDenseLayerMLP,
sensor: nn.Module,
feature_extractor: nn.Module,
actor: nn.Module,
critic: nn.Module,
message_passer: nn.Module,
action_low,
action_high,
):
@ -361,10 +363,11 @@ def _rollout_jit(
key,
max_steps,
step_env_fn,
sensor: GenericDenseLayersWithActivation,
feature_extractor: GenericDenseLayersWithActivation,
actor: Actor,
critic: OneDenseLayerMLP,
sensor: nn.Module,
feature_extractor: nn.Module,
actor: nn.Module,
critic: nn.Module,
message_passer: nn.Module,
action_low,
action_high,
):
@ -375,6 +378,7 @@ def _rollout_jit(
feature_extractor=feature_extractor,
actor=actor,
critic=critic,
message_passer=message_passer,
env_step_fn=step_env_fn,
action_low=action_low,
action_high=action_high,
@ -470,7 +474,15 @@ class PPOTrainer:
self.logger.info(f"[INIT]: Used morphology mode {self.morph_mode}")
self.adj = build_adjacency(cfg.morphology.segments_per_arm, self.morph_mode)
self.sensor, self.feature_extractor, self.actor, self.critic = self._init_agent()
(
self.sensor,
self.message_passer,
self.actor,
self.feature_extractor,
self.critic,
self.needed_copies,
) = self._init_agent()
self.sensor.apply = jax.jit(self.sensor.apply)
self.feature_extractor.apply = jax.jit(self.feature_extractor.apply)
self.actor.apply = jax.jit(self.actor.apply)
@ -494,6 +506,7 @@ class PPOTrainer:
feature_extractor=self.feature_extractor,
actor=self.actor,
critic=self.critic,
message_passer=self.message_passer,
action_low=action_low,
action_high=action_high,
adj_matrix=self.adj,
@ -527,10 +540,7 @@ class PPOTrainer:
def _init_agent(self):
self.logger.info("[AGENT]: Initializing agent...")
sensors = []
actors = []
message_passers = []
needed_copies = 1
if (self.morph_mode == MorphMode.FULLY_CONNECTED) or (self.morph_mode == MorphMode.RING):
needed_copies = sum(1 for s in self.segments_per_arm if s > 0)
else:
@ -538,22 +548,19 @@ class PPOTrainer:
1 for s in self.segments_per_arm if s > 0
)
for _ in range(needed_copies):
actor = Actor(action_dim=self.env.single_action_space.shape[0])
sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
sensors.append(sensor)
actors.append(actor)
message_passers.append(OneDenseLayerMLP())
actor = Actor(action_dim=self.env.single_action_space.shape[0])
sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
message_passer = OneDenseLayerMLP()
feature_extractor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
critic = OneDenseLayerMLP()
return sensors, message_passers, actors, feature_extractor, critic
return sensor, message_passer, actor, feature_extractor, critic, needed_copies
def _init_agent_state(self) -> TrainState:
self.logger.info("[AGENT STATE]: Initializing agent state...")
self.key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(
self.key, 5
self.key, sensor_key, actor_key, critic_key, feature_extractor_key, message_passer_key = (
jax.random.split(self.key, 6)
)
dummy_reset = self.env.reset(seed=0)
@ -564,12 +571,28 @@ class PPOTrainer:
self.morph_mode,
self.segments_per_arm,
)[0] # take first env
self.obs_mean = jnp.zeros((len(sample_obs),))
self.obs_var = jnp.ones((len(sample_obs),))
self.obs_count = 1e-4
sensor_params = self.sensor.init(sensor_key, sample_obs)
sensor_keys = jax.random.split(sensor_key, self.needed_copies)
actor_keys = jax.random.split(actor_key, self.needed_copies)
message_passer_keys = jax.random.split(message_passer_key, self.needed_copies)
sensor_params = jax.vmap(lambda k: self.sensor.init(k, sample_obs))(sensor_keys)
single_sensor_param = jax.tree.map(lambda x: x[0], sensor_params)
sensor_params_sample = self.sensor.apply(single_sensor_param, sample_obs)
actor_params = jax.vmap(lambda k: self.actor.init(k, sensor_params_sample))(actor_keys)
message_passer_params = jax.vmap(
lambda k: self.message_passer.init(
k, self.sensor.apply(single_sensor_param, sample_obs)
)
)(message_passer_keys)
feature_extractor_params = self.feature_extractor.init(feature_extractor_key, sample_obs)
actor_params = self.actor.init(actor_key, self.sensor.apply(sensor_params, sample_obs))
critic_params = self.critic.init(
critic_key, self.feature_extractor.apply(feature_extractor_params, sample_obs)
)
@ -577,7 +600,13 @@ class PPOTrainer:
return TrainState.create(
apply_fn=None,
params=asdict(
AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params)
AgentParams(
sensor_params,
actor_params,
critic_params,
feature_extractor_params,
message_passer_params,
)
),
tx=optax.chain(
optax.clip_by_global_norm(self.ppo.max_grad_norm),