Transductive Trainer
- class hivegraph.engine.transductive.TransductiveTrainer(model: Module, dataset: Dataset, device: str = 'cpu', random_state: int = 42)[source]
Bases:
BaseTrainer- fit(epochs: int, batch_size: int, optimizer: Optimizer, verbose: bool, log_to_wandb: bool) 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