1
Fork 0

fix(testing): target direction updated to work with dim >= 2

This commit is contained in:
Robin Meersman 2026-05-05 11:49:32 +02:00
parent b244761db1
commit 3ad1ad8e39
3 changed files with 30 additions and 5 deletions

View file

@ -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)

View file

@ -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

View file

@ -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()