diff --git a/src/brittle_star_project/configs/config_experiment.py b/src/brittle_star_project/configs/config_experiment.py new file mode 100644 index 0000000..69ceca4 --- /dev/null +++ b/src/brittle_star_project/configs/config_experiment.py @@ -0,0 +1,9 @@ +from dataclasses import dataclass + + +@dataclass +class ExperimentConfig: + exp_name: str = "brittle_star_ppo" + seed: int = 1 + torch_deterministic: bool = True + cuda: bool = True diff --git a/src/brittle_star_project/configs/config_networks.py b/src/brittle_star_project/configs/config_networks.py new file mode 100644 index 0000000..347936b --- /dev/null +++ b/src/brittle_star_project/configs/config_networks.py @@ -0,0 +1,23 @@ +from dataclasses import dataclass, field +from typing import List, Optional + + +@dataclass +class LayerConfig: + hidden_dims: List[int] = field(default_factory=lambda: [64, 64]) + activation: str = "relu" + + +@dataclass +class MLPConfig: + actor: LayerConfig = field(default_factory=LayerConfig) + critic: LayerConfig = field(default_factory=LayerConfig) + sensor: LayerConfig = field(default_factory=LayerConfig) + feature_extractor: LayerConfig = field(default_factory=LayerConfig) + + +@dataclass +class NetworksConfig: + mlp: MLPConfig = field(default_factory=MLPConfig) + # Future placeholder for message_passing + # message_passing: Optional[MessagePassingConfig] = None diff --git a/src/brittle_star_project/configs/config_ppo.py b/src/brittle_star_project/configs/config_ppo.py new file mode 100644 index 0000000..d0a29bf --- /dev/null +++ b/src/brittle_star_project/configs/config_ppo.py @@ -0,0 +1,22 @@ +from dataclasses import dataclass +from typing import Optional + + +@dataclass +class PPOConfig: + learning_rate: float = 2.5e-4 + total_timesteps: int = 10000000 + num_envs: int = 100 + num_steps: int = 128 + anneal_lr: bool = True + gamma: float = 0.99 + gae_lambda: float = 0.95 + num_minibatches: int = 4 + update_epochs: int = 4 + norm_adv: bool = True + clip_coef: float = 0.1 + clip_vloss: bool = True + ent_coef: float = 0.01 + vf_coef: float = 0.5 + max_grad_norm: float = 0.5 + target_kl: Optional[float] = None diff --git a/src/brittle_star_project/configs/main_config.py b/src/brittle_star_project/configs/main_config.py new file mode 100644 index 0000000..941a5f2 --- /dev/null +++ b/src/brittle_star_project/configs/main_config.py @@ -0,0 +1,18 @@ +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_networks import NetworksConfig +from brittle_star_project.environment.env_config import MorphologyConfig, ArenaConfig, EnvConfig + + +@dataclass +class BrittleStarConfig: + experiment: ExperimentConfig = field(default_factory=ExperimentConfig) + logging: LoggingConfig = field(default_factory=LoggingConfig) + ppo: PPOConfig = field(default_factory=PPOConfig) + networks: NetworksConfig = field(default_factory=NetworksConfig) + 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/configs/register_configs.py b/src/brittle_star_project/configs/register_configs.py new file mode 100644 index 0000000..22d93ca --- /dev/null +++ b/src/brittle_star_project/configs/register_configs.py @@ -0,0 +1,27 @@ +from hydra.core.config_store import ConfigStore + +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_networks import NetworksConfig, MLPConfig, LayerConfig +from brittle_star_project.environment.env_config import MorphologyConfig, ArenaConfig, EnvConfig +from brittle_star_project.configs.main_config import BrittleStarConfig + + +def register_configs(): + """Register dataclasses with Hydra's ConfigStore.""" + cs = ConfigStore.instance() + + # Store the main config schema + cs.store(name="brittle_star_config", node=BrittleStarConfig) + + # Store individual structured configs for validation + cs.store(group="experiment", name="base_experiment", node=ExperimentConfig) + cs.store(group="logging", name="base_logging", node=LoggingConfig) + cs.store(group="ppo", name="base_ppo", node=PPOConfig) + cs.store(group="networks", name="base_networks", node=NetworksConfig) + + # Environment configs (reusing existing environment config objects) + cs.store(group="morphology", name="base_morphology", node=MorphologyConfig) + cs.store(group="arena", name="base_arena", node=ArenaConfig) + cs.store(group="environment", name="base_environment", node=EnvConfig) diff --git a/src/experiment_logger/config_logger.py b/src/experiment_logger/config_logger.py new file mode 100644 index 0000000..5d39b8e --- /dev/null +++ b/src/experiment_logger/config_logger.py @@ -0,0 +1,14 @@ +from dataclasses import dataclass +from typing import Optional + + +@dataclass +class LoggingConfig: + track: bool = False + wandb_project_name: str = "PPO-Modularity" + wandb_entity: Optional[str] = "SEL3-2026-Groep-4" + capture_video: bool = False + save_model: bool = True + checkpoint_frequency: int = 100 + upload_model: bool = False + hf_entity: str = ""