diff --git a/src/brittle_star_project/environment/obs_processing.py b/src/brittle_star_project/environment/obs_processing.py index 41afe57..216b008 100644 --- a/src/brittle_star_project/environment/obs_processing.py +++ b/src/brittle_star_project/environment/obs_processing.py @@ -6,8 +6,6 @@ from brittle_star_project.environment.env_config import MorphMode from experiment_logger import get_logger -logger11 = get_logger() - _JOINT_SCALED_KEYS = frozenset( { "joint_position", @@ -57,6 +55,8 @@ def create_obs_processor( segments_per_arm=[4, 4, 4, 4, 4], agent_indices=[0, 1, 2, 3, 4], ): + logger = get_logger() + # made a set to allow O(1) search ordered_keys = frozenset( [ @@ -87,6 +87,13 @@ def create_obs_processor( 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: normalized = {} for key, arr in obs.items(): @@ -124,37 +131,29 @@ def create_obs_processor( return padded def _split_to_agents(obs: dict, morph_mode) -> dict: - total = 0 - 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 = {} + key_to_agents = {} num_agents = needed_copies # IMPORTANT: number of MLPs + 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 - logger11.debug(f"[INPUT] {key}: {arr.shape}") + + 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) + key_to_agents[key] = arr.reshape(1, -1) continue - # -------- SEGMENTS -------- if key in _SEGMENT_SCALED_KEYS: per_agent = [] - for i, agent_id in enumerate(agent_indices): - idx = segment_indices[i] - taken = jnp.take(arr, idx, axis=0) # (segs, ...) - logger11.debug(f"WHY {taken.shape}") + for agent_id in agent_indices: + taken = jnp.take(arr, agent_id, axis=0) # (segs, ...) + logger.debug(f"WHY {taken.shape}") # pad to 4 pad_len = 4 - taken.shape[0] 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) - # -------- JOINTS -------- elif key in _JOINT_SCALED_KEYS: per_agent = [] - for i, agent_id in enumerate(agent_indices): - idx = joint_indices[i] - taken = jnp.take(arr, idx, axis=0) # (joint_n, ...) + for agent_id in agent_indices: + taken = jnp.take(arr, agent_id, axis=0) # (joint_n, ...) # pad to 8 pad_len = 8 - taken.shape[0] @@ -182,10 +179,10 @@ def create_obs_processor( else: out = jnp.repeat(arr[None, :], num_agents, axis=0) - logger11.debug(f"[OUTPUT] {key}: {out.shape}") - output[key] = out + logger.debug(f"[OUTPUT] {key}: {out.shape}") + key_to_agents[key] = out - return output + return key_to_agents def _flatten_features(obs: dict) -> jnp.ndarray: """ @@ -217,11 +214,14 @@ def create_obs_processor( def _process_single(obs_dict: dict) -> jnp.ndarray: processed = _add_derived_features(obs_dict) + processed = _prune_features(processed) processed = _normalize_features(processed) processed = _split_to_agents(processed, morph_mode) 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}") - return _flatten_features(processed) # (agents, feat) + + logger.debug(f"[FLATTENED FINAL] shape: {flat.shape}") + logger.debug(f"[PER AGENT] example row 0 shape: {flat[0].shape}") + + return flat # (agents, feat) return jax.jit(jax.vmap(_process_single)) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index f7c6489..f82ba83 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -31,8 +31,6 @@ from brittle_star_project.ppo import PPO from brittle_star_project.environment import MorphMode from brittle_star_project.utils import logged_jit -logger11 = get_logger() - @logged_jit def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray: @@ -121,14 +119,16 @@ def _step_once( action_low, action_high, ) - logger11.debug(f"[_step_once] raw_action: {raw_action.shape}") - logger11.debug(f"[_step_once] clipped_action: {flat_clipped_action.shape}") + logger = get_logger() + + 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) - logger11.debug(f"[_step_once] logprob: {logprob.shape}") - logger11.debug(f"[_step_once] value: {value.shape}") - logger11.debug(f"[_step_once] mean: {mean.shape}") - logger11.debug(f"[_step_once] std: {std.shape}") + logger.debug(f"[_step_once] logprob: {logprob.shape}") + logger.debug(f"[_step_once] value: {value.shape}") + logger.debug(f"[_step_once] mean: {mean.shape}") + logger.debug(f"[_step_once] std: {std.shape}") key, reset_key = jax.random.split(key) reset_rngs = jax.random.split(reset_key, num_envs) @@ -147,9 +147,9 @@ def _step_once( terminated_any = terminated_any | terminated truncated_any = truncated_any | truncated - logger11.debug(f"[_step_once] next_obs: {next_obs.shape}") - logger11.debug(f"[_step_once] reward: {reward.shape}") - logger11.debug(f"[_step_once] next_done: {next_done.shape}") + logger.debug(f"[_step_once] next_obs: {next_obs.shape}") + logger.debug(f"[_step_once] reward: {reward.shape}") + logger.debug(f"[_step_once] next_done: {next_done.shape}") storage = Storage( obs=obs, @@ -547,7 +547,6 @@ class PPOTrainer: case MorphMode.SEGMENT: agent_mask = self.segments_per_arm > 0 agent_indices = jnp.where(agent_mask)[0] - needed_copies = jnp.where(self.segments_per_arm > 0, 1, 0).sum().item() needed_copies = ( self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum() ).item() @@ -627,17 +626,17 @@ class PPOTrainer: message_passer_params = {} 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_key, self.sensor.apply(single_sensor_param, sample_obs), ) - self.logger.debug( - f"[_init_agent_state] message_passer_params: { - jax.tree.map(lambda x: x.shape, message_passer_params) - }" - ) + self.logger.debug( + f"[_init_agent_state] message_passer_params: { + jax.tree.map(lambda x: x.shape, message_passer_params) + }" + ) flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic self.logger.debug(f"[_init_agent_state] flat_obs: {flat_obs.shape}")