ml4co_kit.learning.train
Trainer for ML4CO models.
Classes
|
Save the model periodically by monitoring a quantity. |
|
Logger for Wandb. |
|
Display progress-bar metrics with fixed decimal precision. |
|
Trainer for ML4CO models. |
- class ml4co_kit.learning.train.Checkpoint(dirpath: str = 'wandb/checkpoints', monitor: str = 'val/loss', every_n_epochs: int = 1, every_n_train_steps=None, filename=None, save_top_k: int = -1, mode: str = None)[source]
Bases:
ModelCheckpointSave the model periodically by monitoring a quantity. Look at the above link for more detailed information.
- class ml4co_kit.learning.train.Logger(name: str = 'wandb', project: str = 'project', entity: str | None = None, save_dir: str = 'log', id: str | None = None, resume_id: str | None = None)[source]
Bases:
WandbLoggerLogger for Wandb.
- class ml4co_kit.learning.train.MetricProgressBar(metric_precision: int = 4, **kwargs)[source]
Bases:
TQDMProgressBarDisplay progress-bar metrics with fixed decimal precision.
- get_metrics(trainer: Trainer, pl_module: LightningModule) dict[str, Union[int, str, float]][source]
Combines progress bar metrics collected from the trainer with standard metrics from get_standard_metrics. Implement this to override the items displayed in the progress bar.
Here is an example of how to override the defaults:
def get_metrics(self, trainer, model): # don't show the version number items = super().get_metrics(trainer, model) items.pop("v_num", None) return items
- Return:
Dictionary with the items to be displayed in the progress bar.
- class ml4co_kit.learning.train.Trainer(model: Module, logger: Logger | None = None, wandb_logger_name: str = 'wandb', resume_id: str | None = None, ckpt_save_path: str | None = None, ckpt_monitor: str = 'val/loss', save_top_k: int = -1, mode: str = 'min', ckpt_every_n_epochs: int = 1, ckpt_every_n_train_steps: int | None = None, ckpt_filename: str = None, accelerator: str = 'auto', strategy: str | Strategy = None, devices: List[int] | str | int = 'auto', fp16: bool = False, max_epochs: int = 100, max_steps: int = -1, val_check_interval: int | None = None, log_every_n_steps: int | None = 50, gradient_clip_val: int = 1, inference_mode: bool = False, reload_dataloaders_every_n_epochs: int = 0, progress_bar_refresh_rate: int = 20, metric_precision: int = 4, disable_profiling_executor: bool = True, ckpt_path: str | None = None, weight_path: str | None = None)[source]
Bases:
TrainerTrainer for ML4CO models.