From eff0c7c1df16f4647e1d55e114022bc60034f9d6 Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Wed, 15 Apr 2026 16:47:17 +0200 Subject: [PATCH] test: additional testing and visual confirmations --- .../configs/config_experiment.py | 1 + .../configs/main_config.py | 5 +- .../environment/env_config.py | 12 ++-- .../trainers/PPOTrainer.py | 7 ++ tests/test_config.py | 33 --------- tests/test_configs.py | 42 ++++++++++++ tests/test_morphology_render.py | 58 ++++++++++++++++ tests/test_network_shapes.py | 67 +++++++++++++++++++ 8 files changed, 184 insertions(+), 41 deletions(-) delete mode 100644 tests/test_config.py create mode 100644 tests/test_configs.py create mode 100644 tests/test_morphology_render.py create mode 100644 tests/test_network_shapes.py diff --git a/src/brittle_star_project/configs/config_experiment.py b/src/brittle_star_project/configs/config_experiment.py index 69ceca4..716d3b6 100644 --- a/src/brittle_star_project/configs/config_experiment.py +++ b/src/brittle_star_project/configs/config_experiment.py @@ -7,3 +7,4 @@ class ExperimentConfig: seed: int = 1 torch_deterministic: bool = True cuda: bool = True + debug_sanity: bool = False diff --git a/src/brittle_star_project/configs/main_config.py b/src/brittle_star_project/configs/main_config.py index aab864b..cd54473 100644 --- a/src/brittle_star_project/configs/main_config.py +++ b/src/brittle_star_project/configs/main_config.py @@ -18,8 +18,9 @@ class BrittleStarConfig: experiment: ExperimentConfig = field(default_factory=ExperimentConfig) logging: LoggingConfig = field(default_factory=LoggingConfig) ppo: PPOConfig = field(default_factory=PPOConfig) - # Default to centralized; swap with architecture=decentralized on the CLI. - architecture: ArchitectureConfig = field(default_factory=CentralizedConfig) + # This field is polymorphic; defaults to the base class to allow subclasses + # (CentralizedConfig, DecentralizedConfig) to be merged in via Hydra. + architecture: ArchitectureConfig = field(default_factory=ArchitectureConfig) morphology: MorphologyConfig = field(default_factory=MorphologyConfig) arena: ArenaConfig = field(default_factory=ArenaConfig) environment: EnvConfig = field(default_factory=EnvConfig) diff --git a/src/brittle_star_project/environment/env_config.py b/src/brittle_star_project/environment/env_config.py index 48f0e34..ee30e87 100644 --- a/src/brittle_star_project/environment/env_config.py +++ b/src/brittle_star_project/environment/env_config.py @@ -5,7 +5,7 @@ from dataclasses import dataclass, field from .env_types import Task -@dataclass(frozen=True, slots=True) +@dataclass class MorphologyConfig: """Brittle star morphology configuration. @@ -17,7 +17,7 @@ class MorphologyConfig: The upstream biorobot library natively supports per-arm segment counts. """ - segments_per_arm: tuple[int, ...] = (4, 4, 4, 4, 4) + segments_per_arm: list[int] = field(default_factory=lambda: [4, 4, 4, 4, 4]) use_p_control: bool = True use_torque_control: bool = False @@ -26,16 +26,16 @@ class MorphologyConfig: return len(self.segments_per_arm) -@dataclass(frozen=True, slots=True) +@dataclass class ArenaConfig: - size: tuple[float, float] = (10.0, 5.0) + size: list[float] = field(default_factory=lambda: [10.0, 5.0]) sand_ground_color: bool = True attach_target: bool = True wall_height: float = 1.5 wall_thickness: float = 0.1 -@dataclass(frozen=True, slots=True) +@dataclass class EnvConfig: """Shared environment settings. @@ -50,7 +50,7 @@ class EnvConfig: camera_ids: list[int] = field(default_factory=lambda: [0, 1]) # (height, width) - render_size: tuple[int, int] = (480, 640) + render_size: list[int] = field(default_factory=lambda: [480, 640]) joint_randomization_noise_scale: float = 0.0 diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index b20c4dd..3dc9666 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -527,6 +527,13 @@ class PPOTrainer: f"ETA {eta_str}" ) + if getattr(self.config.experiment, "debug_sanity", False): + self.logger.log( + "\n[SANITY CHECK] Successfully completed 1 epoch of data collection and gradient updates." + ) + self.logger.log("[SANITY CHECK] Gradients flowed without NaN. Exiting gracefully.") + break + if self.logging_cfg.save_model: model_path = f"{self.run_dir}/{self.experiment.exp_name}.cleanrl_model" self._save_model(model_path=model_path) diff --git a/tests/test_config.py b/tests/test_config.py deleted file mode 100644 index e0ee0ef..0000000 --- a/tests/test_config.py +++ /dev/null @@ -1,33 +0,0 @@ -"""Tests for YAML config loading.""" - -import sys -from pathlib import Path - -import pytest - -# Ensure src is on the path when running from the project root -sys.path.insert(0, str(Path(__file__).parent.parent / "src")) - -CONFIGS_DIR = Path(__file__).parent.parent / "configs" - - -class TestYamlConfig: - def test_load_yaml_config(self): - from experiment_logger.config_utils import load_yaml_config - - config = load_yaml_config(str(CONFIGS_DIR / "default_ppo.yaml")) - assert isinstance(config, dict) - assert "total_timesteps" in config - assert "learning_rate" in config - - def test_load_dev_test_config(self): - from experiment_logger.config_utils import load_yaml_config - - config = load_yaml_config(str(CONFIGS_DIR / "dev_test.yaml")) - assert config["total_timesteps"] == 100000 - - def test_missing_config_raises(self): - from experiment_logger.config_utils import load_yaml_config - - with pytest.raises(FileNotFoundError): - load_yaml_config("nonexistent.yaml") diff --git a/tests/test_configs.py b/tests/test_configs.py new file mode 100644 index 0000000..40ff35e --- /dev/null +++ b/tests/test_configs.py @@ -0,0 +1,42 @@ +import pytest +from hydra import compose, initialize +from omegaconf import OmegaConf +from brittle_star_project.configs.main_config import BrittleStarConfig +from brittle_star_project.configs.register_configs import register_configs + +# Registration must happen before composition to enable validation against schemas +register_configs() + + +def test_config_composition_centralized(): + """Test that the centralized configuration composes and validates correctly.""" + with initialize(version_base="1.3", config_path="../configs"): + # We compose the config; it follows main_config.yaml + cfg = compose(config_name="main_config", overrides=["architecture=centralized"]) + + # Merge with the dataclass class to get a structured DictConfig, + # then convert to a real dataclass instance to verify validation. + structured_cfg = OmegaConf.to_object(OmegaConf.merge(BrittleStarConfig, cfg)) + + # Basic assertions + assert "CentralizedConfig" in str(type(structured_cfg.architecture)) + assert isinstance(structured_cfg.ppo.learning_rate, float) + assert structured_cfg.ppo.learning_rate > 0 + + +def test_config_composition_decentralized(): + """Test that the decentralized configuration composes and validates correctly.""" + with initialize(version_base="1.3", config_path="../configs"): + cfg = compose(config_name="main_config", overrides=["architecture=decentralized"]) + + # Merge and convert to dataclass instance + structured_cfg = OmegaConf.to_object(OmegaConf.merge(BrittleStarConfig, cfg)) + + # Basic assertions + assert "DecentralizedConfig" in str(type(structured_cfg.architecture)) + assert isinstance(structured_cfg.ppo.learning_rate, float) + assert structured_cfg.ppo.learning_rate > 0 + + # Decentralized specifics + assert hasattr(structured_cfg.architecture, "message_passing_steps") + assert structured_cfg.architecture.message_passing_steps > 0 diff --git a/tests/test_morphology_render.py b/tests/test_morphology_render.py new file mode 100644 index 0000000..c127a01 --- /dev/null +++ b/tests/test_morphology_render.py @@ -0,0 +1,58 @@ +import jax +import os +import mujoco +from PIL import Image +from brittle_star_project.environment.env_config import EnvConfig, MorphologyConfig, ArenaConfig +from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper + + +def test_render_morphologies(): + base_dir = "runs/renders" + os.makedirs(base_dir, exist_ok=True) + + # --- 1. Full 5-Arm Morphology --- + morph_full = MorphologyConfig(segments_per_arm=[4, 4, 4, 4, 4]) + env_full = BrittleStarJaxEnvWrapper( + morphology=morph_full, arena=ArenaConfig(), env_config=EnvConfig(), num_envs=1 + ) + state_full = env_full.reset(seed=0) + + model_full = state_full.mj_model + data_full = state_full.mj_data + + # 1. Compute forward kinematics so geoms are correctly positioned + mujoco.mj_forward(model_full, data_full) + + # 2. Render using the environment's primary camera (camera=0) + renderer_full = mujoco.Renderer(model=model_full) + renderer_full.update_scene(data_full, camera=1) + pixels_full = renderer_full.render() + image_path = os.path.join(base_dir, "full_5_arm.png") + Image.fromarray(pixels_full).save(image_path) + print(f"Generated full morphology render: {image_path}") + + # --- 2. Partially Amputated Morphology --- + morph_amp = MorphologyConfig(segments_per_arm=[4, 0, 4, 2, 4]) + env_amp = BrittleStarJaxEnvWrapper( + morphology=morph_amp, arena=ArenaConfig(), env_config=EnvConfig(), num_envs=1 + ) + state_amp = env_amp.reset(seed=0) + + model_amp = state_amp.mj_model + data_amp = state_amp.mj_data + + # Compute forward kinematics + mujoco.mj_forward(model_amp, data_amp) + + renderer_amp = mujoco.Renderer(model=model_amp) + renderer_amp.update_scene(data_amp, camera=1) + pixels_amp = renderer_amp.render() + image_path = os.path.join(base_dir, "amputated_arm.png") + Image.fromarray(pixels_amp).save(image_path) + print(f"Generated amputated morphology render: {image_path}") + + print("Morphology render test successful!") + + +if __name__ == "__main__": + test_render_morphologies() diff --git a/tests/test_network_shapes.py b/tests/test_network_shapes.py new file mode 100644 index 0000000..32e9898 --- /dev/null +++ b/tests/test_network_shapes.py @@ -0,0 +1,67 @@ +import jax +import jax.numpy as jnp +from brittle_star_project.environment.padded_obs_wrapper import ( + compute_padding_masks, + pad_observations_batched, +) + +# We use Actor and OneDenseLayerMLP (as the critic) based on your mlps.py +from brittle_star_project.MLPs.mlps import Actor, OneDenseLayerMLP + + +def test_centralized_forward_pass_with_padding(): + batch_size = 2 + + # 1. Simulate Amputated Observation [4, 0, 4, 2, 4] -> 14 segments total + # 14 segments * 2 = 28 joints + amputated_obs = { + "joint_position": jnp.zeros((batch_size, 28)), + "joint_velocity": jnp.zeros((batch_size, 28)), + "segment_contact": jnp.zeros((batch_size, 14)), + } + + # 2. Pad Observation using the boolean scattering wrapper + 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, + ) + + # 40 + 40 + 20 = 100 dimensions + assert global_state.shape == (batch_size, 100), ( + f"Expected global state shape (2, 100), got {global_state.shape}" + ) + + # 4. Initialize dummy networks (40 actuators for the max morphology output) + actor = Actor(action_dim=40) + critic = OneDenseLayerMLP() # Acts as the centralized critic + + rng = jax.random.PRNGKey(0) + rng_a, rng_c = jax.random.split(rng) + + # Initialize Flax variables + actor_params = actor.init(rng_a, global_state) + critic_params = critic.init(rng_c, global_state) + + # 5. Forward Pass Assertions + action_mean, action_log_std = actor.apply(actor_params, global_state) + value = critic.apply(critic_params, global_state) + + assert action_mean.shape == (batch_size, 40), f"Actor mean shape mismatch: {action_mean.shape}" + assert action_log_std.shape == (40,), f"Actor log_std shape mismatch: {action_log_std.shape}" + assert value.shape == (batch_size, 1) or value.shape == (batch_size,), ( + f"Critic value shape mismatch: {value.shape}" + ) + + +if __name__ == "__main__": + test_centralized_forward_pass_with_padding()