From 143dbbc710ba90496491ad9eaeb045f4ac1f9438 Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Wed, 15 Apr 2026 19:26:33 +0200 Subject: [PATCH] fix: explicit arms length check --- configs/architecture/centralized.yaml | 1 + configs/architecture/decentralized.yaml | 1 + .../environment/padded_obs_wrapper.py | 14 +++++++++++++- tests/test_configs.py | 5 +++-- 4 files changed, 18 insertions(+), 3 deletions(-) diff --git a/configs/architecture/centralized.yaml b/configs/architecture/centralized.yaml index 62ff16d..3b7783d 100644 --- a/configs/architecture/centralized.yaml +++ b/configs/architecture/centralized.yaml @@ -4,6 +4,7 @@ # Default values are defined in CentralizedConfig dataclass. # Use this configuration for standard PPO experiments. +name: "centralized" sensor: hidden_dims: [64, 64] activation: "tanh" diff --git a/configs/architecture/decentralized.yaml b/configs/architecture/decentralized.yaml index d87e144..fbf07bc 100644 --- a/configs/architecture/decentralized.yaml +++ b/configs/architecture/decentralized.yaml @@ -4,6 +4,7 @@ # Default values are defined in DecentralizedConfig dataclass. # Use this configuration for decentralized execution experiments. +name: "decentralized" sensor: hidden_dims: [64, 64] activation: "tanh" diff --git a/src/brittle_star_project/environment/padded_obs_wrapper.py b/src/brittle_star_project/environment/padded_obs_wrapper.py index 95f310b..3f22038 100644 --- a/src/brittle_star_project/environment/padded_obs_wrapper.py +++ b/src/brittle_star_project/environment/padded_obs_wrapper.py @@ -42,10 +42,22 @@ def compute_padding_masks( Returns: A dict containing 1D boolean masks and target sizes. """ + if len(segments_per_arm) != len(reference_segments_per_arm): + raise ValueError( + f"Morphology mismatch: current has {len(segments_per_arm)} arms, " + f"but reference requires {len(reference_segments_per_arm)} arms." + ) + mask_1x = [] mask_2x = [] - for actual, ref in zip(segments_per_arm, reference_segments_per_arm): + for arm_idx, (actual, ref) in enumerate(zip(segments_per_arm, reference_segments_per_arm)): + if not (0 <= actual <= ref): + raise ValueError( + f"Invalid amputation at arm {arm_idx}: " + f"actual segments ({actual}) must be between 0 and reference ({ref})." + ) + # 1x scaling (e.g., contacts: 1 value per segment) # 1x scaling (e.g., contacts: 1 value per segment) mask_1x.extend([True] * actual + [False] * (ref - actual)) # 2x scaling (e.g., joints: 2 values per segment) diff --git a/tests/test_configs.py b/tests/test_configs.py index 84a7cde..c8cad63 100644 --- a/tests/test_configs.py +++ b/tests/test_configs.py @@ -18,7 +18,8 @@ def test_config_composition_centralized(): structured_cfg = OmegaConf.to_object(OmegaConf.merge(BrittleStarConfig, cfg)) # Basic assertions - assert "CentralizedConfig" in str(type(structured_cfg.architecture)) + assert structured_cfg.architecture.name == "centralized" + assert structured_cfg.architecture.propagator is None assert isinstance(structured_cfg.ppo.learning_rate, float) assert structured_cfg.ppo.learning_rate > 0 @@ -32,7 +33,7 @@ def test_config_composition_decentralized(): structured_cfg = OmegaConf.to_object(OmegaConf.merge(BrittleStarConfig, cfg)) # Basic assertions - assert "DecentralizedConfig" in str(type(structured_cfg.architecture)) + assert structured_cfg.architecture.name == "decentralized" assert isinstance(structured_cfg.ppo.learning_rate, float) assert structured_cfg.ppo.learning_rate > 0