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