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."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -93,6 +93,8 @@ class UnifiedLogger:
|
|||
entity: Optional[str] = None,
|
||||
base_dir: str = "runs",
|
||||
use_wandb: bool = True,
|
||||
upload_final_model: bool = False,
|
||||
upload_checkpoints: bool = False,
|
||||
save_code: bool = True,
|
||||
log_level: int = logging.INFO,
|
||||
):
|
||||
|
|
@ -110,6 +112,8 @@ class UnifiedLogger:
|
|||
self.run_name = run_name
|
||||
self.config = config
|
||||
self.use_wandb = use_wandb
|
||||
self.upload_final_model = upload_final_model
|
||||
self.upload_checkpoints = upload_checkpoints
|
||||
self.wandb_available = False
|
||||
self.wandb_run = None
|
||||
self.is_interactive = sys.stdout.isatty()
|
||||
|
|
@ -327,7 +331,7 @@ class UnifiedLogger:
|
|||
self.info(f"Checkpoint saved: {checkpoint_path}")
|
||||
|
||||
# Log to WandB as artifact
|
||||
if self.wandb_run is not None:
|
||||
if self.wandb_run is not None and self.upload_checkpoints:
|
||||
try:
|
||||
import wandb
|
||||
|
||||
|
|
@ -363,7 +367,7 @@ class UnifiedLogger:
|
|||
self.info(f"Final model saved: {final_model_path}")
|
||||
|
||||
# Log to WandB
|
||||
if self.wandb_run is not None:
|
||||
if self.wandb_run is not None and self.upload_final_model:
|
||||
try:
|
||||
import wandb
|
||||
|
||||
|
|
|
|||
Reference in a new issue