Class Trainer
java.lang.Object
io.github.kirstenali.deepj.training.Trainer
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.
-
Nested Class Summary
Nested ClassesModifier and TypeClassDescriptionstatic interfacestatic interface -
Constructor Summary
Constructors -
Method Summary
Modifier and TypeMethodDescriptionTrain until maxSteps or until EMA loss goes below targetEmaLoss (if provided).train(int maxSteps, int batchSize, int logEvery, float emaBeta, Float targetEmaLoss, int releaseEverySteps) 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).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).floattrainStep(int batchSize)
-
Constructor Details
-
Trainer
-
-
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 <= 0disables periodic release, but final release still runs.
-