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

train_step(train_dataloder: Iterable, optimizer: Optimizer, **kwargs) float[source]

Performs a single training step

Parameters:
Returns:

Training Loss

Return type:

float