polymon.exp

Pipeline

Trainer

class polymon.exp.train.Trainer(out_dir: str, model: ModelWrapper, lr: float, num_epochs: int, logger: Logger, device: device = 'cuda', early_stopping_patience: int = 10)[source]

Bases: object

Trainer for the model.

Parameters:
  • out_dir (str) – The directory to save the model and results.

  • model (nn.Module) – The model to train.

  • lr (float) – The learning rate.

  • num_epochs (int) – The number of epochs to train.

  • logger (logging.Logger) – The logger. If not provided, a logger will be created in the out_dir.

  • ema_decay (float) – The decay rate for the EMA. If 0, EMA will not be used. Default is 0.

  • device (torch.device) – The device to train on. Default is cuda.

  • early_stopping_patience (int) – The number of epochs to wait before stopping the training. Default is 10.

build_optimizer() → Optimizer[source]

Build the optimizer.

Returns:

The optimizer.

Return type:

torch.optim.Optimizer

eval(loader: DataLoader, label: str) → Dict[str, float][source]

Evaluate the model on the given data loader.

Parameters:
  • loader (DataLoader) – The data loader.

  • metrics (List[Literal['mae', 'r2']]) – The metrics to evaluate. If None, all metrics will be evaluated.

Returns:

The metrics and their values.

Return type:

Dict[str, float]

train(train_loader: DataLoader, val_loader: DataLoader, test_loader: DataLoader | None = None, label: str = 'Rg', trial: Trial | None = None)[source]

Train the model.

Parameters:
  • train_loader (DataLoader) – The training data loader.

  • val_loader (DataLoader) – The validation data loader.

  • test_loader (DataLoader) – The test data loader.

  • label (str) – The label to train on.

  • trial (optuna.Trial) – The trial object. If provided, the model will be stopped training when the pruning is triggered.

  • n_fold (int) – The number of folds. This is used to report the

train_step(ith_epoch: int, train_loader: DataLoader, val_loader: DataLoader, optimizer: Optimizer, label: str) → float[source]

Train the model for one epoch.

Parameters:
  • ith_epoch (int) – The current epoch.

  • train_loader (DataLoader) – The training data loader.

  • val_loader (DataLoader) – The validation data loader.

  • optimizer (torch.optim.Optimizer) – The optimizer.

Returns:

The F1 score on the validation set.

Return type:

float