From 5e7fd27b61a9927ae2c824c81c736404f4277294 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Wed, 6 May 2026 19:26:32 +0200 Subject: [PATCH] backup --- .../environment/obs_processing.py | 89 +++++++++++++------ .../trainers/PPOTrainer.py | 2 +- tests/test_network_shapes.py | 3 +- tests/test_target_direction.py | 3 +- 4 files changed, 63 insertions(+), 34 deletions(-) diff --git a/src/brittle_star_project/environment/obs_processing.py b/src/brittle_star_project/environment/obs_processing.py index 5733177..96dd4ab 100644 --- a/src/brittle_star_project/environment/obs_processing.py +++ b/src/brittle_star_project/environment/obs_processing.py @@ -22,11 +22,23 @@ _SEGMENT_SCALED_KEYS = frozenset( def create_obs_processor( bounds_dict: Dict[str, Tuple[float, float]], - num_segments: int, num_arms: int, + needed_copies: int, padding_masks: Optional[Dict] = None, morph_mode: MorphMode = MorphMode.CENTRALIZED, ): + # made a set to allow O(1) search + ordered_keys = frozenset( + [ + "disk_z_tilt", + "joint_actuator_force", + "joint_position", + "joint_velocity", + "robot_direction_to_target", + "segment_contact", + ] + ) + def _add_derived_features(obs: dict) -> dict: new_obs = dict(obs) if "disk_rotation" in new_obs: @@ -57,49 +69,61 @@ def create_obs_processor( normalized[key] = arr return normalized - def _pad_features(obs: dict) -> dict: + # ndarray (agents, features) + def _pad_features(obs: dict, agent_count: int) -> dict: assert padding_masks is not None padded = {} + + # arr shape = (agent_count, ...) TODO for key, arr in obs.items(): if key in _JOINT_SCALED_KEYS: - padded_arr = jnp.zeros(padding_masks["target_size_2x"], dtype=arr.dtype) - padded[key] = padded_arr.at[padding_masks["mask_2x"]].set(arr) + padded_arr = jnp.zeros( + (agent_count, padding_masks["target_size_2x"]), dtype=arr.dtype + ) + padded[key] = padded_arr.at[:, padding_masks["mask_2x"]].set(arr) elif key in _SEGMENT_SCALED_KEYS: - padded_arr = jnp.zeros(padding_masks["target_size_1x"], dtype=arr.dtype) - padded[key] = padded_arr.at[padding_masks["mask_1x"]].set(arr) + padded_arr = jnp.zeros( + (agent_count, padding_masks["target_size_1x"]), dtype=arr.dtype + ) + padded[key] = padded_arr.at[:, padding_masks["mask_1x"]].set(arr) else: padded[key] = arr return padded - # TODO - def _split_to_agents() -> dict: - return {} + def _split_to_agents(obs: dict, agent_count: int) -> dict: + """ + Each observation type is now updated to (agent_count, feature_shape), + thus duplicating the observation for each agent + """ + output = {} + + for key, arr in obs.items(): + if key not in ordered_keys or arr.size == 0: + continue + + if arr.ndim == 0: + arr = arr.reshape(1) + + output[key] = jnp.repeat(arr[None, :], agent_count, axis=0) + + return output + + # TODO: update def _flatten_features(obs: dict) -> jnp.ndarray: - ordered_keys = [ - "disk_z_tilt", - "joint_actuator_force", - "joint_position", - "joint_velocity", - "robot_direction_to_target", - "segment_contact", - ] + """ + Collapse all observations into a single array + """ values = [] for key in sorted(obs.keys()): if key not in ordered_keys: continue - v = obs[key] - # skip empty arrays and scalars - if v.size == 0: - continue - - # reshape scalars to 1D array - if v.ndim == 0: - v = v.reshape(1) + # empty arrays handled in split_to_agents + v = obs[key] # (agent_count, feat) # -------- CENTRALIZED -------- if morph_mode == MorphMode.CENTRALIZED: @@ -124,11 +148,18 @@ def create_obs_processor( return jnp.concatenate(values, axis=-1) # (agent, feat) def _process_single(obs_dict: dict) -> jnp.ndarray: - processed = _add_derived_features(obs_dict) # key |--> (feat-count, feat-lengths, ) - processed = _normalize_features(processed) # key |--> (feat-count, feat-lengths, ) - # TODO: split to agents # key |--> (agents, feat-count, feat-lengths) + processed = _add_derived_features(obs_dict) # key |--> (feat-count,) + processed = _normalize_features(processed) # key |--> (feat-count,) + # TODO: split to agents + processed = _split_to_agents(processed, needed_copies) # key |--> (agents, feat-count) if padding_masks is not None: - processed = _pad_features(processed) # key |--> (agents, feat-count, feat-lengths') + processed = _pad_features( + processed, agent_count=needed_copies + ) # key |--> (agents, feat-count') + for k, v in processed.items(): + print(k, v.shape) + + exit(1) return _flatten_features(processed) # (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 43de64b..05761f5 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -363,7 +363,7 @@ class PPOTrainer: # Build the centralized observation processor: derive -> normalize -> pad -> flatten. self.obs_processor = create_obs_processor( bounds_dict=self.cfg.obs_bounds.to_bounds_dict(), - num_segments=self.num_segments, + needed_copies=self.needed_copies, num_arms=self.num_arms, morph_mode=self.morph_mode, padding_masks=self.env.padding_masks, diff --git a/tests/test_network_shapes.py b/tests/test_network_shapes.py index 1183377..3cf4c14 100644 --- a/tests/test_network_shapes.py +++ b/tests/test_network_shapes.py @@ -22,14 +22,13 @@ def test_centralized_forward_pass_with_padding(): } segments_per_arm = jnp.array((4, 0, 4, 2, 4)) - num_segments = segments_per_arm.sum().item() num_arms = jnp.where(segments_per_arm > 0, 1, 0).sum().item() # 2. Process and Pad Observation masks = compute_padding_masks(segments_per_arm=list(segments_per_arm)) obs_processor = create_obs_processor( bounds_dict={}, - num_segments=num_segments, + needed_copies=1, num_arms=num_arms, padding_masks=masks, morph_mode=MorphMode.CENTRALIZED, diff --git a/tests/test_target_direction.py b/tests/test_target_direction.py index 64a6361..f931c8a 100644 --- a/tests/test_target_direction.py +++ b/tests/test_target_direction.py @@ -45,12 +45,11 @@ def test_processor_converts_to_egocentric_direction(): env = BrittleStarJaxEnvWrapper.default(num_envs=1, backend=Backend.MJX) segments_per_arm = jnp.array((4, 4, 4, 4, 4)) - num_segments = segments_per_arm.sum().item() num_arms = jnp.where(segments_per_arm > 0, 1, 0).sum().item() obs_processor = create_obs_processor( bounds_dict=cfg.obs_bounds.to_bounds_dict(), - num_segments=num_segments, + needed_copies=1, num_arms=num_arms, padding_masks=env.padding_masks, morph_mode=MorphMode.CENTRALIZED,