feat: upgraded obs_to_array, now takes care of splitting for agents
This commit is contained in:
parent
39945e8821
commit
a005c0ccad
1 changed files with 30 additions and 24 deletions
|
|
@ -138,7 +138,7 @@ def _normalize_obs(obs, mean, var, eps=1e-8):
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def _convert_obs_dict_to_array(obs_dict, obs_mode, segments_per_arm):
|
def _convert_obs_dict_to_array(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 = len(segments_per_arm)
|
||||||
|
|
@ -155,55 +155,54 @@ def _convert_obs_dict_to_array(obs_dict, obs_mode, segments_per_arm):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# -------- CENTRALIZED --------
|
# -------- CENTRALIZED --------
|
||||||
if obs_mode == 0:
|
if morph_mode == 0:
|
||||||
values.append(v.reshape(v.shape[0], -1))
|
values.append(v.reshape(v.shape[0], -1))
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# -------- SPLIT TO SEGMENTS --------
|
# -------- SPLIT TO SEGMENTS --------
|
||||||
if key in _JOINT_SCALED_KEYS:
|
if key in _JOINT_SCALED_KEYS:
|
||||||
|
if morph_mode == 3:
|
||||||
|
B = v.shape[0]
|
||||||
|
|
||||||
|
center_size = num_arms * 3 * 2 # 3 slots per arm, each slot has 2 values
|
||||||
|
|
||||||
|
v_center = v[:, :center_size]
|
||||||
|
v_center = v_center.reshape(B, num_arms, 3, 2)
|
||||||
|
|
||||||
|
v_segs = v[:, center_size:]
|
||||||
|
v_segs = v_segs.reshape(B, -1, 2)
|
||||||
|
|
||||||
|
v = jnp.concatenate([v_center.reshape(B, -1, 2), v_segs], axis=1)
|
||||||
|
values.append(v)
|
||||||
|
continue
|
||||||
v = v.reshape(v.shape[0], num_segments, 2)
|
v = v.reshape(v.shape[0], num_segments, 2)
|
||||||
|
|
||||||
elif key in _SEGMENT_SCALED_KEYS:
|
elif key in _SEGMENT_SCALED_KEYS:
|
||||||
v = v[..., None] # (env, segments, 1)
|
v = v[..., None] # (env, segments, 1)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
# global → broadcast
|
if morph_mode == 3: # segment
|
||||||
if obs_mode == 2:
|
v = jnp.repeat(v[:, None, :], num_segments + num_arms, axis=1)
|
||||||
v = jnp.repeat(v[:, None, :], num_segments, axis=1)
|
|
||||||
else:
|
else: # ring or fully connect
|
||||||
v = jnp.repeat(v[:, None, :], num_arms, axis=1)
|
v = jnp.repeat(v[:, None, :], num_arms, axis=1)
|
||||||
values.append(v)
|
values.append(v)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# -------- SEGMENT MODE --------
|
# -------- SEGMENT MODE --------
|
||||||
if obs_mode == 2:
|
if morph_mode == 3:
|
||||||
values.append(v)
|
values.append(v)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# -------- ARM MODE --------
|
# -------- ARM MODE --------
|
||||||
# reshape (segments → arms, seg_per_arm)
|
|
||||||
v = v.reshape(v.shape[0], num_arms, -1)
|
v = v.reshape(v.shape[0], num_arms, -1)
|
||||||
|
|
||||||
values.append(v)
|
values.append(v)
|
||||||
|
|
||||||
# -------- CONCAT --------
|
|
||||||
if obs_mode == 0:
|
|
||||||
return jnp.concatenate(values, axis=-1)
|
|
||||||
else:
|
|
||||||
return jnp.concatenate(values, axis=-1)
|
return jnp.concatenate(values, axis=-1)
|
||||||
|
|
||||||
return jax.vmap(_filter_and_flatten)(obs_dict)
|
return jax.vmap(_filter_and_flatten)(obs_dict)
|
||||||
|
|
||||||
|
|
||||||
from enum import Enum
|
|
||||||
|
|
||||||
|
|
||||||
class ObsMode(Enum):
|
|
||||||
CENTRALIZED = 0 # dirty dirty code, todo, mooove outta here
|
|
||||||
ARM = 1
|
|
||||||
SEGMENT = 2
|
|
||||||
|
|
||||||
|
|
||||||
# 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(
|
||||||
{
|
{
|
||||||
|
|
@ -480,7 +479,12 @@ class TrainingMeasurements:
|
||||||
|
|
||||||
class PPOTrainer:
|
class PPOTrainer:
|
||||||
def __init__(
|
def __init__(
|
||||||
self, cfg: BrittleStarConfig, env: BrittleStarJaxEnvWrapper, run_dir: str, run_name: str
|
self,
|
||||||
|
cfg: BrittleStarConfig,
|
||||||
|
env: BrittleStarJaxEnvWrapper,
|
||||||
|
run_dir: str,
|
||||||
|
run_name: str,
|
||||||
|
morph: MorphMode = MorphMode.CENTRALIZED,
|
||||||
):
|
):
|
||||||
self.cfg = cfg
|
self.cfg = cfg
|
||||||
self.ppo = cfg.ppo
|
self.ppo = cfg.ppo
|
||||||
|
|
@ -497,6 +501,8 @@ class PPOTrainer:
|
||||||
|
|
||||||
self.key = jax.random.PRNGKey(self.experiment.seed)
|
self.key = jax.random.PRNGKey(self.experiment.seed)
|
||||||
|
|
||||||
|
self.adj = build_adjacency(cfg.morphology.segments_per_arm, morph)
|
||||||
|
|
||||||
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)
|
||||||
|
|
|
||||||
Reference in a new issue