Class Trainer

java.lang.Object
io.github.kirstenali.deepj.training.Trainer

public final class Trainer extends Object
A small, reusable training loop wrapper.

This library supports different model/data shapes (e.g. supervised Tensor->Tensor models, and causal language models that operate on token ids). Rather than duplicating full trainers, Trainer delegates a single training step to a pluggable Trainer.StepFunction.

  • Constructor Details

  • Method Details

    • trainStep

      public float trainStep(int batchSize)
    • train

      public TrainingResult train(int maxSteps, int batchSize, int logEvery, float emaBeta, Float targetEmaLoss)
      Train until maxSteps or until EMA loss goes below targetEmaLoss (if provided). Uses the default periodic backend release cadence.
    • train

      public TrainingResult train(int maxSteps, int batchSize, int logEvery, float emaBeta, Float targetEmaLoss, int releaseEverySteps)
    • train

      public TrainingResult train(int maxSteps, int batchSize, int logEvery, float emaBeta, Float targetEmaLoss, Trainer.StepHook stepHook)
      Train until maxSteps or until EMA loss goes below targetEmaLoss (if provided). Uses the default periodic backend release cadence.
    • train

      public TrainingResult train(int maxSteps, int batchSize, int logEvery, float emaBeta, Float targetEmaLoss, int releaseEverySteps, Trainer.StepHook stepHook)
      Train until maxSteps or until EMA loss goes below targetEmaLoss (if provided). releaseEverySteps <= 0 disables periodic release, but final release still runs.