1
Fork 0

chore: draft new configs

This commit is contained in:
Tibo De Peuter 2026-04-14 21:53:52 +02:00
parent 67f620d599
commit 157e061f86
Signed by: tdpeuter
SSH key fingerprint: SHA256:u/h/LVoqKF1Iz02uOyxe6hcjmoZASCGV2HM0TG9ZMoU
6 changed files with 113 additions and 0 deletions

View 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

View 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

View 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

View 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)

View 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)

View 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 = ""