diff --git a/src/brittle_star_project/configs/main_config.py b/src/brittle_star_project/configs/main_config.py index cd54473..2164a30 100644 --- a/src/brittle_star_project/configs/main_config.py +++ b/src/brittle_star_project/configs/main_config.py @@ -3,7 +3,7 @@ from dataclasses import dataclass, field from experiment_logger.config_logger import LoggingConfig from brittle_star_project.configs.config_experiment import ExperimentConfig from brittle_star_project.configs.config_ppo import PPOConfig -from brittle_star_project.configs.config_architecture import ArchitectureConfig, CentralizedConfig +from brittle_star_project.configs.config_architecture import ArchitectureConfig from brittle_star_project.environment.env_config import MorphologyConfig, ArenaConfig, EnvConfig diff --git a/src/brittle_star_project/environment/padded_obs_wrapper.py b/src/brittle_star_project/environment/padded_obs_wrapper.py index 8c07248..95f310b 100644 --- a/src/brittle_star_project/environment/padded_obs_wrapper.py +++ b/src/brittle_star_project/environment/padded_obs_wrapper.py @@ -5,23 +5,28 @@ a constant size regardless of how many segments are amputated. This wrapper pads the observation dictionary values with zeros using spatial insertion so that the flattened observation maintains the correct physical mapping to the neural network. """ + from __future__ import annotations from typing import Any import jax.numpy as jnp # Observation keys whose size scales with the number of joints (2 per segment). -_JOINT_SCALED_KEYS = frozenset({ - "joint_position", - "joint_velocity", - "joint_actuator_force", - "actuator_force", -}) +_JOINT_SCALED_KEYS = frozenset( + { + "joint_position", + "joint_velocity", + "joint_actuator_force", + "actuator_force", + } +) # Observation keys whose size scales with the number of segments (1 per segment). -_SEGMENT_SCALED_KEYS = frozenset({ - "segment_contact", -}) +_SEGMENT_SCALED_KEYS = frozenset( + { + "segment_contact", + } +) def compute_padding_masks( diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 1e49cc2..b3321b2 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -528,10 +528,7 @@ class PPOTrainer: ) if getattr(self.cfg.experiment, "debug_sanity", False): - self.logger.info( - "\n[SANITY CHECK] Successfully completed 1 epoch of data collection and gradient updates." - ) - self.logger.info("[SANITY CHECK] Gradients flowed without NaN. Exiting gracefully.") + self.logger.info("\n[SANITY CHECK] Successfully completed 1 epoch") break if self.logging_cfg.save_model: diff --git a/src/experiment_logger/config_utils.py b/src/experiment_logger/config_utils.py index 6017041..80d1748 100644 --- a/src/experiment_logger/config_utils.py +++ b/src/experiment_logger/config_utils.py @@ -1,7 +1,6 @@ """Configuration utilities for loading YAML configs and merging with CLI args.""" import os -import sys from typing import Dict, Any, Type, TypeVar import yaml from dataclasses import fields, is_dataclass diff --git a/src/experiment_logger/unified_logger.py b/src/experiment_logger/unified_logger.py index a02bdb1..2f5d06f 100644 --- a/src/experiment_logger/unified_logger.py +++ b/src/experiment_logger/unified_logger.py @@ -6,9 +6,7 @@ This logger ensures all experimental data is preserved by writing to: 3. stdout (for real-time monitoring) """ -import datetime import logging -import subprocess import yaml import sys import time diff --git a/tests/test_configs.py b/tests/test_configs.py index 40ff35e..84a7cde 100644 --- a/tests/test_configs.py +++ b/tests/test_configs.py @@ -1,4 +1,3 @@ -import pytest from hydra import compose, initialize from omegaconf import OmegaConf from brittle_star_project.configs.main_config import BrittleStarConfig diff --git a/tests/test_morphology_render.py b/tests/test_morphology_render.py index d8648d8..423a575 100644 --- a/tests/test_morphology_render.py +++ b/tests/test_morphology_render.py @@ -1,4 +1,3 @@ -import jax import pytest import os import mujoco