Classification Trainer
- class hivegraph.engine.classification.ClassificationTrainer(model: ~torch.nn.modules.module.Module, dataset: ~torch_geometric.data.dataset.Dataset, criterion: ~typing.Callable = <function nll_loss>, device: str = 'cpu', num_folds: int = 10, random_state: int = 42, test_metric: str | None = None)[source]
Bases:
BaseTrainerCustom Trainer Class for Classification Tasks
- fit(epochs: int, batch_size: int, optimizer: Optimizer, verbose: bool = True, log_to_wandb: bool = False) None[source]
Trains the model on the given dataset
- Parameters:
epochs (int) – Number of epochs
batch_size (int) – batch size
optimizer (torch.optim.Optimizer) – Optimizer
verbose (bool) – Whether to print stats, defaults to True
log_to_wandb (bool) – Whether to log stats to wandb, defaults to False
- train_step(train_dataloder: Iterable, optimizer: Optimizer) float[source]
Performs a single training step
- Parameters:
train_dataloder (Iterable) – Training Dataloader
optimizer (torch.optim.Optimizer) – Optimizer
- Returns:
Training Loss
- Return type: