feat: configure checkpoints saving
This commit is contained in:
parent
4d0f729aee
commit
37a4b59e04
10 changed files with 65 additions and 36 deletions
|
|
@ -5,10 +5,29 @@ from typing import Optional
|
|||
@dataclass
|
||||
class LoggingConfig:
|
||||
track: bool = False
|
||||
wandb_project_name: str = "PPO-Modularity"
|
||||
wandb_project_name: str = "default-project"
|
||||
wandb_entity: Optional[str] = "SEL3-2026-Groep-4"
|
||||
capture_video: bool = False
|
||||
save_model: bool = True
|
||||
|
||||
# Local Saving
|
||||
save_model: bool = True # Final model
|
||||
save_checkpoints: bool = True # Intermediate checkpoints
|
||||
checkpoint_frequency: int = 100
|
||||
upload_model: bool = False
|
||||
|
||||
# Remote Uploading (WandB Artifacts)
|
||||
upload_final_model: bool = False
|
||||
upload_checkpoints: bool = False
|
||||
|
||||
hf_entity: str = ""
|
||||
|
||||
def __post_init__(self):
|
||||
if self.upload_final_model and not (self.track and self.save_model):
|
||||
raise ValueError(
|
||||
"Configuration Error: 'upload_final_model' is True, but it requires "
|
||||
"both 'track' and 'save_model' to also be True."
|
||||
)
|
||||
if self.upload_checkpoints and not (self.track and self.save_checkpoints):
|
||||
raise ValueError(
|
||||
"Configuration Error: 'upload_checkpoints' is True, but it requires "
|
||||
"both 'track' and 'save_checkpoints' to also be True."
|
||||
)
|
||||
|
|
|
|||
Reference in a new issue