feat(obs_processor): rewrote split method, todo: testing + code cleanup
This commit is contained in:
parent
07ee32fd3b
commit
34f0b1b80e
10 changed files with 87 additions and 34 deletions
|
|
@ -102,61 +102,47 @@ def create_obs_processor(
|
|||
return normalized
|
||||
|
||||
def _split_to_agents(obs: dict, morph_mode) -> dict:
|
||||
# TODO: cleanup + MORE testing (works for centralized 5 arms + damaged arms)
|
||||
output = {}
|
||||
num_agents = needed_copies # IMPORTANT: number of MLPs
|
||||
for key, arr in obs.items():
|
||||
if key not in ordered_keys or arr.size == 0:
|
||||
if arr.size == 0:
|
||||
continue
|
||||
|
||||
logger.debug(f"[INPUT] {key}: {arr.shape}")
|
||||
|
||||
if arr.ndim == 0:
|
||||
arr = arr.reshape(1)
|
||||
|
||||
# -------- CENTRALIZED --------
|
||||
if morph_mode == MorphMode.CENTRALIZED:
|
||||
output[key] = arr.reshape(1, -1)
|
||||
# TODO: padding for centralized
|
||||
continue
|
||||
|
||||
# -------- SEGMENTS --------
|
||||
if key in _SEGMENT_SCALED_KEYS:
|
||||
per_agent = []
|
||||
|
||||
for i, agent_id in enumerate(agent_indices):
|
||||
for i, _ in enumerate(agent_indices):
|
||||
idx = segment_indices[i]
|
||||
taken = jnp.take(arr, idx, axis=0) # (segs, ...)
|
||||
logger.debug(f"WHY {taken.shape}")
|
||||
|
||||
# pad to 4 (segments per arm?)
|
||||
taken = jnp.take(arr, idx, axis=0)
|
||||
pad_len = 4 - taken.shape[0]
|
||||
padded = jnp.pad(taken, [(0, pad_len)] + [(0, 0)] * (taken.ndim - 1))
|
||||
padded = jnp.pad(taken, [(9, pad_len)] + [(0, 0)] * (taken.ndim - 1))
|
||||
|
||||
per_agent.append(padded.reshape(-1))
|
||||
|
||||
out = jnp.stack(per_agent)
|
||||
|
||||
# -------- JOINTS --------
|
||||
arr = jnp.stack(per_agent)
|
||||
elif key in _JOINT_SCALED_KEYS:
|
||||
per_agent = []
|
||||
|
||||
for i, _ in enumerate(agent_indices):
|
||||
idx = joint_indices[i]
|
||||
taken = jnp.take(arr, idx, axis=0) # (joint_n, ...)
|
||||
# pad to 8
|
||||
taken = jnp.take(arr, idx, axis=0)
|
||||
pad_len = 8 - taken.shape[0]
|
||||
|
||||
padded = jnp.pad(taken, [(0, pad_len)] + [(0, 0)] * (taken.ndim - 1))
|
||||
|
||||
per_agent.append(padded.reshape(-1))
|
||||
|
||||
out = jnp.stack(per_agent)
|
||||
|
||||
# -------- GLOBAL --------
|
||||
arr = jnp.stack(per_agent)
|
||||
else:
|
||||
out = jnp.repeat(arr[None, :], num_agents, axis=0)
|
||||
arr = jnp.repeat(arr[None, :], num_agents, axis=0)
|
||||
|
||||
logger.debug(f"[OUTPUT] {key}: {out.shape}")
|
||||
output[key] = out
|
||||
if morph_mode == MorphMode.CENTRALIZED:
|
||||
output[key] = arr.reshape(1, -1)
|
||||
elif key in _JOINT_SCALED_KEYS:
|
||||
output[key] = arr.reshape(num_agents, -1)
|
||||
elif key in _SEGMENT_SCALED_KEYS:
|
||||
output[key] = arr[:, None]
|
||||
else:
|
||||
output[key] = arr
|
||||
|
||||
return output
|
||||
|
||||
|
|
|
|||
17
src/brittle_star_project/environment/obs_processing_tmp.py
Normal file
17
src/brittle_star_project/environment/obs_processing_tmp.py
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
obs = {}
|
||||
segs = set()
|
||||
joints = set()
|
||||
|
||||
for key, arr in obs.items():
|
||||
if key in segs:
|
||||
# padding using mask 1x
|
||||
pass
|
||||
elif key in joints:
|
||||
# padding using mask 2x
|
||||
pass
|
||||
else:
|
||||
# no padding
|
||||
pass
|
||||
|
||||
# reshape according to segs, joints, or global
|
||||
pass
|
||||
|
|
@ -277,7 +277,6 @@ def apply_shared(net, params, x):
|
|||
return jax.vmap(lambda xi: net.apply(params, xi))(x_flattened)
|
||||
|
||||
|
||||
# TODO: update to work with extra dimension + message passing
|
||||
def _rollout_jit(
|
||||
agent_state,
|
||||
episode_stats,
|
||||
|
|
|
|||
Reference in a new issue