refactor: cleanups while reviewing
This commit is contained in:
parent
4e3cd0c673
commit
6956c5e853
2 changed files with 49 additions and 50 deletions
|
|
@ -6,8 +6,6 @@ from brittle_star_project.environment.env_config import MorphMode
|
||||||
|
|
||||||
from experiment_logger import get_logger
|
from experiment_logger import get_logger
|
||||||
|
|
||||||
logger11 = get_logger()
|
|
||||||
|
|
||||||
_JOINT_SCALED_KEYS = frozenset(
|
_JOINT_SCALED_KEYS = frozenset(
|
||||||
{
|
{
|
||||||
"joint_position",
|
"joint_position",
|
||||||
|
|
@ -57,6 +55,8 @@ def create_obs_processor(
|
||||||
segments_per_arm=[4, 4, 4, 4, 4],
|
segments_per_arm=[4, 4, 4, 4, 4],
|
||||||
agent_indices=[0, 1, 2, 3, 4],
|
agent_indices=[0, 1, 2, 3, 4],
|
||||||
):
|
):
|
||||||
|
logger = get_logger()
|
||||||
|
|
||||||
# made a set to allow O(1) search
|
# made a set to allow O(1) search
|
||||||
ordered_keys = frozenset(
|
ordered_keys = frozenset(
|
||||||
[
|
[
|
||||||
|
|
@ -87,6 +87,13 @@ def create_obs_processor(
|
||||||
|
|
||||||
return new_obs
|
return new_obs
|
||||||
|
|
||||||
|
def _prune_features(obs: dict) -> dict:
|
||||||
|
pruned = {}
|
||||||
|
for key, arr in obs.items():
|
||||||
|
if key in ordered_keys:
|
||||||
|
pruned[key] = arr
|
||||||
|
return pruned
|
||||||
|
|
||||||
def _normalize_features(obs: dict) -> dict:
|
def _normalize_features(obs: dict) -> dict:
|
||||||
normalized = {}
|
normalized = {}
|
||||||
for key, arr in obs.items():
|
for key, arr in obs.items():
|
||||||
|
|
@ -124,37 +131,29 @@ def create_obs_processor(
|
||||||
return padded
|
return padded
|
||||||
|
|
||||||
def _split_to_agents(obs: dict, morph_mode) -> dict:
|
def _split_to_agents(obs: dict, morph_mode) -> dict:
|
||||||
total = 0
|
key_to_agents = {}
|
||||||
for k, v in obs.items():
|
|
||||||
if hasattr(v, "shape"):
|
|
||||||
size = v.size
|
|
||||||
logger11.debug(f"[RAW] {k}: shape={v.shape}, size={size}")
|
|
||||||
total += size
|
|
||||||
else:
|
|
||||||
logger11.debug(f"[RAW] {k}: non-array")
|
|
||||||
|
|
||||||
logger11.debug(f"[RAW TOTAL FEATURES]: {total}")
|
|
||||||
output = {}
|
|
||||||
num_agents = needed_copies # IMPORTANT: number of MLPs
|
num_agents = needed_copies # IMPORTANT: number of MLPs
|
||||||
|
|
||||||
for key, arr in obs.items():
|
for key, arr in obs.items():
|
||||||
if key not in ordered_keys or arr.size == 0:
|
# TODO Should this still be here?
|
||||||
|
if arr.size == 0:
|
||||||
continue
|
continue
|
||||||
logger11.debug(f"[INPUT] {key}: {arr.shape}")
|
|
||||||
|
logger.debug(f"[INPUT] {key}: {arr.shape}")
|
||||||
|
|
||||||
if arr.ndim == 0:
|
if arr.ndim == 0:
|
||||||
arr = arr.reshape(1)
|
arr = arr.reshape(1)
|
||||||
# -------- CENTRALIZED --------
|
|
||||||
if morph_mode == MorphMode.CENTRALIZED:
|
if morph_mode == MorphMode.CENTRALIZED:
|
||||||
output[key] = arr.reshape(1, -1)
|
key_to_agents[key] = arr.reshape(1, -1)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# -------- SEGMENTS --------
|
|
||||||
if key in _SEGMENT_SCALED_KEYS:
|
if key in _SEGMENT_SCALED_KEYS:
|
||||||
per_agent = []
|
per_agent = []
|
||||||
|
|
||||||
for i, agent_id in enumerate(agent_indices):
|
for agent_id in agent_indices:
|
||||||
idx = segment_indices[i]
|
taken = jnp.take(arr, agent_id, axis=0) # (segs, ...)
|
||||||
taken = jnp.take(arr, idx, axis=0) # (segs, ...)
|
logger.debug(f"WHY {taken.shape}")
|
||||||
logger11.debug(f"WHY {taken.shape}")
|
|
||||||
# pad to 4
|
# pad to 4
|
||||||
pad_len = 4 - taken.shape[0]
|
pad_len = 4 - taken.shape[0]
|
||||||
padded = jnp.pad(taken, [(0, pad_len)] + [(0, 0)] * (taken.ndim - 1))
|
padded = jnp.pad(taken, [(0, pad_len)] + [(0, 0)] * (taken.ndim - 1))
|
||||||
|
|
@ -163,13 +162,11 @@ def create_obs_processor(
|
||||||
|
|
||||||
out = jnp.stack(per_agent)
|
out = jnp.stack(per_agent)
|
||||||
|
|
||||||
# -------- JOINTS --------
|
|
||||||
elif key in _JOINT_SCALED_KEYS:
|
elif key in _JOINT_SCALED_KEYS:
|
||||||
per_agent = []
|
per_agent = []
|
||||||
|
|
||||||
for i, agent_id in enumerate(agent_indices):
|
for agent_id in agent_indices:
|
||||||
idx = joint_indices[i]
|
taken = jnp.take(arr, agent_id, axis=0) # (joint_n, ...)
|
||||||
taken = jnp.take(arr, idx, axis=0) # (joint_n, ...)
|
|
||||||
# pad to 8
|
# pad to 8
|
||||||
pad_len = 8 - taken.shape[0]
|
pad_len = 8 - taken.shape[0]
|
||||||
|
|
||||||
|
|
@ -182,10 +179,10 @@ def create_obs_processor(
|
||||||
else:
|
else:
|
||||||
out = jnp.repeat(arr[None, :], num_agents, axis=0)
|
out = jnp.repeat(arr[None, :], num_agents, axis=0)
|
||||||
|
|
||||||
logger11.debug(f"[OUTPUT] {key}: {out.shape}")
|
logger.debug(f"[OUTPUT] {key}: {out.shape}")
|
||||||
output[key] = out
|
key_to_agents[key] = out
|
||||||
|
|
||||||
return output
|
return key_to_agents
|
||||||
|
|
||||||
def _flatten_features(obs: dict) -> jnp.ndarray:
|
def _flatten_features(obs: dict) -> jnp.ndarray:
|
||||||
"""
|
"""
|
||||||
|
|
@ -217,11 +214,14 @@ def create_obs_processor(
|
||||||
|
|
||||||
def _process_single(obs_dict: dict) -> jnp.ndarray:
|
def _process_single(obs_dict: dict) -> jnp.ndarray:
|
||||||
processed = _add_derived_features(obs_dict)
|
processed = _add_derived_features(obs_dict)
|
||||||
|
processed = _prune_features(processed)
|
||||||
processed = _normalize_features(processed)
|
processed = _normalize_features(processed)
|
||||||
processed = _split_to_agents(processed, morph_mode)
|
processed = _split_to_agents(processed, morph_mode)
|
||||||
flat = _flatten_features(processed) # (num_arms, total_feat)
|
flat = _flatten_features(processed) # (num_arms, total_feat)
|
||||||
logger11.debug(f"[FLATTENED FINAL] shape: {flat.shape}")
|
|
||||||
logger11.debug(f"[PER AGENT] example row 0 shape: {flat[0].shape}")
|
logger.debug(f"[FLATTENED FINAL] shape: {flat.shape}")
|
||||||
return _flatten_features(processed) # (agents, feat)
|
logger.debug(f"[PER AGENT] example row 0 shape: {flat[0].shape}")
|
||||||
|
|
||||||
|
return flat # (agents, feat)
|
||||||
|
|
||||||
return jax.jit(jax.vmap(_process_single))
|
return jax.jit(jax.vmap(_process_single))
|
||||||
|
|
|
||||||
|
|
@ -31,8 +31,6 @@ from brittle_star_project.ppo import PPO
|
||||||
from brittle_star_project.environment import MorphMode
|
from brittle_star_project.environment import MorphMode
|
||||||
from brittle_star_project.utils import logged_jit
|
from brittle_star_project.utils import logged_jit
|
||||||
|
|
||||||
logger11 = get_logger()
|
|
||||||
|
|
||||||
|
|
||||||
@logged_jit
|
@logged_jit
|
||||||
def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray:
|
def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray:
|
||||||
|
|
@ -121,14 +119,16 @@ def _step_once(
|
||||||
action_low,
|
action_low,
|
||||||
action_high,
|
action_high,
|
||||||
)
|
)
|
||||||
logger11.debug(f"[_step_once] raw_action: {raw_action.shape}")
|
logger = get_logger()
|
||||||
logger11.debug(f"[_step_once] clipped_action: {flat_clipped_action.shape}")
|
|
||||||
|
logger.debug(f"[_step_once] raw_action: {raw_action.shape}")
|
||||||
|
logger.debug(f"[_step_once] clipped_action: {flat_clipped_action.shape}")
|
||||||
|
|
||||||
# Supporting signals (often where mismatch originates)
|
# Supporting signals (often where mismatch originates)
|
||||||
logger11.debug(f"[_step_once] logprob: {logprob.shape}")
|
logger.debug(f"[_step_once] logprob: {logprob.shape}")
|
||||||
logger11.debug(f"[_step_once] value: {value.shape}")
|
logger.debug(f"[_step_once] value: {value.shape}")
|
||||||
logger11.debug(f"[_step_once] mean: {mean.shape}")
|
logger.debug(f"[_step_once] mean: {mean.shape}")
|
||||||
logger11.debug(f"[_step_once] std: {std.shape}")
|
logger.debug(f"[_step_once] std: {std.shape}")
|
||||||
|
|
||||||
key, reset_key = jax.random.split(key)
|
key, reset_key = jax.random.split(key)
|
||||||
reset_rngs = jax.random.split(reset_key, num_envs)
|
reset_rngs = jax.random.split(reset_key, num_envs)
|
||||||
|
|
@ -147,9 +147,9 @@ def _step_once(
|
||||||
terminated_any = terminated_any | terminated
|
terminated_any = terminated_any | terminated
|
||||||
truncated_any = truncated_any | truncated
|
truncated_any = truncated_any | truncated
|
||||||
|
|
||||||
logger11.debug(f"[_step_once] next_obs: {next_obs.shape}")
|
logger.debug(f"[_step_once] next_obs: {next_obs.shape}")
|
||||||
logger11.debug(f"[_step_once] reward: {reward.shape}")
|
logger.debug(f"[_step_once] reward: {reward.shape}")
|
||||||
logger11.debug(f"[_step_once] next_done: {next_done.shape}")
|
logger.debug(f"[_step_once] next_done: {next_done.shape}")
|
||||||
|
|
||||||
storage = Storage(
|
storage = Storage(
|
||||||
obs=obs,
|
obs=obs,
|
||||||
|
|
@ -547,7 +547,6 @@ class PPOTrainer:
|
||||||
case MorphMode.SEGMENT:
|
case MorphMode.SEGMENT:
|
||||||
agent_mask = self.segments_per_arm > 0
|
agent_mask = self.segments_per_arm > 0
|
||||||
agent_indices = jnp.where(agent_mask)[0]
|
agent_indices = jnp.where(agent_mask)[0]
|
||||||
needed_copies = jnp.where(self.segments_per_arm > 0, 1, 0).sum().item()
|
|
||||||
needed_copies = (
|
needed_copies = (
|
||||||
self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum()
|
self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum()
|
||||||
).item()
|
).item()
|
||||||
|
|
@ -627,17 +626,17 @@ class PPOTrainer:
|
||||||
|
|
||||||
message_passer_params = {}
|
message_passer_params = {}
|
||||||
if self.morph_mode != MorphMode.CENTRALIZED:
|
if self.morph_mode != MorphMode.CENTRALIZED:
|
||||||
assert self.message_passer is not None, "MessagePasser is None"
|
assert self.message_passer is not None, "decentralized modes require a message passer"
|
||||||
|
|
||||||
message_passer_params = self.message_passer.init(
|
message_passer_params = self.message_passer.init(
|
||||||
message_passer_key,
|
message_passer_key,
|
||||||
self.sensor.apply(single_sensor_param, sample_obs),
|
self.sensor.apply(single_sensor_param, sample_obs),
|
||||||
)
|
)
|
||||||
self.logger.debug(
|
self.logger.debug(
|
||||||
f"[_init_agent_state] message_passer_params: {
|
f"[_init_agent_state] message_passer_params: {
|
||||||
jax.tree.map(lambda x: x.shape, message_passer_params)
|
jax.tree.map(lambda x: x.shape, message_passer_params)
|
||||||
}"
|
}"
|
||||||
)
|
)
|
||||||
|
|
||||||
flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic
|
flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic
|
||||||
self.logger.debug(f"[_init_agent_state] flat_obs: {flat_obs.shape}")
|
self.logger.debug(f"[_init_agent_state] flat_obs: {flat_obs.shape}")
|
||||||
|
|
|
||||||
Reference in a new issue