diff --git a/src/brittle_star_project/environment/obs_processing.py b/src/brittle_star_project/environment/obs_processing.py index 14b48f2..73c2e94 100644 --- a/src/brittle_star_project/environment/obs_processing.py +++ b/src/brittle_star_project/environment/obs_processing.py @@ -105,6 +105,7 @@ def create_obs_processor( # -------- SPLIT TO SEGMENTS -------- if key in _JOINT_SCALED_KEYS: if morph_mode == MorphMode.SEGMENT: + # TODO: remove magic constants center_size = num_arms * 3 * 2 v_center = v[:center_size].reshape(num_arms, 3 * 2) # (arms, 6) v_segs = v[center_size:].reshape(-1, 2) # (segs, 2) diff --git a/tests/test_network_shapes.py b/tests/test_network_shapes.py index e1a7a1a..aad79fc 100644 --- a/tests/test_network_shapes.py +++ b/tests/test_network_shapes.py @@ -20,9 +20,18 @@ def test_centralized_forward_pass_with_padding(): "segment_contact": jnp.zeros((batch_size, 14)), } + 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=(4, 0, 4, 2, 4)) - obs_processor = create_obs_processor(bounds_dict={}, padding_masks=masks) + masks = compute_padding_masks(segments_per_arm=list(segments_per_arm)) + obs_processor = create_obs_processor( + bounds_dict={}, + num_segments=num_segments, + num_arms=num_arms, + padding_masks=masks, + ) global_state = obs_processor(amputated_obs) # 40 + 40 + 20 = 100 dimensions diff --git a/tests/test_target_direction.py b/tests/test_target_direction.py index a81cdf1..825c50c 100644 --- a/tests/test_target_direction.py +++ b/tests/test_target_direction.py @@ -1,6 +1,7 @@ import jax.numpy as jnp from brittle_star_project.configs.main_config import BrittleStarConfig +from brittle_star_project.environment.env_config import MorphMode from brittle_star_project.environment.env_types import Backend from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.environment.obs_processing import create_obs_processor @@ -43,8 +44,16 @@ def test_processor_converts_to_egocentric_direction(): cfg = BrittleStarConfig() 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(), padding_masks=env.padding_masks + bounds_dict=cfg.obs_bounds.to_bounds_dict(), + num_segments=num_segments, + num_arms=num_arms, + padding_masks=env.padding_masks, + morph_mode=MorphMode.CENTRALIZED, ) env_state = env.reset(seed=42) @@ -66,13 +75,19 @@ def test_processor_converts_to_egocentric_direction(): processed_2 = obs_processor(dummy_obs_2) # Find the indices of the elements that changed - diff_array = jnp.abs(processed_1[0] - processed_2[0]) + diff_array = jnp.abs(processed_1[0, 0] - processed_2[0, 0]) changed_indices = jnp.where(diff_array > 1e-4)[0] - local_target = processed_1[0, changed_indices] + # (143,) + local_target = processed_1[0, 0, changed_indices] + + # (2,) expected_local_target = jnp.array([0.0, -1.0]) assert jnp.sum(jnp.abs(local_target - expected_local_target)) < 1e-4, ( f"The obs_processor did not correctly rotate the vector to egocentric. " f"Expected {expected_local_target}, but got {local_target}." ) + + +test_processor_converts_to_egocentric_direction()