feat: init agents array
This commit is contained in:
parent
db5d4c3dc2
commit
bc3408b59a
2 changed files with 49 additions and 85 deletions
|
|
@ -70,4 +70,4 @@ class EnvConfig:
|
||||||
# Per docs in upstream env config: integer factors of 200.
|
# Per docs in upstream env config: integer factors of 200.
|
||||||
light_perlin_noise_scale: int = 0
|
light_perlin_noise_scale: int = 0
|
||||||
|
|
||||||
obs_mode: int = 0 # TODO: ????
|
morph_mode: MorphMode = MorphMode.CENTRALIZED
|
||||||
|
|
|
||||||
|
|
@ -43,7 +43,7 @@ _ALLOWED_OBS_KEYS = {
|
||||||
|
|
||||||
|
|
||||||
def build_adjacency(segments_per_arm, mode: MorphMode):
|
def build_adjacency(segments_per_arm, mode: MorphMode):
|
||||||
num_arms = len(segments_per_arm)
|
num_arms = sum(1 for s in segments_per_arm if s > 0)
|
||||||
num_segments = sum(segments_per_arm)
|
num_segments = sum(segments_per_arm)
|
||||||
|
|
||||||
# FOR NOW SEMI HARDCODE:
|
# FOR NOW SEMI HARDCODE:
|
||||||
|
|
@ -131,9 +131,8 @@ def _normalize_obs(obs, mean, var, eps=1e-8):
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm):
|
def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm):
|
||||||
|
|
||||||
num_segments = sum(segments_per_arm)
|
num_segments = sum(segments_per_arm)
|
||||||
num_arms = len(segments_per_arm)
|
num_arms = sum(1 for s in segments_per_arm if s > 0)
|
||||||
|
|
||||||
def _filter_and_flatten(o):
|
def _filter_and_flatten(o):
|
||||||
values = []
|
values = []
|
||||||
|
|
@ -154,10 +153,12 @@ def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm
|
||||||
# -------- SPLIT TO SEGMENTS --------
|
# -------- SPLIT TO SEGMENTS --------
|
||||||
if key in _JOINT_SCALED_KEYS:
|
if key in _JOINT_SCALED_KEYS:
|
||||||
if morph_mode == 3:
|
if morph_mode == 3:
|
||||||
|
# special case: segment lvl and scale with joints,
|
||||||
|
# needs to split logic for center mlps and arms
|
||||||
|
# center:
|
||||||
|
# 3 joints per arm, each joint has 2 values
|
||||||
B = v.shape[0]
|
B = v.shape[0]
|
||||||
|
center_size = num_arms * 3 * 2
|
||||||
center_size = num_arms * 3 * 2 # 3 slots per arm, each slot has 2 values
|
|
||||||
|
|
||||||
v_center = v[:, :center_size]
|
v_center = v[:, :center_size]
|
||||||
v_center = v_center.reshape(B, num_arms, 3, 2)
|
v_center = v_center.reshape(B, num_arms, 3, 2)
|
||||||
|
|
||||||
|
|
@ -173,6 +174,7 @@ def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm
|
||||||
v = v[..., None] # (env, segments, 1)
|
v = v[..., None] # (env, segments, 1)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
|
# global key, share with all
|
||||||
if morph_mode == 3: # segment
|
if morph_mode == 3: # segment
|
||||||
v = jnp.repeat(v[:, None, :], num_segments + num_arms, axis=1)
|
v = jnp.repeat(v[:, None, :], num_segments + num_arms, axis=1)
|
||||||
|
|
||||||
|
|
@ -197,7 +199,7 @@ def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm
|
||||||
|
|
||||||
# Observation keys whose size scales with the number of joints (2 per segment).
|
# Observation keys whose size scales with the number of joints (2 per segment).
|
||||||
_JOINT_SCALED_KEYS = frozenset(
|
_JOINT_SCALED_KEYS = frozenset(
|
||||||
{
|
{ # TODO CODE SMELL
|
||||||
"joint_position",
|
"joint_position",
|
||||||
"joint_velocity",
|
"joint_velocity",
|
||||||
"joint_actuator_force",
|
"joint_actuator_force",
|
||||||
|
|
@ -213,64 +215,6 @@ _SEGMENT_SCALED_KEYS = frozenset(
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@jax.jit(static_argnums=(1,))
|
|
||||||
def _convert_obs_dict_to_array(obs_dict, obs_mode, segments_per_arm):
|
|
||||||
|
|
||||||
num_segments = sum(segments_per_arm)
|
|
||||||
num_arms = len(segments_per_arm)
|
|
||||||
|
|
||||||
def _filter_and_flatten(o):
|
|
||||||
values = []
|
|
||||||
|
|
||||||
for key in sorted(o.keys()):
|
|
||||||
if key not in _ALLOWED_OBS_KEYS:
|
|
||||||
continue
|
|
||||||
|
|
||||||
v = o[key]
|
|
||||||
if v.size == 0:
|
|
||||||
continue
|
|
||||||
|
|
||||||
# -------- CENTRALIZED --------
|
|
||||||
if obs_mode == 0:
|
|
||||||
values.append(v.reshape(v.shape[0], -1))
|
|
||||||
continue
|
|
||||||
|
|
||||||
# -------- SPLIT TO SEGMENTS --------
|
|
||||||
if key in _JOINT_SCALED_KEYS:
|
|
||||||
v = v.reshape(v.shape[0], num_segments, 2)
|
|
||||||
|
|
||||||
elif key in _SEGMENT_SCALED_KEYS:
|
|
||||||
v = v[..., None] # (env, segments, 1)
|
|
||||||
|
|
||||||
else:
|
|
||||||
# global → broadcast
|
|
||||||
if obs_mode == 2:
|
|
||||||
v = jnp.repeat(v[:, None, :], num_segments, axis=1)
|
|
||||||
else:
|
|
||||||
v = jnp.repeat(v[:, None, :], num_arms, axis=1)
|
|
||||||
values.append(v)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# -------- SEGMENT MODE --------
|
|
||||||
if obs_mode == 2:
|
|
||||||
values.append(v)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# -------- ARM MODE --------
|
|
||||||
# reshape (segments → arms, seg_per_arm)
|
|
||||||
v = v.reshape(v.shape[0], num_arms, -1)
|
|
||||||
|
|
||||||
values.append(v)
|
|
||||||
|
|
||||||
# -------- CONCAT --------
|
|
||||||
if obs_mode == 0:
|
|
||||||
return jnp.concatenate(values, axis=-1)
|
|
||||||
else:
|
|
||||||
return jnp.concatenate(values, axis=-1)
|
|
||||||
|
|
||||||
return jax.vmap(_filter_and_flatten)(obs_dict)
|
|
||||||
|
|
||||||
|
|
||||||
# TODO: update to work with extra dimension + message passing
|
# TODO: update to work with extra dimension + message passing
|
||||||
def _get_action_and_value_noise(
|
def _get_action_and_value_noise(
|
||||||
sensor: GenericDenseLayersWithActivation,
|
sensor: GenericDenseLayersWithActivation,
|
||||||
|
|
@ -284,12 +228,12 @@ def _get_action_and_value_noise(
|
||||||
action_high,
|
action_high,
|
||||||
adj_matrix: jnp.ndarray,
|
adj_matrix: jnp.ndarray,
|
||||||
):
|
):
|
||||||
hidden = sensor.apply(agent_state.params["sensor_params"], next_obs)
|
hidden = apply_per_node(sensor, agent_state.params["sensor_params"], next_obs)
|
||||||
hidden_critic = feature_extractor.apply(
|
hidden_critic = apply_per_node(
|
||||||
agent_state.params["feature_extractor_params"], next_obs
|
feature_extractor, agent_state.params["feature_extractor_params"], next_obs
|
||||||
)
|
)
|
||||||
|
|
||||||
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
|
mean, log_std = apply_per_node(actor, agent_state.params["actor_params"], hidden)
|
||||||
log_std = jnp.clip(log_std, -5, 2)
|
log_std = jnp.clip(log_std, -5, 2)
|
||||||
key, subkey = jax.random.split(key)
|
key, subkey = jax.random.split(key)
|
||||||
noise = jax.random.normal(subkey, shape=mean.shape)
|
noise = jax.random.normal(subkey, shape=mean.shape)
|
||||||
|
|
@ -367,7 +311,7 @@ def _reward_fn(env_state, next_env_state):
|
||||||
|
|
||||||
|
|
||||||
# TODO: update to work with extra dimension + message passing
|
# TODO: update to work with extra dimension + message passing
|
||||||
def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, obs_mode, segments_per_arm):
|
def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, morph_mode, segments_per_arm):
|
||||||
next_env_state = env_step_fn(env_state, action)
|
next_env_state = env_step_fn(env_state, action)
|
||||||
|
|
||||||
reward = _reward_fn(env_state, next_env_state)
|
reward = _reward_fn(env_state, next_env_state)
|
||||||
|
|
@ -392,13 +336,20 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, obs_mode, s
|
||||||
episode_stats,
|
episode_stats,
|
||||||
next_env_state,
|
next_env_state,
|
||||||
(
|
(
|
||||||
_convert_obs_dict_to_array(next_env_state.observations, obs_mode, segments_per_arm),
|
_convert_obs_dict_to_array_morphology(
|
||||||
|
next_env_state.observations, morph_mode, segments_per_arm
|
||||||
|
),
|
||||||
reward,
|
reward,
|
||||||
done,
|
done,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def apply_per_node(net, params, x):
|
||||||
|
# x: (batch, nodes, feat)
|
||||||
|
return jax.vmap(lambda node_x: net.apply(params, node_x), in_axes=1, out_axes=1)(x)
|
||||||
|
|
||||||
|
|
||||||
# TODO: update to work with extra dimension + message passing
|
# TODO: update to work with extra dimension + message passing
|
||||||
def _rollout_jit(
|
def _rollout_jit(
|
||||||
agent_state,
|
agent_state,
|
||||||
|
|
@ -514,18 +465,16 @@ class PPOTrainer:
|
||||||
|
|
||||||
self.key = jax.random.PRNGKey(self.experiment.seed)
|
self.key = jax.random.PRNGKey(self.experiment.seed)
|
||||||
|
|
||||||
morph = cfg.morphology.morph_mode
|
self.morph_mode = self.cfg.morphology.morph_mode
|
||||||
|
|
||||||
self.logger.info(f"[INIT]: Used morphology mode {morph}")
|
self.logger.info(f"[INIT]: Used morphology mode {self.morph_mode}")
|
||||||
self.adj = build_adjacency(cfg.morphology.segments_per_arm, morph)
|
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.feature_extractor, self.actor, self.critic = self._init_agent()
|
||||||
self.sensor.apply = jax.jit(self.sensor.apply)
|
self.sensor.apply = jax.jit(self.sensor.apply)
|
||||||
self.feature_extractor.apply = jax.jit(self.feature_extractor.apply)
|
self.feature_extractor.apply = jax.jit(self.feature_extractor.apply)
|
||||||
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.obs_mode = self.cfg.environment.obs_mode
|
|
||||||
self.segments_per_arm = jnp.asarray(self.cfg.morphology.segments_per_arm, dtype=jnp.int32)
|
self.segments_per_arm = jnp.asarray(self.cfg.morphology.segments_per_arm, dtype=jnp.int32)
|
||||||
|
|
||||||
action_low = jnp.asarray(self.env.single_action_space.low, dtype=jnp.float32)
|
action_low = jnp.asarray(self.env.single_action_space.low, dtype=jnp.float32)
|
||||||
|
|
@ -538,7 +487,7 @@ class PPOTrainer:
|
||||||
step_env_fn=partial(
|
step_env_fn=partial(
|
||||||
_step_env_wrapped,
|
_step_env_wrapped,
|
||||||
env_step_fn=self.env.step,
|
env_step_fn=self.env.step,
|
||||||
obs_mode=self.obs_mode,
|
morph_mode=self.morph_mode,
|
||||||
segments_per_arm=self.segments_per_arm,
|
segments_per_arm=self.segments_per_arm,
|
||||||
),
|
),
|
||||||
sensor=self.sensor,
|
sensor=self.sensor,
|
||||||
|
|
@ -578,12 +527,27 @@ class PPOTrainer:
|
||||||
|
|
||||||
def _init_agent(self):
|
def _init_agent(self):
|
||||||
self.logger.info("[AGENT]: Initializing agent...")
|
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:
|
||||||
|
needed_copies = sum(self.segments_per_arm) + sum(
|
||||||
|
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())
|
||||||
|
|
||||||
sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
|
|
||||||
feature_extractor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
|
feature_extractor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
|
||||||
actor = Actor(action_dim=self.env.single_action_space.shape[0])
|
|
||||||
critic = OneDenseLayerMLP()
|
critic = OneDenseLayerMLP()
|
||||||
return sensor, feature_extractor, actor, critic
|
return sensors, message_passers, actors, feature_extractor, critic
|
||||||
|
|
||||||
def _init_agent_state(self) -> TrainState:
|
def _init_agent_state(self) -> TrainState:
|
||||||
self.logger.info("[AGENT STATE]: Initializing agent state...")
|
self.logger.info("[AGENT STATE]: Initializing agent state...")
|
||||||
|
|
@ -595,9 +559,9 @@ class PPOTrainer:
|
||||||
dummy_reset = self.env.reset(seed=0)
|
dummy_reset = self.env.reset(seed=0)
|
||||||
for k, v in dummy_reset.observations.items():
|
for k, v in dummy_reset.observations.items():
|
||||||
self.logger.debug(k, v.shape)
|
self.logger.debug(k, v.shape)
|
||||||
sample_obs = _convert_obs_dict_to_array(
|
sample_obs = _convert_obs_dict_to_array_morphology(
|
||||||
dummy_reset.observations,
|
dummy_reset.observations,
|
||||||
self.obs_mode,
|
self.morph_mode,
|
||||||
self.segments_per_arm,
|
self.segments_per_arm,
|
||||||
)[0] # take first env
|
)[0] # take first env
|
||||||
self.obs_mean = jnp.zeros((len(sample_obs),))
|
self.obs_mean = jnp.zeros((len(sample_obs),))
|
||||||
|
|
@ -826,9 +790,9 @@ class PPOTrainer:
|
||||||
self.logger.log_non_interactive(f"Initial reset started: {time.ctime()}")
|
self.logger.log_non_interactive(f"Initial reset started: {time.ctime()}")
|
||||||
|
|
||||||
env_state = self.env.reset(seed=self.experiment.seed)
|
env_state = self.env.reset(seed=self.experiment.seed)
|
||||||
next_obs = _convert_obs_dict_to_array(
|
next_obs = _convert_obs_dict_to_array_morphology(
|
||||||
env_state.observations,
|
env_state.observations,
|
||||||
self.obs_mode,
|
self.morph_mode,
|
||||||
self.segments_per_arm,
|
self.segments_per_arm,
|
||||||
)
|
)
|
||||||
next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_)
|
next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_)
|
||||||
|
|
|
||||||
Reference in a new issue