Keras Stochastic Weight Averaging Save

Keras callback function for stochastic weight averaging

Project README

Stochastic Weight Averaging with Keras callback function

Stochastic Weight Averaging following paper Averaging Weights Leads to Wider Optima and Better Generalization

The file swa.py contains an implementation for stochastic weight averaging (SWA) with a constant learning rate for a user defined amount of epochs.

Callback is instantiated with filename for saving the final weights of the model after SWA and the number of epochs to average.

Example

The total number of training epochs 150, SWA to start from epoch 140 to average last 10 epochs.

from swa import SWA

# specify number of training epochs
number_of_epochs = 150

# specify the start epoch of stochastic weight averaging
swa = SWA(140, filepath = None)

# call SWA during model fitting
model.fit(..., callbacks = [swa])
Open Source Agenda is not affiliated with "Keras Stochastic Weight Averaging" Project. README Source: kristpapadopoulos/keras-stochastic-weight-averaging

Open Source Agenda Badge

Open Source Agenda Rating