diff --git a/src/brittle_star_project/dataclasses/PPOArgs.py b/src/brittle_star_project/dataclasses/PPOArgs.py index edaf318..6d8b9b4 100644 --- a/src/brittle_star_project/dataclasses/PPOArgs.py +++ b/src/brittle_star_project/dataclasses/PPOArgs.py @@ -8,7 +8,7 @@ class PPOArgs: """ # path to environment config file, if None, use default config - config_path: str | None = None + env_config_path: str | None = None # the name of this experiment exp_name: str = "brittle_star_ppo" diff --git a/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py b/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py index 259f1a5..7c86d51 100644 --- a/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py +++ b/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py @@ -8,7 +8,7 @@ from brittle_star_project import ( ArenaConfig, Backend, ) -from brittle_star_project.environment import from_json +from brittle_star_project.environment import from_file class BrittleStarJaxEnvWrapper: @@ -82,7 +82,7 @@ class BrittleStarJaxEnvWrapper: def from_config( config_path: str, num_envs: int, backend: Backend = Backend.MJX ) -> "BrittleStarJaxEnvWrapper": - morphology_cfg, arena_cfg, env_cfg = from_json(config_path) + morphology_cfg, arena_cfg, env_cfg = from_file(config_path) return BrittleStarJaxEnvWrapper( morphology_cfg, arena_cfg, env_cfg, num_envs=num_envs, backend=backend ) diff --git a/src/brittle_star_project/environment/__init__.py b/src/brittle_star_project/environment/__init__.py index f87eae5..aad8c35 100644 --- a/src/brittle_star_project/environment/__init__.py +++ b/src/brittle_star_project/environment/__init__.py @@ -1,4 +1,4 @@ -from .env_config import ArenaConfig, EnvConfig, MorphologyConfig, from_json +from .env_config import ArenaConfig, EnvConfig, MorphologyConfig, from_file from .env_types import Backend, Task from .env_wrapper import BrittleStarEnv, StepResult from .factory import BrittleStarEnvFactory @@ -12,5 +12,5 @@ __all__ = [ "BrittleStarEnv", "StepResult", "BrittleStarEnvFactory", - "from_json", + "from_file", ] diff --git a/src/brittle_star_project/environment/env_config.py b/src/brittle_star_project/environment/env_config.py index 6f55f35..1420e19 100644 --- a/src/brittle_star_project/environment/env_config.py +++ b/src/brittle_star_project/environment/env_config.py @@ -50,10 +50,16 @@ class EnvConfig: light_perlin_noise_scale: int = 0 -def from_json(path: str) -> tuple[MorphologyConfig, ArenaConfig, EnvConfig]: +def from_file(path: str) -> tuple[MorphologyConfig, ArenaConfig, EnvConfig]: + """Load configurations from a JSON or YAML file.""" with open(path, "r") as f: - config_json = json.load(f) - morphology = MorphologyConfig(**config_json.get("morphology", {})) - arena = ArenaConfig(**config_json.get("arena", {})) - env = EnvConfig(**config_json.get("env", {})) + if path.endswith(".yaml") or path.endswith(".yml"): + import yaml + config_dict = yaml.safe_load(f) + else: + config_dict = json.load(f) + + morphology = MorphologyConfig(**config_dict.get("morphology", {})) + arena = ArenaConfig(**config_dict.get("arena", {})) + env = EnvConfig(**config_dict.get("env", {})) return morphology, arena, env diff --git a/src/train.py b/src/train.py index 3c9bdea..3995a36 100644 --- a/src/train.py +++ b/src/train.py @@ -31,11 +31,11 @@ def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: ) -def make_env(config_path: str | None, num_envs: int) -> Callable: +def make_env(env_config_path: str | None, num_envs: int) -> Callable: def thunk(): - if config_path is None: + if env_config_path is None: return BrittleStarJaxEnvWrapper.default(num_envs=num_envs) - return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs) + return BrittleStarJaxEnvWrapper.from_config(env_config_path, num_envs=num_envs) return thunk @@ -88,7 +88,7 @@ def train(args: PPOArgs): print(f"Running on device: {device}") print("Creating the environment...") - env = make_env(config_path=args.config_path, num_envs=args.num_envs)() + env = make_env(env_config_path=args.env_config_path, num_envs=args.num_envs)() print(f"Environment: {env}") episode_stats = EpisodeStatistics(