chore: draft new configs
This commit is contained in:
parent
67f620d599
commit
157e061f86
6 changed files with 113 additions and 0 deletions
9
src/brittle_star_project/configs/config_experiment.py
Normal file
9
src/brittle_star_project/configs/config_experiment.py
Normal file
|
|
@ -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
|
||||
23
src/brittle_star_project/configs/config_networks.py
Normal file
23
src/brittle_star_project/configs/config_networks.py
Normal file
|
|
@ -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
|
||||
22
src/brittle_star_project/configs/config_ppo.py
Normal file
22
src/brittle_star_project/configs/config_ppo.py
Normal file
|
|
@ -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
|
||||
18
src/brittle_star_project/configs/main_config.py
Normal file
18
src/brittle_star_project/configs/main_config.py
Normal file
|
|
@ -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)
|
||||
27
src/brittle_star_project/configs/register_configs.py
Normal file
27
src/brittle_star_project/configs/register_configs.py
Normal file
|
|
@ -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)
|
||||
14
src/experiment_logger/config_logger.py
Normal file
14
src/experiment_logger/config_logger.py
Normal file
|
|
@ -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 = ""
|
||||
Reference in a new issue