From ba8184b03423a0e2e4825224dd035d41b29235a5 Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Tue, 28 Apr 2026 18:03:53 +0200 Subject: [PATCH] test: fix outdated imports etc. --- tests/test_network_shapes.py | 19 ++++--------------- 1 file changed, 4 insertions(+), 15 deletions(-) diff --git a/tests/test_network_shapes.py b/tests/test_network_shapes.py index 32e9898..e1a7a1a 100644 --- a/tests/test_network_shapes.py +++ b/tests/test_network_shapes.py @@ -2,8 +2,8 @@ import jax import jax.numpy as jnp from brittle_star_project.environment.padded_obs_wrapper import ( compute_padding_masks, - pad_observations_batched, ) +from brittle_star_project.environment.obs_processing import create_obs_processor # We use Actor and OneDenseLayerMLP (as the critic) based on your mlps.py from brittle_star_project.MLPs.mlps import Actor, OneDenseLayerMLP @@ -20,21 +20,10 @@ def test_centralized_forward_pass_with_padding(): "segment_contact": jnp.zeros((batch_size, 14)), } - # 2. Pad Observation using the boolean scattering wrapper + # 2. Process and Pad Observation masks = compute_padding_masks(segments_per_arm=(4, 0, 4, 2, 4)) - padded_obs = pad_observations_batched(amputated_obs, masks) - - # Assertions to ensure padding sizes are correct (40 joints, 20 segments) - assert padded_obs["joint_position"].shape == (batch_size, 40), "Padding failed for joint keys" - assert padded_obs["segment_contact"].shape == (batch_size, 20), ( - "Padding failed for segment keys" - ) - - # 3. Concatenate for Centralized MLP (simulating the global state vector) - global_state = jnp.concatenate( - [padded_obs["joint_position"], padded_obs["joint_velocity"], padded_obs["segment_contact"]], - axis=-1, - ) + obs_processor = create_obs_processor(bounds_dict={}, padding_masks=masks) + global_state = obs_processor(amputated_obs) # 40 + 40 + 20 = 100 dimensions assert global_state.shape == (batch_size, 100), (