ml4co_kit.learning.train

Trainer for ML4CO models.

Classes

Checkpoint([dirpath, monitor, ...])

Save the model periodically by monitoring a quantity.

Logger([name, project, entity, save_dir, ...])

Logger for Wandb.

MetricProgressBar([metric_precision])

Display progress-bar metrics with fixed decimal precision.

Trainer(model[, logger, wandb_logger_name, ...])

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: ModelCheckpoint

Save 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: WandbLogger

Logger for Wandb.

class ml4co_kit.learning.train.MetricProgressBar(metric_precision: int = 4, **kwargs)[source]

Bases: TQDMProgressBar

Display 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: Trainer

Trainer for ML4CO models.

model_test()[source]
model_train()[source]